diff --git a/kernels/conv/conv3d_implicit.py b/kernels/conv/conv3d_implicit.py index 1bb9abbb0..6b8f3c01a 100644 --- a/kernels/conv/conv3d_implicit.py +++ b/kernels/conv/conv3d_implicit.py @@ -15,9 +15,9 @@ import flydsl.compiler as flyc import flydsl.expr as fx -from flydsl._mlir import ir from flydsl._mlir.dialects import llvm from flydsl.expr import arith, const_expr, range_constexpr, rocdl +from flydsl.expr.rocdl.universal import make_buffer_ptr from flydsl.expr.typing import T from kernels.common import buffer_ops from kernels.common.mem_ops import buffer_atomic_add @@ -246,10 +246,16 @@ def compile_conv3d_implicit( @flyc.kernel(known_block_size=[BLOCK_THREADS, 1, 1]) def conv3d_implicit_kernel(y: fx.Tensor, x: fx.Tensor, weight: fx.Tensor, bias: fx.Tensor): - w_rsrc = buffer_ops.create_buffer_resource(weight, num_records_bytes=W_BYTES) - if const_expr(not BIG_IN): - x_rsrc = buffer_ops.create_buffer_resource(x, num_records_bytes=X_BYTES) y_rsrc = buffer_ops.create_buffer_resource(y) + # Buffer tensors for the im2col gather, flattened to a 1-D element view so the + # per-thread flat gather offset indexes elements (multi-dim views would not). + w_buf0 = fx.rocdl.make_buffer_tensor(weight, max_size=False) + w_buf = fx.Tensor(fx.make_view(fx.get_iter(w_buf0), fx.make_layout(k * crs, 1))) + w_div = fx.logical_divide(w_buf, fx.make_layout(1, 1)) + if const_expr(not BIG_IN): + x_buf0 = fx.rocdl.make_buffer_tensor(x, max_size=False) + x_buf = fx.Tensor(fx.make_view(fx.get_iter(x_buf0), fx.make_layout(n * c * d * h * w, 1))) + x_div = fx.logical_divide(x_buf, fx.make_layout(1, 1)) if const_expr(has_bias): bias_rsrc = buffer_ops.create_buffer_resource(bias) @@ -276,6 +282,17 @@ def conv3d_implicit_kernel(y: fx.Tensor, x: fx.Tensor, weight: fx.Tensor, bias: else: k_off = 0 + # BIG_IN (>2GB): flat buffer tensor from a rebased address with explicit ~2GB + # num_records (same mechanism as x_div, replacing create_buffer_resource_from_addr). + GXPtrTy = fx.PointerType.get(elem_ty.ir_type, 1, BF16_BYTES) if const_expr(BIG_IN) else None + + def _x_div_from_addr(addr_i64): + gptr = fx.inttoptr(GXPtrTy, addr_i64) + buf_ptr = make_buffer_ptr(gptr, num_records_bytes=BIG_IN_NR) + # 1-D element view so the flat im2col gather offset indexes elements. + buf = fx.Tensor(fx.make_view(buf_ptr, fx.make_layout(BIG_IN_NR // BF16_BYTES, 1))) + return fx.logical_divide(buf, fx.make_layout(1, 1)) + if const_expr(BIG_IN_N1): nbase = m_offset // dhw ot_base0 = (m_offset % dhw) // hw_o @@ -283,7 +300,7 @@ def conv3d_implicit_kernel(y: fx.Tensor, x: fx.Tensor, weight: fx.Tensor, bias: base_t = arith.select(base_t < fx.Index(0), fx.Index(0), base_t) x_base_elem = ((nbase * fx.Index(d) + base_t) * fx.Index(h) + fx.Index(0)) * fx.Index(w) * fx.Index(c) x_addr = fx.Int64(buffer_ops.extract_base_index(x)) + fx.Int64(x_base_elem) * fx.Int64(2) - x_rsrc = buffer_ops.create_buffer_resource_from_addr(x_addr, num_records_bytes=BIG_IN_NR) + x_div_big = _x_div_from_addr(x_addr) if const_expr(BIG_IN_NM): x_base_addr = fx.Int64(buffer_ops.extract_base_index(x)) @@ -420,27 +437,24 @@ def _b_addr(i, k_base): return g_off, col_valid # ---- global -> LDS DMA copy, masking via OOB routing ---- - DMA_BYTES = LDG_VEC * BF16_BYTES # 16 - OOB_ELEM = fx.Int32(OOB_SENTINEL_ELEM) - - def _lds_dma_ptr(lds_array, stage_tile, i): - off_elems = fx.Index(stage_tile) + (fx.Index(tid) + fx.Index(i * BLOCK_THREADS)) * fx.Index(LDG_VEC) - base_bytes = off_elems * fx.Index(BF16_BYTES) - addr = fx.Int64(fx.ptrtoint(lds_array.ptr)) + fx.Int64(base_bytes) - addr = rocdl.readfirstlane(T.i64, arith.index_cast(T.i64, addr.ir_value())) - return llvm.inttoptr(ir.Type.parse("!llvm.ptr<3>"), addr) - - def _dma_to_lds(rsrc, lds_ptr, voff_elem): - voff_b = (voff_elem * fx.Int32(BF16_BYTES)).ir_value() - rocdl.raw_ptr_buffer_load_lds( - rsrc, - lds_ptr, - arith.constant(DMA_BYTES, type=T.i32), - voff_b, - arith.constant(0, type=T.i32), - arith.constant(0, type=T.i32), - arith.constant(0, type=T.i32), - ) + # OOB sentinel element indices: a gather past num_records makes the buffer + # hardware return 0, reproducing the padding/halo zeroing. + X_ELEMS = fx.Int32(n * c * d * h * w) + W_ELEMS = fx.Int32(k * c * kt * kh * kw) + BIG_OOB_ELEM = fx.Int32(BIG_IN_NR // BF16_BYTES) + + # global->LDS DMA (BufferCopyLDS128b): gather offset as a flat element index; + # OOB padding routes past-end. Used for both non-BIG and BIG_IN gathers. + g2s_atom = fx.make_copy_atom(fx.rocdl.BufferCopyLDS128b(), 128) + LdsPtrTy = fx.PointerType.get(elem_ty.ir_type, 2, 512) + + def _copy_g2s(src_div, lds_array, stage_tile, i, src_elem): + off_elems = fx.Int32(stage_tile) + (fx.Int32(tid) + fx.Int32(i * BLOCK_THREADS)) * fx.Int32(LDG_VEC) + lds_byte_addr = fx.Int32(fx.ptrtoint(lds_array.ptr)) + off_elems * fx.Int32(BF16_BYTES) + lds_ptr = fx.inttoptr(LdsPtrTy, lds_byte_addr) + dst = fx.make_view(lds_ptr, fx.make_layout(1, 1)) + src = fx.slice(src_div, (None, fx.Int32(src_elem))) + fx.copy(g2s_atom, src, dst) def _load_a(stage, k_base): kbase_i = fx.Index(k_base) @@ -451,26 +465,32 @@ def _load_a(stage, k_base): stage_tile = fx.Index(stage) * TILE_M * TILE_K for i in range_constexpr(LDG_A_COUNT): if const_expr(BIG_IN_NM): + # Rebase the buffer tensor per load to the sample base. addr_ret = _a_addr(i, kbase_i, cc_base, ckk_base) g_off_i, valid, n_idx_i = addr_ret sample_addr = x_base_addr + fx.Int64(n_idx_i) * fx.Int64(X_SAMPLE_BYTES) - x_rsrc_i = buffer_ops.create_buffer_resource_from_addr(sample_addr, num_records_bytes=BIG_IN_NR) - voff = fx.Int32(arith.select(valid, g_off_i, OOB_ELEM)) - _dma_to_lds(x_rsrc_i, _lds_dma_ptr(a_lds, stage_tile, i), voff) + x_div_i = _x_div_from_addr(sample_addr) + voff = fx.Int32(arith.select(valid, g_off_i, BIG_OOB_ELEM)) + _copy_g2s(x_div_i, a_lds, stage_tile, i, voff) + elif const_expr(BIG_IN): + # BIG_IN_N1: rebased buffer tensor (built once above). + g_off_i, valid = _a_addr(i, kbase_i, cc_base, ckk_base) + voff = fx.Int32(arith.select(valid, g_off_i, BIG_OOB_ELEM)) + _copy_g2s(x_div_big, a_lds, stage_tile, i, voff) else: g_off_i, valid = _a_addr(i, kbase_i, cc_base, ckk_base) - voff = fx.Int32(arith.select(valid, g_off_i, OOB_ELEM)) - _dma_to_lds(x_rsrc, _lds_dma_ptr(a_lds, stage_tile, i), voff) + voff = fx.Int32(arith.select(valid, g_off_i, X_ELEMS)) + _copy_g2s(x_div, a_lds, stage_tile, i, voff) def _load_b(stage, k_base): stage_tile = fx.Index(stage) * TILE_N * TILE_K for i in range_constexpr(LDG_B_COUNT): g_off, col_valid = _b_addr(i, k_base) if const_expr(n_tail): - voff = fx.Int32(arith.select(col_valid, g_off, OOB_ELEM)) + voff = fx.Int32(arith.select(col_valid, g_off, W_ELEMS)) else: voff = g_off - _dma_to_lds(w_rsrc, _lds_dma_ptr(b_lds, stage_tile, i), voff) + _copy_g2s(w_div, b_lds, stage_tile, i, voff) # ---- single-vec ds_read (LDS -> register), indexed by per-wave MFMA row ---- def read_a_vec(stage, mi): diff --git a/kernels/conv/conv3d_implicit_fp8.py b/kernels/conv/conv3d_implicit_fp8.py index f31d94808..5a2414f56 100644 --- a/kernels/conv/conv3d_implicit_fp8.py +++ b/kernels/conv/conv3d_implicit_fp8.py @@ -16,7 +16,7 @@ from flydsl.expr.typing import T from kernels.common import buffer_ops from kernels.common.mem_ops import buffer_atomic_add -from kernels.gemm.fp8_gemm_utils import Mfma16x16x128, make_fp8_buffer_tensor, pack_i32x4_i32x8 +from kernels.gemm.fp8_gemm_utils import make_fp8_buffer_tensor, pack_i32x4_i32x8 TILE_M = 128 TILE_N = 128 @@ -254,13 +254,15 @@ def conv3d_8wave_fp8_kernel(y: fx.Tensor, x: fx.Tensor, weight: fx.Tensor, bias: c_m_vec = lane_div_16 * MFMA_C_VALUES c_n = lane_mod_16 - mfma = Mfma16x16x128(QM_STEPS, QN_STEPS) - acc00 = [mfma.zero_value for _ in range_constexpr(N_SUB)] - acc01 = [mfma.zero_value for _ in range_constexpr(N_SUB)] - acc10 = [mfma.zero_value for _ in range_constexpr(N_SUB)] - acc11 = [mfma.zero_value for _ in range_constexpr(N_SUB)] - + # 16x16x128 FP8 MFMA via the layout API, built in-kernel (a tiled_mma + # kernel-arg compiles warm but fails cold-compile). + mma_atom = fx.make_mma_atom(fx.rocdl.cdna4.MFMA_Scale(16, 16, 128, elem_ty)) Vec = fx.Vector + mfma_zero = Vec.filled(MFMA_C_VALUES, 0.0, fx.Float32) + acc00 = [mfma_zero for _ in range_constexpr(N_SUB)] + acc01 = [mfma_zero for _ in range_constexpr(N_SUB)] + acc10 = [mfma_zero for _ in range_constexpr(N_SUB)] + acc11 = [mfma_zero for _ in range_constexpr(N_SUB)] class Vec16U8Ty: ir_type = Vec.make_type(16, fx.Uint8) @@ -358,8 +360,19 @@ def read_b_vec(stage, n_half, wn): def setprio(level): llvm.InlineAsmOp(None, [], f"s_setprio {level}", "", has_side_effects=True) + def _do_mma(a, b, c): + # split-16@64 packed i32x8 fragments feed the MFMA_Scale atom via fx.gemm. + a_frag = fx.make_rmem_tensor(8, fx.Int32) + a_frag.store(Vec(a)) + b_frag = fx.make_rmem_tensor(8, fx.Int32) + b_frag.store(Vec(b)) + c_frag = fx.make_rmem_tensor(MFMA_C_VALUES, fx.Float32) + c_frag.store(Vec(c)) + fx.gemm(mma_atom, c_frag, a_frag, b_frag, c_frag) + return c_frag.load().ir_value() + def mfma_one(a, b, c_acc): - out = mfma._do_mma(a, b, c_acc) + out = _do_mma(a, b, c_acc) fx.rocdl.sched_mfma(1) return out diff --git a/kernels/gemm/fp8_gemm_4wave.py b/kernels/gemm/fp8_gemm_4wave.py index 6c830543d..af1cf6eeb 100644 --- a/kernels/gemm/fp8_gemm_4wave.py +++ b/kernels/gemm/fp8_gemm_4wave.py @@ -58,6 +58,76 @@ def _do_mma(self, a, b, c): ) +def _make_swz_layout(): + """Composed layout carrying the XOR bank-avoidance swizzle. + + ``swizzle_128`` (BLOCK_K == 128) was verified byte-identical to + ``SwizzleType.get(3, 4, 4)`` and to ``crd2idx`` over this layout, so offsets + computed here match ``compute_global_swizzle`` / ``_compute_lds_swizzle`` / + ``S2RLoader`` exactly and the emitted addressing is unchanged. + """ + return fx.make_composed_layout( + fx.static(fx.SwizzleType.get(3, 4, 4)), + fx.make_ordered_layout((128, 128), (1, 0)), + ) + + +def _layout_global_swizzle(lane_id, wave_id, K, n_rounds, block_dim_x): + """Layout-API form of ``compute_global_swizzle(preshuffled=False)``. + + swizzle_128 keeps r == row (the XOR only permutes columns), so the global + element is ``row * K + swizzled_col``. Reads the swizzled column with + ``crd2idx`` on concrete coords (``get_scalar`` -> compile-time constant). + """ + swz_layout = _make_swz_layout() + n_waves = block_dim_x // 64 + offsets = [] + for rnd in range_constexpr(n_rounds): + row = lane_id // 8 + wave_id * 8 + rnd * (n_waves * 8) + col = (lane_id % 8) * 16 + swz_col = fx.get_scalar(fx.crd2idx((row % 16, col), swz_layout)) % 128 + offsets.append(row * K + swz_col) + return offsets + + +class LayoutS2R: + """Shared->register reader with the XOR swizzle carried in a composed layout. + + Mirrors ``S2RLoader.load(preshuffled=False)`` but computes the swizzled LDS + byte offset with ``crd2idx`` over ``_make_swz_layout`` instead of the + hand-written ``swizzle_128`` math. The split-16@64 ``pack_i32x4_i32x8`` + packing is kept -- it is the MFMA operand ABI. Preshuffled B stays on + ``S2RLoader`` (a different affine, non-XOR LDS map). + """ + + def __init__(self, wave_idx, n_tiles): + self.lane_id = fx.thread_idx.x % 64 + self.wave_idx = wave_idx + self.n_tiles = n_tiles + self.swz_layout = _make_swz_layout() + + def _vec_load_16xf8(self, lds_src, offset): + ptr_off = fx.add_offset(lds_src.ptr, fx.make_int_tuple(offset)) + i8_iter = fx.recast_iter(fx.Uint8, ptr_off) + return fx.make_view(i8_iter, fx.make_layout(16, 1)).load() + + def load(self, lds_src, preshuffled=False): + assert not preshuffled, "LayoutS2R only handles the non-preshuffled (XOR) LDS layout" + frag = [] + for i in range_constexpr(self.n_tiles): + halves = [] + row = self.wave_idx * (self.n_tiles * 16) + i * 16 + self.lane_id % 16 + for step in range_constexpr(2): + col = (self.lane_id // 16) * 16 + step * 64 + offset = fx.get_scalar(fx.crd2idx((row, col), self.swz_layout)) + halves.append(self._vec_load_16xf8(lds_src, offset).bitcast(fx.Int32)) + frag.append(pack_i32x4_i32x8(halves[0], halves[1])) + return frag + + def load_one(self, lds_src, lds_offset): + return self._vec_load_16xf8(lds_src, lds_offset).bitcast(fx.Int32) + + def _min(a, b): return arith.select(a < b, a, b) @@ -170,6 +240,8 @@ def kernel_gemm( gb_div = fx.logical_divide(gB, fx.make_layout(1, 1)) def _compute_lds_swizzle(s2r, preshuffled=False): + # Manual swizzle_128 XOR kept (not crd2idx): the interleaved AGPR MMA path + # is register-schedule sensitive; crd2idx regresses it ~1.3%. lds_swz = [] for row_offset in range_constexpr(s2r.n_tiles): row = s2r.wave_idx * (s2r.n_tiles * 16) + row_offset * 16 + lane_id % 16 @@ -309,13 +381,22 @@ def _compute_block( c10_frag = [mfma.zero_value] * N_ACCUMS c11_frag = [mfma.zero_value] * N_ACCUMS - gl_off_a = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=False) - gl_off_b = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=b_preshuffled) + gl_off_a = _layout_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, fx.block_dim.x) + if const_expr(b_preshuffled): + gl_off_b = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=True) + else: + gl_off_b = _layout_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, fx.block_dim.x) a_g2s = G2SLoader(ga_div, gl_off_a, N_TILES_A, F8_IR_t, wave_id) b_g2s = G2SLoader(gb_div, gl_off_b, N_TILES_B, F8_IR_t, wave_id) - a_s2r = S2RLoader(wave_i, N_TILES_A) - b_s2r = S2RLoader(wave_j, N_TILES_B) + # crd2idx s2r on the non-interleaved path only; interleaved keeps S2RLoader + # (AGPR accumulator is register-schedule sensitive, ~1.3%). + if const_expr(_use_interleaved_block): + a_s2r = S2RLoader(wave_i, N_TILES_A) + b_s2r = S2RLoader(wave_j, N_TILES_B) + else: + a_s2r = LayoutS2R(wave_i, N_TILES_A) + b_s2r = S2RLoader(wave_j, N_TILES_B) if b_preshuffled else LayoutS2R(wave_j, N_TILES_B) store_c = StoreC(A_scale, B_scale, C, c_m, c_n, mfma.idx, N_TILES_A, N_TILES_B) # Prologue: 8-buffer LDS pipeline pre-fill. diff --git a/kernels/gemm/fp8_gemm_8wave.py b/kernels/gemm/fp8_gemm_8wave.py index c57c178c3..b90567390 100644 --- a/kernels/gemm/fp8_gemm_8wave.py +++ b/kernels/gemm/fp8_gemm_8wave.py @@ -9,10 +9,10 @@ import flydsl.compiler as flyc import flydsl.expr as fx -from flydsl.expr import range_constexpr, rocdl +from flydsl.expr import const_expr, range_constexpr, rocdl +from flydsl.expr.typing import Vector as Vec from kernels.gemm.fp8_gemm_utils import ( G2SLoader, - Mfma16x16x128, S2RLoader, StoreC, ceildiv, @@ -23,6 +23,128 @@ ) +def _layout_global_swizzle(lane_id, wave_id, K, n_rounds, block_dim_x): + """Layout-API form of ``compute_global_swizzle(preshuffled=False)``. + + Carries the XOR bank-avoidance swizzle in a + ``make_composed_layout(SwizzleType.get(3, 4, 4), ...)`` and reads the + swizzled column back with ``crd2idx`` on concrete (compile-time) coords + (``get_scalar`` -> a constant), instead of the hand-written ``swizzle_128`` + bit math. ``swizzle_128`` was verified equal to ``SwizzleType.get(3, 4, 4)`` + and this offset list is byte-identical to the helper's, so the emitted + ``buffer_load_lds`` addressing is unchanged. + """ + swz_layout = fx.make_composed_layout( + fx.static(fx.SwizzleType.get(3, 4, 4)), + fx.make_ordered_layout((16, 128), (1, 0)), + ) + n_waves = block_dim_x // 64 + offsets = [] + for rnd in range_constexpr(n_rounds): + row = lane_id // 8 + wave_id * 8 + rnd * (n_waves * 8) + col = (lane_id % 8) * 16 + # swizzle_128 permutes only columns (r == row), periodic in 16 rows. + swz_col = fx.get_scalar(fx.crd2idx((row % 16, col), swz_layout)) % 128 + offsets.append(row * K + swz_col) + return offsets + + +class LayoutS2R: + """Shared->register reader with the XOR swizzle carried in a composed layout. + + Mirrors ``S2RLoader.load(preshuffled=False)`` but computes the swizzled LDS + byte offset with ``crd2idx`` over + ``make_composed_layout(SwizzleType.get(3, 4, 4), ...)`` (``get_scalar`` -> + compile-time constant) instead of the hand-written ``swizzle_128`` math. + ``swizzle_128`` was verified equal to ``SwizzleType.get(3, 4, 4)`` and the + resulting offsets are byte-identical, so the emitted ``ds_read`` addressing + is unchanged. The split-16@64 i32x8 packing (``pack_i32x4_i32x8``) is kept + as-is -- it is the MFMA operand ABI. Preshuffled B stays on ``S2RLoader`` + (a different affine, non-XOR LDS map). + """ + + def __init__(self, wave_idx, n_tiles): + self.lane_id = fx.thread_idx.x % 64 + self.wave_idx = wave_idx + self.n_tiles = n_tiles + self.swz_layout = fx.make_composed_layout( + fx.static(fx.SwizzleType.get(3, 4, 4)), + fx.make_ordered_layout((128, 128), (1, 0)), + ) + + def _vec_load_16xf8(self, lds_src, offset): + ptr_off = fx.add_offset(lds_src.ptr, fx.make_int_tuple(offset)) + i8_iter = fx.recast_iter(fx.Uint8, ptr_off) + return fx.make_view(i8_iter, fx.make_layout(16, 1)).load() + + def load(self, lds_src, preshuffled=False): + assert not preshuffled, "LayoutS2R only handles the non-preshuffled (XOR) LDS layout" + frag = [] + for i in range_constexpr(self.n_tiles): + halves = [] + row = self.wave_idx * (self.n_tiles * 16) + i * 16 + self.lane_id % 16 + for step in range_constexpr(2): + col = (self.lane_id // 16) * 16 + step * 64 + offset = fx.get_scalar(fx.crd2idx((row, col), self.swz_layout)) + v = self._vec_load_16xf8(lds_src, offset) + halves.append(v.bitcast(fx.Int32)) + frag.append(_pack_i32x4_i32x8(halves[0], halves[1])) + return frag + + +def _pack_i32x4_i32x8(lo, hi): + return lo.shuffle(hi, list(range(8))) + + +class TiledMmaDriver: + """MMA driver in the example-04 layout-API idiom. + + Replaces ``Mfma16x16x128``'s bare-atom ``fx.gemm(atom, ...)`` calls with + ``fx.gemm(tiled_mma, ...)`` over a ``make_tiled_mma``. The 16x16x128 fp8 + operands stay the split-16@64-packed i32x8 / f32x4 register fragments + produced by ``S2RLoader`` and consumed by ``StoreC``; only the MMA + construction moves to the layout API. + """ + + def __init__(self, tiled_mma, n_tiles_a, n_tiles_b): + self.tiled_mma = tiled_mma + self.zero_value = Vec.filled(4, 0.0, fx.Float32) + self.n_tiles_a = n_tiles_a + self.n_tiles_b = n_tiles_b + + def idx(self, i, j): + return i * self.n_tiles_b + j + + def _make_operand_frag(self, value): + frag = fx.make_rmem_tensor(8, fx.Int32) + frag.store(Vec(value)) + return frag + + def _make_accum_frag(self, value): + frag = fx.make_rmem_tensor(4, fx.Float32) + frag.store(Vec(value)) + return frag + + def call(self, a, b, c, *, set_prio=True): + 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 + + a_frags = [self._make_operand_frag(a[idx]) for idx in range_constexpr(self.n_tiles_a)] + b_frags = [self._make_operand_frag(b[idx]) for idx in range_constexpr(self.n_tiles_b)] + c_frags = [self._make_accum_frag(c[idx]) for idx in range_constexpr(self.n_tiles_a * self.n_tiles_b)] + if const_expr(set_prio): + rocdl.s_setprio(1) + 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.tiled_mma, cf, a_frags[i], b_frags[j], cf) + if const_expr(set_prio): + rocdl.s_setprio(0) + rocdl.s_barrier() + return [c_frags[idx].load().ir_value() for idx in range_constexpr(self.n_tiles_a * self.n_tiles_b)] + + def compile_fp8_gemm_8w(*, K: int, BLOCK_M: int = 256, BLOCK_N: int = 256, b_preshuffled: bool = False): BLOCK_K = 128 @@ -70,6 +192,13 @@ def kernel_gemm( ): F8_IR_t = fx.Float8E4M3FN.ir_type + # Single 16x16x128 fp8 atom, built in-kernel (raw i32x8 operands need the + # concrete tiled_mma; a tiled_mma kernel-arg fails cold-compile). + tiled_mma = fx.make_tiled_mma( + fx.make_mma_atom(fx.rocdl.cdna4.MFMA_Scale(16, 16, 128, fx.Float8E4M3FN)), + fx.make_layout((1, 1, 1), (0, 0, 0)), + ) + n_blocks = ceildiv(c_n, BLOCK_N) lds = fx.SharedAllocator().allocate(SharedStorage).peek() @@ -99,15 +228,21 @@ def kernel_gemm( a_div = fx.logical_divide(gA, fx.make_layout(1, 1)) b_div = fx.logical_divide(gB, fx.make_layout(1, 1)) - gl_off_a = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=False) - gl_off_b = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=b_preshuffled) + gl_off_a = _layout_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, fx.block_dim.x) + if const_expr(b_preshuffled): + gl_off_b = compute_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, preshuffled=True) + else: + gl_off_b = _layout_global_swizzle(lane_id, wave_id, K, N_LDS_ROUNDS, fx.block_dim.x) - mfma = Mfma16x16x128(N_TILES_A, N_TILES_B) + mfma = TiledMmaDriver(tiled_mma, N_TILES_A, N_TILES_B) a_g2s = G2SLoader(a_div, gl_off_a, N_LDS_STEPS_A, F8_IR_t, wave_id) b_g2s = G2SLoader(b_div, gl_off_b, N_LDS_STEPS_B, F8_IR_t, wave_id) - a_s2r = S2RLoader(wave_m, N_TILES_A) - b_s2r = S2RLoader(wave_n, N_TILES_B) + a_s2r = LayoutS2R(wave_m, N_TILES_A) + if const_expr(b_preshuffled): + b_s2r = S2RLoader(wave_n, N_TILES_B) + else: + b_s2r = LayoutS2R(wave_n, N_TILES_B) store_c = StoreC(A_scale, B_scale, C, c_m, c_n, mfma.idx, N_TILES_A, N_TILES_B) # 2x2 config of 4x2 (instead of 4x4 in 4wave) 16x16 sub-tiles diff --git a/kernels/gemm/preshuffle_gemm.py b/kernels/gemm/preshuffle_gemm.py index 41e35f83a..67b9a20fd 100644 --- a/kernels/gemm/preshuffle_gemm.py +++ b/kernels/gemm/preshuffle_gemm.py @@ -13,7 +13,6 @@ from flydsl.expr.typing import BFloat16, Float8E4M3FN, Float8E4M3FNUZ, Float16, Float32, Int8, Int32, T from flydsl.expr.typing import Vector as Vec from flydsl.runtime.device import get_rocm_arch -from kernels.common import buffer_ops from kernels.common.mma.mfma_preshuffle_pipeline import xcd_remap_bx_by # (dsrd_preload, dvmem_preload) per (tile_m, tile_n, tile_k). @@ -563,44 +562,50 @@ def two_tiles(k_base): lane_div_16 = lane_id // 16 lane_mod_16 = lane_id % 16 + # Epilogue scalar/vec4 gathers via buffer_copy atoms over a make_buffer_tensor + # (element-index addressing; same OOB-checked descriptor as the legacy buffer_load). + epi_copy_32b = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), Float32) + epi_copy_128b = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), Float32) + epi_copy_16b = fx.make_copy_atom(fx.rocdl.BufferCopy16b(), out_elem_cls) + def load_epi_operands(): s_a = s_b = bias = None if const_expr(is_8bit): # Per-row(scale_a) × per-col(scale_b) scaling, applied in the epilogue. - scale_b_rsrc = buffer_ops.create_buffer_resource(arg_scale_b, max_size=True) - s_b = [ - buffer_ops.buffer_load( - scale_b_rsrc, - fx.Int32(by_n + (ni * num_waves + wave_id) * 16 + lane_mod_16), - vec_width=1, - dtype=T.f32, + sb_buf = fx.logical_divide( + fx.rocdl.make_buffer_tensor(arg_scale_b, max_size=True), fx.make_layout(1, 1) + ) + s_b = [] + for ni in range_constexpr(num_acc_n): + f = fx.make_rmem_tensor(1, Float32) + fx.copy_atom_call( + epi_copy_32b, + sb_buf[None, fx.Int32(by_n + (ni * num_waves + wave_id) * 16 + lane_mod_16)], + f, ) - for ni in range_constexpr(num_acc_n) - ] - scale_a_rsrc = buffer_ops.create_buffer_resource(arg_scale_a, max_size=True) - s_a = [ - Vec( - buffer_ops.buffer_load( - scale_a_rsrc, fx.Int32(bx_m + mi * 16 + lane_div_16 * 4), vec_width=4, dtype=T.f32 - ) - ).bitcast(fx.Float32) - for mi in range_constexpr(m_repeat) - ] + s_b.append(fx.Float32(f.load()[0])) + # scale_a: vec4 f32 per m-block (BufferCopy128b). + sa_buf = fx.logical_divide( + fx.rocdl.make_buffer_tensor(arg_scale_a, max_size=True), fx.make_layout(4, 1) + ) + s_a = [] + for mi in range_constexpr(m_repeat): + f = fx.make_rmem_tensor(4, Float32) + grp = (bx_m + mi * 16 + lane_div_16 * 4) // 4 + fx.copy_atom_call(epi_copy_128b, sa_buf[None, fx.Int32(grp)], f) + s_a.append(Vec(f.load())) if const_expr(_has_bias): # Per-column bias (out_dtype), one scalar per N-block, shared across rows. - bias_rsrc = buffer_ops.create_buffer_resource(arg_bias, max_size=True) - bias_elem_ty = T.bf16 if out_dtype == "bf16" else T.f16 - bias = [ - fx.Float32( - buffer_ops.buffer_load( - bias_rsrc, - fx.Int32(by_n + (ni * num_waves + wave_id) * 16 + lane_mod_16), - vec_width=1, - dtype=bias_elem_ty, - ) + bias_buf = fx.logical_divide(fx.rocdl.make_buffer_tensor(arg_bias, max_size=True), fx.make_layout(1, 1)) + bias = [] + for ni in range_constexpr(num_acc_n): + f = fx.make_rmem_tensor(1, out_elem_cls) + fx.copy_atom_call( + epi_copy_16b, + bias_buf[None, fx.Int32(by_n + (ni * num_waves + wave_id) * 16 + lane_mod_16)], + f, ) - for ni in range_constexpr(num_acc_n) - ] + bias.append(fx.Float32(f.load()[0])) return s_a, s_b, bias overlap_epi_load = acc_size <= 64 # small enough accumulator to keep operands live over the MMA diff --git a/kernels/moe/mxfp_moe/gemm1.py b/kernels/moe/mxfp_moe/gemm1.py index 23ee69047..559c09bdf 100644 --- a/kernels/moe/mxfp_moe/gemm1.py +++ b/kernels/moe/mxfp_moe/gemm1.py @@ -268,6 +268,8 @@ def issue_a_ds_read(slot): t.store(lo.shuffle(hi, list(range(8)))) a[i][k] = t else: + # Manual XOR swizzle kept: the crd2idx form (see gemm2) is ISA-identical + # but ~5% slower here from scheduling sensitivity in this tuned loop. lds_col = (lane_div_16 * fx.Int32(16) + fx.Int32(k * 64)) ^ mask for i in range_constexpr(kMChunks): lds_row = lane_mod_16 + fx.Int32(i * 16) diff --git a/kernels/moe/mxfp_moe/gemm2.py b/kernels/moe/mxfp_moe/gemm2.py index 20e4f7fc2..6168cce8a 100644 --- a/kernels/moe/mxfp_moe/gemm2.py +++ b/kernels/moe/mxfp_moe/gemm2.py @@ -9,6 +9,8 @@ from flydsl.expr.typing import Vector as Vec from .mxfp4_gemm_common import ( + _a_lds_swz_block_idx, + _a_lds_swz_block_layout, _e8m0_from_amax, _fabs_f32, _gep1, @@ -499,18 +501,20 @@ def issue_a_load_lds(slot, kt): k_half=_K_HALF, ) + # A-LDS read block-swizzle via layout algebra (crd2idx over composed swizzle). + _a_lds_swz = _a_lds_swz_block_layout(_aStages * BM) + def issue_a_ds_read(slot): lane_row = lane_mod_16 - lane_col = lane_div_16 * fx.Int32(16) - mask = _lds_swizzle_mask(lane_row) + lane_col_block = lane_div_16 # 16-byte block column within the 128-byte row a = [[None, None] for _ in range(_kMChunks)] for k in range_constexpr(2): - lds_col = (lane_col + fx.Int32(k * 64)) ^ mask + block_col = lane_col_block + fx.Int32(k * 4) # +64 bytes == +4 blocks for i in range_constexpr(_kMChunks): - lds_row = lane_row + fx.Int32(i * 16) - byte_off = fx.Int32(slot * _slot_bytes) + lds_row * fx.Int32(KH_TILE) + lds_col + global_row = fx.Int32(slot * BM) + lane_row + fx.Int32(i * 16) + block_idx = _a_lds_swz_block_idx(_a_lds_swz, global_row, block_col) r = fx.make_rmem_tensor(lds_a_read_lay, fx.Int32) - fx.copy_atom_call(lds_a_read_atom, fx.slice(s_aq_i32x4_tiles, (None, byte_off // fx.Int32(16))), r) + fx.copy_atom_call(lds_a_read_atom, fx.slice(s_aq_i32x4_tiles, (None, block_idx)), r) a[i][k] = r return a diff --git a/kernels/moe/mxfp_moe/mxfp4_gemm_common.py b/kernels/moe/mxfp_moe/mxfp4_gemm_common.py index 82a2d9dac..071e36e38 100644 --- a/kernels/moe/mxfp_moe/mxfp4_gemm_common.py +++ b/kernels/moe/mxfp_moe/mxfp4_gemm_common.py @@ -137,6 +137,27 @@ def _lds_swizzle_mask(row): return (row & fx.Int32(14)) << fx.Int32(3) +# A-LDS bank-conflict swizzle as layout algebra: over the flat 16-byte-block index +# (row*8 + block_col) the manual XOR mask equals SwizzleType(3, 0, 4) (verified bit-for-bit). +_A_LDS_BLOCKS_PER_ROW = 8 # 128-byte fp4 row / 16-byte block + + +def _a_lds_swz_block_layout(rows): + """Composed layout over 16-byte A-LDS blocks: (rows, 8) row-major, swizzled S<3,0,4>. + + crd2idx((row, block_col), ) == the manual `(byte_off ^ mask) // 16` block index. + """ + return fx.make_composed_layout( + fx.static(fx.SwizzleType.get(3, 0, 4)), + fx.make_layout((rows, _A_LDS_BLOCKS_PER_ROW), (_A_LDS_BLOCKS_PER_ROW, 1)), + ) + + +def _a_lds_swz_block_idx(swz_layout, row, block_col): + """Swizzled 16-byte-block index for (row, block_col) via layout algebra.""" + return fx.Int32(crd2idx([fx.Int64(row), fx.Int64(block_col)], swz_layout)) + + def _fabs_f32(x): return fx.Float32(llvm.call_intrinsic(T.f32, "llvm.fabs.f32", [_raw(x)], [], []))