From bc802ffe76937191b3566dac894399c0f2efa01e Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 05:32:43 +0000 Subject: [PATCH 1/8] gemm_ar: MXFP8 (32-wide ue8m0), the quantisation V4.1-Flash actually ships `GemmAllReduceOp` served 1x128 / 128x128 fp8 block scales. DeepSeek-V4.1-Flash quantises 32-wide ue8m0 on both operands (`weight_block_size [32, 32]`, `scale_fmt "ue8m0"`), which is a different operand contract and a different kernel, not a parameter of the old one. CDNA4's `v_mfma_scale_f32_16x16x128_f8f6f4` takes the ue8m0 scales as *instruction operands*, so this form needs no dequantisation arithmetic at all -- no second accumulator, no running rescale. `_BlockScaleK`'s promote/rescale chain exists precisely because a 128-wide block cannot be expressed that way. Two layout results came out of making it fast, both about addresses rather than bytes, and both carry the counters that settle them: * **Scales must be K-block major.** A block group's sixteen lanes want sixteen consecutive rows of one K block. K-block major coalesces them; the quantiser's own row-major `[M, K/32]` spreads them K/32 bytes apart -- same instruction count, 3.2x the cache accesses, +50-67% on the GEMM. * **Four M tiles pack into one scale dword.** `opsel_b` on the scaled MFMA is an atom-time attribute naming which byte of a 32-bit scale operand to read, so one load serves four tiles and the select is free. -29.1% `SQ_INSTS_VMEM` with VALU flat, worth -5 to -7.5%. A quantiser can emit the layout directly, at no extra traffic. `_PinnedLaunch` resolves each compiled kernel's dispatch once. `JitFunction` re-derived its cache key per call -- an inspect bind, a 35-global snapshot, a drift check -- which measured 181.6us of Python per fused call against 5.9us of actual `hipModuleLaunchKernel`, enough to make the layer host-bound. It verifies the argument tuple element by element before pinning and falls back to the normal path on any surprise, so a FlyDSL change costs performance, not correctness. The fp8 all-gather leg gets its two knobs settled by measurement: pull it over LSA rather than pushing over SDMA (2.0-4.0%), and do not fuse the quantise into the reduce (1.2-1.9%). Both hold at every M and in all three modes. `_gemm_a8w8_8wave.py` is vendored from aiter and now diverges in two places rather than one: `Mfma16x16x128` grew the scaled form of the atom. Both new arguments default to off and the unscaled path is byte-for-byte what it was, so a re-vendor stays a merge. The file's header says so. Co-Authored-By: Claude Opus 5 (1M context) --- python/mori/ops/gemm_ar/_gemm_a8w8_8wave.py | 69 ++++- python/mori/ops/gemm_ar/_shuffle.py | 44 +++ python/mori/ops/gemm_ar/kernels_fused.py | 317 ++++++++++++++++++-- python/mori/ops/gemm_ar/op.py | 252 ++++++++++++++-- tests/python/cco/test_gemm_ar.py | 94 +++++- 5 files changed, 706 insertions(+), 70 deletions(-) diff --git a/python/mori/ops/gemm_ar/_gemm_a8w8_8wave.py b/python/mori/ops/gemm_ar/_gemm_a8w8_8wave.py index c4e09cf93..a5c9062ff 100644 --- a/python/mori/ops/gemm_ar/_gemm_a8w8_8wave.py +++ b/python/mori/ops/gemm_ar/_gemm_a8w8_8wave.py @@ -34,9 +34,16 @@ ``kernels_fused.py`` builds its fused GEMM out of these pieces, and mori does not depend on aiter, so the file is copied rather than imported. -Verbatim except for one thing: ``split_row_major_2d`` is inlined below instead -of imported from ``mfma_preshuffle_pipeline`` -- three lines against a -1230-line module. +Two deliberate divergences from upstream, and nothing else: + +* ``split_row_major_2d`` is inlined below instead of imported from + ``mfma_preshuffle_pipeline`` -- three lines against a 1230-line module. +* ``Mfma16x16x128`` grew the scaled form of the atom: ``opsel_b_per_tile`` on + the constructor and ``scale_a`` / ``scale_b`` on ``call``. Both default to + off and the unscaled path is byte-for-byte what it was, so a future re-vendor + is a merge rather than a rewrite. The 32-wide ue8m0 GEMM needs it and cannot + reach it any other way -- the scale is an operand of the instruction, not + arithmetic around it. Note that ``kernels_fused.py`` does *not* call the pipeline in this file. It re-implements the main loop so it can fuse a scatter into the epilogue, and in @@ -281,10 +288,32 @@ def store(self, c_frag, base_row, base_col): class Mfma16x16x128: - def __init__(self, n_tiles_a, n_tiles_b): + """``opsel_b_per_tile`` packs the B-tile scales four-to-a-dword. + + The scale operand is a 32-bit register and ``opsel_b`` is an atom-time + attribute naming which of its four bytes the instruction reads. So if the + four B tiles' ue8m0 bytes are packed into one dword, one load serves all + four and the byte select costs nothing -- it is the instruction's own field, + not a shift. That needs one atom per tile, since the attribute is baked in. + CK does the same thing (``preShuffleScaleBuffer_gfx950``'s ``MNXdlPack``). + """ + + def __init__(self, n_tiles_a, n_tiles_b, *, opsel_b_per_tile=False): self.atom = fx.make_mma_atom( fx.rocdl.cdna4.MFMA_Scale(16, 16, 128, fx.Float8E4M3FN) ) + self.atoms_b = [ + fx.make_mma_atom( + fx.rocdl.cdna4.MFMA_Scale( + 16, + 16, + 128, + fx.Float8E4M3FN, + opsel_b=(j if opsel_b_per_tile else 0), + ) + ) + for j in range_constexpr(n_tiles_b) + ] self.zero_value = Vec.filled(4, 0.0, fx.Float32) self.n_tiles_a = n_tiles_a self.n_tiles_b = n_tiles_b @@ -309,10 +338,27 @@ def _do_mma(self, a, b, c): fx.gemm(self.atom, c_frag, a_frag, b_frag, c_frag) return c_frag.load().ir_value() - def call(self, a, b, c, *, set_prio=True): + def call(self, a, b, c, *, set_prio=True, scale_a=None, scale_b=None): + """``scale_a`` / ``scale_b``, when given, are per-tile ue8m0 operands. + + ``v_mfma_scale_f32_16x16x128_f8f6f4`` carries one ue8m0 scale per 32 K + per row and gathers the four of them *across lanes*: lane ``16*s + r`` + supplies block ``s`` of row ``r`` at op_sel 0. So one MFMA consumes a + whole K=128 step with its four 32-blocks already dequantised in + hardware -- there is no promote arithmetic to schedule, which is the + entire reason the 32-wide form is cheaper than ``_BlockScaleK``'s + 128-wide rescale chain. + + Leaving both None keeps the unscaled call byte-for-byte as it was. + """ assert len(a) == self.n_tiles_a assert len(b) == self.n_tiles_b assert len(c) == self.n_tiles_a * self.n_tiles_b + scaled = scale_a is not None + assert scaled == (scale_b is not None), "pass both scales or neither" + if scaled: + assert len(scale_a) == self.n_tiles_a + assert len(scale_b) == self.n_tiles_b a_frags = [ self._make_operand_frag(a[idx]) for idx in range_constexpr(self.n_tiles_a) @@ -329,7 +375,18 @@ def call(self, a, b, c, *, set_prio=True): for i in range_constexpr(self.n_tiles_a): for j in range_constexpr(self.n_tiles_b): cf = c_frags[self.idx(i, j)] - fx.gemm(self.atom, cf, a_frags[i], b_frags[j], cf) + if const_expr(scaled): + fx.gemm( + self.atoms_b[j], + cf, + a_frags[i], + b_frags[j], + cf, + scale_a=scale_a[i], + scale_b=scale_b[j], + ) + else: + fx.gemm(self.atom, cf, a_frags[i], b_frags[j], cf) if const_expr(set_prio): rocdl.s_setprio(0) rocdl.s_barrier() diff --git a/python/mori/ops/gemm_ar/_shuffle.py b/python/mori/ops/gemm_ar/_shuffle.py index b5fb20349..3e760d561 100644 --- a/python/mori/ops/gemm_ar/_shuffle.py +++ b/python/mori/ops/gemm_ar/_shuffle.py @@ -63,3 +63,47 @@ def preshuffle_b(w: torch.Tensor) -> torch.Tensor: out = w.view(n // n_lane, n_lane, k // bk, bk // k_pack, k_pack) out = out.permute(0, 2, 3, 1, 4).contiguous() return out.view(n, k).view(dtype) + + +#: ue8m0 block size along K, fixed by the checkpoint and by the MFMA. +MXFP8_BLOCK = 32 +#: M rows a wave's four 16-row A tiles span, and so the packing group. +_A_SCALE_GROUP = 64 + + +def preshuffle_a_scale(exps: torch.Tensor) -> torch.Tensor: + """Put the ue8m0 A scales in the layout ``--quant mxfp8`` reads. + + Takes the exponent bytes as ``[M, K/32]`` -- one per 32-wide K block, the + orientation every quantiser emits -- and returns a flat int32 buffer to hand + the kernel as its ``A_scale`` argument. + + Two things happen, for two different reasons. + + **K-block major.** A block group's sixteen lanes want sixteen consecutive + rows of one K block, so K major puts them at consecutive addresses and the + load coalesces. Row major spreads them ``K/32`` bytes apart, which measured + +50-67% on the whole GEMM -- the addresses, not the bytes: same instruction + count, 3.2x the cache accesses. + + **Four M tiles to a dword.** A lane's four A tiles differ only by sixteen + rows and share the K block, and ``opsel_b`` on the scaled MFMA names which + byte of the 32-bit scale operand the instruction reads. So packing the four + into one dword turns four loads into one and the byte select costs nothing + -- it is the instruction's own field, not a shift. Worth -5 to -7.5%, and + it also shrinks the buffer 4x by storing bytes rather than int32. CK does + the same thing in ``preShuffleScaleBuffer_gfx950`` (its ``MNXdlPack``). + + So within each 64-row group the order goes from ``ti*16 + r`` to + ``r*4 + ti``, which is a ``(4, 16) -> (16, 4)`` transpose and nothing else. + """ + if exps.ndim != 2: + raise ValueError(f"expected a 2-D [M, K/32] scale, got {exps.ndim}-D") + m, kb = exps.shape + if m % _A_SCALE_GROUP: + raise ValueError(f"M={m} must be a multiple of {_A_SCALE_GROUP}") + if (m * kb) % 4: + raise ValueError(f"M*K/32 = {m * kb} must be a multiple of 4") + out = exps.to(torch.uint8).t().contiguous() # [K/32, M], K-block major + out = out.view(kb, m // _A_SCALE_GROUP, 4, 16) # split M into (group, ti, r) + return out.permute(0, 1, 3, 2).contiguous().reshape(-1).view(torch.int32) diff --git a/python/mori/ops/gemm_ar/kernels_fused.py b/python/mori/ops/gemm_ar/kernels_fused.py index 08fb2874d..61a760059 100644 --- a/python/mori/ops/gemm_ar/kernels_fused.py +++ b/python/mori/ops/gemm_ar/kernels_fused.py @@ -184,8 +184,14 @@ def __init__(self, inner): def idx(self, i, j): return self._inner.idx(j, i) - def call(self, a, b, c, *, set_prio=True): - return self._inner.call(b, a, c, set_prio=set_prio) + def call(self, a, b, c, *, set_prio=True, scale_a=None, scale_b=None): + # The scales follow their operands: with the exchange the instruction's + # A is our B, so its scale_a must be our B scale. The lane mapping is + # unaffected -- row is still ``lane % 16`` and the 32-block still + # ``lane // 16`` -- only which matrix that row indexes changes. + return self._inner.call( + b, a, c, set_prio=set_prio, scale_a=scale_b, scale_b=scale_a + ) class _BlockScaleK: @@ -348,6 +354,198 @@ def final_scale(self, accs, prev, *, idx_fn, n_tiles_b): ) +class _Mxfp8ScaleK: + """DeepSeek-V4.1-Flash's 32-wide ue8m0 scales, fed to the MFMA as operands. + + The checkpoint quantises A per 32 K and B per 32x32 with ue8m0 -- exponent + only, so every scale is exactly a power of two and applying it is lossless. + That is precisely the operand format of + ``v_mfma_scale_f32_16x16x128_f8f6f4``, so unlike ``_BlockScaleK`` there is + **no arithmetic here at all**: no second accumulator, no running rescale, no + division. One MFMA consumes a whole K=128 step with its four 32-blocks + already dequantised in hardware. ``_BlockScaleK`` exists only because a + 128-wide block cannot be expressed this way. + + Lane mapping (from sglang's ``mxfp8_gemv_gfx95.cuh``, whose header notes it + was measured with a lane probe and is not documented): lane ``16*s + r`` + supplies the ue8m0 scale of 32-block ``s`` of row ``r``, at op_sel 0. So for + a K step ``ks`` this lane wants block ``4*ks + lane//16`` of row + ``lane % 16`` -- one scalar per tile, and the same shape of index + ``_BlockScaleK.a_scales`` already uses under ``swap_ab``. + + Scales arrive as int32 with the exponent byte in the low 8 bits rather than + as packed uint8. The MFMA's scale operand is a 32-bit register read at + op_sel 0, so the widening is free at the instruction; doing it on the host + costs 4x on an array that is 1 MB against the weight's hundreds, and buys a + plain dword load here instead of sub-dword addressing. + """ + + #: ue8m0 block size along K (and along N for B). Fixed by the checkpoint + #: and by the instruction. + BLOCK = 32 + + def __init__( + self, + A_scale, + B_scale, + m, + n, + k, + *, + n_tiles_a, + n_tiles_b, + row_major=False, + packed_a=False, + ): + self.kb_count = k // self.BLOCK + self.m = m + self.n_groups = n // self.BLOCK + self.n_tiles_a = n_tiles_a + self.n_tiles_b = n_tiles_b + self.row_major = row_major + self.packed_a = packed_a + self.lane = fx.thread_idx.x % 64 + # Row major packs four 32-blocks of one row into a dword, so a K step is + # one dword wide and the arrays are a quarter the size. + self.k_dwords = k // 128 + if packed_a: + # Four bytes to a dword, so exactly four M tiles. M % 4 == 0 is also + # required for the dword index to be exact, but M is dynamic here; + # BLOCK_M >= 128 already forces it at the call site. + if n_tiles_a != 4: + raise ValueError( + f"packed A scales need exactly 4 M tiles (BLOCK_M=256), got " + f"n_tiles_a={n_tiles_a}" + ) + # one byte per scale, [K/32, M/64, 16, 4] + a_bytes = m * self.kb_count + else: + a_bytes = m * (self.k_dwords if row_major else self.kb_count) * 4 + b_bytes = (n // self.BLOCK) * (self.k_dwords if row_major else self.kb_count) * 4 + gSA = fx.rocdl.make_buffer_tensor( + A_scale, max_size=False, num_records_bytes=a_bytes + ) + gSB = fx.rocdl.make_buffer_tensor( + B_scale, max_size=False, num_records_bytes=b_bytes + ) + # Byte select for row major: lane 16*s+r wants byte s of the dword. + # Loop invariant, so it is hoisted out of the mainloop by the compiler. + self.byte_shift = (self.lane // 16) * fx.Int32(8) + self.sa_div = fx.logical_divide(gSA, fx.make_layout(1, 1)) + self.sb_div = fx.logical_divide(gSB, fx.make_layout(1, 1)) + self.atom_1 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) + self.reg_1 = fx.make_rmem_tensor(fx.make_layout(1, 1), fx.Int32) + + def _load1(self, div, index): + fx.copy(self.atom_1, fx.slice(div, (None, fx.Int32(index))), self.reg_1) + return Vec(fx.memref_load_vec(self.reg_1))[0] + + def _load_byte(self, div, dword_index): + """One scale out of a row-major dword: load, shift down, mask. + + The four 32-blocks a single MFMA consumes are four consecutive bytes of + one row, so the dword address does not depend on ``s`` at all -- the + four block groups read the same sixteen addresses and only the shift + differs. Two VALU ops per scale, both cheap; what row major costs is in + the addresses, not here. + """ + v = self._load1(div, dword_index) + return arith.andi(arith.shrui(v, self.byte_shift), fx.Int32(0xFF)) + + def _kb(self, ks): + """This lane's 32-block for K step ``ks``: block ``4*ks + lane//16``.""" + return fx.Int32(ks * 4) + self.lane // 16 + + def a_scales(self, base_row, ks): + """Per M-tile A scale operand. ``A_scale`` is **K-block major**, ``[K/32, M]``. + + The layout is the point. A lane wants row ``lane % 16`` of block + ``4*ks + lane//16``, so the sixteen lanes of one block group differ only + in the row. K-block major puts those sixteen at consecutive addresses -- + one coalesced load. Row major ``[M, K/32]`` puts them ``K/32`` dwords + apart instead, which is sixteen separate transactions per group and + measured 19% off the pace at M=4096 even after matching BLOCK_M. + ``_BlockScaleK`` takes the same column-major A scale for the same + reason. + """ + if const_expr(self.packed_a): + # ``[K/32, M/64, 16, 4]`` bytes: the four M tiles of one lane sit in + # one dword, because they differ only by 16 rows and share the K + # block. One load feeds all four MFMAs and ``opsel_b`` picks the + # byte, so the four returned entries are deliberately the same + # register. Sixteen lanes of a block group still read sixteen + # consecutive dwords, so this keeps the coalescing and divides the + # A-scale load count by four. base_row is a multiple of 64 and M a + # multiple of 4, so both shifts are exact. + kb = self._kb(ks) + v = self._load1( + self.sa_div, + kb * fx.Int32(self.m // 4) + base_row // fx.Int32(4) + self.lane % 16, + ) + return [v for _ in range_constexpr(self.n_tiles_a)] + row = base_row + self.lane % 16 + if const_expr(self.row_major): + # [M, K/32] uint8 read as [M, K/128] dwords: the sixteen lanes of a + # group land K/128 dwords apart, sixteen addresses instead of one + # coalesced run. + return [ + self._load_byte( + self.sa_div, (row + ti * 16) * fx.Int32(self.k_dwords) + ks + ) + for ti in range_constexpr(self.n_tiles_a) + ] + kb = self._kb(ks) + return [ + self._load1(self.sa_div, kb * fx.Int32(self.m) + row + ti * 16) + for ti in range_constexpr(self.n_tiles_a) + ] + + def b_scales(self, base_col, ks): + """Per N-tile B scale operand. ``B_scale`` is K-block major, ``[K/32, N/32]``. + + A 16-column tile never straddles a 32-column group (``base_col`` is a + multiple of 16), so the group index is constant across the tile's rows: + all sixteen lanes of a block group read the same address and the load is + a broadcast. + """ + col = base_col + self.lane % 16 + if const_expr(self.row_major): + # [N/32, K/32] uint8 as [N/32, K/128] dwords; still a broadcast + # within the group, so B is the cheap side either way. + return [ + self._load_byte( + self.sb_div, + ((col + tj * 16) // fx.Int32(self.BLOCK)) * fx.Int32(self.k_dwords) + + ks, + ) + for tj in range_constexpr(self.n_tiles_b) + ] + kb = self._kb(ks) + return [ + self._load1( + self.sb_div, + kb * fx.Int32(self.n_groups) + (col + tj * 16) // fx.Int32(self.BLOCK), + ) + for tj in range_constexpr(self.n_tiles_b) + ] + + def step(self, base_row, base_col, ks, lds_block_m, lds_block_n): + """The A/B scale operands of K step ``ks`` for the four LDS halves. + + Returned in the order the mainloop consumes them: ``(a0, a1, b0, b1)``, + pairing as c00=(a0,b0), c01=(a0,b1), c10=(a1,b0), c11=(a1,b1). A method + rather than a closure for the reason ``_BlockScaleK.rescale_for`` + documents: FlyDSL rewrites the AST of every ``def`` nested in a kernel + and that turns captured instances into locals. + """ + return ( + self.a_scales(base_row + 0 * lds_block_m, ks), + self.a_scales(base_row + 1 * lds_block_m, ks), + self.b_scales(base_col + 0 * lds_block_n, ks), + self.b_scales(base_col + 1 * lds_block_n, ks), + ) + + class _SwapABStoreC(StoreC): """C store for an A/B-swapped MFMA: 4 consecutive N per lane -> 64-bit store. @@ -877,20 +1075,39 @@ def compile_fused_gemm_scatter( "--lane-transpose for the real thing on --mode fused-lsa" ) direct_lsa = fuse and transport == "lsa" - if quant not in ("ptpc", "blockscale"): - raise ValueError(f"quant must be ptpc or blockscale, got {quant!r}") + if quant not in ("ptpc", "blockscale", "mxfp8", "mxfp8_unpacked", "mxfp8_row"): + raise ValueError( + f"quant must be ptpc, blockscale, mxfp8, mxfp8_unpacked or " + f"mxfp8_row, got {quant!r}" + ) blockscale = quant == "blockscale" + # ``mxfp8`` takes ``preshuffle_a_scale``'s layout: K-block major with a + # lane's four M tiles packed into one dword, the byte picked by the MFMA's + # opsel. The other two are the layouts it was chosen over, kept so the + # choice stays measurable rather than asserted: + # mxfp8_unpacked K-block major, one int32 per scale (+5 to +7.5%) + # mxfp8_row the quantiser's own [M, K/32] bytes (+50 to +67%) + mxfp8_row = quant == "mxfp8_row" + mxfp8_pack = quant == "mxfp8" + mxfp8 = quant in ("mxfp8", "mxfp8_unpacked") or mxfp8_row + if (blockscale or mxfp8) and not swap_ab: + raise ValueError( + f"--quant {quant} needs --swap-ab: the unswapped epilogue " + "applies its scales inside the stock StoreC, which this copy does " + "not override" + ) if blockscale: - if not swap_ab: - raise ValueError( - "--quant blockscale needs --swap-ab: the unswapped epilogue " - "applies its scales inside the stock StoreC, which this copy does " - "not override" - ) if K % 128: raise ValueError(f"blockscale needs K % 128 == 0, got K={K}") if N % 128: raise ValueError(f"blockscale needs N % 128 == 0, got N={N}") + if mxfp8: + # K by the MFMA's own step, N by the 32-column scale group. Both hold + # for DeepSeek-V4.1-Flash's wo_b: N=5120, K=8192/tp. + if K % 128: + raise ValueError(f"mxfp8 needs K % 128 == 0, got K={K}") + if N % _Mxfp8ScaleK.BLOCK: + raise ValueError(f"mxfp8 needs N % {_Mxfp8ScaleK.BLOCK} == 0, got N={N}") rotated = fuse if rotated is None else rotated # n_stripe spans 1..N//BLOCK_N; the top of the range *is* chunk-major, so # 0 ("per-mode default") resolves to it rather than to a separate branch. @@ -909,8 +1126,17 @@ def compile_fused_gemm_scatter( if n_stripe < N // BLOCK_N and not rotated: raise ValueError("a striped tile order needs --tile-order rotated") - assert BLOCK_M >= 128 and BLOCK_N >= 256 - assert BLOCK_M % 128 == 0 and BLOCK_N % 256 == 0 + # BLOCK_N's floor belongs to the *store*, not the mainloop: the mainloop + # builds N_TILES_B = BLOCK_N//128 accumulators and is happy with one, while + # _LaneTransposeStoreC's permlane mapping pairs exactly two N-tiles. So 128 + # is legal without permlane and not with it, which is what the store's own + # assert says a few frames deeper and less usefully. + n_floor = 256 if permlane else 128 + assert BLOCK_M >= 128 and BLOCK_N >= n_floor, ( + f"BLOCK_N={BLOCK_N} is below {n_floor}" + + (" (permlane's store pairs two N-tiles)" if permlane else "") + ) + assert BLOCK_M % 128 == 0 and BLOCK_N % 128 == 0 assert K % BLOCK_K == 0 if N % BLOCK_N: raise ValueError( @@ -1082,7 +1308,11 @@ def kernel_gemm_scatter( # Tile counts swap with the operands so Mfma's own asserts and its # idx() line up; the store then addresses the accumulator as # idx(tj, ti). - mfma_raw = Mfma16x16x128(N_TILES_B, N_TILES_A) + # With the swap the instruction's B is our A, so the per-tile + # opsel that selects a packed A byte is opsel_b on the raw atom. + mfma_raw = Mfma16x16x128( + N_TILES_B, N_TILES_A, opsel_b_per_tile=mxfp8_pack + ) mfma = _SwappedMfma(mfma_raw) else: mfma = Mfma16x16x128(N_TILES_A, N_TILES_B) @@ -1219,6 +1449,24 @@ def kernel_gemm_scatter( # visible to a nested def afterwards (a `nonlocal` on one raised "no # binding" and reading one raised UnboundLocalError). bsk = nb0 = base_row_pre = None + msk = base_col_pre = None + if mxfp8: + # Same deal as blockscale: the scales land in the accumulator during + # the mainloop, so the epilogue must not apply them again. + store_c._preapplied = True + msk = _Mxfp8ScaleK( + A_scale, + B_scale, + c_m, + N, + K, + n_tiles_a=N_TILES_A, + n_tiles_b=N_TILES_B, + row_major=mxfp8_row, + packed_a=mxfp8_pack, + ) + base_row_pre = block_m * BLOCK_M + wave_m * (N_TILES_A * 16) + base_col_pre = block_n * BLOCK_N + wave_n * (N_TILES_B * 16) if blockscale: # Set after construction rather than threading a kwarg through # thirteen call sites: it is a trace-time Python bool read by @@ -1271,24 +1519,29 @@ def kernel_gemm_scatter( n_tiles_b=N_TILES_B, lds_block_m=LDS_BLOCK_M, ) + msa0 = msa1 = msb0 = msb1 = None + if mxfp8: + msa0, msa1, msb0, msb1 = msk.step( + base_row_pre, base_col_pre, k, LDS_BLOCK_M, LDS_BLOCK_N + ) b0_frag = b_s2r.load(b_cur0, preshuffled=b_preshuffled) a0_frag = a_s2r.load(a_cur0) a_g2s.load(a_next1, A1_gl_offset + (k + 1) * BLOCK_K) rocdl.s_barrier() - c00_frag = mfma.call(a0_frag, b0_frag, c00_frag) + c00_frag = mfma.call(a0_frag, b0_frag, c00_frag, scale_a=msa0, scale_b=msb0) b1_frag = b_s2r.load(b_cur1, preshuffled=b_preshuffled) b_g2s.load(b_cur0, B0_gl_offset + (k + 2) * B_K_STEP) rocdl.s_barrier() - c01_frag = mfma.call(a0_frag, b1_frag, c01_frag) + c01_frag = mfma.call(a0_frag, b1_frag, c01_frag, scale_a=msa0, scale_b=msb1) a1_frag = a_s2r.load(a_cur1) a_g2s.load(a_cur0, A0_gl_offset + (k + 2) * BLOCK_K) rocdl.s_barrier() - c10_frag = mfma.call(a1_frag, b0_frag, c10_frag) + c10_frag = mfma.call(a1_frag, b0_frag, c10_frag, scale_a=msa1, scale_b=msb0) b_g2s.load(b_cur1, B1_gl_offset + (k + 2) * B_K_STEP) # aiter has `2 * N_LDS_STEPS_A + N_LDS_STEPS_B` here, which lets @@ -1312,7 +1565,7 @@ def kernel_gemm_scatter( # measurably correct here, not a proven-general formula. wait_barrier(N_LDS_STEPS_A + N_LDS_STEPS_B - 1) - c11_frag = mfma.call(a1_frag, b1_frag, c11_frag) + c11_frag = mfma.call(a1_frag, b1_frag, c11_frag, scale_a=msa1, scale_b=msb1) a_cur0, a_next0 = a_next0, a_cur0 a_cur1, a_next1 = a_next1, a_cur1 @@ -1331,27 +1584,32 @@ def kernel_gemm_scatter( n_tiles_b=N_TILES_B, lds_block_m=LDS_BLOCK_M, ) + msa0 = msa1 = msb0 = msb1 = None + if mxfp8: + msa0, msa1, msb0, msb1 = msk.step( + base_row_pre, base_col_pre, K_ITERS - 2, LDS_BLOCK_M, LDS_BLOCK_N + ) b0_frag = b_s2r.load(b_cur0, preshuffled=b_preshuffled) a0_frag = a_s2r.load(a_cur0) rocdl.s_barrier() - c00_frag = mfma.call(a0_frag, b0_frag, c00_frag) + c00_frag = mfma.call(a0_frag, b0_frag, c00_frag, scale_a=msa0, scale_b=msb0) b1_frag = b_s2r.load(b_cur1, preshuffled=b_preshuffled) rocdl.s_barrier() - c01_frag = mfma.call(a0_frag, b1_frag, c01_frag) + c01_frag = mfma.call(a0_frag, b1_frag, c01_frag, scale_a=msa0, scale_b=msb1) a1_frag = a_s2r.load(a_cur1) a_g2s.load(a_next1, A1_gl_offset + (K_ITERS - 1) * BLOCK_K) rocdl.s_barrier() - c10_frag = mfma.call(a1_frag, b0_frag, c10_frag) + c10_frag = mfma.call(a1_frag, b0_frag, c10_frag, scale_a=msa1, scale_b=msb0) b0_frag = b_s2r.load(b_next0, preshuffled=b_preshuffled) rocdl.s_barrier() - c11_frag = mfma.call(a1_frag, b1_frag, c11_frag) + c11_frag = mfma.call(a1_frag, b1_frag, c11_frag, scale_a=msa1, scale_b=msb1) a_cur0, a_next0 = a_next0, a_cur0 a_cur1, a_next1 = a_next1, a_cur1 @@ -1370,22 +1628,31 @@ def kernel_gemm_scatter( n_tiles_b=N_TILES_B, lds_block_m=LDS_BLOCK_M, ) + msa0 = msa1 = msb0 = msb1 = None + if mxfp8: + msa0, msa1, msb0, msb1 = msk.step( + base_row_pre, base_col_pre, K_ITERS - 1, LDS_BLOCK_M, LDS_BLOCK_N + ) a0_frag = a_s2r.load(a_cur0) wait_barrier(0) - c00_frag = mfma.call(a0_frag, b0_frag, c00_frag) + c00_frag = mfma.call(a0_frag, b0_frag, c00_frag, scale_a=msa0, scale_b=msb0) b1_frag = b_s2r.load(b_cur1, preshuffled=b_preshuffled) rocdl.s_barrier() - c01_frag = mfma.call(a0_frag, b1_frag, c01_frag) + c01_frag = mfma.call(a0_frag, b1_frag, c01_frag, scale_a=msa0, scale_b=msb1) a1_frag = a_s2r.load(a_cur1) rocdl.s_barrier() rocdl.s_setprio(1) - c10_frag = mfma.call(a1_frag, b0_frag, c10_frag, set_prio=False) - c11_frag = mfma.call(a1_frag, b1_frag, c11_frag, set_prio=False) + c10_frag = mfma.call( + a1_frag, b0_frag, c10_frag, set_prio=False, scale_a=msa1, scale_b=msb0 + ) + c11_frag = mfma.call( + a1_frag, b1_frag, c11_frag, set_prio=False, scale_a=msa1, scale_b=msb1 + ) rocdl.s_setprio(0) rocdl.s_barrier() diff --git a/python/mori/ops/gemm_ar/op.py b/python/mori/ops/gemm_ar/op.py index 134aab8a5..81b4142da 100644 --- a/python/mori/ops/gemm_ar/op.py +++ b/python/mori/ops/gemm_ar/op.py @@ -23,12 +23,12 @@ from __future__ import annotations -from typing import NamedTuple, Optional +from typing import NamedTuple +import flydsl.expr as fx import torch -import flydsl.expr as fx -from mori.cco import CCODevCommRequirements, GDA_CONNECTION_NONE +from mori.cco import GDA_CONNECTION_NONE, CCODevCommRequirements from mori.tensor_utils import from_gpu_ptr from .kernels_fused import BLOCK_K, compile_fused_gemm_scatter @@ -55,6 +55,44 @@ # The block-scale group, on both operands: A is 1x128, B is 128x128. SCALE_BLOCK_K = 128 +# The ue8m0 group, on both operands: A is 1x32, B is 32x32. DeepSeek-V4.1-Flash +# quantises this way (`weight_block_size [32, 32]`, `scale_fmt "ue8m0"`). +MXFP8_BLOCK = 32 + +#: What the GEMM's operands are quantised as. +#: +#: ``"blockscale"`` A 1x128 fp32 scales, B 128x128. DeepSeek-V4-Pro. +#: ``"mxfp8"`` 32-wide ue8m0 on both, fed to the scaled MFMA as +#: instruction operands. DeepSeek-V4.1-Flash. +#: +#: They differ in more than a group size: mxfp8 needs BLOCK_M=256 (see +#: MXFP8_BLOCK_M) and takes its A scale through ``preshuffle_a_scale`` as +#: int32, where blockscale takes fp32. Both are checked, not assumed. +QUANTS = ("blockscale", "mxfp8") + + +#: mxfp8's BLOCK_M is not a preference. The kernel packs a lane's four M tiles +#: into one dword and picks the byte with the MFMA's ``opsel``, and there are +#: four tiles only when ``BLOCK_M // 64 == 4``. blockscale has no such rule -- +#: its 128 is the measured tile (256x256 spills the doubled accumulator VGPRs), +#: so other granule-legal tiles stay expressible there. +MXFP8_BLOCK_M = 256 + + +def _default_block_m(quant: str) -> int: + return MXFP8_BLOCK_M if quant == "mxfp8" else DEFAULT_BLOCK_M + + +def _quant_tile_constraint(quant: str, block_m: int) -> str | None: + """Why this BLOCK_M cannot serve this quantisation, or None.""" + if quant == "mxfp8" and block_m != MXFP8_BLOCK_M: + return ( + f"quant='mxfp8' requires block_m={MXFP8_BLOCK_M}, got {block_m}: the " + f"packed A scale puts a lane's four M tiles in one dword, which is " + f"four tiles only at BLOCK_M//64 == 4" + ) + return None + # Chunks are how many separate pushes a destination receives, and so how early # the first bytes leave. More is better until the pieces get small enough that # the SDMA per-packet cost shows; 8 is the measured knee. @@ -112,7 +150,7 @@ def counter_chunks(m_pad: int, world_size: int, block_m: int = DEFAULT_BLOCK_M) return max(c for c in range(1, min(MAX_CHUNKS, bands) + 1) if bands % c == 0) -def _tile_constraints(block_m: int, block_n: int) -> Optional[str]: +def _tile_constraints(block_m: int, block_n: int, n_granule: int | None = None) -> str | None: """Why this tile is not one the kernel can build, or None. Checked before anything divides by a tile size. ``padded_m`` and @@ -123,7 +161,7 @@ def _tile_constraints(block_m: int, block_n: int) -> Optional[str]: """ for name, value, granule in ( ("block_m", block_m, TILE_M_GRANULE), - ("block_n", block_n, TILE_N_GRANULE), + ("block_n", block_n, n_granule or TILE_N_GRANULE), ): if value < granule or value % granule: return ( @@ -134,10 +172,15 @@ def _tile_constraints(block_m: int, block_n: int) -> Optional[str]: return None -def _gemm_constraints(n: int, k: int, block_n: int) -> Optional[str]: +def _gemm_constraints( + n: int, k: int, block_n: int, quant: str = "blockscale" +) -> str | None: """Why the GEMM cannot take this shape, or None. See :func:`supports`.""" + # K % 128 holds for both, for different reasons: it is blockscale's group, + # and it is the scaled MFMA's K step. if k % SCALE_BLOCK_K: - return f"K={k} must be a multiple of {SCALE_BLOCK_K} (the block-scale group)" + why = "the block-scale group" if quant == "blockscale" else "the MFMA's K step" + return f"K={k} must be a multiple of {SCALE_BLOCK_K} ({why})" if k < MIN_K: return ( f"K={k} is below the minimum {MIN_K}: the mainloop prefetches a " @@ -145,8 +188,9 @@ def _gemm_constraints(n: int, k: int, block_n: int) -> Optional[str]: ) if n % block_n: return f"N={n} must be a multiple of block_n={block_n}" - if n % SCALE_BLOCK_K: - return f"N={n} must be a multiple of {SCALE_BLOCK_K} (the B scale group)" + b_group = MXFP8_BLOCK if quant == "mxfp8" else SCALE_BLOCK_K + if n % b_group: + return f"N={n} must be a multiple of {b_group} (the B scale group)" return None @@ -192,10 +236,11 @@ def supports( k: int, world_size: int, *, - block_m: int = DEFAULT_BLOCK_M, + block_m: int | None = None, block_n: int = DEFAULT_BLOCK_N, gather_dtype: str = "bf16", gather_transport: str = "lsa", + quant: str = "blockscale", ) -> bool: """Whether this shape is *expressible*, which is not whether it is faster. @@ -206,10 +251,17 @@ def supports( rather than by a second copy of its rules: a predicate that says yes where construction raises, or vice versa, is worse than no predicate. """ + if quant not in QUANTS: + return False + # block_m follows the quantisation unless the caller pins it; passing the + # wrong one is a rejection rather than a silent reinterpretation. + block_m = _default_block_m(quant) if block_m is None else block_m # Tile and world size first: everything below divides by their product. if world_size < 1 or _tile_constraints(block_m, block_n) is not None: return False - if m <= 0 or _gemm_constraints(n, k, block_n) is not None: + if _quant_tile_constraint(quant, block_m) is not None: + return False + if m <= 0 or _gemm_constraints(n, k, block_n, quant) is not None: return False try: _build_cfg( @@ -301,6 +353,91 @@ def _flatten_a_scale(a_scale: torch.Tensor, m: int, kb: int) -> torch.Tensor: ) +def _flatten_mxfp8_a_scale(a_scale: torch.Tensor, m: int, k: int) -> torch.Tensor: + """A's ue8m0 scales as ``preshuffle_a_scale`` emits them. + + Deliberately stricter than :func:`_flatten_a_scale`: there is no logical 2-D + spelling to disambiguate here, because the layout is not a transpose of + anything. Inside each 64-row group the order goes from ``ti*16 + r`` to + ``r*4 + ti`` so a lane's four M tiles land in one dword, which no 2-D shape + describes. Callers build it with ``preshuffle_a_scale`` (or have their + quantiser write it directly) and pass it flat; anything else is rejected + rather than reinterpreted. + + int32 rather than the uint8 the scales really are: the MFMA's scale operand + is a 32-bit register, and four packed bytes are one element of it. + """ + want = m * k // (4 * MXFP8_BLOCK) # four ue8m0 bytes to an int32 + if a_scale.dim() != 1: + raise ValueError( + f"a_scale must be the flat buffer preshuffle_a_scale returns; got " + f"shape {tuple(a_scale.shape)}. The packed layout is a permutation " + f"within each 64-row group, so no 2-D shape describes it." + ) + if a_scale.numel() != want: + raise ValueError( + f"a_scale has {a_scale.numel()} elements, expected M*K/128 = {want} " + f"(M*K/32 ue8m0 bytes, four to an int32)" + ) + return a_scale + + +class _PinnedLaunch: + """A compiled FlyDSL kernel with its dispatch resolved once. + + ``JitFunction.__call__`` re-derives the cache key on every call before it + reaches the ``CallState`` that does the work: an ``inspect`` signature bind, + a snapshot of ~35 referenced globals, a drift check against them, a + compile-hint resolve and a runtime-pairing check. Measured on the fused wo_b + at M=5120: **31.9us for a phase and 85.9us for the GEMM**, against 5.9us of + actual ``hipModuleLaunchKernel``. Four launches per call is 181.6us of + Python, and the phases are sequential, so it sits on the critical path + between them -- enough to make a layer whose GPU work is 526us take 707us. + + So let the first call go through the normal path, then keep the + ``CallState`` it produced and hand it the argument tuple directly. Safe + because ``CallState`` re-reads every argument each call -- it pre-allocates + the ctypes storage, not the values. + + The tuple order is *verified*, not assumed: at capture time this binds the + same arguments through FlyDSL's own signature and checks element-by-element + that the tuple it would build is identical. Getting that order wrong would + not fail, it would feed the kernel the wrong pointers. + """ + + __slots__ = ("_fn", "_state") + + def __init__(self, fn): + self._fn = fn + self._state = None + + def __call__(self, *args, stream): + if self._state is not None: + return self._state(args + (stream,)) + out = self._fn(*args, stream=stream) + self._pin(args, stream) + return out + + def _pin(self, args, stream): + """Capture the CallState, or give up and stay on the normal path.""" + try: + cache = self._fn._call_state_cache + # One JitFunction per compiled shape here, so one entry. More than + # one means it is shared and the entry cannot be chosen blind. + if len(cache) != 1: + return + state = next(iter(cache.values())) + bound = self._fn._sig.bind(*args, stream=stream) + bound.apply_defaults() + want = tuple(bound.arguments.values()) + ours = args + (stream,) + if len(want) != len(ours) or any(a is not b for a, b in zip(want, ours)): + return + except Exception: # noqa: BLE001 - any surprise means stay on the slow path + return + self._state = state + + class _Plan(NamedTuple): """One M's compiled pipeline. @@ -335,12 +472,13 @@ def __init__( n: int, k: int, m_max: int, - block_m: int = DEFAULT_BLOCK_M, + block_m: int | None = None, block_n: int = DEFAULT_BLOCK_N, sdma_queues: int = 1, gather_dtype: str = "bf16", gather_transport: str = "lsa", - max_shapes: Optional[int] = None, + quant: str = "blockscale", + max_shapes: int | None = None, ): # Cleanup state before anything can raise. The rollback below calls # close(), which clears these; initialising them after the try block @@ -349,8 +487,18 @@ def __init__( self._closed = True # so a failed constructor leaves close() a no-op self.mem = self.win = self.dev_comm = None self._cache: dict[int, _Plan] = {} - self._pad_in: Optional[torch.Tensor] = None - + self._pad_in: torch.Tensor | None = None + + if quant not in QUANTS: + raise ValueError(f"quant={quant!r} is not one of {sorted(QUANTS)}") + # Each quantisation has exactly one legal BLOCK_M (see _quant_block_m). + # Default to it, and reject a mismatch loudly rather than compiling a + # kernel whose scale indexing does not match the operand handed in. + if block_m is None: + block_m = _default_block_m(quant) + why = _quant_tile_constraint(quant, block_m) + if why is not None: + raise ValueError(why) if gather_dtype not in WIRE_DTYPES: raise ValueError( f"gather_dtype={gather_dtype!r} is not one of {sorted(WIRE_DTYPES)}" @@ -381,7 +529,7 @@ def __init__( ) if max_shapes < 1: raise ValueError(f"max_shapes must be >= 1; got {max_shapes}") - why = _gemm_constraints(n, k, block_n) + why = _gemm_constraints(n, k, block_n, quant) if why is not None: raise ValueError(f"unsupported shape: {why}") self.comm = comm @@ -390,6 +538,7 @@ def __init__( # world_size. Normalise here so the op's own surface has one name. self.world_size = comm.nranks self.n, self.k = n, k + self.quant = quant self.block_m, self.block_n = block_m, block_n self.sdma_queues = sdma_queues #: fp8 halves the all-gather's bytes, which is ~40% of a fused layer. @@ -445,7 +594,7 @@ def window_bytes_for( block_m: int = DEFAULT_BLOCK_M, gather_dtype: str = "bf16", gather_transport: str = "lsa", - max_shapes: Optional[int] = None, + max_shapes: int | None = None, ) -> int: """Symmetric-window bytes an op for this shape will allocate. @@ -521,7 +670,7 @@ def _compiled(self, m: int): b_preshuffled=True, fuse=True, transport="sdma", - quant="blockscale", + quant=self.quant, sdma_queues=self.sdma_queues, # The three C-store stages are off by default in # compile_fused_gemm_scatter, and blockscale requires swap_ab. @@ -533,9 +682,11 @@ def _compiled(self, m: int): # fused_order rather than a literal list: the fp8 wire inserts a # quantise and a dequantise around the gather, and a hardcoded # drain/reduce/gather would skip them and reduce into zeros. - tail = tuple(parts[name] for name in parts["fused_order"]) + # Pinned launchers on the hot path; `parts` keeps the raw ones, which is + # what self_test uses (it runs a different order, once). + tail = tuple(_PinnedLaunch(parts[name]) for name in parts["fused_order"]) hit = _Plan( - gemm=gemm, + gemm=_PinnedLaunch(gemm), tail=tail, parts=parts, input=from_gpu_ptr( @@ -548,7 +699,7 @@ def _compiled(self, m: int): self._cache[m] = hit return hit - def self_test(self, m: Optional[int] = None) -> None: + def self_test(self, m: int | None = None) -> None: """Verify that the collective actually moves bytes. Raises if it does not. The failure this exists for is silent. A mori built without @@ -645,7 +796,7 @@ def close(self) -> None: handle.close() self.win = self.mem = None - def __enter__(self) -> "GemmAllReduceOp": + def __enter__(self) -> GemmAllReduceOp: return self def __exit__(self, *exc) -> None: @@ -699,10 +850,22 @@ def __call__( block_m`` (see :meth:`padded_m`), ``b_preshuffled`` is ``[N, K]`` through :func:`preshuffle_b`. - ``a_scale`` may be ``[M, K/128]`` (the logical shape, either memory - order), ``[K/128, M]``, or already flat in physical order -- see - :func:`_flatten_a_scale`, which is where the distinction is made rather - than assumed. ``b_scale`` is ``[N/128, K/128]`` row-major. + The scales follow ``quant``: + + ``blockscale`` + ``a_scale`` may be ``[M, K/128]`` (the logical shape, either memory + order), ``[K/128, M]``, or already flat in physical order -- see + :func:`_flatten_a_scale`, which is where the distinction is made + rather than assumed. ``b_scale`` is ``[N/128, K/128]`` row-major. + Both fp32. + + ``mxfp8`` + ``a_scale`` is what :func:`~mori.ops.gemm_ar.preshuffle_a_scale` + returns: flat int32, ``M*K/128`` elements, and **only** flat -- the + packed layout is a permutation inside each 64-row group, so no 2-D + shape describes it. ``b_scale`` is the ``[N/32, K/32]`` ue8m0 bytes + K-block major and widened, i.e. ``e.t().contiguous().to(int32)``, + flat. Both int32. The returned tensor aliases the window and is overwritten by the next call. Clone it to keep it. One instance is not usable concurrently: the @@ -727,7 +890,11 @@ def __call__( a_fp8.contiguous().view(torch.int8).view(-1), b_preshuffled.contiguous().view(torch.int8).view(-1), plan.input.view(-1), - _flatten_a_scale(a_scale, m, self.k // SCALE_BLOCK_K), + ( + _flatten_mxfp8_a_scale(a_scale, m, self.k) + if self.quant == "mxfp8" + else _flatten_a_scale(a_scale, m, self.k // SCALE_BLOCK_K) + ), b_scale.reshape(-1), m, self.n, @@ -749,6 +916,7 @@ def _check_operands(self, a_fp8, b_preshuffled, a_scale, b_scale, m: int) -> Non """ if self.mem is None: raise RuntimeError("this GemmAllReduceOp has been closed") + mxfp8 = self.quant == "mxfp8" kb = self.k // SCALE_BLOCK_K for name, t, shape in ( ("a_fp8", a_fp8, (m, self.k)), @@ -764,12 +932,30 @@ def _check_operands(self, a_fp8, b_preshuffled, a_scale, b_scale, m: int) -> Non ) if t.device.type != "cuda": raise ValueError(f"{name} must be on a GPU, got {t.device}") - for name, t, numel in ( - ("a_scale", a_scale, m * kb), - ("b_scale", b_scale, (self.n // SCALE_BLOCK_K) * kb), - ): - if t.dtype != torch.float32: - raise ValueError(f"{name} must be float32, got {t.dtype}") + # mxfp8's scales are ue8m0 exponent bytes read as int32, four packed to + # an element on A and one per 32x32 block on B. blockscale's are fp32. + # The dtype is load-bearing on both: the kernel indexes dwords, so a + # tensor of the right element count in the wrong dtype spans the wrong + # bytes and returns finite nonsense. + if mxfp8: + want_dtype = torch.int32 + sizes = ( + ("a_scale", a_scale, m * self.k // (4 * MXFP8_BLOCK)), + ( + "b_scale", + b_scale, + (self.n // MXFP8_BLOCK) * (self.k // MXFP8_BLOCK), + ), + ) + else: + want_dtype = torch.float32 + sizes = ( + ("a_scale", a_scale, m * kb), + ("b_scale", b_scale, (self.n // SCALE_BLOCK_K) * kb), + ) + for name, t, numel in sizes: + if t.dtype != want_dtype: + raise ValueError(f"{name} must be {want_dtype}, got {t.dtype}") if t.numel() != numel: raise ValueError(f"{name} must have {numel} elements, got {t.numel()}") if t.device.type != "cuda": diff --git a/tests/python/cco/test_gemm_ar.py b/tests/python/cco/test_gemm_ar.py index 882d7c0fc..5f3439be2 100644 --- a/tests/python/cco/test_gemm_ar.py +++ b/tests/python/cco/test_gemm_ar.py @@ -29,14 +29,16 @@ import json import os -from pathlib import Path import subprocess import sys +from pathlib import Path import pytest import torch - -from mori.ops.gemm_ar import layout +from mori.ops.gemm_ar import ( + layout, + preshuffle_a_scale, +) REPO_ROOT = Path(__file__).resolve().parents[3] BENCH = REPO_ROOT / "benchmark" / "cco" / "flydsl" / "gemm_ar" / "bench_gemm_ar.py" @@ -230,7 +232,6 @@ def test_ptpc_matches_an_fp32_reference(m, n, k): pytest.skip("requires a GPU") pytest.importorskip("flydsl") import flydsl.expr as fx - from mori.ops.gemm_ar import compile_fused_gemm_scatter, preshuffle_b g = torch.Generator(device="cuda").manual_seed(7) @@ -290,7 +291,6 @@ def test_blockscale_lands_at_the_fp8_floor(m, n, k): pytest.skip("requires a GPU") pytest.importorskip("flydsl") import flydsl.expr as fx - from mori.ops.gemm_ar import compile_fused_gemm_scatter, preshuffle_b BK = 128 @@ -353,6 +353,89 @@ def test_blockscale_lands_at_the_fp8_floor(m, n, k): assert rel < 3e-3, f"blockscale relL2 {rel:.3e} against the fp32 reference" +@pytest.mark.parametrize("m,n,k", [(1024, 512, 512), (2048, 5120, 2048)]) +def test_mxfp8_lands_at_the_fp8_floor(m, n, k): + """``quant="mxfp8"`` -- DeepSeek-V4.1-Flash's quantisation -- against fp32. + + 32-wide ue8m0 on both sides, which is the scaled MFMA's own operand format: + the four scales of a K=128 step are gathered *across lanes* (lane ``16*s+r`` + supplies block ``s`` of row ``r``) and applied in hardware, so unlike + ``blockscale`` there is no promote arithmetic in the mainloop at all. + + The scales must be exact powers of two -- ue8m0 is an exponent with no + mantissa -- so the only error here is the fp8 rounding of the operands, + same floor as the other two quantisations. A wrong lane mapping does not + land near the floor: it misses by the ratio of two random powers of two. + + ``n=5120, k=2048`` is ``wo_b`` per rank at TP4 on V4.1-Flash. + """ + if not torch.cuda.is_available(): + pytest.skip("requires a GPU") + pytest.importorskip("flydsl") + import flydsl.expr as fx + from mori.ops.gemm_ar import compile_fused_gemm_scatter, preshuffle_b + + BK = 32 + g = torch.Generator(device="cuda").manual_seed(11) + a = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + b = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + # ue8m0 exponent bytes around 2**-7; the kernel reads them as int32 at + # op_sel 0, so the byte sits in the low 8 bits. + ea = torch.randint( + 120, 123, (m, k // BK), generator=g, device="cuda", dtype=torch.int32 + ) + eb = torch.randint( + 120, 123, (n // BK, k // BK), generator=g, device="cuda", dtype=torch.int32 + ) + sa, sb = torch.exp2(ea.float() - 127.0), torch.exp2(eb.float() - 127.0) + b_shuf = preshuffle_b(b) + + af, bf = a.float(), b.float() + ref = torch.zeros(m, n, device="cuda", dtype=torch.float32) + for i in range(k // BK): + ks = slice(i * BK, (i + 1) * BK) + ref += ( + (af[:, ks] @ bf[:, ks].T) + * sa[:, i][:, None] + * sb[:, i].repeat_interleave(BK)[None, :] + ) + rn = ref.norm() + + cfg = layout.ArConfig(world_size=2, m=m, n=n) + gemm = compile_fused_gemm_scatter( + cfg, + 0, + K=k, + BLOCK_M=256, + BLOCK_N=256, + b_preshuffled=True, + fuse=False, + swap_ab=True, + permlane=True, + lane_transpose=True, + quant="mxfp8", + ) + ours = torch.zeros(m, n, device="cuda", dtype=torch.bfloat16) + gemm( + a.contiguous().view(torch.int8).view(-1), + b_shuf.contiguous().view(torch.int8).view(-1), + ours.view(-1), + # A goes through preshuffle_a_scale (K-block major, four M tiles to a + # dword); B stays K-block major [K/32, N/32] int32. + preshuffle_a_scale(ea), + eb.t().reshape(-1).contiguous(), + m, + n, + 0, + 0, + stream=fx.Stream(torch.cuda.current_stream()), + ) + torch.cuda.synchronize() + + rel = ((ours.float() - ref).norm() / rn).item() + assert rel < 3e-3, f"mxfp8 relL2 {rel:.3e} against the fp32 reference" + + @pytest.mark.parametrize("m,n,k", [(512, 512, 256), (4096, 7168, 1024)]) def test_swap_ab_is_bitwise_identical(m, n, k): """Exchanging the MFMA operands must not change a single bit. @@ -368,7 +451,6 @@ def test_swap_ab_is_bitwise_identical(m, n, k): pytest.skip("requires a GPU") pytest.importorskip("flydsl") import flydsl.expr as fx - from mori.ops.gemm_ar import compile_fused_gemm_scatter, preshuffle_b g = torch.Generator(device="cuda").manual_seed(7) From 54e74a02e90bebf3e66e64b8be5712a29a488184 Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 05:33:11 +0000 Subject: [PATCH 2/8] gemm_ar: the GEMM without a collective, for prefill M and for decode's Most of what this integration is worth turns out to be the multiply, not the overlap. At V4.1-Flash's shapes the GEMM is ~180us against 1130-1410us of communication, so hiding it entirely would be worth 11-18% and fusing collects about half of that -- while the GEMM alone is worth ~30% against the route SGLang would otherwise take. And a `ColumnParallelLinear` has no all-reduce to fuse with at all, so `wq_b` was unreachable. `Mxfp8GemmOp` (`gemm.py`) is the same kernel with the epilogue tail compiled out: no communicator, no window, an ordinary tensor out. Two things about it are not obvious. M only has to be a multiple of 64, not of BLOCK_M, because the grid is ceildiv and the tail block masks -- 64 is the packed scale's group. And there is exactly one compile per (N, K) per tile width, because `c_m` is a runtime argument; a per-M cache looks harmless until a server drives it, where every prefill batch pays a fresh multi-second FlyDSL compile. Which N tile to use is decided by the resulting **grid**, not by M. The tile exists to double the grid when the wide one is launch-starved, and the grid is `ceildiv(M,256) * (N/256)`, so a threshold on M alone is only ever right for the N it was fitted to. Across 110 measured points the separation is total: below 128 workgroups the 128-wide tile wins every single time by 9-17%, from 129 the 256-wide one wins every single time by 26-29%. `_WIDE_TILE_MIN_GRID = 140` sits in the gap at zero regret. `Mxfp8GemvOp` (`gemv.py`, `kernels_gemv.py`) is a different kernel for decode's token counts, where a 256-row tile has nothing to fill it. It inverts the same scaled MFMA -- the weight is the A operand and the tokens are B, because the weight is what there is a lot of -- streams the weight from global with no LDS staging, and splits K across a workgroup's waves, reducing through LDS in a fixed order so a row stays batch-invariant. Against SGLang's `mxfp8_gemv` on the same fp8 bytes it is 3-8% faster on the two tuned shapes and **bit-identical** at every M, which is the check that matters: same instruction, same operands, so a layout disagreement shows up as a mismatch rather than a tolerance. Its margin is tuning, and that is recorded next to the table: shapes falling back to the heuristic land inside +-2%, except `wo_a`, which loses 12-16%. The operands are deliberately not uniform. Both ops take `preshuffle_b`'s weight, so a server shuffles once, but the GEMM wants the A scale through `preshuffle_a_scale` while the GEMV takes both scales exactly as the checkpoint stores them. That is not an inconsistency: the GEMM's sixteen lanes read sixteen rows of one K block and have to coalesce, where the GEMV's are sixteen tokens, M is at most 32, and the whole A scale is under a kilobyte. Tests cover both ops at both real shapes, the ragged and unpadded M contracts, and the GEMV across its configuration space -- each configuration in its own process, because a FlyDSL compile failure takes the interpreter with it. Every masked address in the GEMV is pushed out of its buffer's records or clamped: 0xFF is NaN in both ue8m0 and e4m3, and NaN times a zeroed weight is still NaN. Co-Authored-By: Claude Opus 5 (1M context) --- python/mori/ops/gemm_ar/__init__.py | 18 +- python/mori/ops/gemm_ar/gemm.py | 355 ++++++++++++++++++ python/mori/ops/gemm_ar/gemv.py | 267 +++++++++++++ python/mori/ops/gemm_ar/kernels_gemv.py | 454 +++++++++++++++++++++++ tests/python/cco/gemm_ar_op_worker.py | 77 +++- tests/python/cco/test_gemm_ar_op.py | 190 ++++++++++ tests/python/cco/test_mxfp8_gemv_grid.py | 201 ++++++++++ 7 files changed, 1555 insertions(+), 7 deletions(-) create mode 100644 python/mori/ops/gemm_ar/gemm.py create mode 100644 python/mori/ops/gemm_ar/gemv.py create mode 100644 python/mori/ops/gemm_ar/kernels_gemv.py create mode 100644 tests/python/cco/test_mxfp8_gemv_grid.py diff --git a/python/mori/ops/gemm_ar/__init__.py b/python/mori/ops/gemm_ar/__init__.py index ea81b00a8..2d188eeb1 100644 --- a/python/mori/ops/gemm_ar/__init__.py +++ b/python/mori/ops/gemm_ar/__init__.py @@ -29,8 +29,8 @@ # touching the op or the kernels raises. import importlib -from .layout import ArConfig, MAX_WORLD, ar_config, select_stage -from ._shuffle import preshuffle_b +from ._shuffle import preshuffle_a_scale, preshuffle_b +from .layout import MAX_WORLD, ArConfig, ar_config, select_stage _LAZY = { "GemmAllReduceOp": "op", @@ -41,6 +41,11 @@ "DEFAULT_BLOCK_N": "op", "MAX_CHUNKS": "op", "SCALE_BLOCK_K": "op", + "Mxfp8GemmOp": "gemm", + "supports_gemm": "gemm", + "Mxfp8GemvOp": "gemv", + "supports_gemv": "gemv", + "select_config": "gemv", "compile_fused_gemm_scatter": "kernels_fused", "build_sdma_phases": "kernels_sdma", "build_sdma_ar": "kernels_sdma", @@ -48,19 +53,24 @@ } __all__ = [ + "MAX_WORLD", "ArConfig", - "ar_config", "GemmAllReduceOp", - "MAX_WORLD", + "Mxfp8GemmOp", + "Mxfp8GemvOp", + "ar_config", "build_lsa_ar", "build_sdma_ar", "build_sdma_phases", "compile_fused_gemm_scatter", "counter_chunks", "padded_m", + "preshuffle_a_scale", "preshuffle_b", "select_stage", "supports", + "supports_gemm", + "supports_gemv", ] diff --git a/python/mori/ops/gemm_ar/gemm.py b/python/mori/ops/gemm_ar/gemm.py new file mode 100644 index 000000000..9f284946f --- /dev/null +++ b/python/mori/ops/gemm_ar/gemm.py @@ -0,0 +1,355 @@ +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +"""The mxfp8 GEMM on its own, with no collective attached. + +`GemmAllReduceOp` compiles this same kernel and then fuses a scatter into its +epilogue. Plenty of callers want only the GEMM: a `ColumnParallelLinear` has no +all-reduce to overlap with at all, and a `RowParallelLinear` below the fusion's +profitability floor still wants the multiply. Until now the only way to reach it +was `compile_fused_gemm_scatter(..., fuse=False)`, whose signature asks for an +`ArConfig`, a rank and two window handles -- every one of them an all-reduce +concept that a plain GEMM has no answer for. + +Measured against the two routes SGLang has for the same operands, bf16 in and +bf16 out including the quantisation, at `N=5120 K=2048 M=16384` on MI355X: + + bf16 (fake_quant + hipBLASLt) 281.3 us + sglang mxfp8 (tl.dot_scaled) 274.9 + this 196.0 -30.3% + +**M only has to be a multiple of 64 here**, not of `BLOCK_M`: the grid is +`ceildiv(M, BLOCK_M)` and the tail block masks its stores. 64 is the packed +scale's group -- a lane's four M tiles are four 16-row tiles -- and +`preshuffle_a_scale` enforces it. Anything else is zero-extended internally, +which costs at most 63 rows of GEMM. That is a far looser constraint than the +fused op's `world_size * BLOCK_M`, where a whole row band has to belong to one +destination. +""" + +from __future__ import annotations + +import flydsl.expr as fx +import torch + +from .kernels_fused import BLOCK_K, compile_fused_gemm_scatter +from .layout import ArConfig +from .op import ( + DEFAULT_BLOCK_N, + FP8_DTYPES, + MIN_K, + MXFP8_BLOCK, + MXFP8_BLOCK_M, + _flatten_mxfp8_a_scale, + _PinnedLaunch, + _tile_constraints, +) + +#: Rows one packed A-scale group spans: four M tiles of sixteen rows. M is +#: zero-extended to a multiple of this, not to BLOCK_M. +SCALE_GROUP_M = 64 + +#: Workgroups the 256-wide tile's grid needs before it beats the 128-wide one. +#: +#: The narrow tile exists to double the grid when the wide one is launch-starved, +#: so what decides between them is the wide grid's *size*, not M -- and the grid +#: is `ceildiv(M, BLOCK_M) * (N / BLOCK_N)`, which depends on N as much as on M. +#: An earlier version of this was a bare `M < 2048`, measured on an M grid that +#: jumped 1024 -> 2048 and so never looked between them. It is wrong on both +#: shapes. Cold, GEMM only: +#: +#: M wq_b 8192x1280 wo_b 5120x2048 +#: BN=256 BN=128 BN=256 BN=128 +#: 1024 21.8 20.9 narrow 27.6 25.2 narrow +#: 1280 22.4 33.2 WIDE 28.1 25.7 narrow +#: 1536 23.4 33.8 WIDE 28.4 26.8 narrow +#: 1792 24.1 34.7 WIDE 28.5 44.0 WIDE +#: 2048 25.2 36.2 WIDE 29.9 44.7 WIDE +#: +#: wq_b turns over between 1024 and 1280, wo_b between 1536 and 1792 -- one +#: constant cannot express that, and at M=1280 the old one cost wq_b 48%. In +#: wide-grid workgroups those four transitions are 128 -> 160 and 120 -> 140, so +#: this threshold sits at 140 and fits all ten rows. Fitted to two shapes; treat +#: it as measured rather than derived, and re-measure if a third shape appears. +_WIDE_TILE_MIN_GRID = 140 + +#: Kept for callers that ask "is this M small": the largest M at which *some* +#: shape still prefers the narrow tile. +#: +#: The wide tile is launch-starved at small M. At M=64 with N=8192 its grid is +#: ceildiv(64,256) * ceildiv(8192,256) = 32 workgroups on a 256-CU part, so most +#: of the GPU is idle and the time is flat from M=64 to M=512 -- it is doing one +#: tile-row of work either way. Halving BLOCK_N doubles the grid. GEMM only, +#: both of V4.1-Flash's attention shapes, cold (weights rotated past the LLC, +#: which is what a forward pass does) over hot: +#: +#: M wq_b 8192x1280 wo_b 5120x2048 +#: BN=256 BN=128 BN=256 BN=128 +#: 64 19.3 16.4 -15.1% 25.6 21.7 -15.5% +#: 256 20.3 18.9 -6.8% 26.1 24.1 -7.7% +#: 1024 21.8 20.9 -3.9% 27.6 25.2 -8.7% +#: 2048 25.2 36.2 +43.6% 29.9 44.7 +49.6% +#: 16384 167.2 239.1 +43.0% 148.6 212.0 +42.7% +#: +#: An earlier table here read 23.7 / 19.9 at M=64 and was measured with one call +#: per CUDA-graph capture, which on this box has a 13.4us replay floor -- so the +#: small-M rows were mostly harness. See +#: `benchmark/cco/flydsl/gemm_ar/timing.py`. +#: +#: Above the crossover the wide tile wins by more than the narrow one ever wins +#: below, because it is also the one that keeps `lane_transpose` -- that store +#: pairs exactly two N-tiles, so it needs BLOCK_N=256 and the narrow tile gives +#: it up. +NARROW_N_BELOW_M = 2048 +#: What the narrow tile costs to get: its store cannot use the permlane lane +#: transpose, so this is only worth it where the grid gain is bigger. +NARROW_BLOCK_N = 128 + + +def _gemm_shape_constraint(n: int, k: int, block_n: int) -> str | None: + """Why this shape cannot be compiled, or None.""" + if k % BLOCK_K: + return f"K={k} must be a multiple of {BLOCK_K} (the scaled MFMA's K step)" + if k < MIN_K: + return ( + f"K={k} is below the minimum {MIN_K}: the mainloop prefetches a " + f"second K block and runs two tail steps, so K/{BLOCK_K} must be >= 2" + ) + if n % block_n: + return f"N={n} must be a multiple of block_n={block_n}" + if n % MXFP8_BLOCK: + return f"N={n} must be a multiple of {MXFP8_BLOCK} (the B scale group)" + return None + + +def supports_gemm(n: int, k: int, *, block_n: int = DEFAULT_BLOCK_N) -> bool: + """Whether this shape is *expressible*, which is not whether it is faster. + + M is deliberately not an argument: any M is servable, because one below a + multiple of 64 is zero-extended. Profitability is the caller's call -- at + small M a GEMV-shaped kernel will win, and this says nothing about that. + """ + return ( + _tile_constraints(MXFP8_BLOCK_M, block_n) is None + and _gemm_shape_constraint(n, k, block_n) is None + # Small M runs the narrow tile, so N has to divide by that too. + and _gemm_shape_constraint(n, k, NARROW_BLOCK_N) is None + ) + + +class Mxfp8GemmOp: + """A compiled mxfp8 GEMM for one ``(N, K)``, serving every M. + + :: + + op = Mxfp8GemmOp(n=8192, k=1280) + out = op(a_fp8, b_preshuffled, a_scale, b_scale) # [M, N] bf16 + + ``b_preshuffled`` is a weight through :func:`~mori.ops.gemm_ar.preshuffle_b` + and ``a_scale`` is through :func:`~mori.ops.gemm_ar.preshuffle_a_scale`; + ``b_scale`` is the ``[N/32, K/32]`` ue8m0 bytes K-block major and widened to + int32, flat. Those are the same operands ``GemmAllReduceOp`` takes, and for + the same reasons -- see its docstring. + + No communicator and no symmetric window: the output is an ordinary tensor, + allocated here or supplied by the caller. + """ + + def __init__( + self, + *, + n: int, + k: int, + block_n: int = DEFAULT_BLOCK_N, + ): + why = _tile_constraints(MXFP8_BLOCK_M, block_n) + if why is not None: + raise ValueError(f"unsupported tile: {why}") + why = _gemm_shape_constraint(n, k, block_n) + if why is not None: + raise ValueError(f"unsupported shape: {why}") + self.n, self.k = n, k + self.block_m, self.block_n = MXFP8_BLOCK_M, block_n + self._launch: dict[bool, _PinnedLaunch] = {} + self._pad_in: torch.Tensor | None = None + + def padded_m(self, m: int) -> int: + """``m`` rounded up to a whole packed A-scale group.""" + return (m + SCALE_GROUP_M - 1) // SCALE_GROUP_M * SCALE_GROUP_M + + def wants_narrow_n(self, m: int) -> bool: + """Whether this M runs the 128-wide N tile. See `_WIDE_TILE_MIN_GRID`.""" + wide_grid = -(-m // self.block_m) * (self.n // self.block_n) + return wide_grid < _WIDE_TILE_MIN_GRID + + def _compiled(self, m: int) -> _PinnedLaunch: + """The kernel for this M -- the narrow-N tile below the crossover. + + **Not a per-M cache**, unlike `GemmAllReduceOp`: there are exactly two + kernels ever, chosen by `NARROW_N_BELOW_M`. `c_m` is a runtime argument, + the launch computes its own grid as + `ceildiv(c_m, BLOCK_M) * ceildiv(c_n, BLOCK_N)`, and the tail block + masks -- verified by compiling at M=4096 and calling at 64, 1024, 7040, + 14080 and 16384, all exact. + + That is not a micro-optimisation. A per-M cache looks harmless until a + server drives it: every prefill batch has a different token count, so + each one paid a fresh multi-second FlyDSL compile, and a bounded cache + then made the path decline for good once it filled. + """ + narrow = self.wants_narrow_n(m) + hit = self._launch.get(narrow) + if hit is not None: + return hit + block_n = NARROW_BLOCK_N if narrow else self.block_n + # permlane's store pairs exactly two N-tiles, so it needs BLOCK_N=256. + permlane = not narrow + # A throwaway ArConfig purely to satisfy compile_fused_gemm_scatter's + # signature. With fuse=False the emitted kernel never touches the + # window -- no counters, no puts, no peer addresses -- and nothing + # downstream reads cfg.m, which is why any M compiles the same kernel. + # world_size=2 is the smallest ArConfig.validate accepts; nothing here + # is distributed, and the window handles go in as 0 for the same reason. + cfg = ArConfig(world_size=2, m=self.block_m, n=self.n) + hit = _PinnedLaunch( + compile_fused_gemm_scatter( + cfg, + 0, + K=self.k, + BLOCK_M=self.block_m, + BLOCK_N=block_n, + b_preshuffled=True, + fuse=False, + quant="mxfp8", + swap_ab=True, + permlane=permlane, + lane_transpose=permlane, + ) + ) + self._launch[narrow] = hit + return hit + + def pad_rows(self, x: torch.Tensor, m_pad: int) -> torch.Tensor: + """Zero-extend ``x`` to ``m_pad`` rows, in a buffer reused across calls. + + A GEMM row depends only on the same input row, so the added rows produce + zeros the caller slices off. Pad *before* quantising, so the A scale is + built over the padded M and its packed layout needs no repair. + + The result views a buffer this op reuses; a later call overwrites it. + """ + if x.dim() != 2: + raise ValueError(f"expected a 2-D [M, K] tensor, got {tuple(x.shape)}") + m, k = x.shape + if m_pad < m: + raise ValueError(f"m_pad={m_pad} is smaller than the input's M={m}") + buf = self._pad_in + if ( + buf is None + or buf.shape[1] != k + or buf.shape[0] < m_pad + or buf.dtype != x.dtype + or buf.device != x.device + ): + buf = torch.empty((m_pad, k), dtype=x.dtype, device=x.device) + self._pad_in = buf + out = buf[:m_pad] + out[:m].copy_(x) + out[m:].zero_() + return out + + def _check_operands(self, a_fp8, b_preshuffled, a_scale, b_scale, m: int) -> None: + """Reject what the launch boundary would otherwise reinterpret as bytes. + + ``a_fp8``/``b_preshuffled`` go in as int8 and the scales as dwords, so a + tensor of the right shape in the wrong dtype reaches the kernel and comes + back finite and wrong. All metadata, no device work. + """ + for name, t, shape in ( + ("a_fp8", a_fp8, (m, self.k)), + ("b_preshuffled", b_preshuffled, (self.n, self.k)), + ): + if t.dim() != 2 or tuple(t.shape) != shape: + raise ValueError(f"{name} must be {shape}, got {tuple(t.shape)}") + if t.dtype not in FP8_DTYPES: + raise ValueError( + f"{name} must be one of {[str(d) for d in FP8_DTYPES]}, got " + f"{t.dtype}; the launch reinterprets it as raw bytes, so a " + f"wider dtype runs and returns nonsense rather than failing" + ) + if t.device.type != "cuda": + raise ValueError(f"{name} must be on a GPU, got {t.device}") + for name, t, numel in ( + ("a_scale", a_scale, m * self.k // (4 * MXFP8_BLOCK)), + ("b_scale", b_scale, (self.n // MXFP8_BLOCK) * (self.k // MXFP8_BLOCK)), + ): + if t.dtype != torch.int32: + raise ValueError(f"{name} must be torch.int32, got {t.dtype}") + if t.numel() != numel: + raise ValueError(f"{name} must have {numel} elements, got {t.numel()}") + if t.device.type != "cuda": + raise ValueError(f"{name} must be on a GPU, got {t.device}") + + def __call__( + self, + a_fp8: torch.Tensor, + b_preshuffled: torch.Tensor, + a_scale: torch.Tensor, + b_scale: torch.Tensor, + out: torch.Tensor | None = None, + ) -> torch.Tensor: + """``a_fp8 @ b_preshuffled.T`` with 32-wide ue8m0 scales, into bf16. + + ``M`` must already be a multiple of 64 and the operands must agree with + it -- this does not pad, because the A scale is built over the padded M + and repairing its packed layout afterwards is not a reshape. Use + :meth:`padded_m` and :meth:`pad_rows` on the *unquantised* input. + """ + m = a_fp8.shape[0] + if m % SCALE_GROUP_M: + raise ValueError( + f"M={m} must be a multiple of {SCALE_GROUP_M} (the packed A " + f"scale's group); pad the bf16 input with pad_rows() before " + f"quantising, not the fp8 afterwards" + ) + self._check_operands(a_fp8, b_preshuffled, a_scale, b_scale, m) + if out is None: + out = torch.empty( + (m, self.n), dtype=torch.bfloat16, device=a_fp8.device + ) + elif tuple(out.shape) != (m, self.n) or out.dtype != torch.bfloat16: + raise ValueError( + f"out must be {(m, self.n)} bfloat16, got {tuple(out.shape)} " + f"{out.dtype}" + ) + self._compiled(m)( + a_fp8.contiguous().view(torch.int8).view(-1), + b_preshuffled.contiguous().view(torch.int8).view(-1), + out.view(-1), + _flatten_mxfp8_a_scale(a_scale, m, self.k), + b_scale.reshape(-1), + m, + self.n, + 0, # dev_comm: unused with fuse=False + 0, # win + stream=fx.Stream(torch.cuda.current_stream()), + ) + return out diff --git a/python/mori/ops/gemm_ar/gemv.py b/python/mori/ops/gemm_ar/gemv.py new file mode 100644 index 000000000..bbeaf1dfd --- /dev/null +++ b/python/mori/ops/gemm_ar/gemv.py @@ -0,0 +1,267 @@ +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +"""The public entry to the skinny mxfp8 matmul -- decode's M, not prefill's. + +Same relationship to `kernels_gemv.py` as `Mxfp8GemmOp` has to +`kernels_fused.py`: pick the compile-time configuration, hold the compiled +launchers, and check the operands before the launch boundary reinterprets them +as bytes. + +The handover between this and `Mxfp8GemmOp` is by M and is the caller's to make. +The GEMM's floor is where its 256-row tile stops being launch-starved; this +kernel's ceiling is 32 tokens, because a token is an MFMA row and two tiles of +them is where the register file runs out. They do not meet: M in (32, 1024) is +served well by neither, which is what `mxfp8_native_blockscaled_linear`'s +`dot_scaled` route is for. +""" + +from __future__ import annotations + +import flydsl.expr as fx +import torch + +from .kernels_gemv import MXFP8_BLOCK, STEP_K, TILE, compile_mxfp8_gemv +from .op import FP8_DTYPES, _PinnedLaunch + +#: Most tokens a config can carry: two 16-row MFMA tiles on the B operand. +MAX_TOKENS = 32 + +#: Token counts a config is tuned at, each covering M up to its own value. +M_BUCKETS = (1, 2, 4, 8, 16, 32) + +#: Tuned per (N, K, M bucket) by `benchmark/cco/flydsl/gemm_ar/sweep.py gemv-tune`, +#: on the cold column -- a decode step reads each layer's weight once, so a +#: hot-loop number is measuring the LLC rather than the kernel. +#: +#: Keys are `(N, K, bucket)`; anything absent falls back to `_HEURISTIC`. +#: +#: Against sglang's `mxfp8_gemv` on the same fp8 operands -- the two agree bit +#: for bit, so this is the same arithmetic done faster. Cold us, MI355X: +#: +#: M wq_b 8192x1280 wo_b 5120x2048 +#: sglang mori sglang mori +#: 1 4.19 3.91 -6.8% 4.29 3.93 -8.4% +#: 2 4.14 3.92 -5.4% 4.52 3.99 -11.8% +#: 4 4.20 4.07 -3.0% 4.54 4.17 -8.1% +#: 8 4.29 4.11 -4.1% 4.76 4.27 -10.1% +#: 16 4.57 4.54 -0.7% 5.28 4.94 -6.4% +#: 32 5.58 5.27 -5.7% 6.97 6.32 -9.3% +#: +#: `wq_b` is the thinner margin and it is the shape, not the tuning: K=1280 is +#: ten 128-wide steps and none of 4, 8, 16 waves divides ten, so a K-split wave +#: always issues a step it masks off and throws away. `wo_b`'s K=2048 is sixteen +#: and every wave count divides it -- `EXACT` in `kernels_gemv.py` -- which is +#: where its flat 6-12% comes from. M=16 on `wq_b` is a tie within run-to-run +#: noise and is reported as one. +#: +#: **Both rows are tuned shapes, and the margin does not survive without that.** +#: Measured across all twelve of the checkpoint's shapes, everything that falls +#: back to `_HEURISTIC` lands inside +-2% -- run-to-run noise -- except `wo_a`, +#: which loses 12-16% at M >= 8 on both TP degrees, and `wq_b` at TP8, which +#: loses 14% at M=32. `wo_a` is the only shape here with N <= 2048 *and* +#: K >= 4096, so the heuristic's 4-wave 16x16 tile has both few N tiles to +#: spread over and a long K to walk. Sweep a shape that matters +#: (`sweep.py gemv-tune`, about two minutes a bucket) rather than assuming it +#: inherits this. +#: +#: This is also the fp8-input comparison. sglang quantises a bf16 activation +#: inside its kernel where mori needs a separate ~2us pass, so on these two +#: shapes a bf16 caller is better off with sglang; on shapes where sglang's +#: fusion has to redo that work per workgroup, it is not. The caller decides -- +#: see `mori_mxfp8_gemm.py`'s `_GEMV_MAX_M`. +_TUNED: dict[tuple[int, int, int], dict] = {} + + +def _tune(n: int, k: int, **by_bucket: str) -> None: + import re + + for bucket, key in by_bucket.items(): + g = re.fullmatch(r"w(\d+)s(\d+)r(\d+)t(\d+)([kn])", key) + _TUNED[(n, k, int(bucket[1:]))] = { + "waves": int(g[1]), + "steps": int(g[2]), + "rows": int(g[3]), + "tokens": int(g[4]), + "ksplit": g[5] == "k", + } + + +# V4.1-Flash's two attention shapes, per rank at TP4. +_tune( + 8192, + 1280, # wq_b, ColumnParallel + m1="w4s4r16t16k", + m2="w4s4r16t16k", + m4="w4s4r16t16k", + m8="w4s4r16t16k", + m16="w4s4r32t16k", + m32="w16s2r32t32k", +) +_tune( + 5120, + 2048, # wo_b, RowParallel + m1="w16s1r32t16k", + m2="w16s2r32t16k", + m4="w16s4r32t16k", + m8="w16s4r32t16k", + m16="w16s2r32t16k", + m32="w16s4r32t32k", +) + +#: What to run for an untuned shape. Four waves splitting K, one 16x16 tile per +#: wave: the sweep's best at every bucket it has covered so far, and the only +#: shape of config that does not either starve the grid (no ksplit) or carry +#: token tiles nothing fills (tokens=32 below M=17). +_HEURISTIC = {"waves": 4, "steps": 2, "rows": 16, "tokens": 16, "ksplit": True} + + +def m_bucket(m: int) -> int: + """The smallest tuned bucket that covers ``m``.""" + for b in M_BUCKETS: + if m <= b: + return b + raise ValueError(f"M={m} exceeds the skinny kernel's {MAX_TOKENS} tokens") + + +def supports_gemv(n: int, k: int) -> bool: + """Whether this shape is expressible. Profitability is the caller's call.""" + return k % STEP_K == 0 and n % MXFP8_BLOCK == 0 and n % TILE == 0 + + +def select_config(n: int, k: int, m: int) -> dict: + cfg = dict(_TUNED.get((n, k, m_bucket(m)), _HEURISTIC)) + # A config's token tile is also the most tokens it can serve, so a bucket + # above 16 has to widen it whatever the table says. + if cfg["tokens"] < m: + cfg["tokens"] = MAX_TOKENS + return cfg + + +class Mxfp8GemvOp: + """A compiled skinny mxfp8 matmul for one ``(N, K)``, M up to 32. + + :: + + op = Mxfp8GemvOp(n=8192, k=1280) + out = op(x_fp8, w_preshuffled, x_scale, w_scale) # [M, N] bf16 + + ``w_preshuffled`` is a weight through :func:`~mori.ops.gemm_ar.preshuffle_b`, + shared with :class:`~mori.ops.gemm_ar.Mxfp8GemmOp` so a server shuffles once + and both paths read it. **The scales are not shared**: both go in as the + checkpoint stores them, row-major ue8m0 bytes, ``[M, K/32]`` and + ``[N/32, K/32]``. The GEMM's `preshuffle_a_scale` exists to coalesce sixteen + lanes reading sixteen *rows* of one K block; here those lanes are sixteen + tokens, M is at most 32, and the whole A scale is well under a kilobyte. + + One compiled kernel per M bucket, at most six. Unlike the GEMM, M is *not* + a pure runtime argument -- the token tile is a compile-time tile width -- but + the buckets are a fixed ladder, so the compiles are bounded and a server + cannot drive an unbounded cache with its batch size. + """ + + def __init__(self, *, n: int, k: int): + if not supports_gemv(n, k): + raise ValueError( + f"unsupported shape N={n} K={k}: K must be a multiple of " + f"{STEP_K} and N of {MXFP8_BLOCK}" + ) + self.n, self.k = n, k + self._launch: dict[int, _PinnedLaunch] = {} + + def _compiled(self, m: int) -> _PinnedLaunch: + bucket = m_bucket(m) + hit = self._launch.get(bucket) + if hit is not None: + return hit + cfg = select_config(self.n, self.k, m) + hit = _PinnedLaunch( + compile_mxfp8_gemv(n=self.n, k=self.k, m_max=cfg["tokens"], **cfg) + ) + self._launch[bucket] = hit + return hit + + def _check(self, x_fp8, w, x_scale, w_scale, m: int) -> None: + """Reject what the launch would otherwise reinterpret as raw bytes.""" + for name, t, shape in ( + ("x_fp8", x_fp8, (m, self.k)), + ("w_preshuffled", w, (self.n, self.k)), + ): + if t.dim() != 2 or tuple(t.shape) != shape: + raise ValueError(f"{name} must be {shape}, got {tuple(t.shape)}") + if t.dtype not in FP8_DTYPES: + raise ValueError( + f"{name} must be one of {[str(d) for d in FP8_DTYPES]}, got " + f"{t.dtype}; the launch reads it as bytes, so a wider dtype " + f"runs and returns nonsense rather than failing" + ) + if t.device.type != "cuda": + raise ValueError(f"{name} must be on a GPU, got {t.device}") + for name, t, shape in ( + ("x_scale", x_scale, (m, self.k // MXFP8_BLOCK)), + ("w_scale", w_scale, (self.n // MXFP8_BLOCK, self.k // MXFP8_BLOCK)), + ): + if t.dtype != torch.uint8: + raise ValueError(f"{name} must be torch.uint8 ue8m0, got {t.dtype}") + if t.dim() != 2 or tuple(t.shape) != shape: + raise ValueError(f"{name} must be {shape}, got {tuple(t.shape)}") + if t.device.type != "cuda": + raise ValueError(f"{name} must be on a GPU, got {t.device}") + + def __call__( + self, + x_fp8: torch.Tensor, + w_preshuffled: torch.Tensor, + x_scale: torch.Tensor, + w_scale: torch.Tensor, + out: torch.Tensor | None = None, + ) -> torch.Tensor: + """``x_fp8 @ w_preshuffled.T`` with 32-wide ue8m0 scales, into bf16. + + No padding and none needed: M is a runtime argument, the token tile + masks its own stores, and a token row past M is clamped to row 0 rather + than read out of bounds. + """ + m = x_fp8.shape[0] + if m > MAX_TOKENS: + raise ValueError( + f"M={m} exceeds the skinny kernel's {MAX_TOKENS} tokens; use " + f"Mxfp8GemmOp above it" + ) + self._check(x_fp8, w_preshuffled, x_scale, w_scale, m) + if out is None: + out = torch.empty((m, self.n), dtype=torch.bfloat16, device=x_fp8.device) + elif tuple(out.shape) != (m, self.n) or out.dtype != torch.bfloat16: + raise ValueError( + f"out must be {(m, self.n)} bfloat16, got {tuple(out.shape)} " + f"{out.dtype}" + ) + self._compiled(m)( + w_preshuffled.contiguous().view(torch.int32).view(-1), + w_scale.contiguous().view(torch.int32).view(-1), + x_fp8.contiguous().view(torch.int32).view(-1), + x_scale.contiguous().view(torch.int32).view(-1), + out.view(-1), + m, + self.n, + stream=fx.Stream(torch.cuda.current_stream()), + ) + return out diff --git a/python/mori/ops/gemm_ar/kernels_gemv.py b/python/mori/ops/gemm_ar/kernels_gemv.py new file mode 100644 index 000000000..1fd39256e --- /dev/null +++ b/python/mori/ops/gemm_ar/kernels_gemv.py @@ -0,0 +1,454 @@ +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +"""The mxfp8 matmul for a handful of tokens -- decode's shape, not prefill's. + +`kernels_fused.py`'s GEMM is built for a full prefill batch and stops paying at +small M for two structural reasons, neither of which a tile-size knob reaches. +Its smallest tile is 256 rows by 128 columns, so at M=64 on `N=8192` the grid is +`ceildiv(64,256) * ceildiv(8192,128) = 64` workgroups on a 256-CU part -- most +of the machine idle -- and of the 256 rows each workgroup computes, 64 are real. +And it stages both operands through LDS, which buys nothing when A is 64 rows +and fits in L1 outright. + +So this is a different kernel rather than a tuning of that one, and it inverts +the same instruction: **the weight is the MFMA's A operand and the tokens are +its B operand**, because the weight is what there is a lot of. A wave owns +`ROWS` columns of the output and all `TOKENS` of them, streams the weight from +global straight into registers, and never touches LDS except to reduce. The +memory traffic is then exactly the weight, once, which is the floor for this +shape -- at M=1 the arithmetic is 2 FLOP per weight byte. + +Two ways to fill the machine, both compile-time (`ksplit`): + +- off: one workgroup covers `waves * ROWS` output columns, each wave the whole + of K. No reduction, no LDS, but the grid is `N / (16*AT*waves)`, which on + `wq_b` with four waves is 128 workgroups -- still half the part. +- on: one workgroup covers `ROWS` columns and its waves split K between them, + reducing through LDS. The grid multiplies by `waves` and each wave does + `1/waves` of the work. + +Operands are the **natural** layouts, not the GEMM's. The weight goes through +`preshuffle_b` (shared with the GEMM, so a server shuffles once), but both +scales are read as the checkpoint stores them -- `[N/32, K/32]` and +`[M, K/32]`, row-major ue8m0 bytes. The GEMM needs `preshuffle_a_scale` because +its sixteen lanes want sixteen *rows* of one K block; here the sixteen lanes of +a block group are sixteen *tokens*, M is at most 32, and the whole A scale is +under a kilobyte, so the layout cannot pay for a pass over it. + +Lane mapping of `v_mfma_scale_f32_16x16x128_f8f6f4` is the one measured in +sglang's `mxfp8_gemv_gfx95.cuh` and already relied on by `_Mxfp8ScaleK`: +lane ``16*s + r`` supplies the ue8m0 scale of 32-block ``s`` of row ``r``, and +``acc[r]`` of lane ``l`` is ``D[4*(l//16) + r][l % 16]``. +""" + +# NOTE: no `from __future__ import annotations` here, deliberately, for the +# reason kernels_fused.py states at the same spot: `fx.struct` reads the +# `SharedStorage` field annotations as live objects, and PEP 563 would hand it +# the string "fx.Array[fx.Float32, RED_FLOATS, 16]", whose size operand is a +# local of this factory and so cannot be resolved from the module globals. + +import flydsl.compiler as flyc +import flydsl.expr as fx +from flydsl.expr import arith, const_expr, range_constexpr +from flydsl.expr.typing import Vector as Vec + +from ._gemm_a8w8_8wave import Mfma16x16x128, ceildiv, pack_i32x4_i32x8 + +#: K per MFMA, and so per step. +STEP_K = 128 +#: ue8m0 block size along K, on both operands. +MXFP8_BLOCK = 32 +#: Rows of one MFMA tile, on both operands. +TILE = 16 +#: `preshuffle_b` emits 16 rows x 64 K per 1024-byte block -- see its docstring. +SHUF_BLOCK_K = 64 + + +class _Gemv: + """One wave's share of the work, as loads and MFMAs over a step range. + + Split out of the kernel body only so the two `ksplit` shapes share it; it + holds the buffer views and the lane's fixed offsets, which are every + address in the mainloop bar the step. + """ + + def __init__(self, W, WS, X, XS, *, n, k, at, bt, m_max): + self.k, self.at, self.bt = k, at, bt + self.nsteps = k // STEP_K + #: Scale rows are K/32 bytes; four consecutive blocks are one dword, and + #: one MFMA step consumes exactly those four. So a step is one dword and + #: the byte select is the lane's block index. + self.k_dwords = k // STEP_K + self.lane = fx.thread_idx.x % 64 + self.row = self.lane % TILE # output column within a tile / token + self.g = self.lane // TILE # this lane's 32-block within the step + + w_bytes = n * k + x_bytes = m_max * k + ws_bytes = (n // MXFP8_BLOCK) * (k // MXFP8_BLOCK) + xs_bytes = m_max * (k // MXFP8_BLOCK) + self.w = self._div(W, w_bytes) + self.x = self._div(X, x_bytes) + self.ws = self._div(WS, ws_bytes) + self.xs = self._div(XS, xs_bytes) + #: A dword index at or past this reads outside the buffer descriptor's + #: records, and `buffer_load` returns zero. That is how a masked step + #: costs a `v_cndmask` instead of a branch: zero data times any scale + #: adds zero to the accumulator. + self.w_oob = w_bytes // 4 + self.ws_oob = ws_bytes // 4 + self.x_oob = x_bytes // 4 + + self.atom_1 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) + self.atom_4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Int32) + self.reg_1 = fx.make_rmem_tensor(fx.make_layout(1, 1), fx.Int32) + self.reg_4 = fx.make_rmem_tensor(fx.make_layout(4, 1), fx.Int32) + self.byte_shift = self.g * fx.Int32(8) + + @staticmethod + def _div(t, nbytes): + gt = fx.rocdl.make_buffer_tensor(t, max_size=False, num_records_bytes=nbytes) + return fx.logical_divide(gt, fx.make_layout(1, 1)) + + def _load1(self, div, index): + fx.copy(self.atom_1, fx.slice(div, (None, fx.Int32(index))), self.reg_1) + return Vec(fx.memref_load_vec(self.reg_1))[0] + + def _load4(self, div, index): + fx.copy(self.atom_4, fx.slice(div, (None, fx.Int32(index))), self.reg_4) + return Vec(fx.memref_load_vec(self.reg_4)) + + def _byte(self, div, index): + """One ue8m0 scale out of a row-major dword of four.""" + v = self._load1(div, index) + return arith.andi(arith.shrui(v, self.byte_shift), fx.Int32(0xFF)) + + def w_frag(self, tile, t, step_dw): + """The 32 weight bytes of tile ``tile + t`` this lane feeds one MFMA. + + `preshuffle_b` lays a 16x64 block out as ``KLane(4) NLane(16) KPack(16)`` + so within a 1024-byte block this lane's 16 bytes sit at + ``g*256 + row*16`` -- sixteen lanes of a block group covering 256 + contiguous bytes, which is the coalescing the layout exists for. A + 128-wide MFMA step is two such blocks, hence the second load 1024 bytes + on. + """ + base = (tile + t) * fx.Int32(self.k // SHUF_BLOCK_K * 256) + step_dw + off = self.g * fx.Int32(64) + self.row * fx.Int32(4) + lo = self._load4(self.w, base + off) + hi = self._load4(self.w, base + fx.Int32(256) + off) + return pack_i32x4_i32x8(lo, hi) + + def x_frag(self, tok, step_dw): + """The 32 activation bytes of token ``tok`` for this step. + + The MFMA wants K ``[32*(g//2) + 16*(g%2), +16)`` and the same 64 on: + both simplify to ``g*16``, so the fp8 activation is read where it lies, + row-major, no shuffle. + """ + base = tok * fx.Int32(self.k // 4) + step_dw + self.g * fx.Int32(4) + lo = self._load4(self.x, base) + hi = self._load4(self.x, base + fx.Int32(16)) + return pack_i32x4_i32x8(lo, hi) + + def w_scale_base(self, tile, t): + """Dword index of this lane's weight-scale row. Loop-invariant.""" + n_group = ((tile + t) * fx.Int32(TILE) + self.row) // fx.Int32(MXFP8_BLOCK) + return n_group * fx.Int32(self.k_dwords) + + def x_scale_base(self, tok): + """Dword index of this token's activation-scale row. Loop-invariant.""" + return tok * fx.Int32(self.k_dwords) + + def scale_at(self, div, base, step): + """The ue8m0 scale of this lane's 32-block at ``step``. + + Four consecutive blocks are one dword and one MFMA step consumes exactly + those four, so the step *is* the dword index and the byte select is the + lane's block within the step. + """ + return self._byte(div, base + step) + + +def compile_mxfp8_gemv( + *, + n: int, + k: int, + m_max: int = 32, + waves: int = 4, + steps: int = 2, + rows: int = 16, + tokens: int = 32, + ksplit: bool = True, +): + """A skinny mxfp8 matmul: ``out[M, N] = X[M, K] @ W[N, K].T``, M <= ``m_max``. + + ``rows`` is the output columns one wave owns and ``tokens`` the token rows; + both are 16 or 32, being whole MFMA tiles. ``steps`` is how many 128-wide K + steps are loaded before any of them is multiplied, which is the only knob + that trades registers for latency hiding. + """ + if k % STEP_K: + raise ValueError(f"K={k} must be a multiple of {STEP_K}") + if n % TILE: + raise ValueError(f"N={n} must be a multiple of {TILE}") + if n % MXFP8_BLOCK: + raise ValueError(f"N={n} must be a multiple of {MXFP8_BLOCK} (the B scale group)") + if rows not in (16, 32) or tokens not in (16, 32): + raise ValueError(f"rows/tokens must be 16 or 32, got {rows}/{tokens}") + if tokens < m_max: + raise ValueError(f"tokens={tokens} cannot cover m_max={m_max}") + if waves not in (4, 8, 16): + raise ValueError(f"waves must be 4, 8 or 16, got {waves}") + + AT = rows // TILE + BT = tokens // TILE + NSTEPS = k // STEP_K + #: Steps one wave runs. With ksplit the waves divide K; without, each runs + #: all of it and they divide N instead. + CHUNK = ceildiv(NSTEPS, waves) if ksplit else NSTEPS + N_ITER = ceildiv(CHUNK, steps) + #: Whether every emitted step is a real one. It is whenever `waves` divides + #: `nsteps` -- `wo_b` at K=2048 is 16 steps and all three wave counts divide + #: it -- and then the out-of-range test is a compile-time constant and the + #: mainloop carries no masking at all. `wq_b` at K=1280 is 10 steps, which + #: none of 4, 8, 16 divides, so it pays one select a step. + EXACT = (waves * CHUNK == NSTEPS) if ksplit else True + #: Output columns one workgroup covers. + WG_ROWS = rows if ksplit else rows * waves + RED_WAVES = waves if ksplit else 1 + RED_FLOATS = RED_WAVES * AT * BT * TILE * TILE + BLOCK = waves * 64 + + _kname = ( + f"mori_mxfp8_gemv_n{n}k{k}_w{waves}s{steps}r{rows}t{tokens}" + f"{'_ks' if ksplit else ''}" + ) + + @fx.struct + class SharedStorage: + red: fx.Array[fx.Float32, RED_FLOATS, 16] + + @flyc.kernel(name=_kname, known_block_size=[BLOCK, 1, 1]) + def kernel_gemv( + W: fx.Tensor, + WS: fx.Tensor, + X: fx.Tensor, + XS: fx.Tensor, + C: fx.Tensor, + c_m: fx.Int32, + c_n: fx.Int32, + ): + wave = fx.thread_idx.x // 64 + gv = _Gemv(W, WS, X, XS, n=n, k=k, at=AT, bt=BT, m_max=m_max) + lane_row, g = gv.row, gv.g + + if const_expr(ksplit): + tile = fx.block_idx.x * fx.Int32(AT) + step0 = wave * fx.Int32(CHUNK) + else: + tile = (fx.block_idx.x * fx.Int32(waves) + wave) * fx.Int32(AT) + step0 = fx.Int32(0) + + # A token row past M reads token 0 and is simply never stored. Clamping + # rather than branching keeps the wave converged, and the extra MFMA is + # already being issued for the full 16-token tile whatever M is. + toks = [ + arith.select( + lane_row + fx.Int32(TILE * b) < c_m, + lane_row + fx.Int32(TILE * b), + fx.Int32(0), + ) + for b in range_constexpr(BT) + ] + + mfma = Mfma16x16x128(AT, BT) + acc = [mfma.zero_value] * (AT * BT) + + # Scale rows are fixed for the whole mainloop -- only the step moves -- + # so the row arithmetic is hoisted here rather than redone each step. + w_sc_base = [gv.w_scale_base(tile, t) for t in range_constexpr(AT)] + x_sc_base = [gv.x_scale_base(toks[b]) for b in range_constexpr(BT)] + + for it in range_constexpr(N_ITER): + w_frags, x_frags, w_sc, x_sc = [], [], [], [] + for s in range_constexpr(steps): + # Two separate ways a step can be out of range, and they are not + # the same kind of condition. Past this wave's `CHUNK` share is + # known at compile time (the share is a constant and the loop is + # unrolled), so the step is simply not emitted -- waves must + # *partition* K, and one that runs a step belonging to its + # neighbour double-counts it into the reduction. Past K itself + # is runtime and only arises when `waves * CHUNK` overshoots + # `nsteps`, which `EXACT` decides at compile time. + local = it * steps + s + if const_expr(local >= CHUNK): + continue + step = step0 + fx.Int32(local) + sc_step = step + if const_expr(EXACT): + step_dw = step * fx.Int32(512) + x_step_dw = step * fx.Int32(STEP_K // 4) + else: + # **Every** address this step forms has to be dealt with, + # not just the ones that would give a wrong number. Past K + # an index walks off the end of its row, and the operand + # buffers are sized for `m_max` tokens while the caller + # passes M of them -- so the read can be inside the + # descriptor's records and outside the allocation. At M=1 + # there is no next row at all. Whatever comes back then is + # not merely wrong, it is *poisonous*: 0xFF is NaN in both + # ue8m0 and e4m3, and NaN times the zeroed weight is still + # NaN. Leaving the two operand addresses unmasked because + # "the weight is zero anyway" made every `wq_b` output at + # M=1 NaN, and left M=2, 3, 7 and 17 depending on what + # happened to be in memory past the tensor. + in_k = step < fx.Int32(NSTEPS) + step_dw = arith.select( + in_k, step * fx.Int32(512), fx.Int32(gv.w_oob) + ) + x_step_dw = arith.select( + in_k, step * fx.Int32(STEP_K // 4), fx.Int32(gv.x_oob) + ) + # The scales are clamped rather than pushed out, because one + # `v_min` covers the step where a select would be per tile + # -- and a real scale from the last step, against a zero + # weight, is just as harmless as a zero one. + sc_step = arith.select(in_k, step, fx.Int32(NSTEPS - 1)) + w_frags.append( + [gv.w_frag(tile, t, step_dw) for t in range_constexpr(AT)] + ) + x_frags.append( + [gv.x_frag(toks[b], x_step_dw) for b in range_constexpr(BT)] + ) + w_sc.append( + [gv.scale_at(gv.ws, w_sc_base[t], sc_step) for t in range_constexpr(AT)] + ) + x_sc.append( + [gv.scale_at(gv.xs, x_sc_base[b], sc_step) for b in range_constexpr(BT)] + ) + for s in range_constexpr(len(w_frags)): + acc = mfma.call( + w_frags[s], + x_frags[s], + acc, + set_prio=False, + scale_a=w_sc[s], + scale_b=x_sc[s], + ) + + oob = fx.Int32(0x7FFFFFF0) + out_atom = fx.make_copy_atom(fx.rocdl.BufferCopy16b(), fx.BFloat16) + out_reg = fx.make_rmem_tensor(fx.make_layout(1, 1), fx.BFloat16) + gC = fx.rocdl.make_buffer_tensor( + C, max_size=False, num_records_bytes=m_max * n * 2 + ) + c_div = fx.logical_divide(gC, fx.make_layout(1, 1)) + + def store(value, tok, col): + in_range = arith.andi(tok < c_m, col < c_n) + idx = arith.select(in_range, tok * c_n + col, oob) + fx.memref_store_vec(Vec.filled(1, value, fx.BFloat16), out_reg) + fx.copy(out_atom, out_reg, fx.slice(c_div, (None, fx.Int32(idx)))) + + if const_expr(ksplit): + lds = fx.SharedAllocator().allocate(SharedStorage).peek() + for t in range_constexpr(AT): + for b in range_constexpr(BT): + vec = Vec(acc[mfma.idx(t, b)]) + for r in range_constexpr(4): + slot = ( + wave * fx.Int32((AT * BT) * 256) + + fx.Int32((t * BT + b) * 256) + + (g * fx.Int32(4) + fx.Int32(r)) * fx.Int32(TILE) + + lane_row + ) + fx.ptr_store( + vec[r].ir_value(), fx.add_offset(lds.red.ptr, slot) + ) + fx.barrier() + # A fixed wave order, so repeated calls sum identically and a row + # stays batch-invariant. + # The workgroup can be larger than the reduction -- a 16-wave config + # on a 32x16 tile is 1024 threads over 512 accumulators -- so clamp + # the LDS index and drop the surplus threads at the store. Without + # the clamp they read past `red` and write a live output element + # with whatever came back. + slots = AT * BT * 256 + per_thread = ceildiv(slots, BLOCK) + for e in range_constexpr(per_thread): + flat = fx.thread_idx.x + fx.Int32(e * BLOCK) + in_lds = flat < fx.Int32(slots) + idx = arith.select(in_lds, flat, fx.Int32(0)) + total = None + for w in range_constexpr(RED_WAVES): + v = fx.ptr_load( + fx.add_offset(lds.red.ptr, idx + fx.Int32(w * slots)) + ) + v = fx.Float32(v) if not hasattr(v, "to") else v + total = v if total is None else total + v + tb = flat // fx.Int32(256) + t_i = tb // fx.Int32(BT) + b_i = tb % fx.Int32(BT) + i = (flat % fx.Int32(256)) // fx.Int32(TILE) + j = flat % fx.Int32(TILE) + col = arith.select( + in_lds, (tile + t_i) * fx.Int32(TILE) + i, fx.Int32(0x7FFFFFF0) + ) + tok = j + b_i * fx.Int32(TILE) + store(total.to(fx.BFloat16), tok, col) + else: + for t in range_constexpr(AT): + col_base = (tile + fx.Int32(t)) * fx.Int32(TILE) + g * fx.Int32(4) + for b in range_constexpr(BT): + vec = Vec(acc[mfma.idx(t, b)]) + tok = lane_row + fx.Int32(TILE * b) + for r in range_constexpr(4): + store(vec[r].to(fx.BFloat16), tok, col_base + fx.Int32(r)) + + @flyc.jit + def launch_gemv( + W: fx.Tensor, + WS: fx.Tensor, + X: fx.Tensor, + XS: fx.Tensor, + C: fx.Tensor, + c_m: fx.Int32, + c_n: fx.Int32, + stream: fx.Stream, + ): + grid_x = ceildiv(c_n, WG_ROWS) + kernel_gemv( + W, + WS, + X, + XS, + C, + c_m, + c_n, + value_attrs={ + "rocdl.waves_per_eu": 1, + "rocdl.flat_work_group_size": f"{BLOCK},{BLOCK}", + }, + ).launch(grid=(grid_x, 1, 1), block=(BLOCK, 1, 1), stream=stream) + + return launch_gemv diff --git a/tests/python/cco/gemm_ar_op_worker.py b/tests/python/cco/gemm_ar_op_worker.py index a758cfcff..d9ef9a787 100644 --- a/tests/python/cco/gemm_ar_op_worker.py +++ b/tests/python/cco/gemm_ar_op_worker.py @@ -39,11 +39,11 @@ import torch import torch.distributed as dist - from mori.cco import Communicator, UniqueId -from mori.ops.gemm_ar import GemmAllReduceOp, preshuffle_b +from mori.ops.gemm_ar import GemmAllReduceOp, preshuffle_a_scale, preshuffle_b SCALE_BK = 128 +MXFP8_BK = 32 def _setup(): @@ -254,6 +254,74 @@ def case_self_test(op, rank, world, n, k): ) +def _mxfp8_operands(rank: int, m: int, n: int, k: int, salt: int = 0): + """ue8m0 operands, in the layouts DeepSeek-V4.1-Flash's loader produces. + + Exponent bytes rather than arbitrary floats because that is what ue8m0 is: + every scale is exactly a power of two, which is why the scaled MFMA can take + them as instruction operands and apply them losslessly. + """ + g = torch.Generator(device="cuda").manual_seed(4321 + rank + 7919 * salt) + a = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + b = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ea = torch.randint( + 120, 123, (m, k // MXFP8_BK), generator=g, device="cuda", dtype=torch.int32 + ) + eb = torch.randint( + 120, + 123, + (n // MXFP8_BK, k // MXFP8_BK), + generator=g, + device="cuda", + dtype=torch.int32, + ) + return a, b, ea, eb + + +def _mxfp8_reference(world: int, m: int, n: int, k: int) -> torch.Tensor: + acc = torch.zeros(m, n, device="cuda", dtype=torch.float32) + for r in range(world): + a, b, ea, eb = _mxfp8_operands(r, m, n, k) + af, bf = a.float(), b.float() + sav, sbv = torch.exp2(ea.float() - 127.0), torch.exp2(eb.float() - 127.0) + for j in range(k // MXFP8_BK): + sl = slice(j * MXFP8_BK, (j + 1) * MXFP8_BK) + acc += ( + (af[:, sl] @ bf[:, sl].t()) + * sav[:, j : j + 1] + * sbv[:, j].repeat_interleave(MXFP8_BK)[None, :n] + ) + return acc + + +def case_mxfp8(comm, rank, world, m, n, k): + """``quant="mxfp8"`` end to end: DeepSeek-V4.1-Flash's quantisation. + + Its own op rather than a case on the shared one: mxfp8 compiles at + BLOCK_M=256 where blockscale needs 128, so the two cannot share a window -- + the padding granule and the counter slots both come from block_m. + + The A scale goes through ``preshuffle_a_scale`` and the B scale is the + ``[N/32, K/32]`` exponent bytes K-block major, which is exactly what + sglang's ``prepare_mxfp8_native_weight`` leaves on the layer. + """ + with GemmAllReduceOp(comm, n=n, k=k, m_max=m, quant="mxfp8") as op: + op.self_test() + a, b, ea, eb = _mxfp8_operands(rank, m, n, k) + got = op( + a, + preshuffle_b(b), + preshuffle_a_scale(ea), + eb.t().reshape(-1).contiguous(), + ).clone() + _emit( + rank, + case="mxfp8", + rel_l2=_rel_l2(got, _mxfp8_reference(world, m, n, k)), + m=m, + ) + + def case_close_is_idempotent(comm, rank, n, k, m_max): """close() releases, twice is a no-op, and a closed op refuses to run.""" op = GemmAllReduceOp(comm, n=n, k=k, m_max=m_max) @@ -325,7 +393,10 @@ def run_all(comm, rank, world, n, k): case_changing_data(op, rank, world, kw["m"], n, k, kw["calls"]) elif case == "self_test": case_self_test(op, rank, world, n, k) - # close() needs its own ops, so it comes after the shared one is released. + # Both of these need their own op, so they come after the shared one is + # released: mxfp8 compiles at a different BLOCK_M, and close() is about the + # lifecycle. + case_mxfp8(comm, rank, world, world * 256, n, k) case_close_is_idempotent(comm, rank, n, k, 512) diff --git a/tests/python/cco/test_gemm_ar_op.py b/tests/python/cco/test_gemm_ar_op.py index 08d7e5599..f51e7b7dc 100644 --- a/tests/python/cco/test_gemm_ar_op.py +++ b/tests/python/cco/test_gemm_ar_op.py @@ -414,6 +414,180 @@ def _run_worker(world_size: int, case: str, *extra: str, timeout: int = 900): requires_two_gpus = pytest.mark.skipif( torch.cuda.device_count() < 2, reason="needs 2 GPUs" ) +requires_gpu = pytest.mark.skipif( + torch.cuda.device_count() < 1, reason="needs a GPU" +) + + +def _mxfp8_gemm_rel_l2(n: int, k: int, m: int, pad: bool = False) -> float: + """One standalone GEMM against an fp32 reference built from the same bytes.""" + from mori.ops.gemm_ar import Mxfp8GemmOp, preshuffle_a_scale, preshuffle_b + + g = torch.Generator(device="cuda").manual_seed(11) + op = Mxfp8GemmOp(n=n, k=k) + m_pad = op.padded_m(m) if pad else m + a_bf16 = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.bfloat16) + a_in = op.pad_rows(a_bf16, m_pad) if m_pad != m else a_bf16 + a = a_in.to(torch.float8_e4m3fn) + b = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ea = torch.randint( + 120, 123, (m_pad, k // 32), generator=g, device="cuda", dtype=torch.int32 + ) + eb = torch.randint( + 120, 123, (n // 32, k // 32), generator=g, device="cuda", dtype=torch.int32 + ) + got = op( + a, preshuffle_b(b), preshuffle_a_scale(ea), eb.t().reshape(-1).contiguous() + )[:m] + + sav, sbv = torch.exp2(ea.float() - 127.0), torch.exp2(eb.float() - 127.0) + af, bf = a.float(), b.float() + ref = torch.zeros(m_pad, n, device="cuda", dtype=torch.float32) + for j in range(k // 32): + sl = slice(j * 32, (j + 1) * 32) + ref += ( + (af[:, sl] @ bf[:, sl].t()) + * sav[:, j : j + 1] + * sbv[:, j].repeat_interleave(32)[None, :n] + ) + ref = ref[:m] + return ( + torch.linalg.vector_norm(got.float() - ref) / torch.linalg.vector_norm(ref) + ).item() + + +@requires_gpu +@pytest.mark.parametrize( + "n,k,label", [(8192, 1280, "wq_b"), (5120, 2048, "wo_b")] +) +@pytest.mark.parametrize("m", [64, 192, 1024]) +def test_standalone_mxfp8_gemm(n, k, label, m): + """``Mxfp8GemmOp`` at DeepSeek-V4.1-Flash's two attention shapes. + + Single process on purpose: this op has no communicator and no symmetric + window, which is the whole reason it exists -- a ColumnParallelLinear like + wq_b has no all-reduce to fuse with. M=192 is not a multiple of BLOCK_M and + still has to be exact: the grid is ceildiv and the tail block masks. + """ + assert _mxfp8_gemm_rel_l2(n, k, m) < FP8_FLOOR + + +@requires_gpu +def test_standalone_mxfp8_gemm_pads_ragged_m(): + """A ragged M is zero-extended to the packed A scale's 64-row group. + + The padded rows must not disturb the real ones, and the check is numerical + because they would not: a GEMM row depends only on its own input row, so + getting this wrong shifts the *scale* layout rather than raising. + """ + from mori.ops.gemm_ar import Mxfp8GemmOp + + op = Mxfp8GemmOp(n=5120, k=2048) + assert op.padded_m(100) == 128 + assert op.padded_m(64) == 64 + assert _mxfp8_gemm_rel_l2(5120, 2048, 100, pad=True) < FP8_FLOOR + + +@requires_gpu +def test_standalone_mxfp8_gemm_refuses_unpadded_m(): + """M not on the group boundary is refused rather than quietly mis-scaled.""" + from mori.ops.gemm_ar import Mxfp8GemmOp + + op = Mxfp8GemmOp(n=5120, k=2048) + a = torch.zeros(100, 2048, device="cuda", dtype=torch.float8_e4m3fn) + b = torch.zeros(5120, 2048, device="cuda", dtype=torch.float8_e4m3fn) + sa = torch.zeros(100 * 2048 // 128, device="cuda", dtype=torch.int32) + sb = torch.zeros(160 * 64, device="cuda", dtype=torch.int32) + with pytest.raises(ValueError, match="multiple of 64"): + op(a, b, sa, sb) + + +def _mxfp8_gemv_rel_l2(n: int, k: int, m: int) -> float: + """One skinny GEMM against an fp32 reference built from the same bytes.""" + from mori.ops.gemm_ar import Mxfp8GemvOp, preshuffle_b + + g = torch.Generator(device="cuda").manual_seed(11) + x = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ex = torch.randint( + 120, 123, (m, k // 32), generator=g, device="cuda", dtype=torch.int32 + ) + ew = torch.randint( + 120, 123, (n // 32, k // 32), generator=g, device="cuda", dtype=torch.int32 + ) + got = Mxfp8GemvOp(n=n, k=k)( + x, preshuffle_b(w), ex.to(torch.uint8), ew.to(torch.uint8) + ) + + sxv, swv = torch.exp2(ex.float() - 127.0), torch.exp2(ew.float() - 127.0) + xf, wf = x.float(), w.float() + ref = torch.zeros(m, n, device="cuda", dtype=torch.float32) + for j in range(k // 32): + sl = slice(j * 32, (j + 1) * 32) + ref += ( + (xf[:, sl] @ wf[:, sl].t()) + * sxv[:, j : j + 1] + * swv[:, j].repeat_interleave(32)[None, :n] + ) + return ( + torch.linalg.vector_norm(got.float() - ref) / torch.linalg.vector_norm(ref) + ).item() + + +@requires_gpu +@pytest.mark.parametrize("n,k,label", [(8192, 1280, "wq_b"), (5120, 2048, "wo_b")]) +@pytest.mark.parametrize("m", [1, 3, 16, 17, 32]) +def test_standalone_mxfp8_gemv(n, k, label, m): + """``Mxfp8GemvOp`` at V4.1-Flash's two attention shapes, over the M ladder. + + M=3 and M=17 are the ones that matter: they are not tile widths, so the + token rows past M read a clamped row and must not reach the output, and + M=17 additionally crosses from a 16-token config to a 32-token one. + """ + assert _mxfp8_gemv_rel_l2(n, k, m) < FP8_FLOOR + + +@requires_gpu +def test_mxfp8_gemv_refuses_m_above_its_tile(): + """Past 32 tokens it declines rather than silently truncating the batch. + + The B operand is two 16-row MFMA tiles and there is no third, so a 33rd + token has nowhere to go; without this it would compute 32 rows and return a + tensor whose remaining rows were never written. + """ + from mori.ops.gemm_ar import Mxfp8GemvOp + + op = Mxfp8GemvOp(n=5120, k=2048) + x = torch.zeros(33, 2048, device="cuda", dtype=torch.float8_e4m3fn) + w = torch.zeros(5120, 2048, device="cuda", dtype=torch.float8_e4m3fn) + sx = torch.zeros(33, 64, device="cuda", dtype=torch.uint8) + sw = torch.zeros(160, 64, device="cuda", dtype=torch.uint8) + with pytest.raises(ValueError, match="exceeds the skinny kernel"): + op(x, w, sx, sw) + + +def test_gemv_config_widens_the_token_tile_for_the_top_bucket(): + """A tuned 16-token config must not be handed an M it cannot hold. + + The table is keyed by bucket and a bucket's config is free to be narrower + than the bucket; the guard is here rather than in the kernel because the + kernel's rejection would be a compile error on a server's hot path. + """ + from mori.ops.gemm_ar.gemv import m_bucket, select_config + + assert m_bucket(1) == 1 and m_bucket(17) == 32 + assert select_config(8192, 1280, 1)["tokens"] == 16 + assert select_config(8192, 1280, 17)["tokens"] == 32 + + +def test_supports_gemm_takes_both_attention_shapes(): + """Neither of V4.1-Flash's attention GEMMs needs a special case.""" + from mori.ops.gemm_ar import supports_gemm + + assert supports_gemm(8192, 1280) is True # wq_b, ColumnParallel + assert supports_gemm(5120, 2048) is True # wo_b, RowParallel + assert supports_gemm(5120, 64) is False # K below MIN_K + assert supports_gemm(5000, 2048) is False # N not a multiple of BLOCK_N @pytest.fixture(scope="module") @@ -488,6 +662,22 @@ def test_changing_operands_between_calls_stays_correct(worker_results): assert r[key] < FP8_FLOOR, (key, r) +@requires_two_gpus +def test_mxfp8_quant_runs_end_to_end(worker_results): + """``quant="mxfp8"`` through the public op: V4.1-Flash's quantisation. + + The op hardcoded ``quant="blockscale"`` until this, so the mxfp8 kernel -- + tuned and tested at the kernel layer -- had no way out to a caller. What + this covers that the kernel tests do not is the operand contract: the A + scale is ``preshuffle_a_scale``'s packed int32 and the B scale is + ``[N/32, K/32]`` exponent bytes K-block major, and both are int32 where + blockscale's are fp32. Getting either wrong returns finite nonsense rather + than failing, which is why the check is numerical. + """ + r = _case(worker_results, "mxfp8") + assert r["rel_l2"] < FP8_FLOOR, r + + @requires_two_gpus def test_self_test_passes_and_can_fail(worker_results): """The guard against a mori whose SDMA puts were compiled out. diff --git a/tests/python/cco/test_mxfp8_gemv_grid.py b/tests/python/cco/test_mxfp8_gemv_grid.py new file mode 100644 index 000000000..6ac6c90de --- /dev/null +++ b/tests/python/cco/test_mxfp8_gemv_grid.py @@ -0,0 +1,201 @@ +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +"""The skinny GEMM across its compile-time configuration space. + +``test_gemm_ar_op.py`` covers ``Mxfp8GemvOp``, which runs whatever the tuned +table selects -- so it exercises six configurations out of seventy-two and says +nothing about the rest. This is the grid, and it exists because the kernel is a +different program per configuration: ``ksplit`` changes how K is partitioned and +whether the reduction runs at all, and ``waves`` x ``steps`` decides how many +masked steps a wave issues when the wave count does not divide K. + +That is where the bugs were. Every masked address has to be pushed out of its +buffer's records or clamped, not just the ones that would give a wrong number: +0xFF is NaN in both ue8m0 and e4m3, and NaN times a zeroed weight is still NaN. +Leaving two operand addresses unmasked made every ``wq_b`` output at M=1 NaN and +left M=2, 3, 7 and 17 depending on what happened to sit past the tensor. + +**Each configuration runs in its own process.** A FlyDSL compile failure takes +the interpreter with it, so a shared one would lose the rest of the grid. + +This replaced a shell script that printed ``fail=N`` and exited 0 -- verified +with a probe rigged to fail every case, the script still reported success. A +correctness sweep that cannot fail its caller is not a check. +""" + +from __future__ import annotations + +import argparse +import json +import os +import subprocess +import sys + +import pytest + +#: DeepSeek-V4.1-Flash's two attention shapes, per rank at TP4. K=1280 is ten +#: 128-wide steps and K=2048 is sixteen, which is the distinction that matters: +#: no wave count divides ten, so `wq_b` exercises the masked-step path on every +#: ksplit configuration and `wo_b` exercises none of it. +SHAPES = {"wq_b": (8192, 1280), "wo_b": (5120, 2048)} + +#: (waves, steps, rows, tokens). A representative slice of the 72-entry space: +#: both `rows` and both `tokens`, every `waves`, and `steps` 1/2/4 so the +#: leftover-step arithmetic is covered at each unroll. +CONFIGS = [ + (4, 1, 16, 16), + (4, 2, 16, 32), + (4, 4, 32, 32), + (8, 2, 32, 32), + (8, 4, 16, 32), + (16, 1, 32, 16), +] + +#: M values that are not tile widths. 3 and 17 are the point: rows past M read a +#: clamped token and must not reach the output, and 17 crosses from a 16-token +#: configuration to a 32-token one. +M_VALUES = [1, 2, 3, 8, 16, 17, 32] + +FP8_FLOOR = 2.4e-3 + +requires_gpu = pytest.mark.skipif( + not os.environ.get("MORI_TEST_GPU", "1") == "1", + reason="needs a gfx950 GPU", +) + + +def _cases(): + for shape in SHAPES: + for waves, steps, rows, tokens in CONFIGS: + for m in M_VALUES: + if m > tokens: + continue # the token tile is also the largest M it can serve + for ksplit in (0, 1): + yield shape, m, waves, steps, rows, tokens, ksplit + + +ALL_CASES = list(_cases()) +#: The full grid is ~150 subprocess compiles. Default to every configuration at +#: the M values that broke before, and take the rest under MORI_TEST_GEMV_FULL. +QUICK = [c for c in ALL_CASES if c[1] in (1, 3, 17, 32)] +CASES = ALL_CASES if os.environ.get("MORI_TEST_GEMV_FULL") == "1" else QUICK + + +def _run_case(shape, m, waves, steps, rows, tokens, ksplit): + """One configuration, in its own interpreter. Returns (rc, payload|text).""" + p = subprocess.run( + [sys.executable, os.path.abspath(__file__), "--worker", + "--shape", shape, "-m", str(m), "--waves", str(waves), + "--steps", str(steps), "--rows", str(rows), "--tokens", str(tokens), + "--ksplit", str(ksplit)], + capture_output=True, text=True, timeout=900, + ) + for line in p.stdout.splitlines(): + if line.startswith("RESULT_JSON"): + return p.returncode, json.loads(line.split(" ", 1)[1]) + return p.returncode, (p.stdout + p.stderr)[-2000:] + + +@requires_gpu +@pytest.mark.parametrize( + "shape,m,waves,steps,rows,tokens,ksplit", CASES, + ids=[f"{s}-m{m}-w{w}s{st}r{r}t{t}{'k' if ks else 'n'}" + for s, m, w, st, r, t, ks in CASES], +) +def test_gemv_config(shape, m, waves, steps, rows, tokens, ksplit): + """Every configuration must agree with an fp32 reference, not just the tuned one.""" + rc, got = _run_case(shape, m, waves, steps, rows, tokens, ksplit) + assert rc == 0, f"worker exited {rc}:\n{got}" + assert isinstance(got, dict), f"worker produced no result:\n{got}" + assert got["finite"], f"non-finite output at {got}" + assert got["rel_l2"] < FP8_FLOOR, got + + +# -------------------------------------------------------------------------- +# worker: one configuration, compiled and checked against fp32 +# -------------------------------------------------------------------------- + + +def _worker(args) -> int: + import torch + from mori.ops.gemm_ar import preshuffle_b + from mori.ops.gemm_ar.kernels_gemv import compile_mxfp8_gemv + import flydsl.expr as fx + + n, k = SHAPES[args.shape] + m, bk = args.m, 32 + g = torch.Generator(device="cuda").manual_seed(1234) + # Allocate exactly M rows, never the config's token tile: the operand + # buffers are sized for `m_max` and a masked step that walks off the last + # row must not reach past the allocation. At M=1 there is no next row. + x = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ex = torch.randint(120, 123, (m, k // bk), generator=g, + device="cuda", dtype=torch.int32) + ew = torch.randint(120, 123, (n // bk, k // bk), generator=g, + device="cuda", dtype=torch.int32) + + gemv = compile_mxfp8_gemv( + n=n, k=k, m_max=args.tokens, waves=args.waves, steps=args.steps, + rows=args.rows, tokens=args.tokens, ksplit=bool(args.ksplit), + ) + out = torch.zeros(m, n, device="cuda", dtype=torch.bfloat16) + gemv( + preshuffle_b(w).contiguous().view(torch.int32).view(-1), + ew.to(torch.uint8).contiguous().view(torch.int32).view(-1), + x.contiguous().view(torch.int32).view(-1), + ex.to(torch.uint8).contiguous().view(torch.int32).view(-1), + out.view(-1), m, n, + stream=fx.Stream(torch.cuda.current_stream()), + ) + torch.cuda.synchronize() + + sx, sw = torch.exp2(ex.float() - 127.0), torch.exp2(ew.float() - 127.0) + xf, wf = x.float(), w.float() + ref = torch.zeros(m, n, device="cuda", dtype=torch.float32) + for i in range(k // bk): + ks = slice(i * bk, (i + 1) * bk) + ref += ((xf[:, ks] @ wf[:, ks].T) * sx[:, i][:, None] + * sw[:, i].repeat_interleave(bk)[None, :]) + got = out.float() + rel = (torch.linalg.vector_norm(got - ref) + / torch.linalg.vector_norm(ref)).item() + print("RESULT_JSON " + json.dumps({ + "shape": args.shape, "n": n, "k": k, "m": m, + "config": f"w{args.waves}s{args.steps}r{args.rows}t{args.tokens}" + f"{'k' if args.ksplit else 'n'}", + "rel_l2": rel, "finite": bool(got.isfinite().all()), + })) + return 0 + + +if __name__ == "__main__": + p = argparse.ArgumentParser() + p.add_argument("--worker", action="store_true") + p.add_argument("--shape", choices=sorted(SHAPES)) + p.add_argument("-m", type=int) + p.add_argument("--waves", type=int) + p.add_argument("--steps", type=int) + p.add_argument("--rows", type=int) + p.add_argument("--tokens", type=int) + p.add_argument("--ksplit", type=int) + raise SystemExit(_worker(p.parse_args())) From ad271d199e38b00761dbaf61caafb937f525d8e5 Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 05:33:43 +0000 Subject: [PATCH 3/8] benchmark, docs: one timer, three entry points, and the traps that cost most The measurements, and the harness that has to be right for them to mean anything. Most of the work here was finding out that it was not. **Three measurement traps, each of which produced a confident wrong number.** A single-call CUDA-graph replay has a ~13.4us floor on this box, which is most of any small-M measurement -- `timing.py` amortises the capture. A repeated call reads the weight out of the 256MB LLC at 1.7x the bandwidth a forward pass gets, so cold is the number that predicts a server -- it rotates copies past the cache. And getting *that* right took two more corrections: the ring advances at capture time, so a graph of `reps` calls only ever touches `reps` copies; and the ring has to be sized so each tensor clears the cache on its own, not so their sum does, because a caller may hand over several weights of which the kernel reads one. Each is documented where it bit, because the symptom is always a plausible number rather than a failure. The worst of them was in the comparison itself: the baseline's *bf16* weight was never rotated, while `native_route_plan` picks `hipblaslt_bf16` -- which reads exactly that tensor -- on most large-M points. Those rows measured a hot baseline against a cold mori, which ran against mori, not for it. Corrected, SGLang is 3.6-18.6% slower there and mori's margin moves up to 28 points in its favour. The rows whose route never reads that tensor are unchanged within a point, which is how the fix was confirmed to touch only what it should. **Failure has to reach the caller.** A correctness sweep that printed `fail=N` and exited 0 -- verified with a probe rigged to fail all 112 cases -- is now a pytest grid. Every benchmark returns non-zero, `sweep.py` exits with the count of failed points, and "mori declines this shape by construction" is recorded as a result rather than a crash, so ten expected declines cannot bury one real break. **Six entry points, not seventeen.** `bench_gemm.py` covers the two scopes that are genuinely different questions -- `--scope kernel` for pre-quantised operands and `--scope linear` for a whole layer -- where three scripts had encoded the distinction implicitly. Six shell sweeps and a tuning script become `sweep.py` presets; two reporters become `report.py`, which keys on every axis a sweep declares varying and so cannot silently overwrite one configuration with another. The SGLang baseline is optional throughout: mori's own numbers need nothing but mori. One historical experiment is kept and renamed: `fused-blockscale-control` ran V4.1-Flash's shapes under a quantisation the checkpoint does not use, while presenting itself as a V4.1 benchmark. `--mode split-lsa --gather-dtype fp8` is refused where it is asked for. The LSA all-reduce has no fp8 leg, and the combination used to run a bf16 collective and report it under the fp8 label; the validation gate's two-sided band caught it, but only after paying for the run. `docs/` keeps how-to-run and a summary; every full table lives in the operator README, next to the code it is about. Co-Authored-By: Claude Opus 5 (1M context) --- benchmark/cco/flydsl/gemm_ar/bench_gemm.py | 391 +++++ benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py | 82 +- benchmark/cco/flydsl/gemm_ar/bench_gemv.py | 230 +++ benchmark/cco/flydsl/gemm_ar/report.py | 171 +++ benchmark/cco/flydsl/gemm_ar/sweep.py | 311 ++++ benchmark/cco/flydsl/gemm_ar/timing.py | 173 +++ docs/MORI-GEMM-AR-BENCHMARK.md | 541 +++---- python/mori/ops/gemm_ar/README.md | 1323 ++++++++++++++--- 8 files changed, 2639 insertions(+), 583 deletions(-) create mode 100644 benchmark/cco/flydsl/gemm_ar/bench_gemm.py create mode 100644 benchmark/cco/flydsl/gemm_ar/bench_gemv.py create mode 100644 benchmark/cco/flydsl/gemm_ar/report.py create mode 100644 benchmark/cco/flydsl/gemm_ar/sweep.py create mode 100644 benchmark/cco/flydsl/gemm_ar/timing.py diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemm.py b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py new file mode 100644 index 000000000..2d6d15f12 --- /dev/null +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py @@ -0,0 +1,391 @@ +#!/usr/bin/env python3 +"""mori's mxfp8/blockscale GEMM, at either scope, optionally against SGLang. + +One entry point for the two questions that are *not* about the collective, and +they are different questions rather than two views of one: + +``--scope kernel`` + The operands are already quantised and already padded. This is the multiply + and nothing else, which is what a tile or a scale layout is chosen on. + +``--scope linear`` + bf16 in, bf16 out: quantisation, padding and the op's own tile dispatch are + all inside the measurement. This is what a linear layer actually costs, and + the only scope in which comparing against SGLang means anything, because + SGLang's route pays its own quantisation too. + +``--impl`` picks what runs. ``gemm256`` and ``gemm128`` pin mori's N tile so the +dispatch itself can be measured rather than trusted; ``auto`` leaves the op to +choose. ``sglang`` is the baseline and is **optional** -- mori's own numbers +need nothing installed but mori. + + python bench_gemm.py --shape wq_b -m 4096 --scope kernel + python bench_gemm.py --shape wq_b -m 4096 --scope linear --impl auto,sglang + python bench_gemm.py -n 5120 -k 2048 -m 4096 --quant blockscale --scope kernel +""" + +from __future__ import annotations + +import argparse +import json +import sys +import traceback +from pathlib import Path + +import flydsl.expr as fx +import torch + +sys.path.insert(0, str(Path(__file__).parent)) +import timing # noqa: E402 + +MXFP8_BK = 32 +SCALE_BK = 128 + +#: Every fp8 linear DeepSeek-V4.1-Flash has, read off the checkpoint and split +#: by the parallelism each is declared with. `shared_down_tp4` (5120x576) is +#: absent because SGLang's own route refuses it -- K is not a multiple of 128 -- +#: so there is nothing to compare against. +SHAPES = { + "wq_b": (8192, 1280), + "wo_b": (5120, 2048), + "wq_a_tp4": (1280, 5120), + "wkv_tp4": (512, 5120), + "wqkv_a_tp4": (1792, 5120), + "wo_a_tp4": (2048, 4096), + "shared_gate_up_tp4": (1152, 5120), + "wq_b_tp8": (4096, 1280), + "wo_b_tp8": (5120, 1024), + "wo_a_tp8": (1024, 4096), + "wq_b_tp1": (32768, 1280), + "wo_b_tp1": (5120, 8192), +} + +IMPLS = ("auto", "gemm256", "gemm128", "sglang") + + +# -------------------------------------------------------------------------- +# operands +# -------------------------------------------------------------------------- + + +class _Layer: + """Just enough of an SGLang linear for its native route to read.""" + + def __init__(self, weight, weight_scale_mx_e8m0, weight_bf16): + self.weight = weight + self.weight_scale_mx_e8m0 = weight_scale_mx_e8m0 + self.weight_bf16 = weight_bf16 + self.mxfp8_native_ready = True + + +def build_mxfp8(n, k, seed=1234, want_sglang=False): + """One weight, in every layout either side needs, from the same bytes.""" + from mori.ops.gemm_ar import preshuffle_b + + g = torch.Generator(device="cuda").manual_seed(seed) + w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + eb = torch.randint( + 120, 123, (n // MXFP8_BK, k // MXFP8_BK), generator=g, + device="cuda", dtype=torch.int32, + ) + out = { + "w_raw": w, + "mori_w": preshuffle_b(w), + # The GEMM wants the B scale K-block major as dwords; the GEMV wants + # the checkpoint's own [N/32, K/32] bytes. + "b_scale": eb.to(torch.uint8).t().contiguous().to(torch.int32).reshape(-1), + "w_exps": eb.to(torch.uint8), + "layer": None, + } + if want_sglang: + from sglang.kernels.ops.quantization.mxfp8_native_amd_gfx95 import ( + prepare_mxfp8_native_weight, + ) + + shuffled, scale_e8m0, weight_bf16 = prepare_mxfp8_native_weight( + w, torch.exp2(eb.float() - 127.0), (32, 32) + ) + out["layer"] = _Layer(shuffled.view(torch.float8_e4m3fn), scale_e8m0, weight_bf16) + return out + + +def build_blockscale(n, k, m, seed=1234): + """mori's other operand contract: A 1x128, B 128x128, fp32 scales.""" + from mori.ops.gemm_ar import preshuffle_b + + g = torch.Generator(device="cuda").manual_seed(seed) + a = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + kb = k // SCALE_BK + sa = torch.rand(m, kb, generator=g, device="cuda", dtype=torch.float32) * 0.01 + 0.01 + sb = ( + torch.rand(n // SCALE_BK, kb, generator=g, device="cuda", dtype=torch.float32) + * 0.01 + 0.01 + ) + return a, w, preshuffle_b(w), sa.t().contiguous().t(), sb + + +# -------------------------------------------------------------------------- +# implementations +# -------------------------------------------------------------------------- + + +class _Unsupported(Exception): + """mori declines this (shape, tile) by construction, not by accident.""" + + +def forced_gemm_op(n, k, block_n): + """A `Mxfp8GemmOp` pinned to one N tile. + + Two things have to be got around, and both are the op behaving correctly. + `wants_narrow_n` *is* the dispatch under test, so it is overridden rather + than consulted; and the narrow tile is not reachable through `block_n`, + because on the wide path `block_n` also turns on the permlane store, which + pairs two N tiles and needs 256. 128 lives behind the dispatch as + `NARROW_BLOCK_N`. + + The constructor also requires N to divide the *wide* tile even when only the + narrow one will run, which is what disqualifies `shared_gate_up` (N=1152). + Whether the narrow tile could serve it is a thing worth measuring, so for + that case the op is built field by field. Benchmark-only: a caller has no + business doing this. + """ + from mori.ops.gemm_ar import Mxfp8GemmOp + from mori.ops.gemm_ar.gemm import NARROW_BLOCK_N, _gemm_shape_constraint + from mori.ops.gemm_ar.op import DEFAULT_BLOCK_N, MXFP8_BLOCK_M + + if block_n is None: # 'auto': let the op dispatch + try: + return Mxfp8GemmOp(n=n, k=k) + except ValueError as err: + raise _Unsupported(str(err)) from None + + narrow = block_n == NARROW_BLOCK_N + why = _gemm_shape_constraint(n, k, block_n) + if why is not None: + raise _Unsupported(why) + if n % DEFAULT_BLOCK_N == 0: + op = Mxfp8GemmOp(n=n, k=k) + else: + op = Mxfp8GemmOp.__new__(Mxfp8GemmOp) + op.n, op.k = n, k + op.block_m, op.block_n = MXFP8_BLOCK_M, DEFAULT_BLOCK_N + op._launch, op._pad_in = {}, None + op.wants_narrow_n = lambda m, _v=narrow: _v + return op + + +def mori_kernel_call(ops, n, k, m, block_n, quant): + """Pre-quantised, pre-padded operands. The multiply and nothing else. + + Rotates the weight only, because that is the only operand large enough for + residency to matter: A is M x K and the scales are kilobytes. + """ + if quant == "blockscale": + from mori.ops.gemm_ar import layout + from mori.ops.gemm_ar.kernels_fused import compile_fused_gemm_scatter + + a, _w, w_shuf, sa, sb = ops + gemm = compile_fused_gemm_scatter( + layout.ArConfig(world_size=2, m=128, n=n), 0, K=k, + BLOCK_M=128, BLOCK_N=block_n or 256, b_preshuffled=True, fuse=False, + swap_ab=True, permlane=True, lane_transpose=True, quant="blockscale", + ) + y = torch.zeros(m, n, device="cuda", dtype=torch.bfloat16) + a_i8 = a.contiguous().view(torch.int8).view(-1) + sa_arg, sb_arg = sa.t().reshape(-1).contiguous(), sb.reshape(-1).contiguous() + + def call(picked): + gemm(a_i8, picked[0], y.view(-1), sa_arg, sb_arg, m, n, 0, 0, + stream=fx.Stream(torch.cuda.current_stream())) + return y + + return call, [w_shuf.contiguous().view(torch.int8).view(-1)] + + from mori.ops.gemm_ar import preshuffle_a_scale + + op = forced_gemm_op(n, k, block_n) + g = torch.Generator(device="cuda").manual_seed(99) + m_pad = op.padded_m(m) + a = (torch.randn(m_pad, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ea = torch.randint( + 120, 123, (m_pad, k // MXFP8_BK), generator=g, device="cuda", dtype=torch.int32 + ) + a_scale = preshuffle_a_scale(ea) + + def call(picked): + return op(a, picked[0], a_scale, ops["b_scale"])[:m] + + return call, [ops["mori_w"]] + + +def mori_linear_call(ops, n, k, m, block_n, x_bf16): + """bf16 in, bf16 out, through mori: quantise, pad, dispatch, multiply.""" + from sglang.srt.layers.mori_mxfp8_common import quantize_packed + + op = forced_gemm_op(n, k, block_n) + m_pad = op.padded_m(m) + + def call(picked): + x_in = x_bf16 if m_pad == m else op.pad_rows(x_bf16, m_pad) + a_fp8, a_scale = quantize_packed(x_in) + return op(a_fp8, picked[0], a_scale, ops["b_scale"])[:m] + + return call, [ops["mori_w"]] + + +def sglang_linear_call(layer, x_bf16): + """SGLang's native mxfp8 linear. + + **Every weight the chosen route may read is rotated**, not just the fp8 + one. `native_route_plan` picks `hipblaslt_bf16` for most of these shapes, + and that route reads `weight_bf16` -- which is twice the size of the fp8 + weight and therefore the tensor that decides residency. Rotating only the + fp8 copy left the baseline hot while mori was cold, on 76 of 120 points, + and every one of those comparisons flattered mori. + """ + from sglang.kernels.ops.quantization.mxfp8_native_amd_gfx95 import ( + mxfp8_native_blockscaled_linear, + ) + + weights = [layer.weight] + if layer.weight_bf16 is not None: + weights.append(layer.weight_bf16) + + def call(picked): + return mxfp8_native_blockscaled_linear( + x_bf16, picked[0].view(torch.uint8), layer.weight_scale_mx_e8m0, + weight_bf16=(picked[1] if len(picked) > 1 else None), + ) + + return call, weights + + +# -------------------------------------------------------------------------- +# driver +# -------------------------------------------------------------------------- + + +def rel_l2(got, ref): + if got is None or ref is None: + return None + d = torch.linalg.vector_norm(got.float() - ref.float()) + return (d / torch.linalg.vector_norm(ref.float()).clamp_min(1e-30)).item() + + +def main() -> int: + p = argparse.ArgumentParser() + p.add_argument("--shape", choices=sorted(SHAPES), default=None) + p.add_argument("-n", type=int, default=None) + p.add_argument("-k", type=int, default=None) + p.add_argument("-m", type=int, required=True) + p.add_argument("--scope", choices=("kernel", "linear"), default="linear") + p.add_argument("--quant", choices=("mxfp8", "blockscale"), default="mxfp8") + p.add_argument("--impl", default="auto,gemm256,gemm128,sglang", + help="comma-separated: " + ", ".join(IMPLS)) + p.add_argument("--reps", type=int, default=32) + p.add_argument("--tol", type=float, default=2.4e-3, + help="rel_l2 above this marks the row invalid") + p.add_argument("--json-out", default="gemm.jsonl") + args = p.parse_args() + + if args.shape is not None: + n, k = SHAPES[args.shape] + elif args.n and args.k: + n, k = args.n, args.k + args.shape = f"{n}x{k}" + else: + p.error("pass --shape, or -n and -k") + m = args.m + impls = [i.strip() for i in args.impl.split(",") if i.strip()] + for i in impls: + if i not in IMPLS: + p.error(f"unknown --impl {i!r}; pick from {IMPLS}") + + if args.quant == "blockscale" and ("sglang" in impls or args.scope == "linear"): + p.error("--quant blockscale is mori-only and kernel-scope only") + # `linear` means "quantise a bf16 activation", which is SGLang's quantiser. + needs_sglang = "sglang" in impls or args.scope == "linear" + + vram_before = timing.vram_used() + if args.quant == "mxfp8": + ops = build_mxfp8(n, k, want_sglang="sglang" in impls) + else: + ops = build_blockscale(n, k, m) + x = (torch.randn(m, k, device="cuda") / 8).to(torch.bfloat16) + + common = { + "bench": "gemm", "scope": args.scope, "quant": args.quant, + "shape": args.shape, "n": n, "k": k, "m": m, + "input": "bf16" if args.scope == "linear" else "fp8", + "includes_quant": args.scope == "linear", + "timing": "amortized-graph-cold-hot", + } + rows, ref, failures = [], None, 0 + + # SGLang first when present, so it is the reference the rest are scored on. + order = ([i for i in impls if i == "sglang"] + + [i for i in impls if i != "sglang"]) + for impl in order: + row = dict(common, impl=impl) + try: + if impl == "sglang": + if ops["layer"] is None: + raise RuntimeError("--impl sglang needs the sglang build path") + from sglang.kernels.ops.quantization.mxfp8_native_amd_gfx95 import ( + native_route_plan, + ) + row["route"] = native_route_plan( + m, n, k, ops["layer"].weight_bf16 is not None, False + ) + call, weights = sglang_linear_call(ops["layer"], x) + else: + block_n = {"auto": None, "gemm256": 256, "gemm128": 128}[impl] + row["route"] = impl + if args.scope == "kernel": + call, weights = mori_kernel_call(ops, n, k, m, block_n, args.quant) + else: + call, weights = mori_linear_call(ops, n, k, m, block_n, x) + + got = call(weights) + torch.cuda.synchronize() + if ref is None: + ref = got.float().clone() + row["rel_l2"] = rel_l2(got, ref) + row.update(timing.cold_hot_us(call, weights, reps=args.reps)) + row["validated"] = row["rel_l2"] is None or row["rel_l2"] <= args.tol + except _Unsupported as err: + # The op declining a shape it documents as out of range is a + # *result*, not a crash: `shared_gate_up` is N=1152, which is 4.5 + # tiles of 256, so the wide tile cannot be built for it at all. + # Counting that as a failure makes a sweep of known-good shapes + # exit non-zero and buries a real break in the noise. + row.update(validated=None, supported=False, reason=str(err)) + print(f" {impl:<10} unsupported: {err}", flush=True) + except Exception as err: # noqa: BLE001 - one impl must not stop the rest + row.update(validated=False, error=f"{type(err).__name__}: {err}") + failures += 1 + print(f" {impl:<10} FAILED {row['error']}", flush=True) + traceback.print_exc(limit=2) + else: + flag = "" if row["validated"] else " !! rel_l2 over tol" + print( + f" {impl:<10} hot {row['hot_us']:8.2f} cold {row['cold_us']:8.2f}" + f" relL2 {row['rel_l2']:.2e}{flag}", + flush=True, + ) + if not row["validated"]: + failures += 1 + rows.append(row) + + with open(args.json_out, "a") as f: + for r in rows: + f.write(json.dumps( + dict(r, vram_before=vram_before, vram_after=timing.vram_used()) + ) + "\n") + # A benchmark that fails and exits 0 is how a broken sweep looks green. + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py b/benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py index 806a6178f..b1553ad09 100644 --- a/benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py @@ -97,22 +97,21 @@ import flydsl.expr as fx import torch import torch.distributed as dist - from mori.cco import ( + GDA_CONNECTION_NONE, CCODevCommRequirements, Communicator, - GDA_CONNECTION_NONE, UniqueId, ) -from mori.tensor_utils import from_gpu_ptr - from mori.ops.gemm_ar import ( ArConfig, build_lsa_ar, build_sdma_phases, compile_fused_gemm_scatter, + preshuffle_a_scale, preshuffle_b, ) +from mori.tensor_utils import from_gpu_ptr VMM_SLACK = 512 * 1024 * 1024 MODES = ("gemm-only", "split-sdma", "fused-sdma", "fused-lsa", "split-lsa") @@ -133,6 +132,28 @@ def _setup_distributed(): #: (``aiter_per1x128_quant``) and by the kernel, whose BLOCK_K is already 128. SCALE_BK = 128 +#: MXFP8 block size, along K for A and along both N and K for B. Fixed by the +#: scaled MFMA, which carries one ue8m0 scale per 32 K per row, and by +#: DeepSeek-V4.1-Flash's checkpoint (``weight_block_size [32, 32]``, +#: ``scale_fmt "ue8m0"``). +MXFP8_BK = 32 + + +def _ue8m0_bytes(shape, g, lo=120, hi=123): + """Random ue8m0 exponent bytes, i.e. scales 2**(byte-127) around 1e-2. + + ue8m0 *is* the exponent: there is no mantissa, so every scale is exactly a + power of two and applying it is lossless. That is the whole reason the + scaled MFMA can take it as an operand. + """ + e = torch.randint(lo, hi, shape, generator=g, device="cuda", dtype=torch.int32) + return e.to(torch.uint8) + + +def _ue8m0_value(e: torch.Tensor) -> torch.Tensor: + """ue8m0 exponent bytes -> the fp32 powers of two they denote.""" + return torch.exp2(e.to(torch.float32) - 127.0) + def make_operands(rank: int, m: int, n: int, k: int, quant: str = "ptpc"): """Deterministic per-rank fp8 operands and their scales. @@ -164,8 +185,16 @@ def make_operands(rank: int, m: int, n: int, k: int, quant: str = "ptpc"): torch.rand(n, generator=g, device="cuda", dtype=torch.float32) * 0.01 + 0.01 ) return a, b, sa, sb + if quant == "mxfp8": + # DeepSeek-V4.1-Flash's dense form: A per-32-K, B per-32x32, both ue8m0. + # Scales stay as exponent bytes all the way to the MFMA, which is what + # the instruction's scale operand reads. + kb = k // MXFP8_BK + sa = _ue8m0_bytes((m, kb), g) + sb = _ue8m0_bytes((n // MXFP8_BK, kb), g) + return a, b, sa.contiguous(), sb.contiguous() if quant != "blockscale": - raise ValueError(f"quant must be ptpc or blockscale, got {quant!r}") + raise ValueError(f"quant must be ptpc, blockscale or mxfp8, got {quant!r}") kb = k // SCALE_BK sa = ( torch.rand(m, kb, generator=g, device="cuda", dtype=torch.float32) * 0.01 + 0.01 @@ -183,6 +212,17 @@ def reference_partial(a, b, sa, sb, quant: str) -> torch.Tensor: af, bf = a.float(), b.float() if quant == "ptpc": return (af @ bf.T) * sa[:, None] * sb[None, :] + if quant == "mxfp8": + sav, sbv = _ue8m0_value(sa), _ue8m0_value(sb) + out = torch.zeros(a.shape[0], b.shape[0], device=a.device, dtype=torch.float32) + for i in range(a.shape[1] // MXFP8_BK): + ks = slice(i * MXFP8_BK, (i + 1) * MXFP8_BK) + out += ( + (af[:, ks] @ bf[:, ks].T) + * sav[:, i][:, None] + * sbv[:, i].repeat_interleave(MXFP8_BK)[None, :] + ) + return out out = torch.zeros(a.shape[0], b.shape[0], device=a.device, dtype=torch.float32) for i in range(a.shape[1] // SCALE_BK): ks = slice(i * SCALE_BK, (i + 1) * SCALE_BK) @@ -195,6 +235,12 @@ def reference_partial(a, b, sa, sb, quant: str) -> torch.Tensor: def _median_us(fn, warmup: int, iters: int, *, graph: bool = True) -> float: + # One call per graph capture, which on MI355X carries a ~13.4us replay + # floor. Left alone because everything this file times is a whole fused + # all-reduce at 300-1600us, where that is 1-4% and identical across the + # modes being compared -- but do not copy this into a benchmark of a single + # small kernel. `timing.py` is the one to use there; see the README's + # "Measurement traps". if graph: side = torch.cuda.Stream() side.wait_stream(torch.cuda.current_stream()) @@ -225,6 +271,18 @@ def _median_us(fn, warmup: int, iters: int, *, graph: bool = True) -> float: def run(args) -> int: + # `build_lsa_ar` has no gather_dtype: the LSA 2-stage all-reduce moves bf16 + # and there is no fp8 leg in it. Asking for one used to run a bf16 + # collective and report it under the fp8 label. The two-sided validation + # gate below did catch it -- relL2 landed at the bf16 floor, under the + # 5e-3 lower bound that exists to assert the fp8 wire was taken -- but only + # after paying for the run, and with an error that says "VALIDATION FAILED" + # rather than what is actually wrong. + if args.mode == "split-lsa" and args.gather_dtype == "fp8": + raise SystemExit( + "--mode split-lsa has no fp8 gather leg (build_lsa_ar takes no " + "gather_dtype); use split-sdma or fused-sdma for --gather-dtype fp8" + ) local_rank, rank, world_size, uid = _setup_distributed() # blockscale keeps a second fp32 accumulator for the per-K-block promotion, # which doubles the accumulator VGPRs; 256x256 needs 256 of them and the @@ -348,6 +406,18 @@ def run(args) -> int: if args.quant == "blockscale": sa_arg = sa.t().reshape(-1).contiguous() sb_arg = sb.reshape(-1).contiguous() + elif args.quant == "mxfp8": + # ue8m0 exponent bytes widened to int32 with the byte in the low 8 + # bits: the MFMA's scale operand is a 32-bit register read at + # op_sel 0, so this is what it wants, and it keeps the kernel on a + # plain dword load. + # + # A goes through preshuffle_a_scale, which documents both moves it + # makes and what each was worth. B stays K-block major [K/32, N/32]: + # a 16-column tile never straddles a 32-column group, so its load is + # already a broadcast and there is nothing to pack. + sa_arg = preshuffle_a_scale(sa) + sb_arg = sb.to(torch.int32).t().reshape(-1).contiguous() else: sa_arg, sb_arg = sa, sb @@ -537,7 +607,7 @@ def build_parser() -> argparse.ArgumentParser: p.add_argument("-k", type=int, default=1024) p.add_argument( "--quant", - choices=("ptpc", "blockscale"), + choices=("ptpc", "blockscale", "mxfp8"), default="ptpc", help="a8w8 per-token/per-channel (the aiter 8wave kernel's native form) " "or the model's 1x128 / 128x128 block scale. blockscale applies the " diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemv.py b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py new file mode 100644 index 000000000..1c1692ce3 --- /dev/null +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py @@ -0,0 +1,230 @@ +#!/usr/bin/env python3 +"""mori's mxfp8 GEMV against sglang's, and a config sweep for mori's. + +Both sides get the same raw weight and the same ue8m0 scales, each in its own +shuffle, so the only difference is the kernel. Timed with `timing.py`: amortised +(a single-call capture has a 13.4us floor on this box, and sglang's kernel is +2-5us, i.e. entirely inside it) and cold as well as hot, because a decode step +reads each layer's weight exactly once and a hot loop reads it out of LLC at +1.7x the bandwidth. + +The SGLang baseline is optional (`--baseline none`); mori itself does not +depend on it. + + python bench_gemv.py --shape wq_b -m 1 --sweep + python bench_gemv.py --shape wq_b -m 1 --config w8s2r16t16k +""" + +from __future__ import annotations + +import argparse +import json +import sys +import traceback +from pathlib import Path + +import flydsl.expr as fx +import torch +from mori.ops.gemm_ar import preshuffle_b +from mori.ops.gemm_ar.kernels_gemv import compile_mxfp8_gemv + +sys.path.insert(0, str(Path(__file__).parent)) +import timing # noqa: E402 + +MXFP8_BK = 32 +SHAPES = {"wq_b": (8192, 1280), "wo_b": (5120, 2048)} + +#: The space sglang tunes over, minus the configs that cannot serve the M. +WAVES = (4, 8, 16) +STEPS = (1, 2, 4) +ROWS = (16, 32) +TOKENS = (16, 32) + + +def configs_for(m: int): + for w in WAVES: + for s in STEPS: + for r in ROWS: + for t in TOKENS: + if t < m: + continue + for ks in (True, False): + yield { + "waves": w, + "steps": s, + "rows": r, + "tokens": t, + "ksplit": ks, + } + + +def key_of(c) -> str: + return ( + f"w{c['waves']}s{c['steps']}r{c['rows']}t{c['tokens']}" + f"{'k' if c['ksplit'] else 'n'}" + ) + + +def parse_key(key: str): + import re + + g = re.fullmatch(r"w(\d+)s(\d+)r(\d+)t(\d+)([kn])", key) + if not g: + raise ValueError(f"bad config key {key!r}") + return { + "waves": int(g[1]), + "steps": int(g[2]), + "rows": int(g[3]), + "tokens": int(g[4]), + "ksplit": g[5] == "k", + } + + +def build(m_max, n, k, seed=1234): + g = torch.Generator(device="cuda").manual_seed(seed) + x = (torch.randn(m_max, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) + ex = torch.randint( + 120, 123, (m_max, k // MXFP8_BK), generator=g, device="cuda", dtype=torch.int32 + ).to(torch.uint8) + ew = torch.randint( + 120, + 123, + (n // MXFP8_BK, k // MXFP8_BK), + generator=g, + device="cuda", + dtype=torch.int32, + ).to(torch.uint8) + return x, w, ex, ew + + +def mori_call(cfg, n, k, m, m_max, x, w, ex, ew, out): + gemv = compile_mxfp8_gemv(n=n, k=k, m_max=m_max, **cfg) + w_shuf = preshuffle_b(w).contiguous().view(torch.int32).view(-1) + ws = ew.contiguous().view(torch.int32).view(-1) + xs = ex.contiguous().view(torch.int32).view(-1) + xi = x.contiguous().view(torch.int32).view(-1) + + def call(picked): + gemv( + picked[0], + ws, + xi, + xs, + out.view(-1), + m, + n, + stream=fx.Stream(torch.cuda.current_stream()), + ) + + return call, w_shuf + + +def sglang_call(n, k, m, x, w, ex, ew, out): + from sglang.kernels.ops.quantization.mxfp8_native_amd_gfx95 import ( + mxfp8_gemv, + shuffle_mxfp8_weight, + ) + + w_shuf = shuffle_mxfp8_weight(w).contiguous() + xv, exv = x[:m].contiguous(), ex[:m].contiguous() + + def call(picked): + mxfp8_gemv(xv, picked[0], ew, x_scale=exv, out=out[:m]) + + return call, w_shuf + + +def main() -> int: + p = argparse.ArgumentParser() + p.add_argument("--shape", choices=sorted(SHAPES), required=True) + p.add_argument("-m", type=int, required=True) + p.add_argument("--sweep", action="store_true") + p.add_argument("--config", default=None) + p.add_argument("--reps", type=int, default=64) + p.add_argument("--baseline", choices=("sglang", "none"), default="sglang", + help="'none' drops the SGLang import; mori's own numbers " + "need nothing but mori") + p.add_argument("--tol", type=float, default=2.4e-3) + p.add_argument("--json-out", default="gemv.jsonl") + args = p.parse_args() + + n, k = SHAPES[args.shape] + m = args.m + m_max = 32 + x, w, ex, ew = build(m_max, n, k) + out = torch.zeros(m_max, n, device="cuda", dtype=torch.bfloat16) + vram_before = timing.vram_used() + + common = { + "bench": "gemv", "scope": "kernel", "quant": "mxfp8", + "shape": args.shape, "n": n, "k": k, "m": m, + "input": "fp8", "includes_quant": False, + "timing": "amortized-graph-cold-hot", + } + rows, base, ref, failures = [], None, None, 0 + + if args.baseline == "sglang": + call, wt = sglang_call(n, k, m, x, w, ex, ew, out) + base = timing.cold_hot_us(call, [wt], reps=args.reps) + ref = out[:m].float().clone() + print( + f"{args.shape} M={m} sglang hot {base['hot_us']:6.2f} " + f"cold {base['cold_us']:6.2f}", + flush=True, + ) + rows.append(dict(common, impl="sglang", route="sglang-gemv", + rel_l2=0.0, validated=True, **base)) + + if args.config: + cfgs = [parse_key(args.config)] + elif args.sweep: + cfgs = list(configs_for(m)) + else: + cfgs = [{"waves": 8, "steps": 2, "rows": 16, "tokens": 32, "ksplit": True}] + + for cfg in cfgs: + try: + cfg_m_max = cfg["tokens"] + call, wt = mori_call(cfg, n, k, m, cfg_m_max, x, w, ex, ew, out) + out.zero_() + call([wt]) + torch.cuda.synchronize() + rel = ( + ((out[:m].float() - ref).norm() / ref.norm().clamp_min(1e-30)).item() + if ref is not None else None + ) + res = timing.cold_hot_us(call, [wt], reps=args.reps) + except Exception as err: # noqa: BLE001 - a bad config must not stop the sweep + print(f" {key_of(cfg):<12} FAILED {type(err).__name__}: {err}", flush=True) + traceback.print_exc(limit=3) + rows.append(dict(common, impl="mori", config=key_of(cfg), + validated=False, + error=f"{type(err).__name__}: {err}")) + failures += 1 + continue + ok = rel is None or rel <= args.tol + failures += 0 if ok else 1 + vs = (f" {(res['cold_us'] / base['cold_us'] - 1) * 100:+6.1f}%" + if base else " " * 8) + rel_s = " n/a " if rel is None else f" relL2 {rel:.2e}" + print( + f" {key_of(cfg):<12} hot {res['hot_us']:6.2f} cold {res['cold_us']:6.2f}" + f"{vs}{rel_s}{'' if ok else ' !! over tol'}", + flush=True, + ) + rows.append(dict(common, impl="mori", config=key_of(cfg), + route=f"mori-gemv-{key_of(cfg)}", + rel_l2=rel, validated=ok, **res)) + + with open(args.json_out, "a") as f: + for r in rows: + f.write(json.dumps(dict( + r, vram_before=vram_before, vram_after=timing.vram_used() + )) + "\n") + # A benchmark that fails and exits 0 is how a broken sweep looks green. + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/cco/flydsl/gemm_ar/report.py b/benchmark/cco/flydsl/gemm_ar/report.py new file mode 100644 index 000000000..e390bfe02 --- /dev/null +++ b/benchmark/cco/flydsl/gemm_ar/report.py @@ -0,0 +1,171 @@ +#!/usr/bin/env python3 +"""Render any of the benchmark JSONL files into tables. + +One reader for every producer, because the two it replaced each knew about one +of them and dropped whatever it did not recognise. + +Two rules it keeps that the originals did not: + +* **The grouping key is every configuration field present**, not a hand-listed + few. `report_fused.py` keyed on `(label, world_size, m)` and stored by mode, + so a sweep that varied `gather_transport` or `fuse_quantize` -- which is + exactly what `sweep.py fused-wire` does -- had every such run overwrite the + one before it, silently, and the table showed the last one as if it were the + only one. +* **Nothing is dropped for being unrecognised.** A label or an impl this file + has never seen still gets a row. The old `show(labels, ...)` matched a fixed + list and printed nothing at all for a `fp8gather` sweep. + + python report.py gemm.jsonl + python report.py fused.jsonl --col cold_us + python report.py gemm-full.jsonl --by impl --baseline sglang +""" + +from __future__ import annotations + +import argparse +import json +import statistics +from collections import defaultdict + +#: Fields that identify *a measurement* rather than a configuration. Everything +#: else in the row is part of the key, which is what keeps distinct runs +#: distinct without this file having to know the axes in advance. +MEASURED = { + "hot_us", "cold_us", "us", "max_rank_time_us", "rel_l2", "validated", + "copies", "cold_reps", "working_set_mb", "vram_before", "vram_after", + "error", "timing", "route", "supported", "reason", + # Resolved by the op, not chosen by the caller: `critical_rank` is elected + # at runtime and `resolved_tile_order` is what `--tile-order auto` became. + # They are never sweep axes, so they can be excluded even when a file + # predates `sweep_axes`. + "critical_rank", "resolved_tile_order", +} +#: Shown as the row label when present, in this order; the rest go in the key. +ROW_KEYS = ("label", "shape", "n", "k", "world_size", "quant", "scope") +#: Preferred column axis, first one present wins. +COL_KEYS = ("impl", "mode", "config") + + +def load(paths): + rows = [] + for p in paths: + with open(p) as f: + for line in f: + line = line.strip() + if not line: + continue + if line.startswith("RESULT_JSON"): + line = line.split(" ", 1)[1] + rows.append(json.loads(line)) + return rows + + +def time_of(row, col): + for k in (col, "cold_us", "us", "max_rank_time_us", "hot_us"): + if k in row and isinstance(row[k], (int, float)): + return row[k] + return None + + +def main() -> int: + p = argparse.ArgumentParser() + p.add_argument("paths", nargs="+") + p.add_argument("--col", default="cold_us", + help="which timing field to table (default cold_us)") + p.add_argument("--by", default=None, + help="column axis; default is the first of " + ", ".join(COL_KEYS)) + p.add_argument("--baseline", default=None, + help="column to show the others as a percentage against") + p.add_argument("--keep-invalid", action="store_true") + args = p.parse_args() + + rows = load(args.paths) + if not rows: + print("no rows") + return 1 + + bad = [r for r in rows if r.get("validated") is False] + if bad and not args.keep_invalid: + rows = [r for r in rows if r.get("validated") is not False] + print(f"!! {len(bad)} row(s) failed validation, excluded " + f"(--keep-invalid to see them)") + for r in bad[:5]: + who = r.get("impl") or r.get("mode") or "?" + print(f" {r.get('shape','?')} M={r.get('m','?')} {who}: " + f"{r.get('error') or f'rel_l2={r.get("rel_l2")}'}") + if not rows: + print("nothing left after validation filter") + return 1 + + col_key = args.by + if col_key is None: + for c in COL_KEYS: + if any(c in r for r in rows): + col_key = c + break + if col_key is None: + print("no column axis found; pass --by") + return 1 + + # Key on what the sweep varied when it said so, and on everything + # configuration-ish otherwise. The distinction matters: bench_gemm_ar.py + # echoes knobs it resolved for itself -- `chunks`, `critical_rank`, the + # tile order -- and keying on those turns one matrix into one table per + # point. `sweep_axes` is how sweep.py says which fields were axes. + def keyof(r): + axes = r.get("sweep_axes") + if axes is not None: + keep = [k for k in axes if k not in (col_key, "m")] + keep += [k for k in ("label", "shape", "n", "k", "world_size", + "quant", "scope") if k in r] + return tuple((k, r[k]) for k in sorted(set(keep))) + return tuple( + (k, r[k]) for k in sorted(r) + if k not in MEASURED and k not in (col_key, "m", "bench", "sweep_axes") + and not isinstance(r[k], (dict, list)) + ) + + table = defaultdict(lambda: defaultdict(list)) + for r in rows: + t = time_of(r, args.col) + if t is not None: + table[keyof(r)][(r.get("m"), r.get(col_key))].append(t) + + for key, cells in sorted(table.items(), key=lambda kv: str(kv[0])): + d = dict(key) + head = " ".join(f"{k}={d[k]}" for k in ROW_KEYS if k in d) + rest = " ".join(f"{k}={v}" for k, v in key + if k not in ROW_KEYS and v not in (None, "")) + print(f"\n### {head}" + (f" [{rest}]" if rest else "")) + ms = sorted({m for m, _ in cells if m is not None}) + cols = sorted({c for _, c in cells if c is not None}, key=str) + if not cols: + continue + base = args.baseline if args.baseline in cols else None + + print(f"{'M':>8}" + "".join(f"{str(c):>16}" for c in cols) + + (" (vs " + base + ")" if base else "")) + for m in ms: + line = f"{m:>8}" + bt = None + if base: + vs = cells.get((m, base)) + bt = statistics.median(vs) if vs else None + for c in cols: + vs = cells.get((m, c)) + if not vs: + line += f"{'-':>16}" + continue + t = statistics.median(vs) + spread = f"*{len(vs)}" if len(vs) > 1 else "" + if bt and c != base: + line += f"{t:9.1f}{(t / bt - 1) * 100:+6.0f}%" + else: + line += f"{t:>12.1f}{spread:>4}" + print(line) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/cco/flydsl/gemm_ar/sweep.py b/benchmark/cco/flydsl/gemm_ar/sweep.py new file mode 100644 index 000000000..12ca2d775 --- /dev/null +++ b/benchmark/cco/flydsl/gemm_ar/sweep.py @@ -0,0 +1,311 @@ +#!/usr/bin/env python3 +"""Run a benchmark matrix, one point per process, and say so when one fails. + +The matrices used to be six shell scripts that each re-implemented process +isolation, timeouts, occupancy capture and label tagging. They are presets here +instead. Process isolation is the one thing none of them could drop: a FlyDSL +compile failure takes the interpreter with it, and a multi-M process was once +caught reporting 697us and 1083us for two M that pad to the same size. + + python sweep.py --list + python sweep.py gemm # the default regression + python sweep.py gemm-full --out full.jsonl + python sweep.py fused-wire + +Exit status is the number of failed points, capped at 125. A sweep that fails +and exits 0 is how a broken run looks green. +""" + +from __future__ import annotations + +import argparse +import itertools +import json +import os +import shlex +import subprocess +import sys +import time +from pathlib import Path + +HERE = Path(__file__).parent +PY = os.environ.get("PY") or sys.executable + +#: The two shapes every default regression covers, and the ten the full sweep +#: does. Kept here rather than in the presets so a new preset cannot quietly +#: disagree with an existing one about what `wq_b` means. +MAIN = ["wq_b", "wo_b"] +ALL_SHAPES = MAIN + [ + "wq_a_tp4", "wkv_tp4", "wqkv_a_tp4", "wo_a_tp4", "shared_gate_up_tp4", + "wq_b_tp8", "wo_b_tp8", "wo_a_tp8", "wq_b_tp1", "wo_b_tp1", +] + +#: M values a default run covers. Deliberately includes both sides of the two +#: thresholds this operator has, because a coarse M grid is how the branch got +#: `NARROW_N_BELOW_M` wrong once already: 1024 -> 2048 skipped the crossover. +M_DEFAULT = [64, 256, 1024, 1280, 1536, 1792, 2048, 4096, 16384] +M_FULL = [1, 8, 32, 64, 256, 1024, 2048, 4096, 8192, 16384] +#: Either side of the 128/256 N-tile switch, which is what `--preset gemm-tile` +#: exists to keep honest. +M_TILE = [512, 1024, 1280, 1536, 1792, 2048, 4096] +#: The GEMV's token buckets and the boundaries between them. +M_GEMV = [1, 2, 3, 4, 8, 16, 17, 32] + + +def _grid(**axes): + """Cartesian product of named axes, as a list of dicts.""" + keys = list(axes) + return [dict(zip(keys, v)) for v in itertools.product(*(axes[k] for k in keys))] + + +PRESETS = { + # ---- single process, no collective --------------------------------- + "gemm": dict( + doc="the default regression: two shapes, both scopes, mori + SGLang", + script="bench_gemm.py", + grid=_grid(shape=MAIN, m=M_DEFAULT, scope=["linear"]), + args=["--impl", "auto,sglang"], + ), + "gemm-full": dict( + doc="every shape the checkpoint has, every M, every mori tile", + script="bench_gemm.py", + grid=_grid(shape=ALL_SHAPES, m=M_FULL, scope=["linear"]), + args=["--impl", "auto,gemm256,gemm128,sglang"], + ), + "gemm-tile": dict( + doc="the 128 vs 256 N-tile switch, kernel scope, either side of it", + script="bench_gemm.py", + grid=_grid(shape=ALL_SHAPES, m=M_TILE, scope=["kernel"]), + args=["--impl", "gemm256,gemm128"], + ), + "gemm-blockscale": dict( + doc="mori's other operand contract, kernel scope, no SGLang baseline", + script="bench_gemm.py", + grid=_grid(shape=MAIN, m=[4096, 8192, 16384], scope=["kernel"]), + args=["--impl", "auto", "--quant", "blockscale"], + ), + "gemv": dict( + doc="the skinny GEMM at its token buckets, default config", + script="bench_gemv.py", + grid=_grid(shape=MAIN, m=M_GEMV), + ), + "gemv-tune": dict( + doc="the whole GEMV config space, per bucket -- this is what tunes the table", + script="bench_gemv.py", + grid=_grid(shape=MAIN, m=[1, 2, 4, 8, 16, 32]), + args=["--sweep"], + timeout=3600, + ), + # ---- multi rank, the collective ------------------------------------ + "fused": dict( + doc="split vs fused at V4.1-Flash's wo_b, under the model's own mxfp8", + script="bench_gemm_ar.py", + world=4, + grid=_grid( + m=[4096, 8192, 16384], + mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"], + ), + args=["-n", "5120", "-k", "2048", "--quant", "mxfp8"], + label="mxfp8", + ), + "fused-fp8": dict( + doc="the same mode matrix on the winning fp8 wire, which is where the " + "fusion is actually deployed", + script="bench_gemm_ar.py", + world=4, + grid=_grid( + m=[4096, 8192, 16384], + mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"], + ), + args=["-n", "5120", "-k", "2048", "--quant", "mxfp8", + "--gather-dtype", "fp8", "--gather-transport", "lsa", + "--no-fuse-quantize"], + label="mxfp8-fp8wire", + ), + "fused-wire": dict( + doc="how to move the fp8 all-gather leg: push or pull, fuse the quantise or not", + script="bench_gemm_ar.py", + world=4, + grid=_grid( + m=[4096, 16384], + mode=["split-sdma", "fused-sdma", "fused-lsa"], + gather_transport=["sdma", "lsa"], + fuse_quantize=[True, False], + ), + args=["-n", "5120", "-k", "2048", "--quant", "mxfp8", "--gather-dtype", "fp8"], + label="fp8gather", + ), + "fused-blockscale-control": dict( + doc=( + "CONTROL, not a V4.1-Flash benchmark: the same shapes under 1x128 " + "blockscale, which is *not* what the checkpoint uses. Kept because " + "the mxfp8 numbers are only interpretable against it -- see the " + "README. Use `fused` for the real thing." + ), + script="bench_gemm_ar.py", + world=4, + grid=_grid( + m=[4096, 8192, 16384], + mode=["gemm-only", "split-sdma", "fused-sdma"], + ), + args=["-n", "5120", "-k", "2048", "--quant", "blockscale"], + label="blockscale-control", + ), +} + + +def require_sdma() -> None: + """Refuse to run a collective benchmark against a build that has no queues. + + With `BUILD_CCO_SDMA=OFF` every put silently does nothing: the all-reduce + returns mostly the local slice, the model still answers, and the fused path + measures *faster* than it is because it is not moving data. + """ + try: + from mori.cco.device._build_flags import BUILD_CCO_SDMA + except ImportError: + BUILD_CCO_SDMA = False + if not BUILD_CCO_SDMA: + sys.exit("BUILD_CCO_SDMA is OFF -- the SDMA path is compiled out, " + "numbers would be fiction. Rebuild with BUILD_CCO_SDMA=ON.") + + +def occupancy() -> str: + try: + out = subprocess.run( + ["rocm-smi", "--showmeminfo", "vram", "--csv"], + capture_output=True, text=True, timeout=30, + ).stdout + except (OSError, subprocess.SubprocessError): + return "?" + gib = [f"{int(p[2]) / 2**30:.0f}" for p in + (l.split(",") for l in out.splitlines()) + if len(p) >= 3 and p[2].isdigit()] + return " ".join(gib) or "?" + + +def point_argv(preset, point, out_path): + """The full argv for one point of the matrix.""" + script = str(HERE / preset["script"]) + extra = list(preset.get("args", [])) + world = preset.get("world") + + if world: + # `python -m torch.distributed.run` rather than the `torchrun` console + # script: the latter is only on PATH if the venv is activated, which a + # subprocess inherits only by luck. + cmd = [PY, "-m", "torch.distributed.run", "--standalone", + f"--nproc_per_node={world}", script] + else: + cmd = [PY, script, "--json-out", str(out_path)] + + for key, val in point.items(): + if isinstance(val, bool): + # bench_gemm_ar.py spells these as --fuse-quantize / --no-fuse-quantize + cmd.append(f"--{key.replace('_', '-')}" if val + else f"--no-{key.replace('_', '-')}") + elif key == "m": + cmd += ["-m", str(val)] + else: + cmd += [f"--{key.replace('_', '-')}", str(val)] + return cmd + extra + + +def run_point(preset, point, out_path, timeout): + """One subprocess. Returns (ok, stdout).""" + cmd = point_argv(preset, point, out_path) + env = dict(os.environ) + if preset.get("world"): + env.setdefault("MORI_SOCKET_IFNAME", "lo") + env["MORI_ENABLE_SDMA"] = "1" + try: + p = subprocess.run(cmd, capture_output=True, text=True, + timeout=timeout, cwd=HERE, env=env) + except subprocess.TimeoutExpired: + return False, f"TIMEOUT after {timeout}s" + # Multi-rank runs print RESULT_JSON rather than writing the file themselves. + if preset.get("world"): + wrote = 0 + with open(out_path, "a") as f: + for line in p.stdout.splitlines(): + if line.startswith("RESULT_JSON"): + d = json.loads(line.split(" ", 1)[1]) + d["label"] = preset.get("label", "") + d.update({k: v for k, v in point.items() if k != "m"}) + # What this sweep *varied*, so a reporter can tell an axis + # from a value the op resolved for itself. bench_gemm_ar.py + # echoes every knob it ran with, including ones it chose + # (chunks, critical_rank, tile order); keying a table on + # those splits one matrix into one table per point. + d["sweep_axes"] = sorted(point) + f.write(json.dumps(d) + "\n") + wrote += 1 + if p.returncode == 0 and wrote == 0: + return False, "no RESULT_JSON emitted" + return p.returncode == 0, p.stdout + p.stderr + + +def main() -> int: + p = argparse.ArgumentParser() + p.add_argument("preset", nargs="?", choices=sorted(PRESETS)) + p.add_argument("--list", action="store_true") + p.add_argument("--out", default=None) + p.add_argument("--timeout", type=int, default=None) + p.add_argument("--dry-run", action="store_true") + args = p.parse_args() + + if args.list or not args.preset: + width = max(len(k) for k in PRESETS) + for name, d in sorted(PRESETS.items()): + print(f" {name:<{width}} {d['doc']}") + return 0 + + preset = PRESETS[args.preset] + out_path = Path(args.out or f"{args.preset}.jsonl").resolve() + timeout = args.timeout or preset.get("timeout", 1800) + + if args.dry_run: + for point in preset["grid"]: + print(" ".join(shlex.quote(c) + for c in point_argv(preset, point, out_path))) + return 0 + + if preset.get("world"): + require_sdma() + out_path.write_text("") + before = occupancy() + print(f"# preset {args.preset}: {len(preset['grid'])} points -> {out_path}") + print(f"# occupancy before: {before} GiB/card", file=sys.stderr) + + failed = [] + for i, point in enumerate(preset["grid"], 1): + desc = " ".join(f"{k}={v}" for k, v in point.items()) + print(f"[{i}/{len(preset['grid'])}] {desc}", flush=True) + ok, out = run_point(preset, point, out_path, timeout) + if not ok: + failed.append(desc) + print(f" FAILED: {out.strip().splitlines()[-1] if out.strip() else '?'}", + flush=True) + else: + for line in out.splitlines(): + if line.startswith(" ") or line.startswith("RESULT_JSON"): + print(line if line.startswith(" ") else " ok", flush=True) + if preset.get("world"): + time.sleep(3) # let the SDMA queues drain before the next spawn + + after = occupancy() + print(f"# occupancy after: {after} GiB/card", file=sys.stderr) + if before != after: + print("# NOTE: occupancy moved across the sweep; a leftover process may " + "have shared the GPU. Re-run before trusting these.", file=sys.stderr) + if failed: + print(f"\n{len(failed)} of {len(preset['grid'])} points FAILED:") + for d in failed: + print(f" {d}") + else: + print(f"\nall {len(preset['grid'])} points ok -> {out_path}") + return min(len(failed), 125) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/cco/flydsl/gemm_ar/timing.py b/benchmark/cco/flydsl/gemm_ar/timing.py new file mode 100644 index 000000000..5fed6ff0e --- /dev/null +++ b/benchmark/cco/flydsl/gemm_ar/timing.py @@ -0,0 +1,173 @@ +# Copyright (c) 2025 Advanced Micro Devices, Inc. All rights reserved. +# +# SPDX-License-Identifier: MIT +"""Timing for kernels small enough that the harness is the measurement. + +Two corrections over the obvious `capture one call, replay, take the median`, +both of which matter only at small M and both of which silently inflated an +earlier round of thresholds on this branch. + +**A single-call graph replay has a floor.** On this box an empty-ish kernel +replays in 13.40us; amortised over 200 calls in one graph the same kernel is +1.51us. So every number below ~15us produced by single-call capture is mostly +the floor. sglang's mxfp8 GEMV at `wq_b` M=1 measured 13.48us that way and is +2.45us -- it was entirely inside the floor, and so was the margin it was being +compared on. `amortized=True` (the default) captures `reps` calls per graph. + +**Repeating a call leaves the weight in LLC.** MI355X has 256MB of it and +`wq_b`'s weight is 10.5MB, so a hot loop reads from cache at 4308 GB/s where +one cold pass gets 2561 GB/s -- a 1.7x overstatement of a memory-bound kernel. +A decode step touches each layer's weight once, so cold is the number that +predicts the server and hot is the ceiling. `cold()` rotates enough copies of +the weight to exceed the cache and reports both. + +Both are properties of the measurement, not of the kernel: a compute-bound GEMM +at M=16384 is unaffected by either, which is why they went unnoticed. +""" + +from __future__ import annotations + +import statistics +import subprocess + +import torch + +#: Past this the LLC cannot hold the working set. MI355X has 256MB; the margin +#: covers the activations and output sharing it. +LLC_BYTES = 256 << 20 +COLD_WORKING_SET = 384 << 20 + + +def _graph(fn, reps: int) -> torch.cuda.CUDAGraph: + """Capture `reps` calls into one graph, warming on a side stream first.""" + fn() + torch.cuda.synchronize() + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + for _ in range(3): + fn() + torch.cuda.current_stream().wait_stream(side) + torch.cuda.synchronize() + g = torch.cuda.CUDAGraph() + with torch.cuda.graph(g): + for _ in range(reps): + fn() + return g + + +def median_us(fn, reps: int = 200, iters: int = 21) -> float: + """Median us per call, amortised over `reps` calls inside one graph. + + `reps=1` reproduces the old single-call number, which is useful only for + showing what the floor was doing to it. + """ + g = _graph(fn, reps) + ts = [] + for _ in range(iters): + s, e = torch.cuda.Event(True), torch.cuda.Event(True) + s.record() + g.replay() + e.record() + torch.cuda.synchronize() + ts.append(s.elapsed_time(e) * 1000.0 / reps) + return statistics.median(ts) + + +def replay_floor_us(iters: int = 21) -> float: + """What one single-call graph replay costs when the kernel does nothing. + + Report this next to any single-call number so the reader can see how much + of it is the harness. + """ + x = torch.zeros(1, device="cuda") + return median_us(lambda: x.add_(1.0), reps=1, iters=iters) + + +def n_copies_for(weight_bytes: int, target: int = COLD_WORKING_SET) -> int: + """How many rotating weight copies push the working set past the LLC.""" + return max(1, -(-target // max(weight_bytes, 1))) + + +class Rotating: + """`n` copies of a tensor, handed out round-robin. + + Enough of them and consecutive calls cannot hit in LLC, which is what a + real forward pass looks like: each layer's weight is read once per step. + """ + + def __init__(self, t: torch.Tensor, n: int): + self.copies = [t] + [t.clone() for _ in range(n - 1)] + self.i = 0 + + def next(self) -> torch.Tensor: + t = self.copies[self.i] + self.i = (self.i + 1) % len(self.copies) + return t + + @property + def bytes(self) -> int: + return self.copies[0].numel() * self.copies[0].element_size() * len(self.copies) + + +def cold_hot_us(make_fn, weights, reps: int = 32, iters: int = 21) -> dict: + """Both numbers for one kernel: cold (rotating weights) and hot (one copy). + + `weights` is the tensors whose residency is in question, largest first; + `make_fn(picked)` returns the closure to time, given one tensor per entry. + Cold rotates them together so a rep never revisits the previous rep's copy. + + **The cold capture must be at least as long as the ring.** The ring advances + at capture time, not at replay time, so a graph of `reps` calls bakes in + `reps` pointers and every replay revisits those same ones -- the working set + is `min(reps, n)` copies, not `n`. At `reps=32` that silently left seven of + this branch's twelve shapes under the 256MB LLC and therefore not cold at + all, and the ones that landed *on* it read worst of any: a working set at + exactly cache capacity thrashes, where one comfortably over it just streams. + That is why `wo_a` (8MB weight, 256MB at 32 reps) measured +10% against + SGLang cold and -2% hot. + + **And the ring is sized so that every tensor clears the cache on its own, + not so their sum does.** A caller may hand over several weights of which the + kernel reads only one -- SGLang's native linear takes both an fp8 weight and + a dequantised bf16 one and reads whichever its route picked. Sizing on the + sum then buys `384MB / (fp8 + bf16)` copies, and the route that reads only + the 10MB fp8 weight sees 130MB of it: back inside the LLC, hot again. Sizing + on the largest member costs more VRAM and is correct either way, since a + kernel that does read all of them gets a working set larger still. + """ + n = max(n_copies_for(w.numel() * w.element_size()) for w in weights) + rings = [Rotating(w, n) for w in weights] + + def cold_call(): + return make_fn([r.next() for r in rings]) + + hot = median_us(lambda: make_fn(list(weights)), reps=reps, iters=iters) + cold_reps = max(reps, n) + cold = median_us(cold_call, reps=cold_reps, iters=iters) + return { + "hot_us": hot, + "cold_us": cold, + "copies": n, + "cold_reps": cold_reps, + "working_set_mb": sum(r.bytes for r in rings) / 2**20, + } + + +def vram_used(device: int = 0) -> int | None: + """Bytes in use per rocm-smi, for the before/after check around a run.""" + try: + out = subprocess.run( + ["rocm-smi", "--showmeminfo", "vram", "--csv"], + capture_output=True, + text=True, + timeout=30, + ).stdout + except (OSError, subprocess.SubprocessError): + return None + for line in out.splitlines(): + if line.startswith(f"card{device},"): + parts = line.strip().split(",") + if len(parts) >= 3 and parts[2].isdigit(): + return int(parts[2]) + return None diff --git a/docs/MORI-GEMM-AR-BENCHMARK.md b/docs/MORI-GEMM-AR-BENCHMARK.md index 10feb44b5..0768cbfe4 100644 --- a/docs/MORI-GEMM-AR-BENCHMARK.md +++ b/docs/MORI-GEMM-AR-BENCHMARK.md @@ -1,34 +1,75 @@ # MORI GEMM + All-Reduce Benchmark -Measurements for `mori.ops.gemm_ar`, the fused fp8 GEMM + all-reduce. The design -and the API live next to the code in -[`python/mori/ops/gemm_ar/README.md`](https://github.com/ROCm/mori/blob/main/python/mori/ops/gemm_ar/README.md); -this file is how to reproduce the numbers and what they were. - -Every number below was taken on **8x MI355X (gfx950)**, one node, at the shape -`[M, 7168]` with `K=2048` — DeepSeek-V4-Pro's `wo_b` under TP8 with -`--chunked-prefill-size 16384`. Kernel timings are the median over 11 -graph-replayed iterations, maximum over ranks, on an otherwise idle box. -Run-to-run spread at this shape is about **2%**, so differences below that are -not differences. +How to run the benchmarks for `mori.ops.gemm_ar` -- the fused fp8 GEMM + +all-reduce, and the mxfp8 GEMM it is built on -- and a summary of what they say. + +The design, the API and **every full measurement table** live next to the code in +[`python/mori/ops/gemm_ar/README.md`](https://github.com/ROCm/mori/blob/main/python/mori/ops/gemm_ar/README.md). +This page is the entry point: how to reproduce, what the knobs mean, and the +headline numbers. + +Everything here is **MI355X (gfx950)**, one node, otherwise idle, with occupancy +recorded before and after each measurement. Kernel timings are the median over +repeated iterations, maximum over ranks. Run-to-run spread is about **2%**, so +differences below that are not differences. ## Table of Contents +- [Summary](#summary) - [Running the benchmark](#running-the-benchmark) -- [Headline](#headline) -- [Where the time goes](#where-the-time-goes) -- [The fp8 wire](#the-fp8-wire) - - [Who moves the gather](#who-moves-the-gather) - - [The pull grid](#the-pull-grid) - - [What fp8 costs, numerically](#what-fp8-costs-numerically) -- [Model-level evaluation](#model-level-evaluation) -- [Negative results](#negative-results) + - [Modes and quantisation](#modes-and-quantisation) + - [The fp8 wire's knobs](#the-fp8-wires-knobs) + - [The GEMM on its own](#the-gemm-on-its-own) +- [Measurement discipline](#measurement-discipline) - [Reproducing the end-to-end numbers](#reproducing-the-end-to-end-numbers) +## Summary + +Two models are covered and **they do not reach the same conclusion**, so they +are never averaged: + +| | DeepSeek-V4-Pro | DeepSeek-V4.1-Flash | +|---|---|---| +| `wo_b` per rank | `[M, 7168]` K=2048, TP8 | `[M, 5120]` K=2048, TP4 | +| quantisation | 1x128 / 128x128 fp8 block scale | 32-wide ue8m0 (mxfp8) | +| `--quant` | `blockscale` | `mxfp8` | +| best configuration | `fused-sdma` + fp8 gather | `fused-sdma` + fp8 gather over LSA | +| at M=16384 | **979.3 us against 1472.9, -33%** | **1192.1 us against 1592.5, -25%** | +| what dominates it | the overlap | the fp8 wire, then the GEMM | +| vs the model as it ships | 1419.5 us -> 979.3, **-31%** | see below | + +**On V4.1-Flash most of the win is not the fusing.** The GEMM is 178us against +1414us of communication, so hiding it entirely would be worth 11-15% and fusing +collects about half of that. The fp8 wire is worth -15.6% on its own and the +mxfp8 GEMM another -30.3% against the route SGLang would otherwise take. Reading +V4-Pro's conclusion across would set every threshold wrong. + +The GEMM is also usable on its own, with no collective attached +(`Mxfp8GemmOp`, and `Mxfp8GemvOp` for decode's token counts). Against SGLang's +native mxfp8 linear across all twelve of V4.1-Flash's fp8 linear shapes, what +decides the outcome is the grid, `ceildiv(M,256) * (N/256)`: + +| grid (workgroups) | median vs SGLang | +|---|---:| +| 1-32 | +246.7% | +| 33-64 | -0.7% | +| 65-128 | **-10.2%** | +| 129-256 | **-26.6%** | +| >256 | **-26.4%** | + +End to end in a server, V4.1-Flash prefill throughput is **+5.1% to +5.7%** at +batch sizes 4-16 on the fp8 wire; V4-Pro is **-5.0% GPU busy** over a profiled +prefill. Per-layer wins are larger than end-to-end ones because `wo_b` is about +12% of the profile. + +Full tables, the numerical cost of the fp8 wire, the model-level quality +evaluation, and the threshold derivations are in +[the operator README](https://github.com/ROCm/mori/blob/main/python/mori/ops/gemm_ar/README.md#measured-results). + ## Running the benchmark Needs a mori built with `BUILD_CCO_SDMA=ON`. Setting `MORI_ENABLE_SDMA` in the -environment only rebuilds the *device* bitcode — a host library built without +environment only rebuilds the *device* bitcode -- a host library built without the flag has no SDMA queues, every put silently does nothing, and the all-reduce quietly produces zeros. @@ -36,132 +77,91 @@ quietly produces zeros. cd /path/to/mori BUILD_CCO_SDMA=ON pip install . +# DeepSeek-V4-Pro, TP8 MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo \ torchrun --standalone --nproc_per_node=8 \ benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py \ --mode fused-sdma --quant blockscale -m 16384 -n 7168 -k 2048 -``` - -`--mode` selects what is measured, all reaching the same end state: - -| mode | what runs | -|---|---| -| `gemm-only` | the GEMM alone, to size the ceiling | -| `split-sdma` | `gemm` then a 4-kernel SDMA all-reduce | -| `fused-sdma` | GEMM with the scatter fused into its epilogue | -| `split-lsa` | `gemm` then the 2-kernel LSA all-reduce | -| `fused-lsa` | GEMM storing straight into peers | -`--quant blockscale` is the model's own quantisation (A 1x128, B 128x128, fp32 -scales) and is the column to read. `--gather-dtype fp8` and -`--gather-transport {sdma,lsa}` select the wire; see [the fp8 wire](#the-fp8-wire). +# DeepSeek-V4.1-Flash, TP4, with the wire that wins there +MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo \ + torchrun --standalone --nproc_per_node=4 \ + benchmark/cco/flydsl/gemm_ar/bench_gemm_ar.py \ + --mode fused-sdma --quant mxfp8 -m 16384 -n 5120 -k 2048 \ + --gather-dtype fp8 --gather-transport lsa --no-fuse-quantize +``` Each run prints a `RESULT_JSON` line with `max_rank_time_us`, `rel_l2` and `validated`, so a sweep can be parsed rather than eyeballed. -> **Between runs, give the SDMA queues time to drain.** Back-to-back 8-rank runs -> hit `hsaKmtCreateQueueExt` failures (`anvil.cpp:237`) if a previous run's ranks -> have not exited. ~10s is enough; a leftover server holding queues is not. +Matrices are presets of `sweep.py`, which runs one point per process, enforces +the `BUILD_CCO_SDMA` guard, records occupancy either side, and **exits non-zero +with the count of failed points**: -## Headline - -`--quant blockscale`, median of 11: +```bash +python benchmark/cco/flydsl/gemm_ar/sweep.py --list +python benchmark/cco/flydsl/gemm_ar/sweep.py fused # the mode matrix +python benchmark/cco/flydsl/gemm_ar/sweep.py fused-wire # the fp8 leg's knobs +python benchmark/cco/flydsl/gemm_ar/report.py fused.jsonl +``` -| M | `split-sdma` | `fused-sdma` | `fused-sdma` + fp8 gather | -|---|---:|---:|---:| -| 4096 | 398.1 us | 351.0 us | **329.5 us** | -| 8192 | 722.0 | 621.1 | **539.3** | -| 16384 | 1472.9 | 1148.8 | **979.3** | +`report.py` keys on every configuration field a row carries, so a matrix that +varies `gather_transport` or `fuse_quantize` renders as separate tables rather +than silently overwriting itself. -Fusing is worth **-22%** at M=16384; the fp8 gather a further **-15%**. +Numerics: -> Both survive to the server, in the same order. See -> [the end-to-end numbers](#reproducing-the-end-to-end-numbers). +```bash +MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo \ + pytest tests/python/cco/test_gemm_ar.py tests/python/cco/test_flydsl_ar.py \ + tests/python/cco/test_gemm_ar_op.py +``` -For scale, the same layer as the model runs it today (a separate GEMM then an -NCCL all-reduce) measures **1419.5 us** at M=16384, and the GEMM alone is -**369.4 us**. +### Modes and quantisation -## Where the time goes +`--mode` selects what is measured, all reaching the same end state: -Per layer, captured in SGLang over one 20000-token prefill (not the standalone -benchmark — this is the pipeline as the model drives it): +| mode | what runs | +|---|---| +| `gemm-only` | the GEMM alone, to size the ceiling | +| `split-sdma` | `gemm` then a 4-kernel SDMA all-reduce | +| `fused-sdma` | GEMM with the scatter fused into its epilogue | +| `split-lsa` | `gemm` then the 2-kernel LSA all-reduce | +| `fused-lsa` | GEMM storing straight into peers | -| phase | bf16 | fp8 / sdma | fp8 / lsa | -|---|---:|---:|---:| -| gemm | 441.9 us | 438.2 us | 441.3 us | -| drain | 180.2 | 170.9 | 204.8 | -| reduce | 41.7 | 42.5 | 43.1 | -| quantize | — | 11.8 | 11.9 | -| gather | 437.8 | 256.4 | 7.5 (barrier only) | -| dequantize | — | 61.0 | — | -| pull | — | — | 230.5 | -| **wo_b layer** | **1101.6** | **980.8** | **939.1** | -| | | -11.0% | **-14.7%** | +`--quant` picks the operand contract, and it is per model -- `blockscale` for +V4-Pro (A 1x128, B 128x128, fp32 scales), `mxfp8` for V4.1-Flash (32-wide ue8m0 +on both, int32). Two more exist and are not for use: `mxfp8_unpacked` and +`mxfp8_row` are the scale layouts `mxfp8` was chosen over, kept so the choice +stays measurable. -Two things to read out of the bf16 column: +> **Pin `--quant` when comparing anything against anything.** The benchmark +> defaults to `ptpc`, which is ~3% faster than `blockscale` and a different +> kernel again from `mxfp8`. Reading one against the other looks exactly like a +> machine that drifts overnight. -* **`gather` is the bottleneck**, not `drain`. It moves 196 MiB at 470 GB/s, - which is 7 xGMI links flat out, so halving its bytes halves its time. -* **`drain`'s apparent 1140 GB/s is not a bandwidth.** Seven links cannot do - that. It is the tell that the scatter's pushes already went out from the GEMM - epilogue and the drain is only waiting for the tail — which is why the scatter - leg has far less to give than its byte count suggests, and why it is still - bf16. +mxfp8 compiles at `BLOCK_M=256` where blockscale needs 128, so M pads to a +multiple of `world_size * 256`. That is not a tuning choice -- the packed scale +puts a lane's four M tiles in one dword, which is four tiles only at +`BLOCK_M//64 == 4`. -## The fp8 wire +### The fp8 wire's knobs `--gather-dtype fp8` sends the all-gather leg as e4m3 with one fp32 scale per -row. The reduce still accumulates in fp32 and `output` is still bf16; only the -wire changes. The scatter leg is unchanged — it carries partial sums that are +row. The reduce still accumulates in fp32 and the output is still bf16; only the +wire changes. The scatter leg is unchanged -- it carries partial sums that are then added across every rank, so its error would compound rather than being a single rounding. -### Who moves the gather - -| `--gather-transport` | how | +| flag | recommendation | |---|---| -| `sdma` | copy engines push, a second kernel widens | -| `lsa` (default) | CUs pull over xGMI and widen on the way in | - -A copy engine has no ALU, so for SDMA the widening *cannot* be the same step: it -is a second kernel that reads the landed fp8 back out of local HBM, 98 MiB a -layer. A CU pull has those bytes in registers already. +| `--gather-transport {sdma,lsa}` | **`lsa`**: CUs pull and widen on the way in, where SDMA has to land the fp8 and read it back out of local HBM | +| `--no-fuse-quantize` | **on**: folding the narrowing into the reduce costs 1-2% | +| `MORI_GEMM_AR_PULL_BLOCKS` | **48-80**; these are xGMI reads, so the grid throttles outstanding remote requests rather than covering HBM latency | -| gather | M=16384 | -|---|---:| -| bf16 / sdma | 1150.7 us | -| fp8 / sdma | 1018.9 | -| **fp8 / lsa** | **957.3** | - -The ordering holds in the server too: 1070.3ms of GPU busy for bf16 against -1052.1 for fp8/sdma and 1041.2 for fp8/lsa. - -The pull also removes fp8's small-M penalty. With the SDMA gather the two -conversion kernels were a fixed cost against a transfer that shrinks with M, so -fp8 measured **+2.3% (slower)** at M=4096. With the pull there is no fixed -conversion cost left and fp8 wins wherever fusing does. - -### The pull grid - -The single most important tuning parameter, and it is not obvious from the -source. These are *xGMI* reads, so the grid throttles outstanding remote -requests rather than covering HBM latency — it wants roughly a **tenth** of what -the local conversion kernels want. - -| blocks | 16 | 24 | 32 | 48 | 64 | 80 | 128 | 256 | 512 | -|---|---:|---:|---:|---:|---:|---:|---:|---:|---:| -| us | 1261 | 1091 | 1006 | 962 | **959** | 963 | 1018 | 1138 | 1184 | - -Flat from 48 to 80, steep either side. The first implementation launched 512 — -the grid the local quantize kernel uses — and **lost to SDMA by 17%**, which -looked like "LSA is the wrong transport" rather than "the grid is wrong". - -Note this is also not `LSA_BLOCK_CAP`'s 24: that cap is for a kernel moving bf16 -with no arithmetic, while this one moves half the bytes and dequantises them, so -it needs more waves in flight to keep the links fed. - -Sweep it with `MORI_GEMM_AR_PULL_BLOCKS`: +The pull grid is the single most important tuning parameter and is not obvious +from the source -- the first implementation launched 512, the grid the local +quantize kernel uses, and lost to SDMA by 17%. Sweep it with: ```bash for b in 16 24 32 48 64 80 128 256; do @@ -174,206 +174,64 @@ for b in 16 24 32 48 64 80 128 256; do done ``` -### What fp8 costs, numerically - -relL2 against an fp32 host reference goes from **2.35e-3** (bf16 wire, which is -bitwise exact through the collective) to **2.49e-2**. - -That is a floor, not a tuning problem. e4m3 carries 3 mantissa bits, and scale -granularity barely moves it — measured in torch on a `[2048, 7168]` standard -normal payload: - -| scale granularity | relL2 | scale bytes | -|---|---:|---:| -| per row (7168) | 2.646e-2 | 0.06% | -| per 512 | 2.631e-2 | 0.78% | -| per 256 | 2.609e-2 | 1.56% | -| per 128 | 2.572e-2 | 3.12% | -| per 32 | 2.399e-2 | 12.5% | - -**200x the scale bytes buys 9%.** Per-row is therefore the right choice, and -~2.5e-2 is what fp8 costs. - -## Model-level evaluation - -The kernel-level cost above is large — 6.7x the bf16 wire. Whether it *matters* -is a different question, and it needs the model, so this section was measured in -SGLang on DeepSeek-V4-Pro at TP8. The fp8 path was verified live throughout: -`SGLANG_DEBUG_FUSED_WO_B_AR=1` logs relL2 per layer call during the very -requests being scored. - -**At the layer**, relL2 against the unfused path, 488 layer calls: - -| wire | min | median | max | -|---|---:|---:|---:| -| bf16 | 3.706e-3 | — | 4.046e-3 | -| fp8 / sdma | 1.435e-2 | **2.505e-2** | 2.654e-2 | -| fp8 / lsa | 2.153e-2 | **2.496e-2** | 2.683e-2 | - -**At the model output**, it is not detectable. Scoring 10941 tokens of real -source text in a single prefill (mean logprob; lower is a worse model): - -| run | mean logprob | ppl | -|---|---:|---:| -| bf16 | -1.184169 | 3.2680 | -| bf16, rerun | -1.182018 | 3.2609 | -| bf16, again | -1.180106 | 3.2547 | -| **fp8** | **-1.183895** | **3.2671** | -| fp8, with debug | -1.183358 | 3.2653 | - -fp8 lands **inside the bf16 run-to-run band**. Paired per token against the same -bf16 run: +### The GEMM on its own -| pair | mean d | sd d | max abs d | -|---|---:|---:|---:| -| CONTROL bf16 rerun | +0.002151 | 0.229 | 3.42 | -| CONTROL bf16 again | +0.004063 | 0.230 | 4.05 | -| **TEST fp8** | **+0.000274** | **0.219** | **2.93** | +Single process, no collective, so no `torchrun`: -On every statistic, fp8 is closer to bf16 than bf16 is to itself. +`bench_gemm.py` covers both scopes, and they are different questions rather +than two views of one. `--scope kernel` takes operands that are already +quantised and already padded, which is what a tile or a scale layout is chosen +on; `--scope linear` starts from bf16 and puts quantisation, padding and the +op's own tile dispatch inside the measurement, which is what a layer costs. -That band is wide because **this model is already strongly non-deterministic**: -two bf16 runs disagree on ~48% of tokens by more than 0.01 logprob, and greedy -decode diverges within 10-20 tokens. The likely cause is the MoE stage-2 -epilogue, which accumulates with `atomic_fadd`. This is also why greedy token -agreement is useless as a metric here — the bf16-vs-bf16 control is as divergent -as bf16-vs-fp8: - -| pair | ~12k tok | ~24k tok | -|---|---:|---:| -| CONTROL bf16 vs bf16 | 35.9% | 15.6% | -| TEST bf16 vs fp8 | 28.1% | 23.4% | - -Needle-in-a-haystack retrieval at 15140 tokens is **24/24 on both wires** — -saturated, so it bounds gross damage without resolving anything finer. - -**What this does and does not say.** It says fp8 causes no gross degradation and -no measurable shift in next-token distribution on one scoring task. It does not -say quality is unaffected on long-chain reasoning, code or maths — that needs a -task benchmark, which has not been run. Note also that short prompts and decode -never reach this path at all (it engages only at M >= 4096), so only -long-prefill workloads are affected. +```bash +B=benchmark/cco/flydsl/gemm_ar -## Negative results +# the multiply alone, either quantisation, any shape +python $B/bench_gemm.py -n 5120 -k 2048 -m 4096 --scope kernel --impl auto -Kept because each reads as obviously right and the reason it is not cannot be -seen from the source. +# the layer, against SGLang. --impl gemm256,gemm128 pins mori's N tile so the +# dispatch can be measured rather than trusted +python $B/bench_gemm.py --shape wq_b -m 1024 --scope linear \ + --impl auto,gemm256,gemm128,sglang -**Folding the narrowing into the reduce** (`fuse_quantize=True`) saves a 28 MiB -re-read and a kernel launch, and costs 20 us: +# the skinny GEMM (M <= 32); --baseline none drops the SGLang import entirely +python $B/bench_gemv.py --shape wq_b -m 1 --sweep +``` -| | us | -|---|---:| -| split reduce + quantize | 957.4 | -| fused, row stashed in registers | 977.1 | -| fused, row re-read | 982.0 | - -Not register pressure — the re-reading variant keeps no stash and is no better. -It is the thread map: a per-row amax cannot be taken by a block holding only part -of a row, so fusing forces one-wave-per-row, where `sdma_reduce` walks packs with -a flat grid stride and streams a block through all 8 source slices at once. - -Re-measured with three repeats each, the cost holds and is if anything larger: -926.6 against 903.2 on `fp8/lsa`, 980.4 against 961.9 on `fp8/sdma`. -`build_sdma_phases` has always defaulted it off, so the op and SGLang take the -fast path; the *benchmark* defaulted it on until `96bd9522`, and any number -produced by an older driver without an explicit `--no-fuse-quantize` reads about -23 us slow. - -**Firing the gather's puts from inside the reduce** (`fuse_reduce_push=True`) -looked like the safest of the three: the push sends this rank's *own* slice, so -unlike the pull it has no cross-rank dependency, and SDMA is a copy engine so it -costs no CU time. It reaches parity and not a win — against 1148.7us unfused: - -| bands | 1 | 4 | 8 | 16 | 32 | -|---|---:|---:|---:|---:|---:| -| `publish="writethrough"` | 1162.1 | **1157.2** | 1160.3 | 1190.1 | 1323.5 | -| `publish="fence"` | 1227.3 | 1380.1 | 1630.3 | 2137.2 | 3026.7 | - -The gap between those rows is the useful result, and it is a lesson about how to -pay for a release rather than whether to. - -Handing a range to a copy engine *does* need one: the engine reads over the -fabric, not through a CU's cache, so `s_waitcnt vmcnt(0)` alone is not enough — -it only retires the stores as far as this XCD's L2. But the release can be paid -two ways. Releasing to system scope **after** the stores is L2-writeback work -charged once per block per band (256 x bands of it, ~61us per band, against a -reduce that is only 42us in total — unrepayable). Storing with `sc0+sc1` so the -bytes never stop in L1 or L2 makes the `waitcnt` itself the release, and that is -**free here**: applying the same store policy to the plain unfused reduce moves -it 1153.4 -> 1148.7us, i.e. nothing. This output is written once and nothing -local reads it again before the gather, so holding it in L2 bought nothing. - -What remains after that fix is small on both sides and nearly cancels. At -bands=1 the publish carries the mechanism's cost with none of its benefit — -1162.1 vs 1148.7, about 13us for the per-band 256-block `wait_barrier`, the -counter atomic and the elected block's locked puts. Four bands buy back about -5us of overlap before the sync cost takes over again. - -Dropping the release entirely is not an option even though it briefly looks like -one: with cached stores and no fence the kernel reaches 1157.9us, but at 32 bands -it produced relL2 1.8e-2 against the 2.35e-3 floor, differing per rank. The same -32 bands are exact under either correct publish mode, which rules out an indexing -bug. - -The contrast with the GEMM's fused scatter is the transferable part: there the -publish is amortised against a 437us transfer hidden behind a compute-bound GEMM; -here against 42us of bandwidth-saturated reduce. The mechanism pays when what is -hidden is much larger than the cost of publishing it. - -Re-measured after the window-geometry work below, with three alternating -repeats rather than a sweep, it is a clearer loss than the table suggests: -1128.4 us on against 1112.0 off, +16.4 us, against spreads of 2.2 and 1.5 us. - -**Hoisting the window geometry out of `lsa_ptr`.** `cco_lsa_ptr` is -`winBase + peer*stride + offset` and loads both fields on every call, through a -*generic* pointer -- which has to be a `flat_load`, since the compiler cannot -rule out LDS, so it counts against `lgkmcnt` as well as `vmcnt`. FlyDSL emits it -as an opaque extern call, and a kernel storing through addresses derived from -that base gives LLVM no way to prove the loads are not clobbered. - -Reading the geometry once and doing the arithmetic in the DSL removes all of -that. It was tried four ways -- hoisting out the band loop, a `lsa_geometry()` -API, `global_load` accessors in C++ (`cco_lsa_win_base` / `cco_lsa_stride`, which -take an `address_space(1)` pointer so each is a single `global_load`), and -finally `cco.CachedWindow`, which reads both in its constructor so a kernel -changes by one line. All four measured nothing on `kernels_sdma` (21 call sites) -and `kernels_fused`, in every wire configuration. - -It pays in exactly one place (`ptpc`, three alternating repeats): - -| | Window | CachedWindow | -|---|---|---| -| `split-lsa` | 1264.69 1264.23 1264.85 | 1250.25 1256.85 1254.53 | -| `fused-sdma` | 1110.45 1109.77 1112.85 | 1110.53 1113.69 1112.61 | - --10.7 us on `split-lsa`, against a 0.6 us spread; nothing on `fused-sdma`. -`ar_1stage`/`ar_2stage` build nine peer addresses in *every block* of a short -kernel; everywhere else the addresses are built once per launch against a body -that runs for a millisecond. **Count address constructions per launch, not -`grep -c lsa_ptr`.** - -The same holds against PR #662's branch in `blockscale`, two alternating -repeats: `split-lsa` 1463.6 -> 1455.6 us, while `split-sdma` (1461.2 -> 1459.9), -`fused-sdma` bf16 (1144.9 -> 1145.8) and fp8/lsa (949.2 -> 949.8) do not move. - -Two things worth carrying. A `CachedWindow` cannot cross an `scf.if` -- FlyDSL -captures every variable an if body reads as state and requires single MLIR -values, which `Window` satisfies only by having exactly one field -- and the way -out is to compute the addresses before the branch, which is what the offsets -usually allow. And pin `--quant` when comparing against anything: the benchmark -defaults to `ptpc`, ~3% faster than the `blockscale` every number on this page is -quoted in, and reading one against the other looks exactly like a machine that -drifts overnight. - -**A CK-shaped 4-wave GEMM**, chasing a 22% gap against CK's block-scale kernel -at the same shape, reached CK's instruction mix and not its speed. Ten hypotheses -were falsified by measurement; hardware counters show identical `SQ_INSTS_MFMA` -(7,340,032) and `SQ_VALU_MFMA_BUSY_CYCLES` (234,881,024), VALU within 1%, -`MemUnitStalled` at approximately zero — but `SQ_WAIT_ANY` 156.0M against 120.4M. -A K-sweep puts the whole difference per-iteration: our fixed cost is *lower* -(51.6 us against 66.0), while each K-block costs 20.9 us against 14.5. The gap is -wait, not work. See commits `cc696762`, `54fef960`, `9f990637`, `1dead3fc`. +The SGLang baseline is optional throughout -- mori's own numbers need nothing +installed but mori -- but `--scope linear` does need it, because the +quantisation inside the measurement is SGLang's quantiser. + +Comparing against SGLang's own route needs its harness, which lives in the +SGLang tree because the baseline does +(`test/registered/perf/models/bench_mori_mxfp8_gemm.py`); it takes `-n`/`-k` for +any shape, or a name for any of V4.1-Flash's layers. + +## Measurement discipline + +These are the ones that have produced a confident wrong number here. The reasons +are in +[Measurement traps](https://github.com/ROCm/mori/blob/main/python/mori/ops/gemm_ar/README.md#measurement-traps). + +- **Amortise the capture.** A single-call CUDA-graph replay has a **13.40 us + floor** on this box, which is most of any small-M measurement. Use + `benchmark/cco/flydsl/gemm_ar/timing.py`. +- **Report cold as well as hot.** A repeated call reads the weight out of the + 256 MB LLC at 1.7x the bandwidth a forward pass gets. Cold is what predicts a + server. +- **One M per process.** A multi-M sweep once reported 697 us and 1083 us for + two M that pad to the same size; isolated, both were 697. +- **Assert the fast path actually ran.** A path that silently declined looks + exactly like one that ran and was not worth it. Both integrations have a shape + log for this (`SGLANG_OPT_FUSED_WO_B_AR_SHAPE_LOG=1`, + `SGLANG_OPT_MORI_MXFP8_GEMM_SHAPE_LOG=1`). +- **Check `BUILD_CCO_SDMA=ON` before believing any end-to-end number.** With it + off the fused path measures *faster* than it is, because it is not moving + data. +- **Let the SDMA queues drain between runs.** Back-to-back 8-rank runs hit + `hsaKmtCreateQueueExt` failures (`anvil.cpp:237`) if a previous run's ranks + have not exited. ~10s is enough; a leftover server holding queues is not. ## Reproducing the end-to-end numbers @@ -385,6 +243,7 @@ export MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo export PYTHONPATH=/path/to/mori-with-sdma export SGLANG_OPT_FUSED_WO_B_AR=1 export SGLANG_OPT_FUSED_WO_B_AR_FP8_GATHER=1 # optional, the fp8 wire +export SGLANG_OPT_MORI_MXFP8_GEMM=1 # optional, the GEMM alone export SGLANG_DEBUG_FUSED_WO_B_AR=1 # optional, logs per-layer relL2 sglang serve --model-path --tp 8 \ @@ -393,13 +252,38 @@ sglang serve --model-path --tp 8 \ --enforce-shared-experts-fusion ``` +V4.1-Flash is the verified MI350X low-latency cell plus mori's variables, at +TP4. `AITER_BF16_FP8_MOE_BOUND=0` is load-bearing and is *not* part of the +published cell: the checkpoint has `swiglu_limit=10.0`, so its MoE takes AITER's +clamped-SwiGLU INTERLEAVE path, where no `ck_moe_stage1` kernel exists for +(bf16 activation x fp4 weight) -- every M below the default bound of 256, which +is every decode step, raises "Unsupported kernel config for moe heuristic +dispatch". + +```bash +export SGLANG_USE_AITER=1 SGLANG_MOE_PADDING=1 +export AITER_FLYDSL_FORCE_REDUCE=1 ROCM_QUICK_REDUCE_QUANTIZATION=NONE +export AITER_BF16_FP8_MOE_BOUND=0 + +sglang serve --model-path --tp 4 --ep-size 4 \ + --disable-radix-cache --mem-fraction-static 0.78 \ + --chunked-prefill-size 16384 \ + --speculative-algorithm DSPARK --speculative-dspark-block-size 5 \ + --cuda-graph-max-bs 64 --cuda-graph-backend-prefill breakable \ + --cuda-graph-max-bs-prefill 4096 +``` + +`--chunked-prefill-size` has to be at least the wire's floor or no prefill chunk +ever reaches it and the fused variants quietly measure the base path. + `--mem-fraction-static` has to leave room for the symmetric window, which is VMM -memory **outside** torch's allocator: 700 MiB on the bf16 wire and 812 MiB on -fp8, which needs the extra staging region. +memory **outside** torch's allocator: 700 MiB on V4-Pro's bf16 wire and 812 MiB +on fp8, which needs the extra staging region; 163 MiB at V4.1-Flash's TP4 shape +with `m_max=5120`. The capture protocol matters more than it looks. Profile a *different* prompt than the one used to warm up, `flush_cache` between them, and compare the same -request on both sides — an earlier A/B without those controls reported a **+8.3% +request on both sides -- an earlier A/B without those controls reported a **+8.3% regression** that did not exist. ```python @@ -408,30 +292,3 @@ main = "The quick brown fox jumps over the lazy dog. " * 2800 gen(warm); post("/flush_cache"); post("/start_profile") gen(main); post("/stop_profile") ``` - -GPU busy time over that capture: - -GPU busy over that capture: - -| | GPU busy | wall | vs unfused | -|---|---:|---:|---:| -| unfused (GEMM + NCCL) | 1096.0 ms | 1.1969 s | — | -| fused, bf16 wire | 1070.3 | 1.1728 | -2.3% | -| fused, fp8 / sdma | 1052.1 | 1.1455 | -4.0% | -| **fused, fp8 / lsa** | **1041.2** | **1.1385** | **-5.0%** | - -An earlier capture of the same four read 1101.9 / 1077.2 / 1050.4 / 1048.2, so -this reproduces to about half a percent. - -> **Check `BUILD_CCO_SDMA=ON` before believing any end-to-end number.** With it -> off every put silently does nothing: the all-reduce returns mostly the local -> slice, the model still answers fluently, every mori kernel still appears in -> the profile, and the fused path measures **faster** than it is because it is -> not moving data -- -6.8% instead of -2.3%, with fp8/lsa appearing *worst* of -> the three rather than best, since the pull is the one leg that does not go -> through SDMA. Perplexity catches it and nothing cheaper does: 862511 against -> 3.26 on the same text. A short prompt cannot catch it either, because fusing -> needs M >= 4096. - -The layer-level win is larger than the end-to-end one because `wo_b` is about -12% of the profile. diff --git a/python/mori/ops/gemm_ar/README.md b/python/mori/ops/gemm_ar/README.md index eaeddf60e..5dcdf1e35 100644 --- a/python/mori/ops/gemm_ar/README.md +++ b/python/mori/ops/gemm_ar/README.md @@ -1,33 +1,80 @@ -# Fused GEMM + all-reduce (`mori.ops.gemm_ar`) +# `mori.ops.gemm_ar` -- an fp8 GEMM, with or without the all-reduce fused in -Fuses a `RowParallelLinear`'s fp8 GEMM with the all-reduce that follows it, over -cco's SDMA copy engines. Built for DeepSeek-V4-Pro's `wo_b` and measured there: -**1146.2us against the model's 1419.5us, -19.4%**, at `[16384, 7168] K=2048` on -8x MI355X. +Three operators over one mxfp8/blockscale GEMM kernel, for a tensor-parallel +attention block: -The public API is `GemmAllReduceOp` in `op.py`. `kernels_fused.py`, -`kernels_sdma.py` and `kernels_lsa.py` hold the kernels, `layout.py` every -window offset and count, and `_gemm_a8w8_8wave.py` / `_shuffle.py` the two -pieces vendored from aiter so mori does not depend on it. +| op | for | what it does | +|---|---|---| +| `GemmAllReduceOp` (`op.py`) | a `RowParallelLinear` at prefill M | the GEMM with the all-reduce's scatter fused into its epilogue | +| `Mxfp8GemmOp` (`gemm.py`) | any mxfp8 linear at prefill M | the same GEMM, nothing fused, no communicator | +| `Mxfp8GemvOp` (`gemv.py`) | any mxfp8 linear at decode M (<= 32) | a skinny GEMM built for a handful of tokens | -Full measurements, how to reproduce them, and the model-level evaluation of the -fp8 wire are in +Built for DeepSeek-V4-Pro's `wo_b` and measured there: **1146.2us against the +model's 1419.5us, -19.4%** at `[16384, 7168] K=2048` on 8x MI355X. On +DeepSeek-V4.1-Flash the same fusing is worth -25% -- but there **most of the win +is the multiply, not the overlap**, which is why the GEMM is also shipped on its +own. See [Measured results](#measured-results). + +`kernels_fused.py`, `kernels_sdma.py`, `kernels_lsa.py` and `kernels_gemv.py` +hold the kernels, `layout.py` every window offset and count, and +`_gemm_a8w8_8wave.py` / `_shuffle.py` the two pieces vendored from aiter so mori +does not depend on it. + +How to *run* the benchmarks, and a one-table summary of what they say, are in [`docs/MORI-GEMM-AR-BENCHMARK.md`](../../../../docs/MORI-GEMM-AR-BENCHMARK.md). +Every full table lives here, next to the code it is about. + +## Contents + +- [What it does](#what-it-does) +- [Using it](#using-it) + - [Fused GEMM + all-reduce](#fused-gemm--all-reduce) + - [The GEMM without a collective](#the-gemm-without-a-collective) + - [fp8 on the wire](#fp8-on-the-wire) + - [Benchmarks and tests](#benchmarks-and-tests) +- [Measured results](#measured-results) + - [The mxfp8 GEMM, on its own](#the-mxfp8-gemm-on-its-own) + - [The mxfp8 GEMM against SGLang, across every shape](#the-mxfp8-gemm-against-sglang-across-every-shape) + - [DeepSeek-V4-Pro (blockscale, TP8)](#deepseek-v4-pro-blockscale-tp8) + - [DeepSeek-V4.1-Flash (mxfp8, TP4)](#deepseek-v41-flash-mxfp8-tp4) + - [End to end, in SGLang](#end-to-end-in-sglang) + - [Measurement traps](#measurement-traps) +- [How it works](#how-it-works) + - [Why the SDMA transport needs no epilogue change](#why-the-sdma-transport-needs-no-epilogue-change) + - [Completion protocol](#completion-protocol) + - [The chunks race: it was the GEMM, and it is gone](#the-chunks-race-it-was-the-gemm-and-it-is-gone) + - [Mode comparison, once the chunks are unblocked](#mode-comparison-once-the-chunks-are-unblocked) + - [The C store: three stages off gcnasm](#the-c-store-three-stages-off-gcnasm) + - [What gcnasm does differently](#what-gcnasm-does-differently) +- [Negative results](#negative-results) + - [A CK-shaped 4-wave GEMM](#a-ck-shaped-4-wave-gemm) + - [Persistent tiles](#persistent-tiles) + - [Direct LSA (`--mode fused-lsa`)](#direct-lsa---mode-fused-lsa) + - [Folding the narrowing into the reduce](#folding-the-narrowing-into-the-reduce) + - [Firing the gather's puts from inside the reduce](#firing-the-gathers-puts-from-inside-the-reduce) + - [Hoisting the window geometry out of `lsa_ptr`](#hoisting-the-window-geometry-out-of-lsa_ptr) ## What it does -aiter's 8-wave fp8 GEMM with the all-reduce's scatter fused into its epilogue. +One GEMM kernel, compiled two ways. With the epilogue's tail switched on it +pushes each destination's slice as it is produced, which is `GemmAllReduceOp`; +with it off it is a plain GEMM that returns an ordinary tensor, which is +`Mxfp8GemmOp`. `Mxfp8GemvOp` is a *different* kernel for the token counts where +a 256-row tile has nothing to fill it. + +The rest of this section, and all of [How it works](#how-it-works), is about the +fused one. -The target is DeepSeek-V4-Pro's ``wo_b``: a RowParallelLinear whose per-rank GEMM -is ``[M,1024] x [7168,1024]`` fp8 -> ``[M,7168]`` bf16, immediately all-reduced. -Split, that is ``gemm(); all_reduce()``. Fused, the GEMM's C lands directly in a +The target is DeepSeek-V4-Pro's `wo_b`: a RowParallelLinear whose per-rank GEMM +is `[M,1024] x [7168,1024]` fp8 -> `[M,7168]` bf16, immediately all-reduced. +Split, that is `gemm(); all_reduce()`. Fused, the GEMM's C lands directly in a registered cco window and each destination's slice is pushed by the copy engine as soon as its last tile is written, so the reduce-scatter transfer overlaps the rest of the GEMM instead of following it. Only the *scatter* half of the all-reduce is absorbed. The reduce and all-gather -phases still run as their own kernels, reused verbatim from ``ar.kernels_sdma`` -(``build_sdma_phases``), with ``scatter`` swapped for its drain-only twin. +phases still run as their own kernels, reused verbatim from `ar.kernels_sdma` +(`build_sdma_phases`), with `scatter` swapped for its drain-only twin. **Nothing here is a production kernel.** aiter's tuned CSV picks a ck/asm/cktile backend for (N=7168, K=1024), not this one, so the fused-vs-split comparison is @@ -35,10 +82,13 @@ internally valid but is not a claim about DSV4 as shipped. ## Using it +### Fused GEMM + all-reduce + Needs a mori built with `BUILD_CCO_SDMA=ON` and `MORI_ENABLE_SDMA=1` in the environment. Setting the variable alone is not enough: a host library built without the flag has no queues, so every put silently does nothing and the -all-reduce quietly produces zeros. +all-reduce quietly produces zeros. (The two standalone ops below need neither -- +they move no data between ranks.) ```python import torch, torch.distributed as dist @@ -121,7 +171,66 @@ communicator. Synchronise first if anything may still be in flight. width but a different exponent bias, and the MMA atom implements OCP's, so it is rejected rather than silently returning a result 4x too large. -## fp8 on the wire +### The GEMM without a collective + +Neither of these takes a communicator, a window or a rank: the output is an +ordinary tensor. They exist because at DeepSeek-V4.1-Flash's shapes the GEMM is +worth more than the fusing, and because a `ColumnParallelLinear` has no +all-reduce to fuse with at all. + +```python +from mori.ops.gemm_ar import ( + Mxfp8GemmOp, Mxfp8GemvOp, preshuffle_a_scale, preshuffle_b, + supports_gemm, supports_gemv, +) + +# Once per weight. w_exps is the checkpoint's [N/32, K/32] ue8m0 bytes. +b_shuffled = preshuffle_b(w_fp8) # [N, K] +b_scale = w_exps.t().contiguous().to(torch.int32).reshape(-1) # K-block major + +# Prefill. M only has to be a multiple of 64 -- the grid is ceildiv(M, BLOCK_M) +# and the tail block masks its stores. +assert supports_gemm(N, K) +op = Mxfp8GemmOp(n=N, k=K) +m_pad = op.padded_m(x.shape[0]) +a_fp8, a_exps = quantize_mxfp8(op.pad_rows(x, m_pad)) # pad before quantising +out = op(a_fp8, b_shuffled, preshuffle_a_scale(a_exps), b_scale)[: x.shape[0]] + +# Decode. No padding and none needed; M is a runtime argument, and both scales +# go in as the checkpoint stores them -- w_exps itself, not b_scale. +assert supports_gemv(N, K) +gemv = Mxfp8GemvOp(n=N, k=K) +out = gemv(x_fp8, b_shuffled, x_exps, w_exps) +``` + +**They share the weight and not the scales.** `preshuffle_b` output serves both, +so a server shuffles once. But `Mxfp8GemmOp` wants the A scale through +`preshuffle_a_scale` -- K-block major, four M tiles to a dword -- while +`Mxfp8GemvOp` takes **both scales exactly as the checkpoint stores them**, +row-major ue8m0 bytes, `[M, K/32]` and `[N/32, K/32]`. That is not an +inconsistency: the GEMM's sixteen lanes want sixteen *rows* of one K block and +have to be coalesced, where the GEMV's are sixteen *tokens*, M is at most 32, and +the whole A scale is under a kilobyte. See +[the layout results](#the-mxfp8-gemm-on-its-own). + +**Shapes.** `supports_gemm(n, k)` and `supports_gemv(n, k)` answer whether a +shape is expressible -- K a multiple of 128, N a multiple of 256 for the GEMM +and of 32 for the GEMV. M is deliberately not an argument to either: any M is +*servable*, and whether it is *profitable* is the caller's call. + +**Which M is profitable is decided by the grid, not by M.** The GEMM's tile is +256x256, so a caller should gate on `ceildiv(M, 256) * (N / 256)` and hand over +at about 80 workgroups; a floor on M alone is only ever right for the N it was +fitted to. The measurement behind that, and what it costs to get it wrong, is +[here](#the-mxfp8-gemm-against-sglang-across-every-shape). + +**The GEMV's ceiling is 32 tokens**, because a token is an MFMA row and two tiles +of them is where the register file runs out. It picks a compile-time +configuration per M bucket from a tuned table in `gemv.py`; an untuned shape gets +a heuristic rather than an error. Between 32 and the GEMM's floor neither wins, +and the caller should use whatever it had. + +### fp8 on the wire `gather_dtype="fp8"` sends the all-gather leg as e4m3 with one fp32 scale per row, halving its bytes. That leg is ~40% of a fused layer and already runs at @@ -134,7 +243,7 @@ Measured, fused-sdma at `[16384, 7168]` K=2048 on 8x MI355X: | bf16 | 1151.5 | 2.35e-3 | | fp8 | **1033.9** (-10.2%) | **2.49e-2** | -### Who moves the fp8 gather +#### Who moves the fp8 gather `gather_transport="lsa"` (the default for fp8) pulls each peer's slice over xGMI into registers and widens it on the way to memory. `"sdma"` pushes with @@ -158,49 +267,13 @@ and it wants roughly a tenth of what the local conversion kernels want. The first version launched 512 -- the quantize grid -- and lost to SDMA by 17%. -### Folding the narrowing into the reduce: measured, and it loses - -`fuse_quantize=True` makes the reduce write its bf16 and narrow to fp8 from the -same accumulators, saving a 28 MiB re-read and a launch. It is **off**, because -it costs 20 us rather than saving 12: - -| | us | -|---|---:| -| split reduce + quantize | 957.4 | -| fused, row stashed in registers | 977.1 | -| fused, row re-read | 982.0 | - -Not register pressure -- the re-reading variant keeps no stash and is no better. -It is the thread map: a per-row amax cannot be taken by a block holding only -part of a row, so fusing forces one-wave-per-row, where `sdma_reduce` walks -packs with a flat grid stride and streams a block through all 8 source slices at -once. Confining a wave to 14 KiB at a time costs the reduce more than the -re-read saves. - -It is not free and the cost is not a tuning problem. e4m3 carries 3 mantissa -bits, so one rounding costs ~2.1e-2 on a normal payload whatever the scale -granularity -- per-row measures 2.65e-2 and per-32 measures 2.40e-2, 9% better -for 200x the scale bytes. Going from 2.35e-3 to 2.49e-2 is the price of the -10%, and whether that is payable is a model-level question, not a kernel one. - -It also only pays at large M. The two conversion kernels are a fixed cost -against a transfer that shrinks with M, so on the standalone all-reduce it is -+2.4% at M=4096 and -11.2% at M=16384. Per phase at M=16384: quantize 12.4us, -dequantize 88.2us, against ~218us saved on the push. - -The scatter leg stays bf16. It carries partial sums that are then added across -every rank, so its fp8 error compounds rather than being a single rounding, and -it is already mostly hidden behind the GEMM -- its 1140 GB/s is not a bandwidth, -it is the tell that the pushes went out from the epilogue and the drain is only -waiting for the tail. `scatter_dtype="fp8"` sizes its regions but raises -`NotImplementedError`. - -## Benchmarks and tests +### Benchmarks and tests ```bash # numerics MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo \ - pytest tests/python/cco/test_gemm_ar.py tests/python/cco/test_flydsl_ar.py + pytest tests/python/cco/test_gemm_ar.py tests/python/cco/test_flydsl_ar.py \ + tests/python/cco/test_gemm_ar_op.py # the full mode comparison at the model's shape MORI_ENABLE_SDMA=1 MORI_SOCKET_IFNAME=lo python -m torch.distributed.run \ @@ -213,42 +286,663 @@ the environment only rebuilds the device bitcode -- a host library built without the flag has no queues, so every put silently does nothing and the all-reduce quietly produces zeros. -## A measured negative result, for the record +Every flag, every sweep driver, and how to drive the whole thing from a server +are in +[`docs/MORI-GEMM-AR-BENCHMARK.md`](../../../../docs/MORI-GEMM-AR-BENCHMARK.md). +Read [Measurement traps](#measurement-traps) before trusting a small-M number +from a harness of your own. + +## Measured results + +MI355X (gfx950), one node, otherwise idle, occupancy recorded before and after +each run. Kernel timings are the median over repeated iterations, maximum over +ranks. Run-to-run spread is about **2%**, so differences below that are not +differences. + +Two models are covered and **they do not reach the same conclusion**, so they +are kept apart rather than averaged: + +| | DeepSeek-V4-Pro | DeepSeek-V4.1-Flash | +|---|---|---| +| `wo_b` per rank | `[M, 7168]` K=2048, TP8 | `[M, 5120]` K=2048, TP4 | +| quantisation | 1x128 / 128x128 fp8 block scale | **32-wide ue8m0 (mxfp8)** | +| `--quant` | `blockscale` | `mxfp8` | +| fusing, best wire | **-25%** at M=16384 | **-25%** at M=16384 | +| what dominates it | the overlap | the fp8 wire and the GEMM | + +### The mxfp8 GEMM, on its own + +CDNA4's `v_mfma_scale_f32_16x16x128_f8f6f4` takes e8m0 scales as instruction +operands, so a 32-wide ue8m0 GEMM needs no dequantisation arithmetic at all -- +unlike `_BlockScaleK`'s promote/rescale chain, which exists precisely because a +128-wide block cannot be expressed that way. Two layout results came out of +making it fast, both about *addresses* rather than bytes. + +**The scales must be K-block major.** Coalescing happens per quarter-wave, which +is exactly the sixteen lanes of one MFMA block group, and those sixteen want +sixteen consecutive rows of one K block. K-block major puts them at consecutive +addresses -- one access per instruction. The quantiser's own `[M, K/32]` row +major puts them `K/32` bytes apart: + +| | K-block major | row major | +|---|---:|---:| +| `SQ_INSTS_VMEM` | 896,000 | 911,360 | +| `TCP_TOTAL_CACHE_ACCESSES` | 8,622,080 | **27,729,920** | +| GEMM, M=4096..16384 | — | **+50% to +67%** | + +The same instruction count and 3.2x the cache accesses. Predicted +19,660,800 +accesses (4 per instruction becoming 64), measured +19,107,840. Stripping the +byte-select arithmetic off the row-major path moved 117.4 -> 110.9us, so 41 of +the 47us is addresses. The data is 16 KB and fully L1-resident; the currency is +requests, not bytes. + +**Four M tiles pack into one scale dword.** A lane's four A tiles differ only by +sixteen rows and share the K block, and `opsel_b` on the scaled MFMA is an +atom-time attribute naming which byte of the 32-bit scale operand to read. So +one load serves all four and the byte select is free: + +| | unpacked | packed | +|---|---:|---:| +| `SQ_INSTS_VMEM` | 896,000 | **634,880** (-29.1%) | +| `SQ_INSTS_VALU` | 2,252,800 | 2,268,160 (+0.7%) | +| GEMM | — | **-5% to -7.5%** | -An earlier revision carried `kernels_preshuffle4w.py`, a port of CK's 4-wave -B-out-of-LDS shape, chasing a 22% gap between this GEMM and CK's at the same -shape. It reached CK's instruction mix and not its speed, and it was deleted -rather than carried (it needed ~1000 lines of further aiter vendoring to serve -a kernel nothing calls). What the investigation ruled out, since the same -ground should not be walked twice: +VALU flat is the check that matters: it confirms `opsel` really is free rather +than degrading to a shift and mask. CK does the same thing in +`preShuffleScaleBuffer_gfx950`; this packs all four M tiles rather than its +`MNXdlPack=2`, which fits opsel's two bits exactly and needs no cross-K-step +state, so the mainloop is untouched. -Ten hypotheses were falsified by measurement -- promote arithmetic, register -spill, store width, VALU scheduling groups, MFMA batch size, scale loads, -address hoisting, load-to-use distance, occupancy, and dependency structure. -Hardware counters (`rocprofv3 --pmc`) show *identical* `SQ_INSTS_MFMA` -(7,340,032) and `SQ_VALU_MFMA_BUSY_CYCLES` (234,881,024), VALU within 1%, and -`MemUnitStalled` at approximately zero -- but `SQ_WAIT_ANY` at 156.0M against -CK's 120.4M. A K-sweep puts the whole difference per-iteration: our fixed cost -is *lower* (51.6us against 66.0us), while each K-block costs 20.9us against -14.5us. The gap is wait, not work, and it is not in any of the places listed -above. See commits `cc696762`, `54fef960`, `9f990637`, `1dead3fc` for the -traces. +`preshuffle_a_scale` is that layout, and a quantiser can write it directly -- +the permutation stays inside the 64 bytes one program already owns, so it is +address arithmetic rather than traffic, and measures slightly *faster* than the +stock kernel. -## Why the SDMA transport needs no epilogue change +**Against the alternatives**, bf16 in and bf16 out, the whole pipeline including +quantisation, M=16384: + +| route | us | +|---|---:| +| bf16 (`fake_quant` + hipBLASLt bf16 GEMM) | 281.3 | +| SGLang's own mxfp8 (`tl.dot_scaled`) | 274.9 | +| **mori mxfp8** | **196.0** | + +**-30.3%**, before any fusion. That is most of what the integration is worth, +and it is not the overlap. + +> Those three came off a harness that predates `timing.py`: hot only, one call +> per graph. At M=16384 each route is 200-280us, so a 13.4us floor is ~5% on all +> three columns and the ratio is close to right -- `bench_gemm.py --scope linear`, on the +> fixed timer, puts the same comparison at -27%. The split between the two +> baseline routes is not reproducible from what is checked in; the mori-vs-best +> -baseline column is, and is the one to quote. + +### The mxfp8 GEMM against SGLang, across every shape + +The table above is two shapes and one kernel. This is every fp8 linear the +checkpoint has -- read off `layers.N.attn.*` and `layers.N.ffn.shared_experts.*` +and split by the parallelism each is declared with in SGLang's +`models/deepseek_v4.py`, at three TP degrees -- against **each of mori's tiles +separately**, rather than against whichever one its own dispatch picks. 120 +points, `sweep.py gemm-full`, cold. + +Baseline is `mxfp8_native_blockscaled_linear`, identical operands. Every cell is +the whole pipeline from bf16 with quantisation included, because that is what +the baseline does. Negative is mori faster. + +> An earlier revision of these tables cooled only the fp8 weight. SGLang's route +> is `hipblaslt_bf16` on most of the large-M points and reads the *dequantised +> bf16* weight, which stayed pinned -- so the baseline was hot where mori was +> cold, and mori's margin was **understated** by up to 28 percentage points. +> `bench_gemm.py` now rotates every weight the chosen route may read, and the +> `gemv` and `dot_scaled` rows, which never read that tensor, are unchanged +> within a point either way. + +**mori's GEMM on the 256x256 tile:** + +| layer | N x K | M=64 | M=256 | M=1024 | M=2048 | M=4096 | M=8192 | M=16384 | +|---|---|---|---|---|---|---|---|---| +| `wq_b` (TP4) | 8192 x 1280 | +113% | +88% | +3% | **-28%** | **-30%** | **-32%** | **-43%** | +| `wo_b` (TP4) | 5120 x 2048 | +126% | +52% | **-5%** | **-34%** | **-26%** | **-31%** | **-27%** | +| `wq_a` (TP4) | 1280 x 5120 | +316% | +221% | +87% | +38% | +5% | **-19%** | **-15%** | +| `wkv` (TP4) | 512 x 5120 | +438% | +309% | +140% | +100% | +36% | +2% | **-21%** | +| `wqkv_a` (TP4) | 1792 x 5120 | +302% | +175% | +62% | +12% | **-16%** | **-28%** | **-24%** | +| `wo_a` (TP4) | 2048 x 4096 | +263% | +155% | +42% | +13% | **-41%** | **-25%** | **-22%** | +| `shared gate_up` (TP4) | 1152 x 5120 | n/a | n/a | n/a | n/a | n/a | n/a | n/a | +| `wq_b` (TP8) | 4096 x 1280 | +131% | +85% | +9% | **-14%** | **-22%** | **-25%** | **-27%** | +| `wo_b` (TP8) | 5120 x 1024 | +85% | +55% | **-8%** | **-28%** | **-23%** | **-26%** | **-26%** | +| `wo_a` (TP8) | 1024 x 4096 | +361% | +253% | +105% | +30% | +4% | **-39%** | **-31%** | +| `wq_b` (TP1) | 32768 x 1280 | +14% | **-22%** | **-34%** | **-34%** | **-31%** | **-23%** | **-17%** | +| `wo_b` (TP1) | 5120 x 8192 | +234% | +140% | +19% | **-19%** | **-25%** | **-31%** | **-21%** | + +**mori's GEMM on the 256x128 tile:** + +| layer | N x K | M=64 | M=256 | M=1024 | M=2048 | M=4096 | M=8192 | M=16384 | +|---|---|---|---|---|---|---|---|---| +| `wq_b` (TP4) | 8192 x 1280 | +85% | +77% | -1% | **-2%** | **-3%** | -1% | **-17%** | +| `wo_b` (TP4) | 5120 x 2048 | +91% | +41% | **-10%** | **-6%** | **-10%** | **-12%** | +0% | +| `wq_a` (TP4) | 1280 x 5120 | +243% | +178% | +60% | +20% | **-8%** | +11% | -1% | +| `wkv` (TP4) | 512 x 5120 | +349% | +256% | +111% | +75% | +20% | **-7%** | **-25%** | +| `wqkv_a` (TP4) | 1792 x 5120 | +238% | +141% | +41% | -2% | **-26%** | **-4%** | +1% | +| `wo_a` (TP4) | 2048 x 4096 | +201% | +123% | +25% | -0% | **-47%** | **-4%** | -1% | +| `shared gate_up` (TP4) | 1152 x 5120 | +243% | +185% | +71% | +30% | **-3%** | +12% | +1% | +| `wq_b` (TP8) | 4096 x 1280 | +93% | +70% | +3% | **-19%** | +4% | +2% | -1% | +| `wo_b` (TP8) | 5120 x 1024 | +58% | +49% | **-10%** | +2% | **-5%** | **-6%** | +2% | +| `wo_a` (TP8) | 1024 x 4096 | +282% | +207% | +87% | +20% | -1% | **-40%** | **-9%** | +| `wq_b` (TP1) | 32768 x 1280 | +0% | **-25%** | **-5%** | +2% | +8% | +33% | +47% | +| `wo_b` (TP1) | 5120 x 8192 | +175% | +102% | -0% | +15% | **-11%** | **-17%** | +19% | + +#### Which tile, and where they cross + +Comparing each tile against the *baseline* is the wrong way to read those two +tables -- it mixes in how fast the baseline happens to be on each shape. +Comparing them against **each other** is unambiguous, and the separation is +total: 110 points where both were measured, and not one on the wrong side of a +single threshold. + +| wide-tile grid | points | 256 faster | 128 faster | mean 256/128 - 1 | +|---|---|---|---|---| +| 1-16 | 33 | 0 | 33 | +16.9% | +| 17-32 | 25 | 0 | 25 | +14.3% | +| 33-64 | 6 | 0 | 6 | +10.6% | +| 65-128 | 15 | 0 | 15 | +8.9% | +| 129-192 | 4 | 4 | 0 | -28.9% | +| >192 | 27 | 27 | 0 | -26.4% | + +The narrow tile is worth a flat 9-17% below the crossover and costs 26-29% above +it, and the crossover is a step rather than a slope. `_WIDE_TILE_MIN_GRID = 140` +sits inside the gap; sweeping the threshold over this data gives **zero regret** +for anything in 129-192 and gets worse either side. This is the third data set +the constant has been re-derived on, including two with measurement bugs in +them, and it has not moved. + +#### Which M is worth serving at all + +Taking the better of the two tiles at each point: + +| wide-tile grid | points | best mori vs baseline | +|---|---|---| +| 1-32 | 65 | +246.7% | +| 33-64 | 7 | -0.7% | +| 65-128 | 16 | -10.2% | +| 129-256 | 10 | -26.6% | +| >256 | 22 | -26.4% | + +Scoring gate thresholds over the 77 points the gate actually decides -- M above +the GEMV's 32 tokens, and a shape `supports_gemm` accepts -- by the percentage +each gets wrong, losses served plus wins declined: + +| gate | served and slower | wins forfeited | total | +|---|---|---|---| +| `grid >= 48` | 2.9% | 0.0% | **2.9%** | +| `grid >= 64` | 2.9% | 0.0% | **2.9%** | +| `grid >= 80` | 0.3% | 7.0% | 7.3% | +| `grid >= 128` | 0.3% | 60.7% | 60.9% | + +**64, where an earlier revision of this said 80.** The move is the cold-baseline +fix: with SGLang measured as cold as mori, it is slower on the `hipblaslt_bf16` +rows, and mori starts winning at a smaller grid than it appeared to. 48 and 64 +score identically because no measured point falls between them, so 64 is the +boundary of the evidence rather than a fitted optimum -- it is where the bin +below it ends. + +A floor on M alone is a different matter and is wrong by an order of magnitude. +`M >= 1280`, read off `wq_b` and `wo_b`, served `wkv` at M=2048 for +100% and +`wq_a` for +38%. + +#### `shared gate_up`, the one shape only the narrow tile can express + +N=1152 is 4.5 tiles of 256 but nine of 128, so the wide tile cannot be built for +it at all and `supports_gemm` refuses the shape. Measured on the narrow tile +anyway, the answer is that it is not worth reaching: **+243% at M=64, +71% at +M=1024, +30% at M=2048, and never better than -3%** -- one point, at M=4096, +inside run-to-run noise. The GEMM's grid there is `ceildiv(M,256) * 4.5`, which +never leaves the starved region for any M a server produces. Relaxing +`supports_gemm` to accept it would buy nothing. + +`bench_gemm.py` records this as `supported: false` rather than as a failure, so +a sweep over known-good shapes still exits zero and a real break is not buried +in ten expected declines. + +The shared expert's *down* projection, `5120 x 576` at TP4, is absent from every +table above because **SGLang refuses it too** -- `native_route_supports` rejects +a K that is not a multiple of 128 -- so there is nothing to compare against and +no gate change on either side can reach it. + +The routed experts are absent throughout for a different reason: they are fp4 +(`expert_dtype` in the checkpoint's `quantization_config`), so they never take +an mxfp8 path at all. + +#### The GEMV, and what its margin actually depends on + +Against SGLang's own `mxfp8_gemv` on the **same fp8 bytes** -- no quantisation on +either side, which is the only form mori's is served in: + +| layer | N x K | M=1 | M=8 | M=32 | +|---|---|---|---|---| +| `wq_b` (TP4) | 8192 x 1280 | **-4.7%** | **-2.7%** | **-4.9%** | +| `wo_b` (TP4) | 5120 x 2048 | **-7.6%** | **-6.0%** | **-7.8%** | +| `wo_b` (TP1) | 5120 x 8192 | **-5.3%** | **-5.9%** | **-6.5%** | +| `wo_b` (TP8) | 5120 x 1024 | **-4.3%** | +0.2% | +0.7% | +| `wq_a` (TP4) | 1280 x 5120 | -0.6% | -1.8% | +1.1% | +| `wkv` (TP4) | 512 x 5120 | -0.2% | -0.1% | +1.0% | +| `wqkv_a` (TP4) | 1792 x 5120 | +1.0% | -0.5% | +2.1% | +| `shared gate_up` (TP4) | 1152 x 5120 | -0.6% | -0.8% | +1.0% | +| `wq_b` (TP1) | 32768 x 1280 | +1.9% | +4.8% | +3.4% | +| `wq_b` (TP8) | 4096 x 1280 | -1.7% | +2.7% | **+13.8%** | +| `wo_a` (TP4) | 2048 x 4096 | +5.4% | **+11.7%** | **+16.1%** | +| `wo_a` (TP8) | 1024 x 4096 | +5.8% | **+13.8%** | **+16.1%** | + +**The win is tuning, and it does not generalise.** `wq_b` and `wo_b` are the two +shapes `gemv.py`'s table was swept for, and they win at every M and every TP. +Everything else falls back to `_HEURISTIC`: six shapes land inside +-2%, which is +run-to-run noise, and **`wo_a` loses outright at both TP degrees, by 12-16% at +M=8 and above.** `wq_b` at TP8 joins it at M=32. + +`wo_a` is the interesting one because it is not a tuning miss in the ordinary +sense -- it is the only shape here with N <= 2048 *and* K >= 4096, so the +heuristic's 4-wave 16x16 tile has both few N tiles to spread over and a long K to +walk. A shape that matters should be swept (`sweep.py gemv-tune`, 72 configs, about +two minutes a bucket) rather than assumed to inherit `wq_b`'s margin. + +None of this reaches the deployed path: `wo_a` is applied through the model's own +batched absorb GEMM, not through the linear's quant method, so the hook never +sees it. + +#### Why a bf16 activation is not simply worse + +mori's GEMV takes fp8, so a bf16 caller pays a separate `mxfp8_e4m3_quantize` +launch. That pass is **1.9-2.2us on every shape**, near enough all of it launch, +since it moves at most 128 KB. It would follow that bf16 is always a loss -- +SGLang's GEMV quantises inside the kernel and pays nothing. + +It does not pay nothing. Subtracting SGLang's own fp8-in GEMV from its bf16 one +prices its in-kernel quantise, and that cost is not fixed at all (M=8): + +| layer | N x K | SGLang's in-kernel quantise | mori's separate pass | bf16 pipeline | +|---|---|---|---|---| +| `wq_b` (TP4) | 8192 x 1280 | 0.31us | 1.97us | +32% | +| `wo_b` (TP4) | 5120 x 2048 | 0.36us | 1.99us | +25% | +| `wo_b` (TP8) | 5120 x 1024 | 0.44us | 1.90us | +35% | +| `wq_b` (TP8) | 4096 x 1280 | 0.83us | 1.94us | +27% | +| `wo_a` (TP8) | 1024 x 4096 | 1.99us | 2.00us | +10% | +| `wo_a` (TP4) | 2048 x 4096 | 2.01us | 2.09us | +10% | +| `wqkv_a` (TP4) | 1792 x 5120 | 2.51us | 2.07us | **-6%** | +| `wkv` (TP4) | 512 x 5120 | 2.64us | 2.09us | **-7%** | +| `wq_a` (TP4) | 1280 x 5120 | 2.65us | 2.02us | **-9%** | +| `shared gate_up` (TP4) | 1152 x 5120 | 2.66us | 2.06us | **-8%** | +| `wq_b` (TP1) | 32768 x 1280 | 4.70us | 2.10us | **-15%** | +| `wo_b` (TP1) | 5120 x 8192 | 5.31us | 2.17us | **-21%** | + +The last column tracks the first two exactly: mori wins the bf16 pipeline at +every point where SGLang's in-kernel quantise costs more than one launch, and +loses at every point where it costs less. Nothing else is needed to explain it. + +The reason SGLang's cost varies 17x is that **its fusion quantises the same +activation once per workgroup.** Each workgroup owns a different N tile and the +same tokens, so the work is redundant across the grid and grows with both N and +K; at `wq_b` TP1 (N=32768) it is 4.7us of the 15us kernel. mori's pass is run +once whatever the grid, which is why a separate launch can beat a free fusion. + +**This is measured, not acted on.** The hook still declines every bf16 call +below 32 tokens, and that stays right for the layers it is enabled on -- `wq_b` +and `wo_b` are the four worst rows in the table. Turning it into a gate would +mean predicting SGLang's redundant quantise cost from the shape, and the +boundary is thin: `wo_a` TP4 sits at 2.01 against 2.09us, a 0.08us margin that +is inside run-to-run noise. + +### DeepSeek-V4-Pro (blockscale, TP8) + +`[M, 7168]` K=2048 on 8 ranks with `--chunked-prefill-size 16384`. +`--quant blockscale`, median of 11: + +| M | `split-sdma` | `fused-sdma` | `fused-sdma` + fp8 gather | +|---|---:|---:|---:| +| 4096 | 398.1 us | 351.0 us | **329.5 us** | +| 8192 | 722.0 | 621.1 | **539.3** | +| 16384 | 1472.9 | 1148.8 | **979.3** | + +Fusing is worth **-22%** at M=16384; the fp8 gather a further **-15%**. For +scale, the same layer as the model runs it today (a separate GEMM then an NCCL +all-reduce) measures **1419.5 us** at M=16384, and the GEMM alone is **369.4 +us**. + +**Where the time goes.** Per layer, captured in SGLang over one 20000-token +prefill -- the pipeline as the model drives it, not the standalone benchmark: + +| phase | bf16 | fp8 / sdma | fp8 / lsa | +|---|---:|---:|---:| +| gemm | 441.9 us | 438.2 us | 441.3 us | +| drain | 180.2 | 170.9 | 204.8 | +| reduce | 41.7 | 42.5 | 43.1 | +| quantize | — | 11.8 | 11.9 | +| gather | 437.8 | 256.4 | 7.5 (barrier only) | +| dequantize | — | 61.0 | — | +| pull | — | — | 230.5 | +| **wo_b layer** | **1101.6** | **980.8** | **939.1** | +| | | -11.0% | **-14.7%** | + +Two things to read out of the bf16 column. **`gather` is the bottleneck**, not +`drain`: it moves 196 MiB at 470 GB/s, which is 7 xGMI links flat out, so +halving its bytes halves its time. And **`drain`'s apparent 1140 GB/s is not a +bandwidth** -- seven links cannot do that. It is the tell that the scatter's +pushes already went out from the GEMM epilogue and the drain is only waiting for +the tail, which is why the scatter leg has far less to give than its byte count +suggests, and why it is still bf16. + +**What fp8 costs, numerically.** relL2 against an fp32 host reference goes from +**2.35e-3** (bf16 wire, bitwise exact through the collective) to **2.49e-2**. +That is a floor, not a tuning problem -- e4m3 carries 3 mantissa bits, and scale +granularity barely moves it, measured in torch on a `[2048, 7168]` standard +normal payload: + +| scale granularity | relL2 | scale bytes | +|---|---:|---:| +| per row (7168) | 2.646e-2 | 0.06% | +| per 512 | 2.631e-2 | 0.78% | +| per 256 | 2.609e-2 | 1.56% | +| per 128 | 2.572e-2 | 3.12% | +| per 32 | 2.399e-2 | 12.5% | + +**200x the scale bytes buys 9%.** Per-row is therefore the right choice, and +~2.5e-2 is what fp8 costs. + +**Whether that matters is a model-level question**, so it was measured in SGLang +on V4-Pro at TP8, with `SGLANG_DEBUG_FUSED_WO_B_AR=1` logging relL2 per layer +call during the very requests being scored. At the layer, against the unfused +path, 488 calls: + +| wire | min | median | max | +|---|---:|---:|---:| +| bf16 | 3.706e-3 | — | 4.046e-3 | +| fp8 / sdma | 1.435e-2 | **2.505e-2** | 2.654e-2 | +| fp8 / lsa | 2.153e-2 | **2.496e-2** | 2.683e-2 | + +At the model output it is not detectable. Scoring 10941 tokens of real source +text in a single prefill (mean logprob; lower is a worse model): + +| run | mean logprob | ppl | +|---|---:|---:| +| bf16 | -1.184169 | 3.2680 | +| bf16, rerun | -1.182018 | 3.2609 | +| bf16, again | -1.180106 | 3.2547 | +| **fp8** | **-1.183895** | **3.2671** | +| fp8, with debug | -1.183358 | 3.2653 | + +fp8 lands **inside the bf16 run-to-run band**. Paired per token against the same +bf16 run: + +| pair | mean d | sd d | max abs d | +|---|---:|---:|---:| +| CONTROL bf16 rerun | +0.002151 | 0.229 | 3.42 | +| CONTROL bf16 again | +0.004063 | 0.230 | 4.05 | +| **TEST fp8** | **+0.000274** | **0.219** | **2.93** | + +On every statistic, fp8 is closer to bf16 than bf16 is to itself. That band is +wide because **this model is already strongly non-deterministic**: two bf16 runs +disagree on ~48% of tokens by more than 0.01 logprob, and greedy decode diverges +within 10-20 tokens -- likely the MoE stage-2 epilogue, which accumulates with +`atomic_fadd`. So greedy token agreement is useless here; the bf16-vs-bf16 +control is as divergent as bf16-vs-fp8 (35.9% / 15.6% against 28.1% / 23.4% at +~12k / ~24k tokens). Needle-in-a-haystack at 15140 tokens is 24/24 on both +wires -- saturated, so it bounds gross damage without resolving anything finer. + +**What this does and does not say.** It says fp8 causes no gross degradation and +no measurable shift in next-token distribution on one scoring task. It does not +say quality is unaffected on long-chain reasoning, code or maths -- that needs a +task benchmark, which has not been run. Short prompts and decode never reach +this path at all (it engages only at M >= 4096), so only long-prefill workloads +are affected. + +### DeepSeek-V4.1-Flash (mxfp8, TP4) + +`[M, 5120]` K=2048 on 4 ranks. **The conclusion is not V4-Pro's.** There, fusing +is the whole story. Here the overlap is the *smallest* of three effects, and +reading the V4-Pro numbers across would set every threshold wrong. + +**What fusing is worth**, `sweep.py fused` and `sweep.py fused-fp8`, max over +ranks: -``C`` is an ordinary tensor argument, so pointing it at a cco window is a host-side -change; ``StoreC`` is untouched. That is the whole reason SDMA is the cheap +| M | wire | `gemm-only` | `split-sdma` | `fused-sdma` | gain | ceiling | +|---|---|---:|---:|---:|---:|---:| +| 4096 | bf16 | 65.6 | 443.6 | 437.3 | +1.4% | 14.8% | +| 8192 | bf16 | 105.0 | 827.3 | 801.8 | +3.1% | 12.7% | +| 16384 | bf16 | 180.2 | 1593.5 | 1500.7 | +5.8% | 11.3% | +| 4096 | fp8 | 65.8 | 364.8 | 358.8 | +1.7% | 18.0% | +| 8192 | fp8 | 102.9 | 669.7 | 643.1 | +4.0% | 15.4% | +| 16384 | fp8 | 176.3 | 1309.2 | 1195.1 | +8.7% | 13.5% | + +The fp8 rows at 4096 and 8192 are new; the rest reproduce an earlier run of the +same matrix to within 2.5% on every cell, with the gains and ceilings identical +to a tenth of a point. + +**The fp8 wire and the fusion are close to independent, and the wire is the +bigger of the two.** At M=16384 it is worth -17.8% on the split path on its own, +before anything is fused. It then *raises* what fusing is worth -- +5.8% to ++8.7% -- because shortening the collective makes the GEMM a larger share of the +layer, which is the same effect the ceiling column reports (11.3% -> 13.5%). +Together: **1195.1us against 1593.5 for split over the bf16 wire, -25.0%.** + +The **ceiling** column is `gemm/split`: what fusing would be worth if it hid the +GEMM entirely. It is 11-18%, because the GEMM is ~180us against 1130-1410us of +communication -- between 1:6 and 1:8. Fusing collects about half of the ceiling. +Making the GEMM faster moved this number the wrong way, which is why V4-Pro's ++21% does not transfer. + +**`fused-lsa` is negative everywhere**, on both TP splits, every M and both +wires: re-measured at TP4 it is 5.3-8.0% worse than `split-sdma` across the six +cells above. `split-lsa` is fine (427.0 against sdma's 443.6 at M=4096), so it +is the fusing that hurts, not the transport: direct-LSA moves the bytes with the +GEMM's own waves, and when the GEMM is 11% of the layer, spending its issue +slots on the transfer is a bad trade. Use `fused-sdma`. `chunks` wants more of +them as M grows (M=16384: c=1 1634.4 -> c=8 1503.6, -8%); lsa is flat and +marginally prefers c=1, and the existing default picks the winner at every point +measured. + +**`split-lsa` has no fp8 leg, and asking for one used to be silent.** +`build_lsa_ar` takes no `gather_dtype` at all -- the LSA 2-stage all-reduce moves +bf16 and that is the whole of it. `--mode split-lsa --gather-dtype fp8` therefore +ran a bf16 collective and reported it under the fp8 label: same time as the bf16 +wire to within 0.2%, and a relL2 at the bf16 quantisation floor rather than +fp8's. The validation gate is two-sided precisely for this -- its lower bound of +5e-3 exists to assert the fp8 wire was *taken*, not just requested -- so it +caught it, but only after paying for the run. `bench_gemm_ar.py` now refuses the +combination where it is asked for. + +**Picking the fp8 gather's transport.** The fp8 leg is worth more than the +fusing here, so its own knobs matter. Three are real; two that look like knobs +are not -- `--gather-bands` and `--fuse-reduce-push` are read only by +`sdma_reduce_push`, which rejects `fp8_gather` outright, and +`scatter_dtype='fp8'` raises `NotImplementedError` (its regions are sized but +nothing writes the payload; that is the other leg and the larger remaining +lever). TP4, M=16384, all 24 runs validated and all at relL2 2.31e-02 -- these +knobs move time, not numbers: + +| | gather=sdma, fq off | fq on | **gather=lsa, fq off** | fq on | +|---|---:|---:|---:|---:| +| `fused-sdma` | 1222.5 | 1244.1 | **1192.1** | 1209.9 | +| `fused-lsa` | 1412.1 | 1433.0 | 1375.6 | 1398.7 | +| `split-sdma` | 1342.1 | 1358.9 | 1310.3 | 1326.8 | + +Every cell says the same two things, at M=4096 as well. **Pull the gather**: +2.4-2.6% faster than pushing at M=16384, 3.6% at M=4096, in all three modes. +**Do not fuse the quantise**: 1.1-2.3% everywhere, +17.8us on the winning +combination. Note this is the opposite of the scatter leg, where `fused-lsa` +loses -- not a contradiction: the scatter runs *inside* the GEMM and competes +with the MFMA for issue slots, the gather runs after it. + +Best: `fused-sdma` + `gather_dtype=fp8` + `gather_transport=lsa` + +`--no-fuse-quantize`. **1192.1us against 1592.5 for split/bf16, -25.1%**, of +which the fp8 wire alone is -15.6% (it needs no fusing) and fusing on top of it +a further -8.8%. + +**Against what the model runs today.** The numbers above compare mori against +mori. The threshold question needs the path being replaced: +`mxfp8_native_blockscaled_linear` followed by an all-reduce, where +`native_route_plan` picks `dot_scaled` below M=8192 and `hipblaslt_bf16` above +it. One M per process, two runs: + +| M | m_pad | fill | today | bf16 wire | fp8 wire | +|---|---:|---:|---:|---:|---:| +| 1024 | 1024 | 1.000 | 164.6 us | +16.4% | +10.3% | +| 2048 | 2048 | 1.000 | 283.3 | +1.8% | -8.8% | +| 4096 | 4096 | 1.000 | 482.8 | -0.9% | -15.4% | +| 4200 | 5120 | 0.820 | 506.1 | +17.0% | -2.3% | +| 4700 | 5120 | 0.918 | 548.3 | +8.4% | -9.3% | +| 5120 | 5120 | 1.000 | 601.5 | -5.8% | -20.7% | +| 7200 | 8192 | 0.879 | 1051.2 | -17.6% | -32.7% | +| 8192 | 8192 | 1.000 | 1042.5 | -19.4% | -34.1% | +| 8200 | 9216 | 0.890 | 1185.1 | -20.8% | -35.7% | +| 9200 | 9216 | 0.998 | 1157.5 | -19.2% | -34.2% | +| 13000 | 13312 | 0.977 | 1547.8 | -8.5% | -25.3% | +| 16384 | 16384 | 1.000 | 1851.3 | -16.2% | -32.9% | + +Two things this settles that a coarser sweep would not. **The baseline is not +monotonic** -- it retunes per M bucket and swings ~20% between them (M=8200 +costs 943us where M=8800 costs 1133), so a threshold has to be read off the +whole curve, not interpolated between two points. And **one `m_pad` is not one +answer**: the fused cost is fixed by the padded size while the baseline follows +the true M, so m_pad 5120 wins at fill 1.000 and loses at 0.918 and 0.820. No +fill threshold separates those from the *winning* low-fill points above (0.879 +at m_pad 8192 wins -17.6%), which is why the floors in SGLang's +`mori_gemm_ar.py` are 8192 (bf16) and 2048 (fp8) rather than a single number +plus a fill guard. + +> These numbers are only meaningful because `_PinnedLaunch` exists -- see +> [Measurement traps](#measurement-traps). + +### End to end, in SGLang + +Three servers, one per wire, GSM8K then `bench_one_batch_server`, two full runs. +Prefill throughput, tok/s, median, on V4.1-Flash: + +| bs | base | bf16 wire | | fp8 wire | | +|---|---:|---:|---:|---:|---:| +| 1 | 27349 | 27632 | +1.0% | 28173 | +3.0% | +| 4 | 34101 | 34990 | +2.6% | 36029 | **+5.7%** | +| 8 | 34948 | 35790 | +2.4% | 36733 | **+5.1%** | +| 16 | 35151 | 35969 | +2.3% | 36930 | **+5.1%** | + +**bs=1 is a built-in control**: M=4096 is below both floors, so no variant fuses +there and all three should agree -- they do, within 3%. That is what makes the +bs>=4 figures a signal rather than drift. + +GSM8K passes on all three (base .915/.916, bf16 .914/.917, fp8 .921/.924) and +the three sit inside each other's noise, so the fp8 leg's relL2 2.3e-2 is below +what 1319 questions resolve. That is not what the gate is for: it catches a +collective that moved nothing, which scores near zero while still answering +fluently. + +The decode column is not reported as evidence. `wo_b` never fuses at decode's M, +so the three variants should be identical, and they scatter -8.6% to +14.1% -- +that is DSPARK's acceptance rate, not this change. + +On V4-Pro, the same comparison as GPU busy time over one profiled prefill: + +| | GPU busy | wall | vs unfused | +|---|---:|---:|---:| +| unfused (GEMM + NCCL) | 1096.0 ms | 1.1969 s | — | +| fused, bf16 wire | 1070.3 | 1.1728 | -2.3% | +| fused, fp8 / sdma | 1052.1 | 1.1455 | -4.0% | +| **fused, fp8 / lsa** | **1041.2** | **1.1385** | **-5.0%** | + +An earlier capture of the same four read 1101.9 / 1077.2 / 1050.4 / 1048.2, so +this reproduces to about half a percent. The layer-level win is larger than the +end-to-end one because `wo_b` is about 12% of the profile. + +### Measurement traps + +Each of these produced a confident wrong number before being caught, and none is +about the kernel. + +**A single-call CUDA-graph capture has a 13.40us floor on this box.** The same +do-nothing kernel amortised over 200 calls in one graph is 1.51us, so anything +under about 15us measured that way is mostly harness. It is invisible at +M=16384, which is why it survived several rounds. Two committed thresholds had +to be re-derived because of it, and one of them was not a win at all. Use +`benchmark/cco/flydsl/gemm_ar/timing.py`. + +**A hot loop reads the weight out of LLC.** MI355X has 256 MB of it and a +`wq_b` weight is 10.5 MB: 4308 GB/s hot against 2561 GB/s with copies rotated +past the cache, a 1.7x overstatement. A forward pass reads each layer's weight +once, so **cold is the number that predicts a server**; report both. + +**And rotating the copies is not enough on its own -- the graph has to be long +enough to reach them.** `cold_hot_us` sized its ring at 384 MB and then captured +`reps=32` calls, but the ring advances at *capture* time, so the graph bakes in +32 pointers and every replay revisits those same ones. The working set was +`min(reps, n)` copies, not `n`, which for seven of these twelve shapes is under +the LLC: `wkv`'s 2.5 MB weight gave 80 MB. Those columns were labelled cold and +were not. + +The shapes that landed *on* 256 MB read worst of all, which is the tell. A +working set at exactly cache capacity thrashes, where one comfortably over it +just streams -- so `wo_a` (8 MB, 256 MB at 32 reps) measured +10% against +SGLang cold and -2% hot, and the two disagreeing in sign is what exposed this. +Fixed by taking `cold_reps = max(reps, n)`. The correction is not cosmetic: on +the GEMV it turned four apparent 9-14% wins into ties and made two losses +larger. Shapes whose weight already exceeded 8 MB -- `wq_b`, `wo_b`, both TP1 +shapes -- moved by less than a point, which is how the diagnosis was confirmed. + +**An eager layer benchmark is biased against the fused path.** It launches four +FlyDSL kernels per call where the baseline launches two Triton ones, and before +`_PinnedLaunch` that was 181.6us of per-dispatch Python on the critical path +(`JitFunction.__call__` re-derived its cache key every launch -- an `inspect` +bind, a 35-global snapshot, a drift check -- 85.9us for the GEMM and 31.9us per +phase against 5.9us of actual `hipModuleLaunchKernel`). The server does not pay +it, since its prefill replays a CUDA graph, so the harness read the bf16 wire at ++11.0% where the server read -2.3%, and a threshold was set from the former. The +tell was arithmetic: 109us of "launch overhead" is impossible when a launch is +2-5us, and checking that rather than accepting it is what found it. + +**Sweeping several M in one process contaminated one point.** M=4200 and M=4700 +pad to the same 5120 and must therefore cost the same; in a multi-M process they +read 697us and 1083us. Isolated, both are 697. The mechanism was never +identified, which is the point -- one M per process is cheap insurance. + +**A silent non-fusing path looks exactly like a slow one.** `layer.weight` is +rebound to its shuffled form `[N/16, K/128, 2048]` after +`prepare_mxfp8_native_weight`, so `shape[0]` is N/16; reading it as N turned +5120 into 320, `supports` rejected it for not being a multiple of `BLOCK_N`, and +the path fell back **without a word**. The server stayed correct and the profile +stayed plausible. Only a harness that asserts fusing *happened* catches this, +which is why both the correctness check and the end-to-end test now do. + +**A threshold fitted on a coarse M grid is fitted to the gap.** `NARROW_N_BELOW_M` +was a bare `M < 2048`, measured on a grid that jumped 1024 -> 2048 and so never +looked between them. It was wrong on both shapes; at `wq_b` M=1280 correcting it +was +19.9% -> -12.8%. + +**With `BUILD_CCO_SDMA=OFF` every put silently does nothing.** The all-reduce +returns mostly the local slice, the model still answers fluently, every mori +kernel still appears in the profile, and the fused path measures **faster** than +it is because it is not moving data -- -6.8% instead of -2.3%, with fp8/lsa +appearing *worst* of the three rather than best, since the pull is the one leg +that does not go through SDMA. Perplexity catches it and nothing cheaper does: +862511 against 3.26 on the same text. A short prompt cannot catch it either, +because fusing needs M >= 4096. + +## How it works + +The fused path, from the epilogue outwards. Everything here is about +`GemmAllReduceOp` -- the two standalone ops compile the same GEMM with this +tail switched off, so only the mainloop and the C store apply to them. + +### Why the SDMA transport needs no epilogue change + +`C` is an ordinary tensor argument, so pointing it at a cco window is a host-side +change; `StoreC` is untouched. That is the whole reason SDMA is the cheap transport to fuse. The LSA alternative -- the epilogue storing straight into a -peer -- is implemented as ``fused-lsa``, and it is the slower of the two for the -reason that was predicted: ``StoreC._store_bf16`` writes one bf16 at a time -through ``BufferCopy16b``, and gcnasm measured that lane-scatter at 0.26x when +peer -- is implemented as `fused-lsa`, and it is the slower of the two for the +reason that was predicted: `StoreC._store_bf16` writes one bf16 at a time +through `BufferCopy16b`, and gcnasm measured that lane-scatter at 0.26x when the destination is a peer. It spends ~560us of its GEMM pushing C over xGMI where the copy engines move the same bytes in ~499. -## Completion protocol +### Completion protocol -One monotonic counter per (destination, chunk) in the window (``cfg.counter_off``). -Every block, after its four ``store_c.store`` calls: +One monotonic counter per (destination, chunk) in the window (`cfg.counter_off`). +Every block, after its four `store_c.store` calls: s_waitcnt vmcnt(0) ; s_barrier -- this block's C tile has retired __threadfence_system() -- ...and is visible to the copy engine @@ -257,9 +951,9 @@ Every block, after its four ``store_c.store`` calls: sdma.put(dest, ...) -- fire and forget, no quiet Counters are never reset, so the modulo test works on every launch and the kernel -stays CUDA-graph-safe, exactly like the barrier flags in ``ar.kernels_lsa``. +stays CUDA-graph-safe, exactly like the barrier flags in `ar.kernels_lsa`. -``quiet`` is deliberately *not* called here: gcnasm measured fusing it into the +`quiet` is deliberately *not* called here: gcnasm measured fusing it into the GEMM as a 1.8-10.7% regression, so the drain kernel does it afterwards. No per-destination spin lock either -- gcnasm needed one because several CTAs could submit for one destination, whereas the modulo test elects exactly one. @@ -268,14 +962,95 @@ Tiles are walked (chunk, destination, n) with the destination rotated by rank, rather than in aiter's linear order. Linear order finishes destination 0's whole slice, then 1's, and so on, so the last destination's link only starts at the end of the GEMM and nothing overlaps. The rotation is gcnasm's -``opus_direct_stripe_tile`` idea. It costs nothing: 49.72us against 49.80us for +`opus_direct_stripe_tile` idea. It costs nothing: 49.72us against 49.80us for the GEMM alone. -## C store: three stages off gcnasm, 6.7% off the GEMM, bit-exact +### The chunks race: it was the GEMM, and it is gone + +`--chunks` > 1 produced wrong output about 3 runs in 10 and was pinned to 1 +for that reason. The cause was not the chunk protocol at all: it was aiter's +8-wave GEMM under-counting one `s_waitcnt` in its main loop, which corrupted +output non-deterministically on large grids whatever the epilogue did (see the +comment at that `wait_barrier` below). Since that fix, at [16384, 7168] +K=2048 on 8 ranks: + +* chunks 1/2/4/8, 10 runs each -- 40/40 correct, no hang; +* chunks=8 alone, 25 more runs -- 25/25 correct, no hang. + +Two hangs were seen at 8 ranks while the sweep was still being set up and never +reproduced in the 100+ runs after. The submit lock is a plain test-and-set spin +(`_acquire_peer_lock`), so a hang is not impossible; use `timeout` when +sweeping and treat one as a finding rather than a flake. -``StoreC`` emits 128 ``buffer_store_short`` per block -- one bf16 at a time -- -because of the MFMA accumulator layout, not a missed vectorization. Lane ``l`` -holds ``D[4*(l/16)+i][l%16]``: four consecutive *rows*, stride ``c_cols``, so its +Two things in this file survived that misdiagnosis and are worth keeping +straight: + +* The half-wave barrier pairing (`if wave_m == 0: s_barrier()` before + `store_c`) is still needed and still right -- gcnasm does it + unconditionally at kernel_template.hpp:693. It took `--chunks 2` from + 1-in-3 failing to 3-in-10, which at the time read as "better but not fixed"; + the residual 3-in-10 was the GEMM. +* Two *other* diagnoses were wrong and are recorded so they are not retried. + **A shared SDMA queue**: gcnasm's per-destination submit lock was ported and + ISA-verified; it fixed nothing, and is kept only because cco's + one-issuing-warp-per-queue rule still applies. **The counter atomic's + ordering**: `acq_rel` appeared to beat `monotonic`, which is how the + acquire half got justified; against a 1-in-3 intermittent failure that + comparison was noise. `acq_rel` stays because release/acquire is right for + a producer handing tiles to a consumer, not because it was measured -- and at + chunks=1 it demonstrably orders nothing, since 20 runs of `monotonic` pass + too. + +Repeated runs are still the only way to judge any of this, and the gate has to +be tight: the corruption landed at 4-9e-3 against an fp8 floor of 2.35e-3, so a +5e-3 threshold reported a corrupt run as validated (it did, at 3.97e-3). The +bench gates at 3e-3, and `test_fused_is_stable_across_repeats` requires three +runs to be *identical* rather than each small. + +### Mode comparison, once the chunks are unblocked + +> These were taken in a different round from the +> [V4-Pro tables above](#deepseek-v4-pro-blockscale-tp8) -- before the fp8 wire +> existed -- and the absolute numbers do not line up with them (split-sdma reads +> 1261.7 here against 1472.9 there). What transfers is the *ordering* of the +> modes and the per-chunk breakdown, which is what this section is for. + +8 ranks, [16384, 7168] K=2048 -- the real prefill shape -- graph replay, all +three C-store stages on (now the benchmark default), median of 31, max over +ranks: + + mode time + fused-sdma chunks=8 1114.1us <- best + split-sdma 1261.7 + split-lsa 1262.4 + fused-lsa 1571.4 + gemm-only 228.9 + +Per-kernel, the overlap is visible directly: + + chunks GEMM drain reduce gather total + 1 257.5 493.8 44.0 496.4 1291.7 + 2 266.5 379.4 44.2 496.4 1186.6 + 8 265.0 303.7 43.8 496.6 1109.1 + +The drain falls 38% while the GEMM grows 7.5us for the extra counter atomics +and the lock. Going finer is worse: at chunks=8 each PUT is 3.5 MiB, and 16 +(via `--block-m 128`) halves that to 1.75 MiB, under the knee in the SDMA +bandwidth curve, for 1130.8us. + +fused-lsa still loses, and for a reason that is not going away: it spends 730us +of its GEMM pushing C over xGMI where the copy engines move the same bytes in +499, and ATT shows 99% of that store time is *stall*, so coalescing the stores +(the three C-store stages) buys it nothing -- 954.5 -> 952.1us. + +The remaining floor is the tail: reduce plus all-gather is 540us of the 1114, +and neither is touched by fusing the GEMM. + +### The C store: three stages off gcnasm + +`StoreC` emits 128 `buffer_store_short` per block -- one bf16 at a time -- +because of the MFMA accumulator layout, not a missed vectorization. Lane `l` +holds `D[4*(l/16)+i][l%16]`: four consecutive *rows*, stride `c_cols`, so its four values are 14336 bytes apart in a row-major C, and the eight bf16 that would make a 16-byte store live in eight different lanes. @@ -286,63 +1061,63 @@ make a 16-byte store live in eight different lanes. --swap-ab --permlane --lane-transpose 35.74us 16 dwordx4 3,852 Median of 42 dispatches, 8 ranks, [4096,7168] K=1024. Bit-identical to aiter at -[512,512,256], [1024,768,512] and wo_b (``test_swap_ab_is_bitwise_identical``), +[512,512,256], [1024,768,512] and wo_b (`test_swap_ab_is_bitwise_identical`), and registers are unchanged throughout (VGPR 128 / SGPR 112 / 0 scratch). -**1. ``--swap-ab``** -- ``mfma_adaptor_swap_ab`` (opus.hpp:2064, literally -``base::operator()(b, a, c)`` with ``dim_c()`` redefined). Computing -``B^T A^T = (A B)^T`` in the accumulator's own layout moves a lane to -``D[l%16][4*(l/16)+k]``: four consecutive *columns*, 8 contiguous bytes. Alone it +**1. `--swap-ab`** -- `mfma_adaptor_swap_ab` (opus.hpp:2064, literally +`base::operator()(b, a, c)` with `dim_c()` redefined). Computing +`B^T A^T = (A B)^T` in the accumulator's own layout moves a lane to +`D[l%16][4*(l/16)+k]`: four consecutive *columns*, 8 contiguous bytes. Alone it is *slower* -- a row still only gets 32 bytes, so the same 16 transactions are squeezed onto a quarter as many instructions and per-instruction address fan-out quadruples (1022 cycles each against 114). -**2. ``--permlane``** -- two ``v_permlane16_swap_b32`` per M-tile. Semantics -measured rather than assumed: ``vdst' = [X.r0, Y.r0, X.r2, Y.r2]``, -``vsrc' = [X.r1, Y.r1, X.r3, Y.r3]``, so:: +**2. `--permlane`** -- two `v_permlane16_swap_b32` per M-tile. Semantics +measured rather than assumed: `vdst' = [X.r0, Y.r0, X.r2, Y.r2]`, +`vsrc' = [X.r1, Y.r1, X.r3, Y.r3]`, so:: (A, B) = permlane16_swap(tile0.d0, tile1.d0) (C, D) = permlane16_swap(tile0.d1, tile1.d1) lane group g stores (A, C, B, D) lands g = 0,1,2,3 on columns 0-7, 16-23, 8-15, 24-31 -- together columns 0..31 -contiguously, 16 bytes per lane and 64 per row. No ``ds_bpermute``, no LDS; the +contiguously, 16 bytes per lane and 64 per row. No `ds_bpermute`, no LDS; the column permutation is absorbed into the address. This does **not** speed the store up (14,608 -> 14,604). What it pays for is everything a 2-byte store drags along: 128 stores need 128 addresses, 128 bounds predicates and 128 scalar multiplies, 16 need 16, and the swapped layout makes -B's scale a vec4 so the scaling packs into ``v_pk_mul_f32``:: +B's scale a vec4 so the scaling packs into `v_pk_mul_f32`:: v_mul_f32_e32 257 -> 1 v_lshlrev_b32_e32 159 -> 45 v_pk_mul_f32 0 -> 128 v_add_u32_e32 133 -> 21 v_cvt_pk_bf16_f32 128 -> 64 v_cndmask_b32_e64 96 -> 8 -**3. ``--lane-transpose``** -- gcnasm's second stage -(kernel_template.hpp:491-510), one ``ds_bpermute`` per dword. After the permlane +**3. `--lane-transpose`** -- gcnasm's second stage +(kernel_template.hpp:491-510), one `ds_bpermute` per dword. After the permlane stage a row's four 8-column chunks sit in lanes 16 apart, so a 16-lane group touches 16 rows at 16 bytes each. Transposing the lane index -- lane -``l' = 4r'+q'`` pulls from the lane holding ``(row r', chunk q')``, with -``g = q``'s two bits swapped, 0,1,2,3 -> 0,2,1,3 -- puts adjacent lanes on one +`l' = 4r'+q'` pulls from the lane holding `(row r', chunk q')`, with +`g = q`'s two bits swapped, 0,1,2,3 -> 0,2,1,3 -- puts adjacent lanes on one row, so lanes 0-3 write 64 contiguous bytes and a group covers 4 rows. Same instructions, same 64 addresses, only which lane holds which. The store's own latency falls **14,604 -> 3,852**, which is the coalescer being sensitive to lane adjacency and not just to the address set -- exactly what gcnasm's -"pair-coalesced" comment is about. The 64 ``ds_bpermute`` cost 9,604 cycles. +"pair-coalesced" comment is about. The 64 `ds_bpermute` cost 9,604 cycles. -### ``--hoist-scales``: a real redundancy that does not pay to remove +#### `--hoist-scales`: a real redundancy that does not pay to remove -The epilogue calls ``store_c.store`` four times, and those calls share base rows +The epilogue calls `store_c.store` four times, and those calls share base rows pairwise and base columns pairwise, so every scale is fetched twice. The -source-attributed thread trace counts it exactly: 16 ``buffer_load_dwordx4`` at -``gemm_a8w8_8wave.py:191`` where only 8 addresses are distinct, and 8 -``buffer_load_dword`` at :199 where only 4 are -- 12 of 24 loads redundant, plus +source-attributed thread trace counts it exactly: 16 `buffer_load_dwordx4` at +`gemm_a8w8_8wave.py:191` where only 8 addresses are distinct, and 8 +`buffer_load_dword` at :199 where only 4 are -- 12 of 24 loads redundant, plus 12 redundant address computations, 4,276 cycles or 0.8% of the kernel. The -compiler cannot merge them because all four calls write the same ``reg_f32_*`` +compiler cannot merge them because all four calls write the same `reg_f32_*` register buffer, which makes them a chain of overwrites rather than pure loads. -``store_all`` loads each scale once. It removes exactly the predicted loads -- +`store_all` loads each scale once. It removes exactly the predicted loads -- A-scale 16 -> 8, B-scale 8 -> 4 -- and is slightly **slower**: variant VGPR scratch instrs dwordx4 dword gemm @@ -354,11 +1129,11 @@ moves and re-materialised addresses than the twelve loads were worth, and it spends the last 2 VGPRs of headroom (254 -> 256). Bit-exact either way. Kept as a switch rather than deleted: the redundant fraction grows with -``N_TILES_B``, so a larger ``BLOCK_N`` would change the arithmetic, and if +`N_TILES_B`, so a larger `BLOCK_N` would change the arithmetic, and if register pressure ever loosens this flips sign. It is also cheaper to re-measure a flag than to re-derive why it was rejected. -Two limits. **It is worth nothing on ``fused-lsa``**, where the store goes to a +Two limits. **It is worth nothing on `fused-lsa`**, where the store goes to a peer over xGMI rather than to local memory:: gemm barrier reduce gather sum @@ -376,16 +1151,69 @@ argument-validation above is there so a missing branch fails loudly instead of silently running a weaker variant.) And 64 bytes is one wave's ceiling here, since a wave owns -``N_TILES_B * 16 = 32`` columns; gcnasm reaches a full 128-byte line only by -having four ``wave_id_n`` waves tile adjacent 16-column runs. +`N_TILES_B * 16 = 32` columns; gcnasm reaches a full 128-byte line only by +having four `wave_id_n` waves tile adjacent 16-column runs. -All three stages are candidates to push back into aiter's ``StoreC``: bit-exact, +All three stages are candidates to push back into aiter's `StoreC`: bit-exact, no register cost, and the 6.7% is on the GEMM itself, independent of any all-reduce. -## Persistent tiles: tried, and there is no register budget for it +### What gcnasm does differently + +`/workspace/gcnasm/opus_gemm_dist/opus_gemm_a2a_lsa` fuses a GEMM with an +all-to-all and gets 16-27% out of it. Reading it explains most of why this does +not, and one of its lessons was worth ~90us here. + +1. **Its collective is one phase; this one is three.** An a2a scatters the GEMM + output once. An all-reduce is scatter + reduce + all-gather, and fusion only + touches the scatter -- reduce + gather is 150us of the 327us baseline, 46%, + untouchable by construction. +2. **Its ratio is inverted.** M=2048 N=18432 K=8192 gives ~518us of GEMM against + ~200us of comm, 2.6:1. wo_b is 39us against 137us per phase, 0.29:1. Overlap + can hide at most the smaller of the two, so theirs hides most of the comm and + this hides at most one GEMM. +3. **Its best mode has no producer/consumer handoff at all.** "Direct LSA" has + the GEMM epilogue store *straight into the destination rank's buffer*. No + staging, no copy engine, nothing to publish mid-kernel -- the only sync is the + barrier at the end. The publication problem this file spends all its time on + simply does not exist there. +4. **Its fused-SDMA path uses no cache fence.** `opus_chunk_sdma_submit` is + `s_waitcnt vmcnt(0)`, `s_barrier`, an `__ATOMIC_ACQ_REL` counter, and a + per-destination submit spin lock -- no `__threadfence_system`, no + `buffer_wbl2`. That is the lesson that transferred: the acq_rel counter *is* + the release, and the explicit fence added here was 90us of pure waste + (fused 440us -> 350us on removing it). Its ISA shows the atomic already emits + its own `buffer_wbl2`/`buffer_inv` pair, on thread 0 only. + +## Negative results + +Kept because each reads as obviously right and the reason it is not cannot be +seen from the source. + +### A CK-shaped 4-wave GEMM + +An earlier revision carried `kernels_preshuffle4w.py`, a port of CK's 4-wave +B-out-of-LDS shape, chasing a 22% gap between this GEMM and CK's at the same +shape. It reached CK's instruction mix and not its speed, and it was deleted +rather than carried (it needed ~1000 lines of further aiter vendoring to serve +a kernel nothing calls). What the investigation ruled out, since the same +ground should not be walked twice: + +Ten hypotheses were falsified by measurement -- promote arithmetic, register +spill, store width, VALU scheduling groups, MFMA batch size, scale loads, +address hoisting, load-to-use distance, occupancy, and dependency structure. +Hardware counters (`rocprofv3 --pmc`) show *identical* `SQ_INSTS_MFMA` +(7,340,032) and `SQ_VALU_MFMA_BUSY_CYCLES` (234,881,024), VALU within 1%, and +`MemUnitStalled` at approximately zero -- but `SQ_WAIT_ANY` at 156.0M against +CK's 120.4M. A K-sweep puts the whole difference per-iteration: our fixed cost +is *lower* (51.6us against 66.0us), while each K-block costs 20.9us against +14.5us. The gap is wait, not work, and it is not in any of the places listed +above. See commits `cc696762`, `54fef960`, `9f990637`, `1dead3fc` for the +traces. -gcnasm builds this kernel family both ways (``PERSISTENT=1|0``) and its README +### Persistent tiles + +gcnasm builds this kernel family both ways (`PERSISTENT=1|0`) and its README has a tail-balance sweep: the win is entirely a function of the remainder after whole 256-CTA batches -- +9.97% at remainder 8, +8.41% at 32, and only +0.77% at 192. wo_b is 16 x 28 = 448 tiles on 256 CUs, i.e. remainder **192**, the benign @@ -416,38 +1244,38 @@ than anything to do with scheduling. From the kernel metadata in the final ISA: pipeline in a runtime loop asks for 38 more -- the K pipeline's accumulators and operand fragments die at the end of a tile in the flat version, but a loop makes the compiler assume they may be live across the back-edge -- and there is nowhere -to put them. ``--amdgpu-num-vgpr 256/512`` changes nothing, because the cap was +to put them. `--amdgpu-num-vgpr 256/512` changes nothing, because the cap was never the constraint. Persistent and non-persistent spill identically, which -confirms the cost is ``scf.for`` itself and not the scheduling idea. +confirms the cost is `scf.for` itself and not the scheduling idea. -Note ``fly-promote-regmem-to-vectorssa`` is **not** the problem, contrary to what -an earlier version of this note claimed. It handles ``scf::ForOp``, it promoted -all 451 register allocas here (zero left afterwards, zero ``llvm.alloca`` in the -final IR), and the emitted loop carries no ``iter_args`` at all. The 156 bytes -are ordinary register-allocator spill: ``.vgpr_spill_count: 38``, 38 -``scratch_store_dword`` / 38 ``scratch_load_dword``. +Note `fly-promote-regmem-to-vectorssa` is **not** the problem, contrary to what +an earlier version of this note claimed. It handles `scf::ForOp`, it promoted +all 451 register allocas here (zero left afterwards, zero `llvm.alloca` in the +final IR), and the emitted loop carries no `iter_args` at all. The 156 bytes +are ordinary register-allocator spill: `.vgpr_spill_count: 38`, 38 +`scratch_store_dword` / 38 `scratch_load_dword`. The pass does have one real inefficiency, just not one we hit: an alloca declared -*outside* a loop is carried as an ``iter_arg`` unconditionally, even when it is -fully overwritten every iteration, because ``collectTouchedRegAllocaInRegion`` +*outside* a loop is carried as an `iter_arg` unconditionally, even when it is +fully overwritten every iteration, because `collectTouchedRegAllocaInRegion` records on any load *or* store with no liveness test. Allocas declared *inside* -the loop are correctly materialised as ``ub.poison`` and not carried, and ours +the loop are correctly materialised as `ub.poison` and not carried, and ours are all inside. Two bugs the attempt surfaced, worth knowing if anyone loops this pipeline: * **The half-wave barrier does not survive a loop.** The prologue's - ``if wave_m == 1: rocdl.s_barrier()`` deliberately runs the two half-waves one + `if wave_m == 1: rocdl.s_barrier()` deliberately runs the two half-waves one barrier out of phase. Fine once; in a loop the offset accumulates by one per tile, so from the second tile on the halves rendezvous at mismatched program points and the LDS double-buffering races -- silently, as partial corruption - inside otherwise-correct tiles. A compensating ``if wave_m == 0: s_barrier()`` + inside otherwise-correct tiles. A compensating `if wave_m == 0: s_barrier()` at the end of each tile fixes it exactly. * **The LDS handles must be rebound per tile.** The pipeline swaps those Python bindings as it advances, so hoisting them makes eight shared-address-space pointers loop-carried, which fails to legalize. -And a FlyDSL gotcha: ``range(...)`` must appear literally in the ``for`` +And a FlyDSL gotcha: `range(...)` must appear literally in the `for` statement or the AST rewriter does not see it and Python evaluates it eagerly ("dynamic 'ArithValue' has no Python integer representation"). Assigning it to a variable first does not work, so a kernel cannot cheaply offer both a looped and @@ -460,7 +1288,7 @@ and it is nowhere near enough. Halving BLOCK_M would free roughly 16 by halving the accumulator count, but moves the tile count to 896 -- remainder 128, where gcnasm measured +0.41%. There is no version of this that pays. -## Direct LSA (``--mode fused-lsa``): correct now, still slower than split +### Direct LSA (`--mode fused-lsa`) gcnasm's best mode has the GEMM epilogue store *straight into the destination rank's window*: no staging buffer, no copy engine, nothing to publish mid-kernel. @@ -472,16 +1300,16 @@ gcnasm measured for a lane-scatter pushed to a peer. It was wrong for a long time, and the fix is one line in the right place. The LSA 2-stage all-reduce publishes from the kernel that produced the data -- -``ar/kernels_lsa.py`` fences right after its ``tmp`` stores, in every block. +`ar/kernels_lsa.py` fences right after its `tmp` stores, in every block. Direct LSA's producer is the GEMM, and the only fence was in the *separate* barrier kernel: one block, therefore one XCD's L2 out of eight. The other seven kept the peer-homed lines dirty, and the LSA flag, being a system-scope atomic, overtook them. Moving the fence into the GEMM fixes it: **10/10 runs bit-correct** where it had been 2-3 in 6. -``--direct-fence leader`` (thread 0 only, the default) rather than every lane is +`--direct-fence leader` (thread 0 only, the default) rather than every lane is worth 80us, 434 -> 355. It is legal only because the half-wave barrier pair is -now closed -- ``wait_barrier(0)`` really does mean every wave's stores have +now closed -- `wait_barrier(0)` really does mean every wave's stores have retired into this CU's L2, so one wave writing it back covers all eight. Measured as incorrect before that fix, which is what made a per-wave release look mandatory. @@ -491,7 +1319,7 @@ Two things that did *not* work, and are worth not re-trying: * **Uncached (sc0|sc1) peer stores**, so there would be nothing to publish. Wrong consistently, ~1.7e-2, at every store width. The first time this was tried the store was one bf16 and the explanation looked like partial-line writes losing - updates over the fabric; with ``--permlane --lane-transpose`` making it 16 + updates over the fabric; with `--permlane --lane-transpose` making it 16 bytes per lane and 64 contiguous per row it is *still* wrong, so that explanation was not it and the real one is unknown. * An uncached recv load in the reduce. No effect, which is what rules out a stale @@ -511,104 +1339,129 @@ peer and then pays a separate 10.9us local reduce pass, its store rate is ~14% under LSA's read rate, and it adds the publishing fence LSA gets from being local. Absorbing the scatter completely still does not cover that. -## What gcnasm does differently, and why its GEMM+a2a wins - -``/workspace/gcnasm/opus_gemm_dist/opus_gemm_a2a_lsa`` fuses a GEMM with an -all-to-all and gets 16-27% out of it. Reading it explains most of why this does -not, and one of its lessons was worth ~90us here. - -1. **Its collective is one phase; this one is three.** An a2a scatters the GEMM - output once. An all-reduce is scatter + reduce + all-gather, and fusion only - touches the scatter -- reduce + gather is 150us of the 327us baseline, 46%, - untouchable by construction. -2. **Its ratio is inverted.** M=2048 N=18432 K=8192 gives ~518us of GEMM against - ~200us of comm, 2.6:1. wo_b is 39us against 137us per phase, 0.29:1. Overlap - can hide at most the smaller of the two, so theirs hides most of the comm and - this hides at most one GEMM. -3. **Its best mode has no producer/consumer handoff at all.** "Direct LSA" has - the GEMM epilogue store *straight into the destination rank's buffer*. No - staging, no copy engine, nothing to publish mid-kernel -- the only sync is the - barrier at the end. The publication problem this file spends all its time on - simply does not exist there. -4. **Its fused-SDMA path uses no cache fence.** ``opus_chunk_sdma_submit`` is - ``s_waitcnt vmcnt(0)``, ``s_barrier``, an ``__ATOMIC_ACQ_REL`` counter, and a - per-destination submit spin lock -- no ``__threadfence_system``, no - ``buffer_wbl2``. That is the lesson that transferred: the acq_rel counter *is* - the release, and the explicit fence added here was 90us of pure waste - (fused 440us -> 350us on removing it). Its ISA shows the atomic already emits - its own ``buffer_wbl2``/``buffer_inv`` pair, on thread 0 only. - -## The chunks race: it was the GEMM, and it is gone - -``--chunks`` > 1 produced wrong output about 3 runs in 10 and was pinned to 1 -for that reason. The cause was not the chunk protocol at all: it was aiter's -8-wave GEMM under-counting one ``s_waitcnt`` in its main loop, which corrupted -output non-deterministically on large grids whatever the epilogue did (see the -comment at that ``wait_barrier`` below). Since that fix, at [16384, 7168] -K=2048 on 8 ranks: - -* chunks 1/2/4/8, 10 runs each -- 40/40 correct, no hang; -* chunks=8 alone, 25 more runs -- 25/25 correct, no hang. - -Two hangs were seen at 8 ranks while the sweep was still being set up and never -reproduced in the 100+ runs after. The submit lock is a plain test-and-set spin -(``_acquire_peer_lock``), so a hang is not impossible; use ``timeout`` when -sweeping and treat one as a finding rather than a flake. - -Two things in this file survived that misdiagnosis and are worth keeping -straight: - -* The half-wave barrier pairing (``if wave_m == 0: s_barrier()`` before - ``store_c``) is still needed and still right -- gcnasm does it - unconditionally at kernel_template.hpp:693. It took ``--chunks 2`` from - 1-in-3 failing to 3-in-10, which at the time read as "better but not fixed"; - the residual 3-in-10 was the GEMM. -* Two *other* diagnoses were wrong and are recorded so they are not retried. - **A shared SDMA queue**: gcnasm's per-destination submit lock was ported and - ISA-verified; it fixed nothing, and is kept only because cco's - one-issuing-warp-per-queue rule still applies. **The counter atomic's - ordering**: ``acq_rel`` appeared to beat ``monotonic``, which is how the - acquire half got justified; against a 1-in-3 intermittent failure that - comparison was noise. ``acq_rel`` stays because release/acquire is right for - a producer handing tiles to a consumer, not because it was measured -- and at - chunks=1 it demonstrably orders nothing, since 20 runs of ``monotonic`` pass - too. - -Repeated runs are still the only way to judge any of this, and the gate has to -be tight: the corruption landed at 4-9e-3 against an fp8 floor of 2.35e-3, so a -5e-3 threshold reported a corrupt run as validated (it did, at 3.97e-3). The -bench gates at 3e-3, and ``test_fused_is_stable_across_repeats`` requires three -runs to be *identical* rather than each small. +### Folding the narrowing into the reduce -## Result: fused-sdma wins, once the chunks are unblocked - -8 ranks, [16384, 7168] K=2048 -- the real prefill shape -- graph replay, all -three C-store stages on (now the benchmark default), median of 31, max over -ranks: +`fuse_quantize=True` makes the reduce write its bf16 and narrow to fp8 from the +same accumulators, saving a 28 MiB re-read and a launch. It is **off**, because +it costs 20 us rather than saving 12: - mode time - fused-sdma chunks=8 1114.1us <- best - split-sdma 1261.7 - split-lsa 1262.4 - fused-lsa 1571.4 - gemm-only 228.9 +| | us | +|---|---:| +| split reduce + quantize | 957.4 | +| fused, row stashed in registers | 977.1 | +| fused, row re-read | 982.0 | -Per-kernel, the overlap is visible directly: +Not register pressure -- the re-reading variant keeps no stash and is no better. +It is the thread map: a per-row amax cannot be taken by a block holding only +part of a row, so fusing forces one-wave-per-row, where `sdma_reduce` walks +packs with a flat grid stride and streams a block through all 8 source slices at +once. Confining a wave to 14 KiB at a time costs the reduce more than the +re-read saves. - chunks GEMM drain reduce gather total - 1 257.5 493.8 44.0 496.4 1291.7 - 2 266.5 379.4 44.2 496.4 1186.6 - 8 265.0 303.7 43.8 496.6 1109.1 +It is not free and the cost is not a tuning problem. e4m3 carries 3 mantissa +bits, so one rounding costs ~2.1e-2 on a normal payload whatever the scale +granularity -- per-row measures 2.65e-2 and per-32 measures 2.40e-2, 9% better +for 200x the scale bytes. Going from 2.35e-3 to 2.49e-2 is the price of the +10%, and whether that is payable is a model-level question, not a kernel one. -The drain falls 38% while the GEMM grows 7.5us for the extra counter atomics -and the lock. Going finer is worse: at chunks=8 each PUT is 3.5 MiB, and 16 -(via ``--block-m 128``) halves that to 1.75 MiB, under the knee in the SDMA -bandwidth curve, for 1130.8us. +It also only pays at large M. The two conversion kernels are a fixed cost +against a transfer that shrinks with M, so on the standalone all-reduce it is ++2.4% at M=4096 and -11.2% at M=16384. Per phase at M=16384: quantize 12.4us, +dequantize 88.2us, against ~218us saved on the push. -fused-lsa still loses, and for a reason that is not going away: it spends 730us -of its GEMM pushing C over xGMI where the copy engines move the same bytes in -499, and ATT shows 99% of that store time is *stall*, so coalescing the stores -(the three C-store stages) buys it nothing -- 954.5 -> 952.1us. +The scatter leg stays bf16. It carries partial sums that are then added across +every rank, so its fp8 error compounds rather than being a single rounding, and +it is already mostly hidden behind the GEMM -- its 1140 GB/s is not a bandwidth, +it is the tell that the pushes went out from the epilogue and the drain is only +waiting for the tail. `scatter_dtype="fp8"` sizes its regions but raises +`NotImplementedError`. -The remaining floor is the tail: reduce plus all-gather is 540us of the 1114, -and neither is touched by fusing the GEMM. +### Firing the gather's puts from inside the reduce + +**Firing the gather's puts from inside the reduce** (`fuse_reduce_push=True`) +looked like the safest of the three: the push sends this rank's *own* slice, so +unlike the pull it has no cross-rank dependency, and SDMA is a copy engine so it +costs no CU time. It reaches parity and not a win — against 1148.7us unfused: + +| bands | 1 | 4 | 8 | 16 | 32 | +|---|---:|---:|---:|---:|---:| +| `publish="writethrough"` | 1162.1 | **1157.2** | 1160.3 | 1190.1 | 1323.5 | +| `publish="fence"` | 1227.3 | 1380.1 | 1630.3 | 2137.2 | 3026.7 | + +The gap between those rows is the useful result, and it is a lesson about how to +pay for a release rather than whether to. + +Handing a range to a copy engine *does* need one: the engine reads over the +fabric, not through a CU's cache, so `s_waitcnt vmcnt(0)` alone is not enough — +it only retires the stores as far as this XCD's L2. But the release can be paid +two ways. Releasing to system scope **after** the stores is L2-writeback work +charged once per block per band (256 x bands of it, ~61us per band, against a +reduce that is only 42us in total — unrepayable). Storing with `sc0+sc1` so the +bytes never stop in L1 or L2 makes the `waitcnt` itself the release, and that is +**free here**: applying the same store policy to the plain unfused reduce moves +it 1153.4 -> 1148.7us, i.e. nothing. This output is written once and nothing +local reads it again before the gather, so holding it in L2 bought nothing. + +What remains after that fix is small on both sides and nearly cancels. At +bands=1 the publish carries the mechanism's cost with none of its benefit — +1162.1 vs 1148.7, about 13us for the per-band 256-block `wait_barrier`, the +counter atomic and the elected block's locked puts. Four bands buy back about +5us of overlap before the sync cost takes over again. + +Dropping the release entirely is not an option even though it briefly looks like +one: with cached stores and no fence the kernel reaches 1157.9us, but at 32 bands +it produced relL2 1.8e-2 against the 2.35e-3 floor, differing per rank. The same +32 bands are exact under either correct publish mode, which rules out an indexing +bug. + +The contrast with the GEMM's fused scatter is the transferable part: there the +publish is amortised against a 437us transfer hidden behind a compute-bound GEMM; +here against 42us of bandwidth-saturated reduce. The mechanism pays when what is +hidden is much larger than the cost of publishing it. + +Re-measured after the window-geometry work below, with three alternating +repeats rather than a sweep, it is a clearer loss than the table suggests: +1128.4 us on against 1112.0 off, +16.4 us, against spreads of 2.2 and 1.5 us. + +### Hoisting the window geometry out of `lsa_ptr` + +**Hoisting the window geometry out of `lsa_ptr`.** `cco_lsa_ptr` is +`winBase + peer*stride + offset` and loads both fields on every call, through a +*generic* pointer -- which has to be a `flat_load`, since the compiler cannot +rule out LDS, so it counts against `lgkmcnt` as well as `vmcnt`. FlyDSL emits it +as an opaque extern call, and a kernel storing through addresses derived from +that base gives LLVM no way to prove the loads are not clobbered. + +Reading the geometry once and doing the arithmetic in the DSL removes all of +that. It was tried four ways -- hoisting out the band loop, a `lsa_geometry()` +API, `global_load` accessors in C++ (`cco_lsa_win_base` / `cco_lsa_stride`, which +take an `address_space(1)` pointer so each is a single `global_load`), and +finally `cco.CachedWindow`, which reads both in its constructor so a kernel +changes by one line. All four measured nothing on `kernels_sdma` (21 call sites) +and `kernels_fused`, in every wire configuration. + +It pays in exactly one place (`ptpc`, three alternating repeats): + +| | Window | CachedWindow | +|---|---|---| +| `split-lsa` | 1264.69 1264.23 1264.85 | 1250.25 1256.85 1254.53 | +| `fused-sdma` | 1110.45 1109.77 1112.85 | 1110.53 1113.69 1112.61 | + +-10.7 us on `split-lsa`, against a 0.6 us spread; nothing on `fused-sdma`. +`ar_1stage`/`ar_2stage` build nine peer addresses in *every block* of a short +kernel; everywhere else the addresses are built once per launch against a body +that runs for a millisecond. **Count address constructions per launch, not +`grep -c lsa_ptr`.** + +The same holds against PR #662's branch in `blockscale`, two alternating +repeats: `split-lsa` 1463.6 -> 1455.6 us, while `split-sdma` (1461.2 -> 1459.9), +`fused-sdma` bf16 (1144.9 -> 1145.8) and fp8/lsa (949.2 -> 949.8) do not move. + +Two things worth carrying. A `CachedWindow` cannot cross an `scf.if` -- FlyDSL +captures every variable an if body reads as state and requires single MLIR +values, which `Window` satisfies only by having exactly one field -- and the way +out is to compute the addresses before the branch, which is what the offsets +usually allow. And pin `--quant` when comparing against anything: the benchmark +defaults to `ptpc`, ~3% faster than the `blockscale` every blockscale number here is +quoted in, and reading one against the other looks exactly like a machine that +drifts overnight. From 0187031f7eb95e5cc8d0387e5085c26f1a7facec Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 07:17:04 +0000 Subject: [PATCH 4/8] docs: re-measure the end-to-end numbers on AMD's own V4.1 branch The SGLang-side numbers were taken against a synthetic snapshot of the container image, because the V4.1 code existed on no branch at the time. It does now -- sgl-project/sglang#39857, `kevin-mii:dsv41-amd-main` -- so all three servers were re-run on top of it. The fused wo_b reproduces and is slightly better: fp8 wire +5.6% to +6.0% at bs 4-16 against +5.1% to +5.7% before, bf16 +1.9% to +2.8% against +2.3% to +2.6%. GSM8K 0.917 / 0.912 / 0.918. Added, because it had never been in this section: the standalone GEMM on its own, +2.2% to +2.5% with nothing fused, GSM8K 0.920 / 0.928. And a note that the two results do not add -- `fused_wo_b` runs first and the linear only sees what it declines, so they split `wo_b` rather than stacking on it. Both together is a fourth configuration nobody has measured. **One claim here was wrong and is corrected.** It called bs=1 a built-in control on the grounds that M=4096 is below both floors. The floors are 8192 for bf16 and 2048 for fp8, and the shape log confirms the consequence: the bf16 wire's smallest served M is 7734, the fp8 wire's is 1792. bs=1 is a control for the bf16 column only, and the fp8 column's +3.4% there is a real win that was being dismissed as drift. Co-Authored-By: Claude Opus 5 (1M context) --- docs/MORI-GEMM-AR-BENCHMARK.md | 10 +++--- python/mori/ops/gemm_ar/README.md | 56 ++++++++++++++++++++++--------- 2 files changed, 46 insertions(+), 20 deletions(-) diff --git a/docs/MORI-GEMM-AR-BENCHMARK.md b/docs/MORI-GEMM-AR-BENCHMARK.md index 0768cbfe4..397305467 100644 --- a/docs/MORI-GEMM-AR-BENCHMARK.md +++ b/docs/MORI-GEMM-AR-BENCHMARK.md @@ -57,10 +57,12 @@ decides the outcome is the grid, `ceildiv(M,256) * (N/256)`: | 129-256 | **-26.6%** | | >256 | **-26.4%** | -End to end in a server, V4.1-Flash prefill throughput is **+5.1% to +5.7%** at -batch sizes 4-16 on the fp8 wire; V4-Pro is **-5.0% GPU busy** over a profiled -prefill. Per-layer wins are larger than end-to-end ones because `wo_b` is about -12% of the profile. +End to end in a server, V4.1-Flash prefill throughput is **+5.6% to +6.0%** at +batch sizes 4-16 on the fp8 wire, or **+2.2% to +2.5%** from the standalone +GEMM with nothing fused; V4-Pro is **-5.0% GPU busy** over a profiled prefill. +Per-layer wins are larger than end-to-end ones because `wo_b` is about 12% of +the profile, and the two SGLang-side results do not add -- they split that +layer rather than stacking on it. Full tables, the numerical cost of the fp8 wire, the model-level quality evaluation, and the threshold derivations are in diff --git a/python/mori/ops/gemm_ar/README.md b/python/mori/ops/gemm_ar/README.md index 5dcdf1e35..eaef66e07 100644 --- a/python/mori/ops/gemm_ar/README.md +++ b/python/mori/ops/gemm_ar/README.md @@ -813,25 +813,49 @@ plus a fill guard. ### End to end, in SGLang -Three servers, one per wire, GSM8K then `bench_one_batch_server`, two full runs. -Prefill throughput, tok/s, median, on V4.1-Flash: +Prefill throughput, tok/s, `bench_one_batch_server`, TP4 on V4.1-Flash, one +server per variant with GSM8K in front of it. + +**The fused wo_b**, three servers: | bs | base | bf16 wire | | fp8 wire | | |---|---:|---:|---:|---:|---:| -| 1 | 27349 | 27632 | +1.0% | 28173 | +3.0% | -| 4 | 34101 | 34990 | +2.6% | 36029 | **+5.7%** | -| 8 | 34948 | 35790 | +2.4% | 36733 | **+5.1%** | -| 16 | 35151 | 35969 | +2.3% | 36930 | **+5.1%** | - -**bs=1 is a built-in control**: M=4096 is below both floors, so no variant fuses -there and all three should agree -- they do, within 3%. That is what makes the -bs>=4 figures a signal rather than drift. - -GSM8K passes on all three (base .915/.916, bf16 .914/.917, fp8 .921/.924) and -the three sit inside each other's noise, so the fp8 leg's relL2 2.3e-2 is below -what 1319 questions resolve. That is not what the gate is for: it catches a -collective that moved nothing, which scores near zero while still answering -fluently. +| 1 | 28148 | 28469 | +1.1% | 29107 | **+3.4%** | +| 4 | 34386 | 35356 | +2.8% | 36442 | **+6.0%** | +| 8 | 35153 | 35920 | +2.2% | 37217 | **+5.9%** | +| 16 | 35423 | 36084 | +1.9% | 37416 | **+5.6%** | + +GSM8K 0.917 / 0.912 / 0.918, all three inside each other's noise, so the fp8 +leg's relL2 2.3e-2 is below what 1319 questions resolve. That is not what the +gate is for: it catches a collective that moved nothing, which scores near zero +while still answering fluently. + +**The two wires do not fuse at the same batch size, and an earlier revision of +this section had that wrong.** It called bs=1 a built-in control on the grounds +that M=4096 is below both floors -- but the floors are 8192 for bf16 and *2048* +for fp8. The shape log settles it: the bf16 wire's smallest served M is 7734, +the fp8 wire's is 1792. So bs=1 is a control for the bf16 column only, and the +fp8 column's +3.4% there is a real win rather than drift. + +**The GEMM on its own**, a separate pair of servers with +`SGLANG_OPT_MORI_MXFP8_GEMM=1` and nothing fused: + +| bs | base | mori GEMM | | +|---|---:|---:|---:| +| 1 | 28598 | 29313 | +2.5% | +| 4 | 34269 | 35047 | +2.3% | +| 8 | 35086 | 35865 | +2.2% | +| 16 | 35412 | 36232 | +2.3% | + +GSM8K 0.920 / 0.928. Its shape log shows `wq_b`, `wo_b` and `wqkv_a` served and +`shared gate_up` declined, which is what the tables above predict. + +**These three results do not add.** At `wo_b` the model calls `fused_wo_b` +first and only reaches the linear when it declines, so the two paths split that +layer rather than stacking on it; what the standalone GEMM adds beyond the +fused path is `wq_b` and `wqkv_a`, which are column-parallel and have no +collective to fuse with. Both enabled together is a fourth configuration and +has not been measured. The decode column is not reported as evidence. `wo_b` never fuses at decode's M, so the three variants should be identical, and they scatter -8.6% to +14.1% -- From 361d385ad23dec4ecc41679c7adc2ab95e25c3a5 Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 08:41:08 +0000 Subject: [PATCH 5/8] docs: the fourth configuration, and one fault that has not reproduced Both SGLang-side paths enabled together, which this section had been calling unmeasured: +7.4% to +7.7% prefill throughput at bs 4-16, against +6.0% for the fp8 wire alone and +2.3% for the standalone GEMM alone. They stack without adding, and the structure says why -- `fused_wo_b` runs first and the linear only sees what it declines, so at `wo_b` they divide the layer by M, and the ~1.7% on top is the column-parallel layers fusing cannot reach. Also recorded: one server out of nine hit `hipErrorIllegalAddress` under sustained load, and re-running the same variant completed cleanly. Written as what it is -- one sample, cause not established, the combination merely the only configuration that has shown it -- rather than as a known defect or as nothing. The five hypotheses already eliminated are listed so the next person does not re-walk them, and so is the reason `AMD_SERIALIZE_KERNEL=3` is being held back: it is worth spending on a repro that reproduces, and serialising perturbs the timing a race would depend on. Co-Authored-By: Claude Opus 5 (1M context) --- docs/MORI-GEMM-AR-BENCHMARK.md | 19 ++++++++---- python/mori/ops/gemm_ar/README.md | 49 +++++++++++++++++++++++++++---- 2 files changed, 56 insertions(+), 12 deletions(-) diff --git a/docs/MORI-GEMM-AR-BENCHMARK.md b/docs/MORI-GEMM-AR-BENCHMARK.md index 397305467..9a36d0428 100644 --- a/docs/MORI-GEMM-AR-BENCHMARK.md +++ b/docs/MORI-GEMM-AR-BENCHMARK.md @@ -57,12 +57,19 @@ decides the outcome is the grid, `ceildiv(M,256) * (N/256)`: | 129-256 | **-26.6%** | | >256 | **-26.4%** | -End to end in a server, V4.1-Flash prefill throughput is **+5.6% to +6.0%** at -batch sizes 4-16 on the fp8 wire, or **+2.2% to +2.5%** from the standalone -GEMM with nothing fused; V4-Pro is **-5.0% GPU busy** over a profiled prefill. -Per-layer wins are larger than end-to-end ones because `wo_b` is about 12% of -the profile, and the two SGLang-side results do not add -- they split that -layer rather than stacking on it. +End to end in a server, V4.1-Flash prefill throughput at batch sizes 4-16 is +**+5.6% to +6.0%** on the fp8 wire, **+2.2% to +2.5%** from the standalone GEMM +with nothing fused, and **+7.4% to +7.7%** with both enabled; V4-Pro is +**-5.0% GPU busy** over a profiled prefill. Per-layer wins are larger than +end-to-end ones because `wo_b` is about 12% of the profile, and the two +SGLang-side results stack without adding -- they split that layer by M rather +than both working on it. + +> One `both`-enabled server out of nine hit an illegal memory access under +> sustained load and has not reproduced. It is not established that the +> combination caused it; see +> [the operator README](../python/mori/ops/gemm_ar/README.md#end-to-end-in-sglang) +> for what has been ruled out. Full tables, the numerical cost of the fp8 wire, the model-level quality evaluation, and the threshold derivations are in diff --git a/python/mori/ops/gemm_ar/README.md b/python/mori/ops/gemm_ar/README.md index eaef66e07..beb6d4d25 100644 --- a/python/mori/ops/gemm_ar/README.md +++ b/python/mori/ops/gemm_ar/README.md @@ -850,12 +850,49 @@ fp8 column's +3.4% there is a real win rather than drift. GSM8K 0.920 / 0.928. Its shape log shows `wq_b`, `wo_b` and `wqkv_a` served and `shared gate_up` declined, which is what the tables above predict. -**These three results do not add.** At `wo_b` the model calls `fused_wo_b` -first and only reaches the linear when it declines, so the two paths split that -layer rather than stacking on it; what the standalone GEMM adds beyond the -fused path is `wq_b` and `wqkv_a`, which are column-parallel and have no -collective to fuse with. Both enabled together is a fourth configuration and -has not been measured. +**Both at once**, which is a fourth configuration rather than the sum of the +other two. Against the same session's base: + +| bs | base | fp8 wire only | both | vs base | vs fp8 only | +|---|---:|---:|---:|---:|---:| +| 1 | 28148 | 29107 | 29504 | +4.8% | +1.4% | +| 4 | 34386 | 36442 | 37023 | **+7.7%** | +1.6% | +| 8 | 35153 | 37217 | 37870 | **+7.7%** | +1.8% | +| 16 | 35423 | 37416 | 38058 | **+7.4%** | +1.7% | + +GSM8K 0.918. + +**They stack, but they do not add.** The fp8 wire alone is +6.0% and the +standalone GEMM alone is +2.3%; together they are +7.7%, not +8.3%. At `wo_b` +the model calls `fused_wo_b` first and the linear only ever sees what it +declines, so the two divide that layer by M rather than both working on it. +What the GEMM adds on top is `wq_b` and `wqkv_a` -- column-parallel, no +collective to fuse with, so fusing could never have reached them. + +> **One server out of nine hit an illegal memory access, and it has not +> reproduced.** It was a `both`-enabled run: it served GSM8K for 13 minutes at +> ~250 concurrent requests and then four ranks failed together with +> `hipErrorIllegalAddress`. Re-running the same variant with the same +> configuration completed cleanly -- 0.918, no faults -- and the numbers above +> are from that second run. +> +> **It is not established that the combination caused it**, and the evidence is +> thin in both directions: the combination is the only configuration that has +> ever shown it, and it has shown it once. What was ruled out, so it is not +> re-walked: the two paths do not share a FlyDSL `JitFunction` or `CallState` +> (two identical compiles return distinct objects); 30 rounds of interleaving +> the two paths on four ranks is clean; so is 40 rounds of driving one layer +> across the fusing threshold; `_PinnedLaunch` survives being pinned inside a +> CUDA graph and reused eagerly, replaying bit-identical to a fresh call; and +> memory is not it -- the two servers' pools are byte-identical and 31.7 GB was +> free at the fault. The standalone GEMM also served the same M distribution in +> the run that did not fault as in the one that did. +> +> `hipErrorIllegalAddress` is reported asynchronously, so the traceback points +> at the scheduler's next synchronisation rather than the faulting kernel. +> Pinning it down needs `AMD_SERIALIZE_KERNEL=3`, which is only worth spending +> on a repro that reproduces -- and serialising changes the timing a race +> depends on. Left open, to be re-tested under longer stress. The decode column is not reported as evidence. `wo_b` never fuses at decode's M, so the three variants should be identical, and they scatter -8.6% to +14.1% -- From 0fda3c1b259c209fcd748b4e29f0bdfad44b2e98 Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 08:56:29 +0000 Subject: [PATCH 6/8] docs: absolute URLs for the links out of the docs tree, as #688 established Rebasing onto main brought in #688, which converted this file's one link into `python/` to an absolute GitHub URL: Sphinx builds `docs/` alone, so a relative path to a file outside that tree is a broken reference and `-W` fails the build. The rewrite in this branch added three more of them. Same rule applied to all four. Links the other way -- from the operator README into `docs/` -- stay relative, because that file is not in the toctree and GitHub resolves them. Verified with the gate itself rather than by eye: `sphinx -E -n -W --keep-going` builds clean, and `tools/check_docs_links.py` reports 8872 local links across 25 pages with 0 errors. Co-Authored-By: Claude Opus 5 (1M context) --- docs/MORI-GEMM-AR-BENCHMARK.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/MORI-GEMM-AR-BENCHMARK.md b/docs/MORI-GEMM-AR-BENCHMARK.md index 9a36d0428..5abeb7603 100644 --- a/docs/MORI-GEMM-AR-BENCHMARK.md +++ b/docs/MORI-GEMM-AR-BENCHMARK.md @@ -68,7 +68,7 @@ than both working on it. > One `both`-enabled server out of nine hit an illegal memory access under > sustained load and has not reproduced. It is not established that the > combination caused it; see -> [the operator README](../python/mori/ops/gemm_ar/README.md#end-to-end-in-sglang) +> [the operator README](https://github.com/ROCm/mori/blob/main/python/mori/ops/gemm_ar/README.md#end-to-end-in-sglang) > for what has been ruled out. Full tables, the numerical cost of the fp8 wire, the model-level quality From f4ec5099d5ee7a4962eec69f007b5ff4cda5b7bb Mon Sep 17 00:00:00 2001 From: xiangch Date: Sun, 20 Sep 2026 09:23:13 +0000 Subject: [PATCH 7/8] gemm_ar: fix the CI break, and satisfy the pre-commit hooks Two unrelated CI failures on this branch. They are in one commit because `black` reflowed lines adjacent to the real fix, so splitting them would mean hunk surgery on the same few lines rather than a cleaner history. **The CCO unit test.** All ten `test_standalone_mxfp8_gemv` cases failed with AttributeError: 'Int32' object has no attribute 'type'. Did you mean: 'dtype'? _arith_ops_gen.py:3163, in ShRUIOp.__init__ while the same tests passed here. The generated MLIR builders infer their result type as `operands[0].type`, so they need an operand that is already an `ArithValue`; `fx.Int32` is a wrapper and carries `dtype`. FlyDSL's operators go through `_make_binop` -> `_extract_arith`, which unwraps first, so `(v >> shift) & 0xFF` is right on any FlyDSL that has the operators at all. Whether the wrapper exposes `.type` differs between the FlyDSL here and the one in CI's container, which is why this was invisible locally. The rule that holds on both: hand wrappers to FlyDSL operators, never to raw builders. `kernels_fused._load_byte` carried the identical expression and the identical latent break; CI never reached it because the GEMV tests failed first. The one `arith.andi` over two predicates became nested `arith.select` -- same result, and `select` is what the rest of mori builds with, so it is the op proven against whatever FlyDSL CI ships. Every `arith.*` call left in this branch is now one `origin/main` also uses. 146 tests pass locally, including the fp32-reference numerics that would catch any change in what the shift and mask compute. **pre-commit.** Formatting from the hooks themselves, plus two they flagged and could not fix: - `needs_sglang` in bench_gemm.py was dead (F841). It computed a condition nothing read; `build_mxfp8(want_sglang=...)` already gates the only thing that needs it, and `--scope linear` imports SGLang's quantiser inside `mori_linear_call`, where a missing build is a reported per-impl failure and a non-zero exit. Deleted rather than wired up: there was no guard to restore. - `l` as a comprehension variable in sweep.py (E741) -> `line`. The license hook appended the full MIT header to timing.py above its SPDX short form, leaving two notices; dropped the short one. The second notice in `_gemm_a8w8_8wave.py` is the deliberate vendored-from-aiter attribution and is left alone. Co-Authored-By: Claude Opus 5 (1M context) --- benchmark/cco/flydsl/gemm_ar/bench_gemm.py | 111 ++++++++++++++++----- benchmark/cco/flydsl/gemm_ar/bench_gemv.py | 100 +++++++++++++++---- benchmark/cco/flydsl/gemm_ar/report.py | 99 ++++++++++++++---- benchmark/cco/flydsl/gemm_ar/sweep.py | 105 ++++++++++++++----- benchmark/cco/flydsl/gemm_ar/timing.py | 22 +++- python/mori/ops/gemm_ar/gemm.py | 4 +- python/mori/ops/gemm_ar/kernels_fused.py | 21 ++-- python/mori/ops/gemm_ar/kernels_gemv.py | 37 +++++-- python/mori/ops/gemm_ar/op.py | 5 +- tests/python/cco/test_gemm_ar_op.py | 14 +-- tests/python/cco/test_mxfp8_gemv_grid.py | 91 ++++++++++++----- 11 files changed, 464 insertions(+), 145 deletions(-) diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemm.py b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py index 2d6d15f12..93a2bd18a 100644 --- a/benchmark/cco/flydsl/gemm_ar/bench_gemm.py +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py @@ -1,4 +1,25 @@ #!/usr/bin/env python3 +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. """mori's mxfp8/blockscale GEMM, at either scope, optionally against SGLang. One entry point for the two questions that are *not* about the collective, and @@ -85,8 +106,12 @@ def build_mxfp8(n, k, seed=1234, want_sglang=False): g = torch.Generator(device="cuda").manual_seed(seed) w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) eb = torch.randint( - 120, 123, (n // MXFP8_BK, k // MXFP8_BK), generator=g, - device="cuda", dtype=torch.int32, + 120, + 123, + (n // MXFP8_BK, k // MXFP8_BK), + generator=g, + device="cuda", + dtype=torch.int32, ) out = { "w_raw": w, @@ -105,7 +130,9 @@ def build_mxfp8(n, k, seed=1234, want_sglang=False): shuffled, scale_e8m0, weight_bf16 = prepare_mxfp8_native_weight( w, torch.exp2(eb.float() - 127.0), (32, 32) ) - out["layer"] = _Layer(shuffled.view(torch.float8_e4m3fn), scale_e8m0, weight_bf16) + out["layer"] = _Layer( + shuffled.view(torch.float8_e4m3fn), scale_e8m0, weight_bf16 + ) return out @@ -117,10 +144,13 @@ def build_blockscale(n, k, m, seed=1234): a = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) kb = k // SCALE_BK - sa = torch.rand(m, kb, generator=g, device="cuda", dtype=torch.float32) * 0.01 + 0.01 + sa = ( + torch.rand(m, kb, generator=g, device="cuda", dtype=torch.float32) * 0.01 + 0.01 + ) sb = ( torch.rand(n // SCALE_BK, kb, generator=g, device="cuda", dtype=torch.float32) - * 0.01 + 0.01 + * 0.01 + + 0.01 ) return a, w, preshuffle_b(w), sa.t().contiguous().t(), sb @@ -187,17 +217,35 @@ def mori_kernel_call(ops, n, k, m, block_n, quant): a, _w, w_shuf, sa, sb = ops gemm = compile_fused_gemm_scatter( - layout.ArConfig(world_size=2, m=128, n=n), 0, K=k, - BLOCK_M=128, BLOCK_N=block_n or 256, b_preshuffled=True, fuse=False, - swap_ab=True, permlane=True, lane_transpose=True, quant="blockscale", + layout.ArConfig(world_size=2, m=128, n=n), + 0, + K=k, + BLOCK_M=128, + BLOCK_N=block_n or 256, + b_preshuffled=True, + fuse=False, + swap_ab=True, + permlane=True, + lane_transpose=True, + quant="blockscale", ) y = torch.zeros(m, n, device="cuda", dtype=torch.bfloat16) a_i8 = a.contiguous().view(torch.int8).view(-1) sa_arg, sb_arg = sa.t().reshape(-1).contiguous(), sb.reshape(-1).contiguous() def call(picked): - gemm(a_i8, picked[0], y.view(-1), sa_arg, sb_arg, m, n, 0, 0, - stream=fx.Stream(torch.cuda.current_stream())) + gemm( + a_i8, + picked[0], + y.view(-1), + sa_arg, + sb_arg, + m, + n, + 0, + 0, + stream=fx.Stream(torch.cuda.current_stream()), + ) return y return call, [w_shuf.contiguous().view(torch.int8).view(-1)] @@ -254,7 +302,9 @@ def sglang_linear_call(layer, x_bf16): def call(picked): return mxfp8_native_blockscaled_linear( - x_bf16, picked[0].view(torch.uint8), layer.weight_scale_mx_e8m0, + x_bf16, + picked[0].view(torch.uint8), + layer.weight_scale_mx_e8m0, weight_bf16=(picked[1] if len(picked) > 1 else None), ) @@ -281,11 +331,18 @@ def main() -> int: p.add_argument("-m", type=int, required=True) p.add_argument("--scope", choices=("kernel", "linear"), default="linear") p.add_argument("--quant", choices=("mxfp8", "blockscale"), default="mxfp8") - p.add_argument("--impl", default="auto,gemm256,gemm128,sglang", - help="comma-separated: " + ", ".join(IMPLS)) + p.add_argument( + "--impl", + default="auto,gemm256,gemm128,sglang", + help="comma-separated: " + ", ".join(IMPLS), + ) p.add_argument("--reps", type=int, default=32) - p.add_argument("--tol", type=float, default=2.4e-3, - help="rel_l2 above this marks the row invalid") + p.add_argument( + "--tol", + type=float, + default=2.4e-3, + help="rel_l2 above this marks the row invalid", + ) p.add_argument("--json-out", default="gemm.jsonl") args = p.parse_args() @@ -304,8 +361,6 @@ def main() -> int: if args.quant == "blockscale" and ("sglang" in impls or args.scope == "linear"): p.error("--quant blockscale is mori-only and kernel-scope only") - # `linear` means "quantise a bf16 activation", which is SGLang's quantiser. - needs_sglang = "sglang" in impls or args.scope == "linear" vram_before = timing.vram_used() if args.quant == "mxfp8": @@ -315,8 +370,13 @@ def main() -> int: x = (torch.randn(m, k, device="cuda") / 8).to(torch.bfloat16) common = { - "bench": "gemm", "scope": args.scope, "quant": args.quant, - "shape": args.shape, "n": n, "k": k, "m": m, + "bench": "gemm", + "scope": args.scope, + "quant": args.quant, + "shape": args.shape, + "n": n, + "k": k, + "m": m, "input": "bf16" if args.scope == "linear" else "fp8", "includes_quant": args.scope == "linear", "timing": "amortized-graph-cold-hot", @@ -324,8 +384,7 @@ def main() -> int: rows, ref, failures = [], None, 0 # SGLang first when present, so it is the reference the rest are scored on. - order = ([i for i in impls if i == "sglang"] - + [i for i in impls if i != "sglang"]) + order = [i for i in impls if i == "sglang"] + [i for i in impls if i != "sglang"] for impl in order: row = dict(common, impl=impl) try: @@ -335,6 +394,7 @@ def main() -> int: from sglang.kernels.ops.quantization.mxfp8_native_amd_gfx95 import ( native_route_plan, ) + row["route"] = native_route_plan( m, n, k, ops["layer"].weight_bf16 is not None, False ) @@ -380,9 +440,12 @@ def main() -> int: with open(args.json_out, "a") as f: for r in rows: - f.write(json.dumps( - dict(r, vram_before=vram_before, vram_after=timing.vram_used()) - ) + "\n") + f.write( + json.dumps( + dict(r, vram_before=vram_before, vram_after=timing.vram_used()) + ) + + "\n" + ) # A benchmark that fails and exits 0 is how a broken sweep looks green. return 1 if failures else 0 diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemv.py b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py index 1c1692ce3..c50c783eb 100644 --- a/benchmark/cco/flydsl/gemm_ar/bench_gemv.py +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py @@ -1,4 +1,25 @@ #!/usr/bin/env python3 +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. """mori's mxfp8 GEMV against sglang's, and a config sweep for mori's. Both sides get the same raw weight and the same ue8m0 scales, each in its own @@ -142,9 +163,13 @@ def main() -> int: p.add_argument("--sweep", action="store_true") p.add_argument("--config", default=None) p.add_argument("--reps", type=int, default=64) - p.add_argument("--baseline", choices=("sglang", "none"), default="sglang", - help="'none' drops the SGLang import; mori's own numbers " - "need nothing but mori") + p.add_argument( + "--baseline", + choices=("sglang", "none"), + default="sglang", + help="'none' drops the SGLang import; mori's own numbers " + "need nothing but mori", + ) p.add_argument("--tol", type=float, default=2.4e-3) p.add_argument("--json-out", default="gemv.jsonl") args = p.parse_args() @@ -157,9 +182,15 @@ def main() -> int: vram_before = timing.vram_used() common = { - "bench": "gemv", "scope": "kernel", "quant": "mxfp8", - "shape": args.shape, "n": n, "k": k, "m": m, - "input": "fp8", "includes_quant": False, + "bench": "gemv", + "scope": "kernel", + "quant": "mxfp8", + "shape": args.shape, + "n": n, + "k": k, + "m": m, + "input": "fp8", + "includes_quant": False, "timing": "amortized-graph-cold-hot", } rows, base, ref, failures = [], None, None, 0 @@ -173,8 +204,16 @@ def main() -> int: f"cold {base['cold_us']:6.2f}", flush=True, ) - rows.append(dict(common, impl="sglang", route="sglang-gemv", - rel_l2=0.0, validated=True, **base)) + rows.append( + dict( + common, + impl="sglang", + route="sglang-gemv", + rel_l2=0.0, + validated=True, + **base, + ) + ) if args.config: cfgs = [parse_key(args.config)] @@ -192,36 +231,57 @@ def main() -> int: torch.cuda.synchronize() rel = ( ((out[:m].float() - ref).norm() / ref.norm().clamp_min(1e-30)).item() - if ref is not None else None + if ref is not None + else None ) res = timing.cold_hot_us(call, [wt], reps=args.reps) except Exception as err: # noqa: BLE001 - a bad config must not stop the sweep print(f" {key_of(cfg):<12} FAILED {type(err).__name__}: {err}", flush=True) traceback.print_exc(limit=3) - rows.append(dict(common, impl="mori", config=key_of(cfg), - validated=False, - error=f"{type(err).__name__}: {err}")) + rows.append( + dict( + common, + impl="mori", + config=key_of(cfg), + validated=False, + error=f"{type(err).__name__}: {err}", + ) + ) failures += 1 continue ok = rel is None or rel <= args.tol failures += 0 if ok else 1 - vs = (f" {(res['cold_us'] / base['cold_us'] - 1) * 100:+6.1f}%" - if base else " " * 8) + vs = ( + f" {(res['cold_us'] / base['cold_us'] - 1) * 100:+6.1f}%" + if base + else " " * 8 + ) rel_s = " n/a " if rel is None else f" relL2 {rel:.2e}" print( f" {key_of(cfg):<12} hot {res['hot_us']:6.2f} cold {res['cold_us']:6.2f}" f"{vs}{rel_s}{'' if ok else ' !! over tol'}", flush=True, ) - rows.append(dict(common, impl="mori", config=key_of(cfg), - route=f"mori-gemv-{key_of(cfg)}", - rel_l2=rel, validated=ok, **res)) + rows.append( + dict( + common, + impl="mori", + config=key_of(cfg), + route=f"mori-gemv-{key_of(cfg)}", + rel_l2=rel, + validated=ok, + **res, + ) + ) with open(args.json_out, "a") as f: for r in rows: - f.write(json.dumps(dict( - r, vram_before=vram_before, vram_after=timing.vram_used() - )) + "\n") + f.write( + json.dumps( + dict(r, vram_before=vram_before, vram_after=timing.vram_used()) + ) + + "\n" + ) # A benchmark that fails and exits 0 is how a broken sweep looks green. return 1 if failures else 0 diff --git a/benchmark/cco/flydsl/gemm_ar/report.py b/benchmark/cco/flydsl/gemm_ar/report.py index e390bfe02..5b5e1f53c 100644 --- a/benchmark/cco/flydsl/gemm_ar/report.py +++ b/benchmark/cco/flydsl/gemm_ar/report.py @@ -1,4 +1,25 @@ #!/usr/bin/env python3 +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. """Render any of the benchmark JSONL files into tables. One reader for every producer, because the two it replaced each knew about one @@ -32,14 +53,28 @@ #: else in the row is part of the key, which is what keeps distinct runs #: distinct without this file having to know the axes in advance. MEASURED = { - "hot_us", "cold_us", "us", "max_rank_time_us", "rel_l2", "validated", - "copies", "cold_reps", "working_set_mb", "vram_before", "vram_after", - "error", "timing", "route", "supported", "reason", + "hot_us", + "cold_us", + "us", + "max_rank_time_us", + "rel_l2", + "validated", + "copies", + "cold_reps", + "working_set_mb", + "vram_before", + "vram_after", + "error", + "timing", + "route", + "supported", + "reason", # Resolved by the op, not chosen by the caller: `critical_rank` is elected # at runtime and `resolved_tile_order` is what `--tile-order auto` became. # They are never sweep axes, so they can be excluded even when a file # predates `sweep_axes`. - "critical_rank", "resolved_tile_order", + "critical_rank", + "resolved_tile_order", } #: Shown as the row label when present, in this order; the rest go in the key. ROW_KEYS = ("label", "shape", "n", "k", "world_size", "quant", "scope") @@ -71,12 +106,19 @@ def time_of(row, col): def main() -> int: p = argparse.ArgumentParser() p.add_argument("paths", nargs="+") - p.add_argument("--col", default="cold_us", - help="which timing field to table (default cold_us)") - p.add_argument("--by", default=None, - help="column axis; default is the first of " + ", ".join(COL_KEYS)) - p.add_argument("--baseline", default=None, - help="column to show the others as a percentage against") + p.add_argument( + "--col", default="cold_us", help="which timing field to table (default cold_us)" + ) + p.add_argument( + "--by", + default=None, + help="column axis; default is the first of " + ", ".join(COL_KEYS), + ) + p.add_argument( + "--baseline", + default=None, + help="column to show the others as a percentage against", + ) p.add_argument("--keep-invalid", action="store_true") args = p.parse_args() @@ -88,12 +130,16 @@ def main() -> int: bad = [r for r in rows if r.get("validated") is False] if bad and not args.keep_invalid: rows = [r for r in rows if r.get("validated") is not False] - print(f"!! {len(bad)} row(s) failed validation, excluded " - f"(--keep-invalid to see them)") + print( + f"!! {len(bad)} row(s) failed validation, excluded " + f"(--keep-invalid to see them)" + ) for r in bad[:5]: who = r.get("impl") or r.get("mode") or "?" - print(f" {r.get('shape','?')} M={r.get('m','?')} {who}: " - f"{r.get('error') or f'rel_l2={r.get("rel_l2")}'}") + print( + f" {r.get('shape','?')} M={r.get('m','?')} {who}: " + f"{r.get('error') or f'rel_l2={r.get("rel_l2")}'}" + ) if not rows: print("nothing left after validation filter") return 1 @@ -117,12 +163,17 @@ def keyof(r): axes = r.get("sweep_axes") if axes is not None: keep = [k for k in axes if k not in (col_key, "m")] - keep += [k for k in ("label", "shape", "n", "k", "world_size", - "quant", "scope") if k in r] + keep += [ + k + for k in ("label", "shape", "n", "k", "world_size", "quant", "scope") + if k in r + ] return tuple((k, r[k]) for k in sorted(set(keep))) return tuple( - (k, r[k]) for k in sorted(r) - if k not in MEASURED and k not in (col_key, "m", "bench", "sweep_axes") + (k, r[k]) + for k in sorted(r) + if k not in MEASURED + and k not in (col_key, "m", "bench", "sweep_axes") and not isinstance(r[k], (dict, list)) ) @@ -135,8 +186,9 @@ def keyof(r): for key, cells in sorted(table.items(), key=lambda kv: str(kv[0])): d = dict(key) head = " ".join(f"{k}={d[k]}" for k in ROW_KEYS if k in d) - rest = " ".join(f"{k}={v}" for k, v in key - if k not in ROW_KEYS and v not in (None, "")) + rest = " ".join( + f"{k}={v}" for k, v in key if k not in ROW_KEYS and v not in (None, "") + ) print(f"\n### {head}" + (f" [{rest}]" if rest else "")) ms = sorted({m for m, _ in cells if m is not None}) cols = sorted({c for _, c in cells if c is not None}, key=str) @@ -144,8 +196,11 @@ def keyof(r): continue base = args.baseline if args.baseline in cols else None - print(f"{'M':>8}" + "".join(f"{str(c):>16}" for c in cols) - + (" (vs " + base + ")" if base else "")) + print( + f"{'M':>8}" + + "".join(f"{str(c):>16}" for c in cols) + + (" (vs " + base + ")" if base else "") + ) for m in ms: line = f"{m:>8}" bt = None diff --git a/benchmark/cco/flydsl/gemm_ar/sweep.py b/benchmark/cco/flydsl/gemm_ar/sweep.py index 12ca2d775..cfb0b6b8f 100644 --- a/benchmark/cco/flydsl/gemm_ar/sweep.py +++ b/benchmark/cco/flydsl/gemm_ar/sweep.py @@ -1,4 +1,25 @@ #!/usr/bin/env python3 +# Copyright © Advanced Micro Devices, Inc. All rights reserved. +# +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. """Run a benchmark matrix, one point per process, and say so when one fails. The matrices used to be six shell scripts that each re-implemented process @@ -36,8 +57,16 @@ #: disagree with an existing one about what `wq_b` means. MAIN = ["wq_b", "wo_b"] ALL_SHAPES = MAIN + [ - "wq_a_tp4", "wkv_tp4", "wqkv_a_tp4", "wo_a_tp4", "shared_gate_up_tp4", - "wq_b_tp8", "wo_b_tp8", "wo_a_tp8", "wq_b_tp1", "wo_b_tp1", + "wq_a_tp4", + "wkv_tp4", + "wqkv_a_tp4", + "wo_a_tp4", + "shared_gate_up_tp4", + "wq_b_tp8", + "wo_b_tp8", + "wo_a_tp8", + "wq_b_tp1", + "wo_b_tp1", ] #: M values a default run covers. Deliberately includes both sides of the two @@ -110,16 +139,26 @@ def _grid(**axes): ), "fused-fp8": dict( doc="the same mode matrix on the winning fp8 wire, which is where the " - "fusion is actually deployed", + "fusion is actually deployed", script="bench_gemm_ar.py", world=4, grid=_grid( m=[4096, 8192, 16384], mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"], ), - args=["-n", "5120", "-k", "2048", "--quant", "mxfp8", - "--gather-dtype", "fp8", "--gather-transport", "lsa", - "--no-fuse-quantize"], + args=[ + "-n", + "5120", + "-k", + "2048", + "--quant", + "mxfp8", + "--gather-dtype", + "fp8", + "--gather-transport", + "lsa", + "--no-fuse-quantize", + ], label="mxfp8-fp8wire", ), "fused-wire": dict( @@ -166,21 +205,27 @@ def require_sdma() -> None: except ImportError: BUILD_CCO_SDMA = False if not BUILD_CCO_SDMA: - sys.exit("BUILD_CCO_SDMA is OFF -- the SDMA path is compiled out, " - "numbers would be fiction. Rebuild with BUILD_CCO_SDMA=ON.") + sys.exit( + "BUILD_CCO_SDMA is OFF -- the SDMA path is compiled out, " + "numbers would be fiction. Rebuild with BUILD_CCO_SDMA=ON." + ) def occupancy() -> str: try: out = subprocess.run( ["rocm-smi", "--showmeminfo", "vram", "--csv"], - capture_output=True, text=True, timeout=30, + capture_output=True, + text=True, + timeout=30, ).stdout except (OSError, subprocess.SubprocessError): return "?" - gib = [f"{int(p[2]) / 2**30:.0f}" for p in - (l.split(",") for l in out.splitlines()) - if len(p) >= 3 and p[2].isdigit()] + gib = [ + f"{int(p[2]) / 2**30:.0f}" + for p in (line.split(",") for line in out.splitlines()) + if len(p) >= 3 and p[2].isdigit() + ] return " ".join(gib) or "?" @@ -194,16 +239,23 @@ def point_argv(preset, point, out_path): # `python -m torch.distributed.run` rather than the `torchrun` console # script: the latter is only on PATH if the venv is activated, which a # subprocess inherits only by luck. - cmd = [PY, "-m", "torch.distributed.run", "--standalone", - f"--nproc_per_node={world}", script] + cmd = [ + PY, + "-m", + "torch.distributed.run", + "--standalone", + f"--nproc_per_node={world}", + script, + ] else: cmd = [PY, script, "--json-out", str(out_path)] for key, val in point.items(): if isinstance(val, bool): # bench_gemm_ar.py spells these as --fuse-quantize / --no-fuse-quantize - cmd.append(f"--{key.replace('_', '-')}" if val - else f"--no-{key.replace('_', '-')}") + cmd.append( + f"--{key.replace('_', '-')}" if val else f"--no-{key.replace('_', '-')}" + ) elif key == "m": cmd += ["-m", str(val)] else: @@ -219,8 +271,9 @@ def run_point(preset, point, out_path, timeout): env.setdefault("MORI_SOCKET_IFNAME", "lo") env["MORI_ENABLE_SDMA"] = "1" try: - p = subprocess.run(cmd, capture_output=True, text=True, - timeout=timeout, cwd=HERE, env=env) + p = subprocess.run( + cmd, capture_output=True, text=True, timeout=timeout, cwd=HERE, env=env + ) except subprocess.TimeoutExpired: return False, f"TIMEOUT after {timeout}s" # Multi-rank runs print RESULT_JSON rather than writing the file themselves. @@ -266,8 +319,7 @@ def main() -> int: if args.dry_run: for point in preset["grid"]: - print(" ".join(shlex.quote(c) - for c in point_argv(preset, point, out_path))) + print(" ".join(shlex.quote(c) for c in point_argv(preset, point, out_path))) return 0 if preset.get("world"): @@ -284,8 +336,10 @@ def main() -> int: ok, out = run_point(preset, point, out_path, timeout) if not ok: failed.append(desc) - print(f" FAILED: {out.strip().splitlines()[-1] if out.strip() else '?'}", - flush=True) + print( + f" FAILED: {out.strip().splitlines()[-1] if out.strip() else '?'}", + flush=True, + ) else: for line in out.splitlines(): if line.startswith(" ") or line.startswith("RESULT_JSON"): @@ -296,8 +350,11 @@ def main() -> int: after = occupancy() print(f"# occupancy after: {after} GiB/card", file=sys.stderr) if before != after: - print("# NOTE: occupancy moved across the sweep; a leftover process may " - "have shared the GPU. Re-run before trusting these.", file=sys.stderr) + print( + "# NOTE: occupancy moved across the sweep; a leftover process may " + "have shared the GPU. Re-run before trusting these.", + file=sys.stderr, + ) if failed: print(f"\n{len(failed)} of {len(preset['grid'])} points FAILED:") for d in failed: diff --git a/benchmark/cco/flydsl/gemm_ar/timing.py b/benchmark/cco/flydsl/gemm_ar/timing.py index 5fed6ff0e..44818aea9 100644 --- a/benchmark/cco/flydsl/gemm_ar/timing.py +++ b/benchmark/cco/flydsl/gemm_ar/timing.py @@ -1,6 +1,24 @@ -# Copyright (c) 2025 Advanced Micro Devices, Inc. All rights reserved. +# Copyright © Advanced Micro Devices, Inc. All rights reserved. # -# SPDX-License-Identifier: MIT +# MIT License +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. """Timing for kernels small enough that the harness is the measurement. Two corrections over the obvious `capture one call, replay, take the median`, diff --git a/python/mori/ops/gemm_ar/gemm.py b/python/mori/ops/gemm_ar/gemm.py index 9f284946f..336a7c197 100644 --- a/python/mori/ops/gemm_ar/gemm.py +++ b/python/mori/ops/gemm_ar/gemm.py @@ -332,9 +332,7 @@ def __call__( ) self._check_operands(a_fp8, b_preshuffled, a_scale, b_scale, m) if out is None: - out = torch.empty( - (m, self.n), dtype=torch.bfloat16, device=a_fp8.device - ) + out = torch.empty((m, self.n), dtype=torch.bfloat16, device=a_fp8.device) elif tuple(out.shape) != (m, self.n) or out.dtype != torch.bfloat16: raise ValueError( f"out must be {(m, self.n)} bfloat16, got {tuple(out.shape)} " diff --git a/python/mori/ops/gemm_ar/kernels_fused.py b/python/mori/ops/gemm_ar/kernels_fused.py index 61a760059..d3233bf44 100644 --- a/python/mori/ops/gemm_ar/kernels_fused.py +++ b/python/mori/ops/gemm_ar/kernels_fused.py @@ -421,7 +421,9 @@ def __init__( a_bytes = m * self.kb_count else: a_bytes = m * (self.k_dwords if row_major else self.kb_count) * 4 - b_bytes = (n // self.BLOCK) * (self.k_dwords if row_major else self.kb_count) * 4 + b_bytes = ( + (n // self.BLOCK) * (self.k_dwords if row_major else self.kb_count) * 4 + ) gSA = fx.rocdl.make_buffer_tensor( A_scale, max_size=False, num_records_bytes=a_bytes ) @@ -448,9 +450,13 @@ def _load_byte(self, div, dword_index): four block groups read the same sixteen addresses and only the shift differs. Two VALU ops per scale, both cheap; what row major costs is in the addresses, not here. + + FlyDSL's own operators rather than ``arith.shrui``/``andi``: the raw MLIR + builders read ``.type`` off the operand, which an ``fx.Int32`` wrapper + does not carry. See the same note in ``kernels_gemv._byte``. """ v = self._load1(div, dword_index) - return arith.andi(arith.shrui(v, self.byte_shift), fx.Int32(0xFF)) + return (v >> self.byte_shift) & fx.Int32(0xFF) def _kb(self, ks): """This lane's 32-block for K step ``ks``: block ``4*ks + lane//16``.""" @@ -1132,9 +1138,10 @@ def compile_fused_gemm_scatter( # is legal without permlane and not with it, which is what the store's own # assert says a few frames deeper and less usefully. n_floor = 256 if permlane else 128 - assert BLOCK_M >= 128 and BLOCK_N >= n_floor, ( - f"BLOCK_N={BLOCK_N} is below {n_floor}" - + (" (permlane's store pairs two N-tiles)" if permlane else "") + assert ( + BLOCK_M >= 128 and BLOCK_N >= n_floor + ), f"BLOCK_N={BLOCK_N} is below {n_floor}" + ( + " (permlane's store pairs two N-tiles)" if permlane else "" ) assert BLOCK_M % 128 == 0 and BLOCK_N % 128 == 0 assert K % BLOCK_K == 0 @@ -1310,9 +1317,7 @@ def kernel_gemm_scatter( # idx(tj, ti). # With the swap the instruction's B is our A, so the per-tile # opsel that selects a packed A byte is opsel_b on the raw atom. - mfma_raw = Mfma16x16x128( - N_TILES_B, N_TILES_A, opsel_b_per_tile=mxfp8_pack - ) + mfma_raw = Mfma16x16x128(N_TILES_B, N_TILES_A, opsel_b_per_tile=mxfp8_pack) mfma = _SwappedMfma(mfma_raw) else: mfma = Mfma16x16x128(N_TILES_A, N_TILES_B) diff --git a/python/mori/ops/gemm_ar/kernels_gemv.py b/python/mori/ops/gemm_ar/kernels_gemv.py index 1fd39256e..e7ad4fb5d 100644 --- a/python/mori/ops/gemm_ar/kernels_gemv.py +++ b/python/mori/ops/gemm_ar/kernels_gemv.py @@ -138,9 +138,19 @@ def _load4(self, div, index): return Vec(fx.memref_load_vec(self.reg_4)) def _byte(self, div, index): - """One ue8m0 scale out of a row-major dword of four.""" + """One ue8m0 scale out of a row-major dword of four. + + Shift and mask through FlyDSL's operators, not ``arith.shrui``/``andi``. + The generated MLIR builders take an operand's ``.type`` directly, which + only works if the value handed in is already an ``ArithValue``; an + ``fx.Int32`` wrapper has ``dtype`` instead and raises. FlyDSL's operators + unwrap the operands first (``_make_binop`` -> ``_extract_arith``), so they + are correct on every FlyDSL that has them -- and that difference is + invisible on one version and fatal on the next, which is how this passed + locally and failed CI on all ten GEMV shapes. + """ v = self._load1(div, index) - return arith.andi(arith.shrui(v, self.byte_shift), fx.Int32(0xFF)) + return (v >> self.byte_shift) & fx.Int32(0xFF) def w_frag(self, tile, t, step_dw): """The 32 weight bytes of tile ``tile + t`` this lane feeds one MFMA. @@ -212,7 +222,9 @@ def compile_mxfp8_gemv( if n % TILE: raise ValueError(f"N={n} must be a multiple of {TILE}") if n % MXFP8_BLOCK: - raise ValueError(f"N={n} must be a multiple of {MXFP8_BLOCK} (the B scale group)") + raise ValueError( + f"N={n} must be a multiple of {MXFP8_BLOCK} (the B scale group)" + ) if rows not in (16, 32) or tokens not in (16, 32): raise ValueError(f"rows/tokens must be 16 or 32, got {rows}/{tokens}") if tokens < m_max: @@ -341,10 +353,16 @@ def kernel_gemv( [gv.x_frag(toks[b], x_step_dw) for b in range_constexpr(BT)] ) w_sc.append( - [gv.scale_at(gv.ws, w_sc_base[t], sc_step) for t in range_constexpr(AT)] + [ + gv.scale_at(gv.ws, w_sc_base[t], sc_step) + for t in range_constexpr(AT) + ] ) x_sc.append( - [gv.scale_at(gv.xs, x_sc_base[b], sc_step) for b in range_constexpr(BT)] + [ + gv.scale_at(gv.xs, x_sc_base[b], sc_step) + for b in range_constexpr(BT) + ] ) for s in range_constexpr(len(w_frags)): acc = mfma.call( @@ -365,8 +383,13 @@ def kernel_gemv( c_div = fx.logical_divide(gC, fx.make_layout(1, 1)) def store(value, tok, col): - in_range = arith.andi(tok < c_m, col < c_n) - idx = arith.select(in_range, tok * c_n + col, oob) + # Nested selects rather than `arith.andi` on the two predicates: + # `select` is the op the rest of mori already builds with, so it is + # the one proven against the FlyDSL that CI ships. Same result -- + # out of range on either axis sends the store past `num_records`. + idx = arith.select( + tok < c_m, arith.select(col < c_n, tok * c_n + col, oob), oob + ) fx.memref_store_vec(Vec.filled(1, value, fx.BFloat16), out_reg) fx.copy(out_atom, out_reg, fx.slice(c_div, (None, fx.Int32(idx)))) diff --git a/python/mori/ops/gemm_ar/op.py b/python/mori/ops/gemm_ar/op.py index 81b4142da..27c1070ad 100644 --- a/python/mori/ops/gemm_ar/op.py +++ b/python/mori/ops/gemm_ar/op.py @@ -93,6 +93,7 @@ def _quant_tile_constraint(quant: str, block_m: int) -> str | None: ) return None + # Chunks are how many separate pushes a destination receives, and so how early # the first bytes leave. More is better until the pieces get small enough that # the SDMA per-packet cost shows; 8 is the measured knee. @@ -150,7 +151,9 @@ def counter_chunks(m_pad: int, world_size: int, block_m: int = DEFAULT_BLOCK_M) return max(c for c in range(1, min(MAX_CHUNKS, bands) + 1) if bands % c == 0) -def _tile_constraints(block_m: int, block_n: int, n_granule: int | None = None) -> str | None: +def _tile_constraints( + block_m: int, block_n: int, n_granule: int | None = None +) -> str | None: """Why this tile is not one the kernel can build, or None. Checked before anything divides by a tile size. ``padded_m`` and diff --git a/tests/python/cco/test_gemm_ar_op.py b/tests/python/cco/test_gemm_ar_op.py index f51e7b7dc..a2be37968 100644 --- a/tests/python/cco/test_gemm_ar_op.py +++ b/tests/python/cco/test_gemm_ar_op.py @@ -414,9 +414,7 @@ def _run_worker(world_size: int, case: str, *extra: str, timeout: int = 900): requires_two_gpus = pytest.mark.skipif( torch.cuda.device_count() < 2, reason="needs 2 GPUs" ) -requires_gpu = pytest.mark.skipif( - torch.cuda.device_count() < 1, reason="needs a GPU" -) +requires_gpu = pytest.mark.skipif(torch.cuda.device_count() < 1, reason="needs a GPU") def _mxfp8_gemm_rel_l2(n: int, k: int, m: int, pad: bool = False) -> float: @@ -457,9 +455,7 @@ def _mxfp8_gemm_rel_l2(n: int, k: int, m: int, pad: bool = False) -> float: @requires_gpu -@pytest.mark.parametrize( - "n,k,label", [(8192, 1280, "wq_b"), (5120, 2048, "wo_b")] -) +@pytest.mark.parametrize("n,k,label", [(8192, 1280, "wq_b"), (5120, 2048, "wo_b")]) @pytest.mark.parametrize("m", [64, 192, 1024]) def test_standalone_mxfp8_gemm(n, k, label, m): """``Mxfp8GemmOp`` at DeepSeek-V4.1-Flash's two attention shapes. @@ -584,9 +580,9 @@ def test_supports_gemm_takes_both_attention_shapes(): """Neither of V4.1-Flash's attention GEMMs needs a special case.""" from mori.ops.gemm_ar import supports_gemm - assert supports_gemm(8192, 1280) is True # wq_b, ColumnParallel - assert supports_gemm(5120, 2048) is True # wo_b, RowParallel - assert supports_gemm(5120, 64) is False # K below MIN_K + assert supports_gemm(8192, 1280) is True # wq_b, ColumnParallel + assert supports_gemm(5120, 2048) is True # wo_b, RowParallel + assert supports_gemm(5120, 64) is False # K below MIN_K assert supports_gemm(5000, 2048) is False # N not a multiple of BLOCK_N diff --git a/tests/python/cco/test_mxfp8_gemv_grid.py b/tests/python/cco/test_mxfp8_gemv_grid.py index 6ac6c90de..a26dc5611 100644 --- a/tests/python/cco/test_mxfp8_gemv_grid.py +++ b/tests/python/cco/test_mxfp8_gemv_grid.py @@ -103,11 +103,28 @@ def _cases(): def _run_case(shape, m, waves, steps, rows, tokens, ksplit): """One configuration, in its own interpreter. Returns (rc, payload|text).""" p = subprocess.run( - [sys.executable, os.path.abspath(__file__), "--worker", - "--shape", shape, "-m", str(m), "--waves", str(waves), - "--steps", str(steps), "--rows", str(rows), "--tokens", str(tokens), - "--ksplit", str(ksplit)], - capture_output=True, text=True, timeout=900, + [ + sys.executable, + os.path.abspath(__file__), + "--worker", + "--shape", + shape, + "-m", + str(m), + "--waves", + str(waves), + "--steps", + str(steps), + "--rows", + str(rows), + "--tokens", + str(tokens), + "--ksplit", + str(ksplit), + ], + capture_output=True, + text=True, + timeout=900, ) for line in p.stdout.splitlines(): if line.startswith("RESULT_JSON"): @@ -117,9 +134,12 @@ def _run_case(shape, m, waves, steps, rows, tokens, ksplit): @requires_gpu @pytest.mark.parametrize( - "shape,m,waves,steps,rows,tokens,ksplit", CASES, - ids=[f"{s}-m{m}-w{w}s{st}r{r}t{t}{'k' if ks else 'n'}" - for s, m, w, st, r, t, ks in CASES], + "shape,m,waves,steps,rows,tokens,ksplit", + CASES, + ids=[ + f"{s}-m{m}-w{w}s{st}r{r}t{t}{'k' if ks else 'n'}" + for s, m, w, st, r, t, ks in CASES + ], ) def test_gemv_config(shape, m, waves, steps, rows, tokens, ksplit): """Every configuration must agree with an fp32 reference, not just the tuned one.""" @@ -149,14 +169,22 @@ def _worker(args) -> int: # row must not reach past the allocation. At M=1 there is no next row. x = (torch.randn(m, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) w = (torch.randn(n, k, generator=g, device="cuda") / 8).to(torch.float8_e4m3fn) - ex = torch.randint(120, 123, (m, k // bk), generator=g, - device="cuda", dtype=torch.int32) - ew = torch.randint(120, 123, (n // bk, k // bk), generator=g, - device="cuda", dtype=torch.int32) + ex = torch.randint( + 120, 123, (m, k // bk), generator=g, device="cuda", dtype=torch.int32 + ) + ew = torch.randint( + 120, 123, (n // bk, k // bk), generator=g, device="cuda", dtype=torch.int32 + ) gemv = compile_mxfp8_gemv( - n=n, k=k, m_max=args.tokens, waves=args.waves, steps=args.steps, - rows=args.rows, tokens=args.tokens, ksplit=bool(args.ksplit), + n=n, + k=k, + m_max=args.tokens, + waves=args.waves, + steps=args.steps, + rows=args.rows, + tokens=args.tokens, + ksplit=bool(args.ksplit), ) out = torch.zeros(m, n, device="cuda", dtype=torch.bfloat16) gemv( @@ -164,7 +192,9 @@ def _worker(args) -> int: ew.to(torch.uint8).contiguous().view(torch.int32).view(-1), x.contiguous().view(torch.int32).view(-1), ex.to(torch.uint8).contiguous().view(torch.int32).view(-1), - out.view(-1), m, n, + out.view(-1), + m, + n, stream=fx.Stream(torch.cuda.current_stream()), ) torch.cuda.synchronize() @@ -174,17 +204,28 @@ def _worker(args) -> int: ref = torch.zeros(m, n, device="cuda", dtype=torch.float32) for i in range(k // bk): ks = slice(i * bk, (i + 1) * bk) - ref += ((xf[:, ks] @ wf[:, ks].T) * sx[:, i][:, None] - * sw[:, i].repeat_interleave(bk)[None, :]) + ref += ( + (xf[:, ks] @ wf[:, ks].T) + * sx[:, i][:, None] + * sw[:, i].repeat_interleave(bk)[None, :] + ) got = out.float() - rel = (torch.linalg.vector_norm(got - ref) - / torch.linalg.vector_norm(ref)).item() - print("RESULT_JSON " + json.dumps({ - "shape": args.shape, "n": n, "k": k, "m": m, - "config": f"w{args.waves}s{args.steps}r{args.rows}t{args.tokens}" - f"{'k' if args.ksplit else 'n'}", - "rel_l2": rel, "finite": bool(got.isfinite().all()), - })) + rel = (torch.linalg.vector_norm(got - ref) / torch.linalg.vector_norm(ref)).item() + print( + "RESULT_JSON " + + json.dumps( + { + "shape": args.shape, + "n": n, + "k": k, + "m": m, + "config": f"w{args.waves}s{args.steps}r{args.rows}t{args.tokens}" + f"{'k' if args.ksplit else 'n'}", + "rel_l2": rel, + "finite": bool(got.isfinite().all()), + } + ) + ) return 0 From ffbaf23037f2d6f3738c74e6f0262abee997c295 Mon Sep 17 00:00:00 2001 From: xiangch Date: Mon, 21 Sep 2026 06:24:23 +0000 Subject: [PATCH 8/8] gemm_ar: close the four gaps from review on the new public surface All four reproduce as reported; none are in the MXFP8 numerics. **`supports_gemm` said yes to a tile the store cannot write.** The permlane guard was `BLOCK_N >= 256`, but `_PermlaneStoreC._emit` pairs *exactly* two N-tiles and asserts `n_tiles_b == 2`. So `block_n=512` passed every predicate and died several frames deeper -- and because the wide tile is only chosen once the grid is large enough, such an instance serves small batches correctly and fails when one grows. That is the worst shape a failure can take. `GemmAllReduceOp` already rejected it at its constructor; `Mxfp8GemmOp` and `supports_gemm` are new here and did not. Rather than restate the rule a third time, it is now `permlane_tile_constraint()` next to `_tile_constraints` in op.py, used by all three, with the kernel's own assert tightened to an equality as the backstop for callers that compile directly. Review suggested the assert alone would cover the predicate; it does not -- `supports_gemm` has to answer without compiling, which is why the rule has to exist as a predicate too. **The benchmarks reported `validated=true` for runs that validated nothing.** `bench_gemm.py` seeded `ref` from the first implementation's output, so that impl was its own reference: `rel_l2` identically 0 and a pass regardless of what it computed. A single-impl run -- which the docs recommend -- checked nothing, and a bug shared by every impl was invisible. `bench_gemv.py` had the same hole via `--baseline none`. Both now score against `reference_partial()`, the fp32 reference `bench_gemm_ar.py` next door already uses, built from the raw operands. The first impl's relL2 is now 1.66e-03 rather than 0, which is the fp8 floor and matches what review measured independently. Where no independent reference exists -- anything quantising inside the timed region -- `validated` is `None` and prints as `n/a`, never `True`. **`--scope kernel --impl sglang` compared two different inputs.** mori's kernel path builds its own pre-quantised A from seed 99 while SGLang's entry point is a linear and quantises the outer bf16 `x`. The baseline was scored against an unrelated result and read as failing validation -- a false accusation -- and the timings were not comparable either, only one side carrying the quantisation. Now refused at argument parsing, which is what this file's own docstring already said: linear is the scope where that comparison means anything. **`sweep.py fused-fp8` shipped 3 points that cannot run.** The preset listed `split-lsa` while appending `--gather-dtype fp8`, and the bench refuses that pair at entry -- a rejection this branch added and did not propagate here. 15 points, 3 of them unrunnable on a healthy box. Dropped from this preset only; the bf16 preset keeps `split-lsa`, where it is legal. Not addressed: the `wo_a` tuning entries review found worth 13.9-17.5%. That is a new performance claim and wants its own reproduction rather than being taken on the review's numbers inside an already-large PR. 66 tests pass; both kernel-scope paths verified against the new reference. Co-Authored-By: Claude Opus 5 (1M context) --- benchmark/cco/flydsl/gemm_ar/bench_gemm.py | 85 ++++++++++++++++++---- benchmark/cco/flydsl/gemm_ar/bench_gemv.py | 17 +++-- benchmark/cco/flydsl/gemm_ar/sweep.py | 6 +- python/mori/ops/gemm_ar/gemm.py | 9 +++ python/mori/ops/gemm_ar/kernels_fused.py | 34 ++++++--- python/mori/ops/gemm_ar/op.py | 28 +++++-- 6 files changed, 140 insertions(+), 39 deletions(-) diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemm.py b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py index 93a2bd18a..c633d087c 100644 --- a/benchmark/cco/flydsl/gemm_ar/bench_gemm.py +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemm.py @@ -58,6 +58,21 @@ sys.path.insert(0, str(Path(__file__).parent)) import timing # noqa: E402 +from bench_gemm_ar import reference_partial # noqa: E402 + +#: Why the reference is built here rather than taken from the first impl. +#: +#: Seeding `ref` from the first implementation's output makes that impl its own +#: reference: its `rel_l2` is identically 0 and it is reported `validated=true` +#: whatever it computed. A single-impl run -- which this benchmark supports and +#: the docs recommend -- then validates nothing at all, and a bug shared by +#: every impl is invisible. +#: +#: `reference_partial` is the fp32 reference `bench_gemm_ar.py` already scores +#: against, reused so the two benchmarks next to each other are equally strong. +#: Where no independent reference exists (anything that quantises inside the +#: timed region), `validated` is `None` -- not `True`. +_REF_NOTE = __doc__ MXFP8_BK = 32 SCALE_BK = 128 @@ -215,7 +230,7 @@ def mori_kernel_call(ops, n, k, m, block_n, quant): from mori.ops.gemm_ar import layout from mori.ops.gemm_ar.kernels_fused import compile_fused_gemm_scatter - a, _w, w_shuf, sa, sb = ops + a, w, w_shuf, sa, sb = ops gemm = compile_fused_gemm_scatter( layout.ArConfig(world_size=2, m=128, n=n), 0, @@ -248,7 +263,10 @@ def call(picked): ) return y - return call, [w_shuf.contiguous().view(torch.int8).view(-1)] + # An fp32 reference built from the same raw operands, not from another + # implementation's output. See `_REF_NOTE`. + ref = reference_partial(a, w, sa, sb, "blockscale") + return call, [w_shuf.contiguous().view(torch.int8).view(-1)], ref from mori.ops.gemm_ar import preshuffle_a_scale @@ -264,7 +282,10 @@ def call(picked): def call(picked): return op(a, picked[0], a_scale, ops["b_scale"])[:m] - return call, [ops["mori_w"]] + # `ea`/`w_exps` are the raw ue8m0 bytes the preshuffles were built from, so + # this reference shares no code with the kernel under test. See `_REF_NOTE`. + ref = reference_partial(a[:m], ops["w_raw"], ea[:m], ops["w_exps"], "mxfp8") + return call, [ops["mori_w"]], ref def mori_linear_call(ops, n, k, m, block_n, x_bf16): @@ -279,7 +300,10 @@ def call(picked): a_fp8, a_scale = quantize_packed(x_in) return op(a_fp8, picked[0], a_scale, ops["b_scale"])[:m] - return call, [ops["mori_w"]] + # No independent reference: the A quantisation happens inside the timed + # region and its packed scales are not the raw exponents `reference_partial` + # takes. Reported as `validated=None` rather than guessed at. + return call, [ops["mori_w"]], None def sglang_linear_call(layer, x_bf16): @@ -308,7 +332,9 @@ def call(picked): weight_bf16=(picked[1] if len(picked) > 1 else None), ) - return call, weights + # SGLang's linear quantises inside the timed region, same as + # `mori_linear_call`; no independent reference. See `_REF_NOTE`. + return call, weights, None # -------------------------------------------------------------------------- @@ -361,6 +387,20 @@ def main() -> int: if args.quant == "blockscale" and ("sglang" in impls or args.scope == "linear"): p.error("--quant blockscale is mori-only and kernel-scope only") + if args.scope == "kernel" and "sglang" in impls: + # The two sides do not share operands: `mori_kernel_call` builds its own + # pre-quantised A from seed 99, while SGLang's entry point is a *linear* + # and quantises the outer bf16 `x` itself. Scoring one against the other + # reads as SGLang failing validation -- a false accusation against the + # baseline -- and the timings are not comparable either, since only one + # side carries the quantisation. `--scope linear` is the scope in which + # the comparison means anything, which is what this file's own docstring + # already says. + p.error( + "--impl sglang needs --scope linear: in kernel scope the two sides " + "consume different operands, so neither the relL2 nor the time is a " + "comparison. Use --scope linear, or drop sglang from --impl." + ) vram_before = timing.vram_used() if args.quant == "mxfp8": @@ -398,22 +438,29 @@ def main() -> int: row["route"] = native_route_plan( m, n, k, ops["layer"].weight_bf16 is not None, False ) - call, weights = sglang_linear_call(ops["layer"], x) + call, weights, impl_ref = sglang_linear_call(ops["layer"], x) else: block_n = {"auto": None, "gemm256": 256, "gemm128": 128}[impl] row["route"] = impl if args.scope == "kernel": - call, weights = mori_kernel_call(ops, n, k, m, block_n, args.quant) + call, weights, impl_ref = mori_kernel_call( + ops, n, k, m, block_n, args.quant + ) else: - call, weights = mori_linear_call(ops, n, k, m, block_n, x) + call, weights, impl_ref = mori_linear_call(ops, n, k, m, block_n, x) got = call(weights) torch.cuda.synchronize() - if ref is None: - ref = got.float().clone() - row["rel_l2"] = rel_l2(got, ref) + # The fp32 reference from the impl's own raw operands, when it has + # one. Never seeded from another impl's output -- see `_REF_NOTE`. + if ref is None and impl_ref is not None: + ref = impl_ref.float() + row["rel_l2"] = rel_l2(got, ref) if ref is not None else None row.update(timing.cold_hot_us(call, weights, reps=args.reps)) - row["validated"] = row["rel_l2"] is None or row["rel_l2"] <= args.tol + # `None` means "nothing to check against", which is not a pass. + row["validated"] = ( + None if row["rel_l2"] is None else row["rel_l2"] <= args.tol + ) except _Unsupported as err: # The op declining a shape it documents as out of range is a # *result*, not a crash: `shared_gate_up` is N=1152, which is 4.5 @@ -428,13 +475,21 @@ def main() -> int: print(f" {impl:<10} FAILED {row['error']}", flush=True) traceback.print_exc(limit=2) else: - flag = "" if row["validated"] else " !! rel_l2 over tol" + if row["validated"] is None: + # Timed but unchecked. Distinguished from a pass in the output + # as well as in the JSON, so a run that validated nothing does + # not read like one that validated everything. + rel = " relL2 n/a (no independent reference)" + else: + rel = f" relL2 {row['rel_l2']:.2e}" + ( + "" if row["validated"] else " !! rel_l2 over tol" + ) print( f" {impl:<10} hot {row['hot_us']:8.2f} cold {row['cold_us']:8.2f}" - f" relL2 {row['rel_l2']:.2e}{flag}", + f"{rel}", flush=True, ) - if not row["validated"]: + if row["validated"] is False: failures += 1 rows.append(row) diff --git a/benchmark/cco/flydsl/gemm_ar/bench_gemv.py b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py index c50c783eb..0a25b6ecc 100644 --- a/benchmark/cco/flydsl/gemm_ar/bench_gemv.py +++ b/benchmark/cco/flydsl/gemm_ar/bench_gemv.py @@ -51,6 +51,7 @@ sys.path.insert(0, str(Path(__file__).parent)) import timing # noqa: E402 +from bench_gemm_ar import reference_partial # noqa: E402 MXFP8_BK = 32 SHAPES = {"wq_b": (8192, 1280), "wo_b": (5120, 2048)} @@ -193,12 +194,18 @@ def main() -> int: "includes_quant": False, "timing": "amortized-graph-cold-hot", } - rows, base, ref, failures = [], None, None, 0 + rows, base, failures = [], None, 0 + + # The fp32 reference comes from the raw operands, never from a baseline's + # output. With `--baseline none` there is no baseline to borrow one from, + # and treating "nothing to compare against" as a pass is how a benchmark + # reports success for a kernel it never checked. `reference_partial` is the + # same reference `bench_gemm_ar.py` scores against. + ref = reference_partial(x[:m], w, ex[:m], ew, "mxfp8").float() if args.baseline == "sglang": call, wt = sglang_call(n, k, m, x, w, ex, ew, out) base = timing.cold_hot_us(call, [wt], reps=args.reps) - ref = out[:m].float().clone() print( f"{args.shape} M={m} sglang hot {base['hot_us']:6.2f} " f"cold {base['cold_us']:6.2f}", @@ -249,8 +256,8 @@ def main() -> int: ) failures += 1 continue - ok = rel is None or rel <= args.tol - failures += 0 if ok else 1 + ok = None if rel is None else rel <= args.tol + failures += 1 if ok is False else 0 vs = ( f" {(res['cold_us'] / base['cold_us'] - 1) * 100:+6.1f}%" if base @@ -259,7 +266,7 @@ def main() -> int: rel_s = " n/a " if rel is None else f" relL2 {rel:.2e}" print( f" {key_of(cfg):<12} hot {res['hot_us']:6.2f} cold {res['cold_us']:6.2f}" - f"{vs}{rel_s}{'' if ok else ' !! over tol'}", + f"{vs}{rel_s}{'' if ok is not False else ' !! over tol'}", flush=True, ) rows.append( diff --git a/benchmark/cco/flydsl/gemm_ar/sweep.py b/benchmark/cco/flydsl/gemm_ar/sweep.py index cfb0b6b8f..17ce93b92 100644 --- a/benchmark/cco/flydsl/gemm_ar/sweep.py +++ b/benchmark/cco/flydsl/gemm_ar/sweep.py @@ -139,12 +139,14 @@ def _grid(**axes): ), "fused-fp8": dict( doc="the same mode matrix on the winning fp8 wire, which is where the " - "fusion is actually deployed", + "fusion is actually deployed. `split-lsa` is absent on purpose: it has " + "no fp8 gather leg (`build_lsa_ar` takes no `gather_dtype`), so the " + "bench refuses the pair at entry and those points can never run", script="bench_gemm_ar.py", world=4, grid=_grid( m=[4096, 8192, 16384], - mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"], + mode=["gemm-only", "split-sdma", "fused-sdma", "fused-lsa"], ), args=[ "-n", diff --git a/python/mori/ops/gemm_ar/gemm.py b/python/mori/ops/gemm_ar/gemm.py index 336a7c197..6bd415a91 100644 --- a/python/mori/ops/gemm_ar/gemm.py +++ b/python/mori/ops/gemm_ar/gemm.py @@ -61,6 +61,7 @@ _flatten_mxfp8_a_scale, _PinnedLaunch, _tile_constraints, + permlane_tile_constraint, ) #: Rows one packed A-scale group spans: four M tiles of sixteen rows. M is @@ -149,6 +150,11 @@ def supports_gemm(n: int, k: int, *, block_n: int = DEFAULT_BLOCK_N) -> bool: """ return ( _tile_constraints(MXFP8_BLOCK_M, block_n) is None + # The wide tile compiles with permlane, which is an equality on + # block_n rather than a floor. Without this, block_n=512 answers + # True here and dies in the store once the grid is wide enough to + # pick that tile -- i.e. only after the batch grows. + and permlane_tile_constraint(block_n) is None and _gemm_shape_constraint(n, k, block_n) is None # Small M runs the narrow tile, so N has to divide by that too. and _gemm_shape_constraint(n, k, NARROW_BLOCK_N) is None @@ -181,6 +187,9 @@ def __init__( block_n: int = DEFAULT_BLOCK_N, ): why = _tile_constraints(MXFP8_BLOCK_M, block_n) + if why is not None: + raise ValueError(f"unsupported tile: {why}") + why = permlane_tile_constraint(block_n) if why is not None: raise ValueError(f"unsupported tile: {why}") why = _gemm_shape_constraint(n, k, block_n) diff --git a/python/mori/ops/gemm_ar/kernels_fused.py b/python/mori/ops/gemm_ar/kernels_fused.py index d3233bf44..32c93fb9a 100644 --- a/python/mori/ops/gemm_ar/kernels_fused.py +++ b/python/mori/ops/gemm_ar/kernels_fused.py @@ -1132,18 +1132,28 @@ def compile_fused_gemm_scatter( if n_stripe < N // BLOCK_N and not rotated: raise ValueError("a striped tile order needs --tile-order rotated") - # BLOCK_N's floor belongs to the *store*, not the mainloop: the mainloop - # builds N_TILES_B = BLOCK_N//128 accumulators and is happy with one, while - # _LaneTransposeStoreC's permlane mapping pairs exactly two N-tiles. So 128 - # is legal without permlane and not with it, which is what the store's own - # assert says a few frames deeper and less usefully. - n_floor = 256 if permlane else 128 - assert ( - BLOCK_M >= 128 and BLOCK_N >= n_floor - ), f"BLOCK_N={BLOCK_N} is below {n_floor}" + ( - " (permlane's store pairs two N-tiles)" if permlane else "" - ) - assert BLOCK_M % 128 == 0 and BLOCK_N % 128 == 0 + # BLOCK_N's constraint belongs to the *store*, not the mainloop: the mainloop + # builds N_TILES_B = BLOCK_N//128 accumulators and is happy with any count, + # while `_PermlaneStoreC._emit` pairs *exactly* two N-tiles and asserts + # `n_tiles_b == 2`. So permlane needs BLOCK_N == 256 -- an equality, not a + # floor. + # + # It was written as `>= 256`, which let BLOCK_N=512 through here and into an + # assert several frames deeper. That is the worst shape for the failure to + # take: the wide tile is only chosen once the grid is large enough, so such + # an instance serves small batches correctly for as long as they stay small + # and dies when one grows. `supports_gemm` asks this same question and was + # wrong in the same way; checking it here fixes that caller too, rather than + # restating the rule at every entry point. + assert BLOCK_M >= 128 and BLOCK_M % 128 == 0, f"BLOCK_M={BLOCK_M}" + if permlane: + assert BLOCK_N == 256, ( + f"BLOCK_N={BLOCK_N}: permlane's store pairs exactly two N-tiles, " + "so it needs BLOCK_N == 256 (pass permlane=False for other widths)" + ) + else: + assert BLOCK_N >= 128, f"BLOCK_N={BLOCK_N} is below 128" + assert BLOCK_N % 128 == 0, f"BLOCK_N={BLOCK_N}" assert K % BLOCK_K == 0 if N % BLOCK_N: raise ValueError( diff --git a/python/mori/ops/gemm_ar/op.py b/python/mori/ops/gemm_ar/op.py index 27c1070ad..044eae7bd 100644 --- a/python/mori/ops/gemm_ar/op.py +++ b/python/mori/ops/gemm_ar/op.py @@ -175,6 +175,26 @@ def _tile_constraints( return None +def permlane_tile_constraint(block_n: int) -> str | None: + """Why ``block_n`` cannot be stored by the permlane epilogue, or None. + + An equality, not a floor: ``_PermlaneStoreC._emit`` pairs *exactly* two + N-tiles and asserts ``n_tiles_b == 2``. The mainloop is happy with any + count, so a wider tile compiles right up to the store. + + Lives here, next to ``_tile_constraints``, because every predicate that + answers "can this shape be served" has to apply it *without* compiling -- + `supports_gemm` cannot reach the kernel's own assert. The assert stays as + the backstop for callers that build a kernel directly. + """ + if block_n != DEFAULT_BLOCK_N: + return ( + f"block_n={block_n} must be {DEFAULT_BLOCK_N}: the permlane " + f"epilogue's lane transpose pairs exactly two N-tiles" + ) + return None + + def _gemm_constraints( n: int, k: int, block_n: int, quant: str = "blockscale" ) -> str | None: @@ -511,11 +531,9 @@ def __init__( f"gather_transport={gather_transport!r} is not one of " f"{sorted(GATHER_TRANSPORTS)}" ) - if block_n != DEFAULT_BLOCK_N: - raise ValueError( - f"block_n must be {DEFAULT_BLOCK_N}: the op always compiles with " - f"permlane, whose lane transpose is written for that width" - ) + why = permlane_tile_constraint(block_n) + if why is not None: + raise ValueError(why) # The tile has to be legal before the padding arithmetic runs: both # padded_m and default_max_shapes divide by world_size * block_m. why = _tile_constraints(block_m, block_n)