Repository navigation
ggml-hrx: fold the MUL_MAT_ID WMMA f16 accumulator into f32 at each 256-wide K tile (engine#284) - #85
Merged
Conversation
…wide dot The wave64 16x16x16 f16 MMA that the rdna3_5 descriptor set offers accumulates in f16; the f32-accumulator variant is rejected (matrix constraint wave_size). Re-rounding the running sum to f16 at every 16-wide step biases it upward once the sum exceeds ~32: a constant 0.0234375 weight gives +17% at K=2816 and +22% at K=4096 instead of the exact K x 0.0234375. Run each MMA from a zero fragment and fold the single 16-wide dot into f32 accumulators, which never chain in f16. No lane/LDS layout changes; the final fptrunc into the existing f16 result fragment is unchanged.
bong-water-water-bong
pushed a commit
that referenced
this pull request
Oct 5, 2026
Conflict in motifs/mul_mat_id_f32_f32_wmma_core.loom (#85 vs this PR): kept this PR's f32 accumulation (vector<4xf32> MMA results), which replaces #85's f16 MMA folded into f32 carries per 256-wide K tile. The other three files merged without conflict. Measured on strixhalo gfx1151 (performance mode), base 86eee89 vs this merge: - The pinned JIT (hrx-system 51b1739) compiles the f32 accumulator: v_wmma_f32_16x16x16_f16 in all 93 mul_mat and 72 mul_mat_id specializations that test-backend-ops compiled (base: v_wmma_f16_16x16x16_f16). - Known-answer cases (MXFP4 0x77, inputs 1.0, exact K x 0.0234375, K = 1152..5120, 6/17-token and 1-token/1-route tiles, iree-test-loom --sanitizer=access, 2 repeats): this merge 84/84; base 28/84 (mul_mat_id passes with #85, the other three kernels fail above K = 1152). - test-backend-ops -b HRX0 MUL_MAT 322/322, MUL_MAT_ID 108/108, MUL_MAT_VEC_FUSION 45/45, 2 runs each. - gpt-oss-20b KLD vs CPU (c512 b512, 8 chunks): 0.029296 -> 0.026932. - llama-bench pp512 (medians of 3): gpt-oss-20b 1045.8 -> 995.7 (-4.8%), Qwen3.8-27B UD-Q4_K_XL 355.3 -> 367.0 (+3.3%), Qwen3-Coder-30B 2060.8 -> 2058.8; tg128 unchanged. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
engine#284 — the HRX WMMA cores accumulate in f16, and that drift is now fixed at the K-tile boundary.
Root. The
amdgpu-rdna3-5descriptor set offers only an f16-accumulator 16x16x16 MMA; an f32 accumulator is rejected (matrix constraint 'wave_size'). Chaining that MMA over the whole K dimension re-rounds the running sum to f16 at every 16-wide step and drifts upward: a constant0.0234375weight gives +17% at K=2816 instead of the exactK × 0.0234375.Fix (
ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_wmma_core.loom). Keep the f16 fragment for exactly one 256-wide K tile, thenvector.extfeach tile's partial into f32 carries (vector<8xf32>) that carry the running sum across tiles. The f16 fragment therefore never spans more than one tile, and the accumulation is exact.Gate (
iree-test-loom, MXFP4 all-0x77= 0.0234375, output 96, 1 token / 1 route,--sanitizer=access, gfx1151):check.expect.close ... atol(0) rtol(0)— bit-exact to the exact f32 result.Ref engine#284. The same f16 accumulation class is the standing hypothesis for the HRX NaN seen above ~4700 tokens (engine#300) and is worth re-measuring at model level with the experts on HRX.