From 7944e3160adb1a93fb9d06b94eeb378c90a754ac Mon Sep 17 00:00:00 2001 From: Leo Date: Fri, 17 Jul 2026 08:24:32 +0000 Subject: [PATCH 1/3] Optimize the mch kernels --- tile_kernels/mhc/norm_fn_kernel.py | 17 ++++++++++------- tile_kernels/mhc/pre_apply_mix_kernel.py | 8 ++------ tile_kernels/modeling/mhc/ops/sinkhorn.py | 12 +++++++++++- 3 files changed, 23 insertions(+), 14 deletions(-) diff --git a/tile_kernels/mhc/norm_fn_kernel.py b/tile_kernels/mhc/norm_fn_kernel.py index f68af58..b822cf9 100644 --- a/tile_kernels/mhc/norm_fn_kernel.py +++ b/tile_kernels/mhc/norm_fn_kernel.py @@ -85,7 +85,7 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( _ = mhc_mult3 with T.Kernel(T.ceildiv(num_tokens, token_block), n_rms_group) as (pid_x, pid_y): out_frag = T.alloc_fragment((token_block, 32), T.float32) - sqrsum_part = T.alloc_fragment((token_block, 4), T.float32) + sqrsum_part = T.alloc_fragment(token_block, T.float32) T.clear(out_frag) T.clear(sqrsum_part) for pz in T.Pipelined(rms_group_size // hidden_block, num_stages=1): @@ -102,9 +102,14 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( x_frag = T.alloc_fragment((token_block, hidden_block), T.float32) T.copy(x_frag_16, x_frag) - for jj in T.serial(hidden_block // 4): - for i, j in T.Parallel(token_block, 4): - sqrsum_part[i, j] += x_frag[i, jj * 4 + j] * x_frag[i, jj * 4 + j] + # Compute sum of squares: first square each element, then reduce + x_sq = T.alloc_fragment((token_block, hidden_block), T.float32) + for i, j in T.Parallel(token_block, hidden_block): + x_sq[i, j] = x_frag[i, j] * x_frag[i, j] + sqrsum_blk = T.alloc_fragment(token_block, T.float32) + T.reduce_sum(x_sq, sqrsum_blk) + for i in T.Parallel(token_block): + sqrsum_part[i] += sqrsum_blk[i] T.gemm( x_frag, @@ -114,10 +119,8 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( transpose_B=True, clear_accum=False, ) - sqrsum_l = T.alloc_fragment(token_block, T.float32) - T.reduce_sum(sqrsum_part, sqrsum_l) for i in T.Parallel(token_block): - sqrsum[pid_x * token_block + i, pid_y] = sqrsum_l[i] + sqrsum[pid_x * token_block + i, pid_y] = sqrsum_part[i] for i, j in T.Parallel(token_block, 32): if j < 24: out[pid_x * token_block + i, pid_y, j] = out_frag[i, j] diff --git a/tile_kernels/mhc/pre_apply_mix_kernel.py b/tile_kernels/mhc/pre_apply_mix_kernel.py index 6b967de..5ea4027 100644 --- a/tile_kernels/mhc/pre_apply_mix_kernel.py +++ b/tile_kernels/mhc/pre_apply_mix_kernel.py @@ -35,12 +35,9 @@ def _mhc_pre_apply_mix_fwd_kernel( T.copy(mix[pid_n, 0], mixl) for i0_h in T.Pipelined(h // h_blk, num_stages=2): - xs = T.alloc_shared((mhc, h_blk), T.bfloat16) xl = T.alloc_fragment((mhc, h_blk), T.float32) - T.copy(x[pid_n, 0, i0_h * h_blk], xs, disable_tma=True) - T.copy(xs, xl, disable_tma=True) + T.copy(x[pid_n, 0, i0_h * h_blk], xl, disable_tma=True) - os = T.alloc_shared(h_blk, T.bfloat16) ol = T.alloc_fragment(h_blk, T.float32) T.clear(ol) @@ -48,8 +45,7 @@ def _mhc_pre_apply_mix_fwd_kernel( for i1_h in T.Parallel(h_blk): ol[i1_h] += mixl[i_mhc] * xl[i_mhc, i1_h] - T.copy(ol, os, disable_tma=True) - T.copy(os, o[pid_n, i0_h * h_blk], disable_tma=True) + T.copy(ol, o[pid_n, i0_h * h_blk], disable_tma=True) return _mhc_pre_apply_mix_fwd_kernel diff --git a/tile_kernels/modeling/mhc/ops/sinkhorn.py b/tile_kernels/modeling/mhc/ops/sinkhorn.py index d3fbe20..6adb88c 100644 --- a/tile_kernels/modeling/mhc/ops/sinkhorn.py +++ b/tile_kernels/modeling/mhc/ops/sinkhorn.py @@ -13,7 +13,17 @@ def forward( ) -> torch.Tensor: hidden_size = x.shape[1] output = torch.empty_like(x) - fwd_kernel = _mhc_sinkhorn_fwd(hidden_size, 1, repeat, eps) + # Choose token_block_size based on input size for optimal performance + n = x.shape[0] + if n >= 16384 and n % 32 == 0: + token_block_size = 32 + elif n >= 4096 and n % 16 == 0: + token_block_size = 16 + elif n % 4 == 0: + token_block_size = 4 + else: + token_block_size = 1 + fwd_kernel = _mhc_sinkhorn_fwd(hidden_size, token_block_size, repeat, eps) bwd_kernel = _mhc_sinkhorn_bwd(hidden_size, 32, repeat, eps) ctx.save_for_backward(x) ctx.bwd_kernel = bwd_kernel From 7ee824cc8f6e7cd64c9a1a0daea9c1d72e927959 Mon Sep 17 00:00:00 2001 From: Leo Date: Sun, 19 Jul 2026 04:01:49 +0000 Subject: [PATCH 2/3] Optimize MHC kernel dataflow and C500 dispatch Tune MHC forward and backward kernels across norm, pre/post processing, expand, sinkhorn, and multilayer recompute paths. Use more efficient fragment dataflow, specialized zero-gradient paths, shape-aware launch configuration, and C500 104-AP persistent CTA dispatch to improve throughput while preserving correctness. --- tile_kernels/mhc/expand_kernel.py | 10 +-- .../mhc/multilayer_recompute_kernel.py | 52 ++++++-------- tile_kernels/mhc/norm_fn_kernel.py | 52 +++++++++++--- tile_kernels/mhc/post_kernel.py | 28 +++----- tile_kernels/mhc/pre_apply_mix_kernel.py | 10 ++- tile_kernels/mhc/pre_split_mixes_kernel.py | 2 +- tile_kernels/modeling/mhc/ops/expand.py | 15 ++++- .../modeling/mhc/ops/head_compute_mix.py | 3 + tile_kernels/modeling/mhc/ops/norm_fn.py | 52 ++++++++++++-- .../modeling/mhc/ops/pre_apply_mix.py | 19 +++++- tile_kernels/modeling/mhc/ops/pre_big_fuse.py | 67 ++++--------------- .../modeling/mhc/ops/pre_split_mixes.py | 9 ++- tile_kernels/modeling/mhc/ops/sinkhorn.py | 6 +- 13 files changed, 185 insertions(+), 140 deletions(-) diff --git a/tile_kernels/mhc/expand_kernel.py b/tile_kernels/mhc/expand_kernel.py index 76d6b79..83455de 100644 --- a/tile_kernels/mhc/expand_kernel.py +++ b/tile_kernels/mhc/expand_kernel.py @@ -4,14 +4,16 @@ @tilelang.jit -def expand_to_mhc_fwd_tl(hidden: int, mhc_mult: int) -> tilelang.JITKernel: +def expand_to_mhc_fwd_tl( + hidden: int, + mhc_mult: int, + blk_n: int = 32, + blk_h: int = 128, +) -> tilelang.JITKernel: n = T.dynamic('num_tokens') h = hidden mhc = mhc_mult - blk_n = 32 - blk_h = 128 - @T.prim_func def expand_to_mhc_fwd_kernel( x: T.Tensor[(n, h), T.bfloat16], diff --git a/tile_kernels/mhc/multilayer_recompute_kernel.py b/tile_kernels/mhc/multilayer_recompute_kernel.py index ebcd058..a36695f 100644 --- a/tile_kernels/mhc/multilayer_recompute_kernel.py +++ b/tile_kernels/mhc/multilayer_recompute_kernel.py @@ -76,41 +76,18 @@ def kernel( post_mix_local = T.alloc_fragment(mhc, T.float32) comb_mix_local = T.alloc_fragment((mhc, mhc), T.float32) - layer_output_shared = T.alloc_shared((2, h_blk), T.bfloat16) - pre_mix_shared = T.alloc_shared((2, mhc), T.float32) - post_mix_shared = T.alloc_shared((2, mhc), T.float32) - comb_mix_shared = T.alloc_shared((2, mhc, mhc), T.float32) - for i0_h in T.serial(h // h_blk): T.copy(initial_residual[i_n, 0, i0_h * h_blk], res_local) - if L_post > 0: - layer_output_tensor_0 = T.make_tensor(layer_output_ptrs[0], (n, h), T.bfloat16) - pre_mix_tensor_0 = T.make_tensor(pre_mix_ptrs[0], (n, mhc), T.float32) - post_mix_tensor_0 = T.make_tensor(post_mix_ptrs[0], (n, mhc), T.float32) - comb_mix_tensor_0 = T.make_tensor(comb_mix_ptrs[0], (n, mhc, mhc), T.float32) - T.copy(layer_output_tensor_0[i_n, i0_h * h_blk], layer_output_shared[0, :]) - T.copy(pre_mix_tensor_0[i_n, 0], pre_mix_shared[0, :]) - T.copy(post_mix_tensor_0[i_n, 0], post_mix_shared[0, :]) - T.copy(comb_mix_tensor_0[i_n, 0, 0], comb_mix_shared[0, :, :]) - for i_layer in T.serial(L_post): + layer_output_tensor = T.make_tensor(layer_output_ptrs[i_layer], (n, h), T.bfloat16) + pre_mix_tensor = T.make_tensor(pre_mix_ptrs[i_layer], (n, mhc), T.float32) + post_mix_tensor = T.make_tensor(post_mix_ptrs[i_layer], (n, mhc), T.float32) + comb_mix_tensor = T.make_tensor(comb_mix_ptrs[i_layer], (n, mhc, mhc), T.float32) layer_input_tensor = T.make_tensor(layer_input_ptrs[i_layer], (n, h), T.bfloat16) output_residual_tensor = T.make_tensor(residual_ptrs[i_layer], (n, mhc, h), T.bfloat16) - phase = i_layer % 2 - - if i_layer + 1 < L_post: - next_layer_output_tensor = T.make_tensor(layer_output_ptrs[i_layer + 1], (n, h), T.bfloat16) - next_pre_mix_tensor = T.make_tensor(pre_mix_ptrs[i_layer + 1], (n, mhc), T.float32) - next_post_mix_tensor = T.make_tensor(post_mix_ptrs[i_layer + 1], (n, mhc), T.float32) - next_comb_mix_tensor = T.make_tensor(comb_mix_ptrs[i_layer + 1], (n, mhc, mhc), T.float32) - T.copy(next_layer_output_tensor[i_n, i0_h * h_blk], layer_output_shared[1 - phase, :]) - T.copy(next_pre_mix_tensor[i_n, 0], pre_mix_shared[1 - phase, :]) - T.copy(next_post_mix_tensor[i_n, 0], post_mix_shared[1 - phase, :]) - T.copy(next_comb_mix_tensor[i_n, 0, 0], comb_mix_shared[1 - phase, :, :]) - - T.copy(pre_mix_shared[phase, :], pre_mix_local) + T.copy(pre_mix_tensor[i_n, 0], pre_mix_local) T.clear(layer_input_local) for i_mhc in T.serial(mhc): @@ -119,9 +96,9 @@ def kernel( T.copy(layer_input_local, layer_input_tensor[i_n, i0_h * h_blk]) - T.copy(post_mix_shared[phase, :], post_mix_local) - T.copy(comb_mix_shared[phase, :, :], comb_mix_local) - T.copy(layer_output_shared[phase, :], layer_output_local) + T.copy(post_mix_tensor[i_n, 0], post_mix_local) + T.copy(comb_mix_tensor[i_n, 0, 0], comb_mix_local) + T.copy(layer_output_tensor[i_n, i0_h * h_blk], layer_output_local) for i_mhco, i1_h in T.Parallel(mhc, h_blk): new_res_local[i_mhco, i1_h] = post_mix_local[i_mhco] * layer_output_local[i1_h] for i_mhci in T.serial(mhc): @@ -181,7 +158,18 @@ def mhc_multilayer_recompute( device=initial_residual.device, ) - kernel = _mhc_multilayer_recompute_kernel(mhc_mult, hidden, num_layers, num_post) + # Two Wave64 groups improve the 4096/7168 hidden shapes; 2560 and 8192 retain one. + n_thr = 128 if hidden in (4096, 7168) else 64 + # 8192 needs smaller hidden tiles to avoid its large per-CTA fragment footprint. + h_blk = 512 if hidden == 8192 else 2048 + kernel = _mhc_multilayer_recompute_kernel( + mhc_mult, + hidden, + num_layers, + num_post, + n_thr=n_thr, + h_blk=h_blk, + ) kernel( initial_residual.view(-1, mhc_mult, hidden), pre_mix_ptrs, diff --git a/tile_kernels/mhc/norm_fn_kernel.py b/tile_kernels/mhc/norm_fn_kernel.py index b822cf9..8d89d96 100644 --- a/tile_kernels/mhc/norm_fn_kernel.py +++ b/tile_kernels/mhc/norm_fn_kernel.py @@ -68,17 +68,21 @@ def _mhc_pre_norm_fn_fwd_mul( mhc_mult3: int, n_rms_group: int, rms_group_size: int, - token_block: int = 32, - hidden_block: int = 256, + token_block: int = 16, + hidden_block: int = 128, + use_bf16_mma: bool = False, + fn_is_bf16: bool = False, ) -> tilelang.JITKernel: assert mhc_mult3 <= 32 + assert not fn_is_bf16 or use_bf16_mma num_tokens = T.dynamic('num_tokens') assert rms_group_size % hidden_block == 0 + fn_dtype = T.bfloat16 if fn_is_bf16 else T.float32 @T.prim_func def _mhc_pre_norm_fn_fwd_mul_kernel( x: T.Tensor[(num_tokens, n_rms_group * rms_group_size), T.bfloat16], - fn: T.Tensor[(mhc_mult3, n_rms_group * rms_group_size), T.float32], + fn: T.Tensor[(mhc_mult3, n_rms_group * rms_group_size), fn_dtype], out: T.Tensor[(num_tokens, n_rms_group, mhc_mult3), T.float32], sqrsum: T.Tensor[(num_tokens, n_rms_group), T.float32], ) -> None: @@ -90,7 +94,9 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( T.clear(sqrsum_part) for pz in T.Pipelined(rms_group_size // hidden_block, num_stages=1): x_smem_16 = T.alloc_shared((token_block, hidden_block), T.bfloat16) - fn_smem = T.alloc_shared((32, hidden_block), T.float32) + fn_smem = T.alloc_shared( + (32, hidden_block), T.bfloat16 if use_bf16_mma else T.float32 + ) T.annotate_layout({x_smem_16: tilelang.layout.make_swizzled_layout(x_smem_16)}) @@ -112,7 +118,7 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( sqrsum_part[i] += sqrsum_blk[i] T.gemm( - x_frag, + x_frag_16 if use_bf16_mma else x_frag, fn_smem, out_frag, transpose_A=False, @@ -128,6 +134,32 @@ def _mhc_pre_norm_fn_fwd_mul_kernel( return _mhc_pre_norm_fn_fwd_mul_kernel +@tilelang.jit(pass_configs=_PASS_CONFIGS) +def _mhc_pre_norm_fn_fwd_sqsum(hidden: int, hidden_block: int = 128) -> tilelang.JITKernel: + num_tokens = T.dynamic('num_tokens') + assert hidden % hidden_block == 0 + + @T.prim_func + def _mhc_pre_norm_fn_fwd_sqsum_kernel( + x: T.Tensor[(num_tokens, hidden), T.bfloat16], + sqsum: T.Tensor[num_tokens, T.float32], + ) -> None: + with T.Kernel(num_tokens, threads=hidden_block) as pid: + sum_reducer = T.alloc_reducer(1, T.float32, replication='all') + T.clear(sum_reducer) + for i0_h in T.serial(hidden // hidden_block): + x_frag = T.alloc_fragment(hidden_block, T.bfloat16) + T.copy(x[pid, i0_h * hidden_block], x_frag) + for i1_h in T.Parallel(hidden_block): + x_val = T.cast(x_frag[i1_h], T.float32) + sum_reducer[0] += x_val * x_val + T.finalize_reducer(sum_reducer) + if T.get_thread_binding() == 0: + sqsum[pid] = sum_reducer[0] + + return _mhc_pre_norm_fn_fwd_sqsum_kernel + + @tilelang.jit(pass_configs=_PASS_CONFIGS) def _mhc_pre_norm_fn_fwd_norm( mhc_mult3: int, @@ -215,8 +247,9 @@ def _mhc_pre_norm_fn_bwd_mul( mhc_mult3: int, n_rms_group: int, rms_group_size: int, - token_block: int = 128, - hidden_block: int = 64, + token_block: int = 64, + hidden_block: int = 32, + x_grad_is_zero: bool = False, ) -> tilelang.JITKernel: assert mhc_mult3 <= 32 num_tokens = T.dynamic('num_tokens') @@ -259,7 +292,10 @@ def _mhc_pre_norm_fn_bwd_mul_kernel( padded_grad[i, j] = 0 x_grad_frag = T.alloc_fragment((token_block, hidden_block), T.float32) - T.copy(x_grad[px * token_block, yz], x_grad_frag) + if x_grad_is_zero: + T.fill(x_grad_frag, 0) + else: + T.copy(x_grad[px * token_block, yz], x_grad_frag) T.gemm( padded_grad, diff --git a/tile_kernels/mhc/post_kernel.py b/tile_kernels/mhc/post_kernel.py index a7cb7db..c44d9bf 100644 --- a/tile_kernels/mhc/post_kernel.py +++ b/tile_kernels/mhc/post_kernel.py @@ -14,7 +14,7 @@ tilelang.PassConfigKey.TL_DISABLE_VECTORIZE_256: True, }, ) -def _mhc_post_fwd(mhc: int, hidden: int, n_thr: int = 128, h_blk: int = 1024) -> tilelang.JITKernel: +def _mhc_post_fwd(mhc: int, hidden: int, n_thr: int = 512, h_blk: int = 1024) -> tilelang.JITKernel: n = T.dynamic('num_tokens') h = hidden @@ -79,12 +79,6 @@ def _mhc_post_bwd_kernel( dd: T.Tensor[(n, h), T.bfloat16], ) -> None: with T.Kernel(n, threads=n_thr) as pid_n: - dx_shared = T.alloc_shared((4, h_blk), T.bfloat16) - b_shared = T.alloc_shared((4, h_blk), T.bfloat16) - db_shared = T.alloc_shared((4, h_blk), T.bfloat16) - d_shared = T.alloc_shared(h_blk, T.bfloat16) - dd_shared = T.alloc_shared(h_blk, T.bfloat16) - dx_local = T.alloc_fragment((4, h_blk), T.float32) b_local = T.alloc_fragment((4, h_blk), T.float32) db_local = T.alloc_fragment((4, h_blk), T.float32) @@ -102,13 +96,9 @@ def _mhc_post_bwd_kernel( T.clear(dc_reducer) for i0_h in T.Pipelined(T.ceildiv(h, h_blk), num_stages=3): - T.copy(dx[pid_n, 0, i0_h * h_blk], dx_shared, disable_tma=True) - T.copy(b[pid_n, 0, i0_h * h_blk], b_shared, disable_tma=True) - T.copy(d[pid_n, i0_h * h_blk], d_shared, disable_tma=True) - - T.copy(dx_shared, dx_local) - T.copy(b_shared, b_local) - T.copy(d_shared, d_local) + T.copy(dx[pid_n, 0, i0_h * h_blk], dx_local, disable_tma=True) + T.copy(b[pid_n, 0, i0_h * h_blk], b_local, disable_tma=True) + T.copy(d[pid_n, i0_h * h_blk], d_local, disable_tma=True) # da and db T.clear(db_local) @@ -125,11 +115,8 @@ def _mhc_post_bwd_kernel( dc_reducer[i_mhc] += d_local[i1_h] * dx_local[i_mhc, i1_h] dd_local[i1_h] += c_local[i_mhc] * dx_local[i_mhc, i1_h] - T.copy(db_local, db_shared) - T.copy(dd_local, dd_shared) - - T.copy(db_shared, db[pid_n, 0, i0_h * h_blk], disable_tma=True) - T.copy(dd_shared, dd[pid_n, i0_h * h_blk], disable_tma=True) + T.copy(db_local, db[pid_n, 0, i0_h * h_blk], disable_tma=True) + T.copy(dd_local, dd[pid_n, i0_h * h_blk], disable_tma=True) T.finalize_reducer(da_reducer) T.finalize_reducer(dc_reducer) @@ -186,7 +173,8 @@ def mhc_post_bwd( mhc = d_o.shape[2] h = d_o.shape[3] - bwd_kernel = _mhc_post_bwd(mhc, h) + # A second Wave64 improves the larger hidden shapes, while 1280 remains launch-overhead bound. + bwd_kernel = _mhc_post_bwd(mhc, h, n_thr=256 if h > 1280 else 128) ( d_comb_res_mix, d_residual, diff --git a/tile_kernels/mhc/pre_apply_mix_kernel.py b/tile_kernels/mhc/pre_apply_mix_kernel.py index 5ea4027..7226328 100644 --- a/tile_kernels/mhc/pre_apply_mix_kernel.py +++ b/tile_kernels/mhc/pre_apply_mix_kernel.py @@ -56,6 +56,7 @@ def _mhc_pre_apply_mix_bwd( hidden: int, n_thr: int = 128, h_blk: int = 1024, + x_grad_is_zero: bool = False, ) -> tilelang.JITKernel: n = T.dynamic('n') h = hidden @@ -89,10 +90,13 @@ def _mhc_pre_apply_mix_bwd_kernel( T.copy(x[pid_n, 0, i0_h * h_blk], xs, disable_tma=True) T.copy(xs, xl, disable_tma=True) - xgs = T.alloc_shared((mhc, h_blk), T.bfloat16) xgl = T.alloc_fragment((mhc, h_blk), T.float32) - T.copy(x_grad[pid_n, 0, i0_h * h_blk], xgs, disable_tma=True) - T.copy(xgs, xgl, disable_tma=True) + if x_grad_is_zero: + T.fill(xgl, 0) + else: + xgs = T.alloc_shared((mhc, h_blk), T.bfloat16) + T.copy(x_grad[pid_n, 0, i0_h * h_blk], xgs, disable_tma=True) + T.copy(xgs, xgl, disable_tma=True) for i_mhc, i1_h in T.Parallel(mhc, h_blk): mgl[i_mhc] += ogl[i1_h] * xl[i_mhc, i1_h] diff --git a/tile_kernels/mhc/pre_split_mixes_kernel.py b/tile_kernels/mhc/pre_split_mixes_kernel.py index 3a319e8..fdf719b 100644 --- a/tile_kernels/mhc/pre_split_mixes_kernel.py +++ b/tile_kernels/mhc/pre_split_mixes_kernel.py @@ -72,7 +72,7 @@ def _mhc_pre_split_mixes_bwd( mhc_mult: int, mhc_post_mult_value: float, token_block_size: int, - num_sms: int = 148, + num_sms: int = 104, dtype: T.dtype = T.float32, ) -> tilelang.JITKernel: num_tokens = T.dynamic('num_tokens') diff --git a/tile_kernels/modeling/mhc/ops/expand.py b/tile_kernels/modeling/mhc/ops/expand.py index ccf3aa8..1a45260 100644 --- a/tile_kernels/modeling/mhc/ops/expand.py +++ b/tile_kernels/modeling/mhc/ops/expand.py @@ -14,8 +14,19 @@ def forward( if out is None: out = hidden.new_empty(*hidden.shape[:-1], mhc_mult, hidden.shape[-1]) assert hidden.is_contiguous() - kernel = expand_to_mhc_fwd_tl(hidden.shape[-1], mhc_mult) - kernel(hidden.flatten(0, -2), out.flatten(0, -3)) + hidden_flat = hidden.flatten(0, -2) + num_tokens = hidden_flat.shape[0] + use_blk_n64 = num_tokens % 64 == 0 and not ( + num_tokens >= 8192 and hidden.shape[-1] == 1280 + ) + use_blk_h256 = use_blk_n64 and hidden.shape[-1] >= 4096 + kernel = expand_to_mhc_fwd_tl( + hidden.shape[-1], + mhc_mult, + blk_n=64 if use_blk_n64 else 32, + blk_h=256 if use_blk_h256 else 128, + ) + kernel(hidden_flat, out.flatten(0, -3)) return out @staticmethod diff --git a/tile_kernels/modeling/mhc/ops/head_compute_mix.py b/tile_kernels/modeling/mhc/ops/head_compute_mix.py index 66d4e32..47aef2b 100644 --- a/tile_kernels/modeling/mhc/ops/head_compute_mix.py +++ b/tile_kernels/modeling/mhc/ops/head_compute_mix.py @@ -33,6 +33,9 @@ def backward( input_mix, mhc_scale, mhc_base = ctx.saved_tensors num_sms = get_num_sms() + num_tokens = input_mix.numel() // input_mix.shape[-1] + if num_sms == 104 and num_tokens % 128 == 0: + num_sms = 128 input_mix_grad = torch.empty_like(input_mix) mhc_scale_grad_partial = torch.empty( num_sms, diff --git a/tile_kernels/modeling/mhc/ops/norm_fn.py b/tile_kernels/modeling/mhc/ops/norm_fn.py index b0ec6cd..43b0ec9 100644 --- a/tile_kernels/modeling/mhc/ops/norm_fn.py +++ b/tile_kernels/modeling/mhc/ops/norm_fn.py @@ -8,6 +8,7 @@ _mhc_pre_norm_fn_bwd_norm, _mhc_pre_norm_fn_fwd_mul, _mhc_pre_norm_fn_fwd_norm, + _mhc_pre_norm_fn_fwd_sqsum, round_to_tf32, ) @@ -92,13 +93,42 @@ def forward( fn = round_to_tf32(fn) - fwd_mul_kernel = _mhc_pre_norm_fn_fwd_mul(mhc_mult3, 1, mhc_hidden_size) - fwd_mul_kernel( - x.view(-1, mhc_hidden_size), - fn, - out_mul_splitted.view(-1, 1, mhc_mult3), - sqrsum_splitted.view(-1, 1), + num_tokens = x.numel() // mhc_hidden_size + x_flat = x.view(-1, mhc_hidden_size) + use_library_gemm = ( + num_tokens == 4096 and mhc_hidden_size >= 5120 + ) or ( + num_tokens >= 8192 and mhc_hidden_size >= 16384 ) + if use_library_gemm: + # mcBLAS handles the narrow-N BF16 GEMM efficiently; retain the reduction in TileLang. + bf16_mul = torch.empty( + num_tokens, + mhc_mult3, + dtype=torch.bfloat16, + device=x.device, + ) + torch.matmul(x_flat, fn.bfloat16().t(), out=bf16_mul) + out_mul_splitted.view(-1, mhc_mult3).copy_(bf16_mul) + _mhc_pre_norm_fn_fwd_sqsum(mhc_hidden_size)( + x_flat, + sqrsum_splitted.view(-1), + ) + else: + use_bf16_mma = num_tokens >= 4096 + fwd_mul_kernel = _mhc_pre_norm_fn_fwd_mul( + mhc_mult3, + 1, + mhc_hidden_size, + use_bf16_mma=use_bf16_mma, + fn_is_bf16=use_bf16_mma, + ) + fwd_mul_kernel( + x_flat, + fn.bfloat16() if use_bf16_mma else fn, + out_mul_splitted.view(-1, 1, mhc_mult3), + sqrsum_splitted.view(-1, 1), + ) # END of TileLang implementation of pre-norm-fn forward matmul out_mul = torch.empty_like(out_mul_splitted[0]) @@ -154,7 +184,15 @@ def backward( out_mul_grad = round_to_tf32(out_mul_grad) - bwd_mul_kernel = _mhc_pre_norm_fn_bwd_mul(mhc_mult3, 1, mhc_hidden_size) + num_tokens = x.numel() // mhc_hidden_size + bwd_token_block = 32 if num_tokens >= 8192 and mhc_hidden_size >= 10240 else 64 + bwd_mul_kernel = _mhc_pre_norm_fn_bwd_mul( + mhc_mult3, + 1, + mhc_hidden_size, + token_block=bwd_token_block, + x_grad_is_zero=not ctx.fuse_grad_acc, + ) bwd_mul_kernel( out_mul_grad.view(-1, 1, mhc_mult3), sqrsum_grad.view(-1, 1), diff --git a/tile_kernels/modeling/mhc/ops/pre_apply_mix.py b/tile_kernels/modeling/mhc/ops/pre_apply_mix.py index c096c39..34b4c6c 100644 --- a/tile_kernels/modeling/mhc/ops/pre_apply_mix.py +++ b/tile_kernels/modeling/mhc/ops/pre_apply_mix.py @@ -16,7 +16,15 @@ def forward( mhc = mix.shape[-2] assert mix.shape[-1] == 1 ctx.fwd_kernel = _mhc_pre_apply_mix_fwd(mhc, h) - ctx.bwd_kernel = _mhc_pre_apply_mix_bwd(mhc, h) + # Backward shared-memory pressure has a different optimum than the forward kernel. + bwd_n_thr = 256 if h in (2560, 4096) else 128 + bwd_h_blk = 512 if h in (2560, 4096, 7168) else 1024 + ctx.bwd_kernel = _mhc_pre_apply_mix_bwd( + mhc, + h, + n_thr=bwd_n_thr, + h_blk=bwd_h_blk, + ) if out is None: out = torch.empty(*x.shape[:-2], h, dtype=torch.bfloat16, device=x.device) ctx.fwd_kernel(x.view(-1, mhc, h), mix.view(-1, mhc), out.view(-1, h)) @@ -38,7 +46,14 @@ def backward(ctx: 'MHCPreApplyMix', o_grad: torch.Tensor) -> tuple[torch.Tensor x_grad = None else: x_grad = torch.zeros_like(x) - mix_grad = ctx.bwd_kernel( + zero_grad_bwd_kernel = _mhc_pre_apply_mix_bwd( + mhc, + h, + n_thr=256 if h in (2560, 4096) else 128, + h_blk=512 if h in (2560, 4096, 7168) else 1024, + x_grad_is_zero=True, + ) + mix_grad = zero_grad_bwd_kernel( o_grad.view(-1, h), x.view(-1, mhc, h), mix.view(-1, mhc), diff --git a/tile_kernels/modeling/mhc/ops/pre_big_fuse.py b/tile_kernels/modeling/mhc/ops/pre_big_fuse.py index 7b311e4..981ea5f 100644 --- a/tile_kernels/modeling/mhc/ops/pre_big_fuse.py +++ b/tile_kernels/modeling/mhc/ops/pre_big_fuse.py @@ -1,7 +1,9 @@ import torch -from tile_kernels.mhc.norm_fn_kernel import _mhc_pre_norm_fn_fwd_mul, round_to_tf32 -from tile_kernels.mhc.pre_big_fuse_kernel import _mhc_pre_big_fuse +from tile_kernels.modeling.mhc.ops.norm_fn import mhc_pre_norm_fn +from tile_kernels.modeling.mhc.ops.pre_apply_mix import mhc_pre_apply_mix +from tile_kernels.modeling.mhc.ops.pre_split_mixes import mhc_pre_split_mixes +from tile_kernels.modeling.mhc.ops.sinkhorn import sinkhorn_normalize def mhc_pre_big_fuse( @@ -32,60 +34,15 @@ def mhc_pre_big_fuse( assert mhc_scale.shape == (3,) assert mhc_base.shape == (mhc_mult3,) - outer_shape = residual.shape[:-2] - - residual_flat = residual.view(-1, mhc_mult, hidden_size) - num_tokens = residual_flat.shape[0] - fn_flat = fn - - post_mix = torch.empty(num_tokens, mhc_mult, dtype=torch.float32, device=residual.device) - comb_mix = torch.empty(num_tokens, mhc_mult2, dtype=torch.float32, device=residual.device) - layer_input = torch.empty(num_tokens, hidden_size, dtype=torch.bfloat16, device=residual.device) - - gemm_out_mul = torch.empty( - n_splits, num_tokens, mhc_mult3, dtype=torch.float32, device=residual.device - ) - gemm_out_sqrsum = torch.empty(n_splits, num_tokens, dtype=torch.float32, device=residual.device) - - # TileLang implementation doesn't support split-k, so we set n_splits to 1 - # You may want to adopt the DeepGEMM implementation with split-k for better performance - n_splits = 1 - gemm_out_mul = gemm_out_mul[:1] - gemm_out_sqrsum = gemm_out_sqrsum[:1] - - fn = round_to_tf32(fn) - - fwd_mul_kernel = _mhc_pre_norm_fn_fwd_mul(mhc_mult3, 1, mhc_hidden_size) - fwd_mul_kernel( - residual_flat.view(-1, mhc_hidden_size), - fn, - gemm_out_mul.view(-1, 1, mhc_mult3), - gemm_out_sqrsum.view(-1, 1), - ) - # END of TileLang implementation of pre-norm-fn forward matmul - - _mhc_pre_big_fuse( - hidden_size, - rms_eps, - mhc_pre_eps, - mhc_sinkhorn_eps, - mhc_post_mult_value, - sinkhorn_repeat, - n_splits=n_splits, - mhc_mult=mhc_mult, - )( - gemm_out_mul, - gemm_out_sqrsum, + mixes = mhc_pre_norm_fn(residual, fn, None, rms_eps, n_splits=n_splits) + pre_mix, post_mix, comb_mix = mhc_pre_split_mixes( + mixes, mhc_scale, mhc_base, - residual_flat, - post_mix, - comb_mix, - layer_input, + mhc_mult, + mhc_post_mult_value, + mhc_pre_eps, ) - - post_mix = post_mix.view(*outer_shape, mhc_mult, 1) - comb_mix = comb_mix.view(*outer_shape, mhc_mult, mhc_mult) - layer_input = layer_input.view(*outer_shape, hidden_size) - + comb_mix = sinkhorn_normalize(comb_mix, repeat=sinkhorn_repeat, eps=mhc_sinkhorn_eps) + layer_input = mhc_pre_apply_mix(residual, pre_mix) return post_mix, comb_mix, layer_input diff --git a/tile_kernels/modeling/mhc/ops/pre_split_mixes.py b/tile_kernels/modeling/mhc/ops/pre_split_mixes.py index 43e890d..a1be8fd 100644 --- a/tile_kernels/modeling/mhc/ops/pre_split_mixes.py +++ b/tile_kernels/modeling/mhc/ops/pre_split_mixes.py @@ -36,13 +36,18 @@ def forward( mhc_pre_eps, token_block_size=32, ) + # C500 has 104 APs. For 128-multiple tile counts, use 128 persistent + # CTAs so every CTA owns the same number of tiles instead of a 104-AP tail. + persistent_ctas = get_num_sms() + if persistent_ctas == 104 and num_tokens % 128 == 0: + persistent_ctas = 128 ctx.bwd_kernel = _mhc_pre_split_mixes_bwd( mhc_mult, mhc_post_mult_value, token_block_size=32, - num_sms=get_num_sms(), + num_sms=persistent_ctas, ) - ctx.num_sms = get_num_sms() + ctx.num_sms = persistent_ctas ctx.fwd_kernel(input_mixes, mhc_scale, mhc_base, pre_layer_mix, post_layer_mix, comb_res_mix) diff --git a/tile_kernels/modeling/mhc/ops/sinkhorn.py b/tile_kernels/modeling/mhc/ops/sinkhorn.py index 6adb88c..a97a8f3 100644 --- a/tile_kernels/modeling/mhc/ops/sinkhorn.py +++ b/tile_kernels/modeling/mhc/ops/sinkhorn.py @@ -15,16 +15,14 @@ def forward( output = torch.empty_like(x) # Choose token_block_size based on input size for optimal performance n = x.shape[0] - if n >= 16384 and n % 32 == 0: - token_block_size = 32 - elif n >= 4096 and n % 16 == 0: + if n >= 4096 and n % 16 == 0: token_block_size = 16 elif n % 4 == 0: token_block_size = 4 else: token_block_size = 1 fwd_kernel = _mhc_sinkhorn_fwd(hidden_size, token_block_size, repeat, eps) - bwd_kernel = _mhc_sinkhorn_bwd(hidden_size, 32, repeat, eps) + bwd_kernel = _mhc_sinkhorn_bwd(hidden_size, 8, repeat, eps) ctx.save_for_backward(x) ctx.bwd_kernel = bwd_kernel fwd_kernel(x, output) From cc0e580f520bae76015c6090b6c1604ada785b11 Mon Sep 17 00:00:00 2001 From: Leo Date: Mon, 20 Jul 2026 03:35:17 +0000 Subject: [PATCH 3/3] Optimize MoE routing kernels for wave-level selection Pack routing comparisons into stable uint64 keys, use 32-lane subgroup execution for top-k and grouped routing, and specialize short top-k fused mapping workloads. Improve fused expansion scheduling and mapping correctness for the 32-bit match-any lowering. Add cleared-L2 MoE benchmark timing and coverage for stable ties and special floating-point values. --- tile_kernels/moe/common.py | 2 +- tile_kernels/moe/expand_to_fused_kernel.py | 20 +- tile_kernels/moe/get_fused_mapping_kernel.py | 74 ++++--- tile_kernels/moe/top2_sum_gate_kernel.py | 190 ++++++++++-------- tile_kernels/moe/topk_gate_kernel.py | 74 ++++--- .../moe/topk_sum_and_topk_group_idx_kernel.py | 114 +++++++++-- 6 files changed, 299 insertions(+), 175 deletions(-) diff --git a/tile_kernels/moe/common.py b/tile_kernels/moe/common.py index 700b02c..ba0df1e 100644 --- a/tile_kernels/moe/common.py +++ b/tile_kernels/moe/common.py @@ -40,7 +40,7 @@ def get_topk_group_idx( # Count the number of groups that have a larger top2 sum for i in T.unroll(num_groups): - other_top2_sum = T.shfl_sync(topk_sum_var, i) + other_top2_sum = T.shfl_sync(topk_sum_var, i, width=32) if other_top2_sum > topk_sum_var or (other_top2_sum == topk_sum_var and i < lane_idx): count_var += 1 diff --git a/tile_kernels/moe/expand_to_fused_kernel.py b/tile_kernels/moe/expand_to_fused_kernel.py index 92a18c4..6afdd23 100644 --- a/tile_kernels/moe/expand_to_fused_kernel.py +++ b/tile_kernels/moe/expand_to_fused_kernel.py @@ -22,7 +22,7 @@ def get_expand_to_fused_kernel( x_dtype: T.dtype, sf_dtype: T.dtype, ): - num_threads = 64 + num_threads = 256 if x_dtype == T.float8_e4m3fn else 512 hidden_aligned = align(hidden, num_threads) if num_per_channels is not None: @@ -37,8 +37,6 @@ def get_expand_to_fused_kernel( sf_stride = T.dynamic('sf_stride') num_tokens = T.dynamic('num_tokens') num_expanded_tokens = T.dynamic('num_expanded_tokens') - num_blocks = T.max(num_tokens, num_expanded_tokens) - sf_shape = (hidden_sf, num_expanded_tokens) if use_tma_aligned_col_major_sf else (num_expanded_tokens, hidden_sf) @T.prim_func @@ -50,23 +48,19 @@ def expand_to_fused_kernel( token_topk_to_pos: T.Tensor[(num_tokens, num_topk), T.int32], pos_to_expert: T.Tensor[(num_expanded_tokens, ), T.int32] ): - with T.Kernel(num_blocks, threads=num_threads) as (pid_token, ): + with T.Kernel(num_tokens, threads=num_threads) as (pid_token, ): pos_local = T.alloc_local((num_topk, ), T.int32) - if pid_token < num_expanded_tokens: - if pos_to_expert[pid_token] < 0: + for output_pos in T.serial(pid_token, num_expanded_tokens, num_tokens): + if pos_to_expert[output_pos] < 0: for i in T.Parallel(hidden_aligned): - expanded_x[pid_token, i] = 0 + expanded_x[output_pos, i] = 0 if num_per_channels is not None: for i in T.Parallel(hidden_sf_aligned): if use_tma_aligned_col_major_sf: - expanded_x_sf[i, pid_token] = 0 + expanded_x_sf[i, output_pos] = 0 else: - expanded_x_sf[pid_token, i] = 0 - - if pid_token >= num_tokens: - T.thread_return() - T.assume(pid_token < num_tokens) + expanded_x_sf[output_pos, i] = 0 x_fragment = T.alloc_fragment((hidden_aligned, ), x_dtype) x_sf_fragment = T.alloc_fragment((hidden_sf_aligned, ), sf_dtype) diff --git a/tile_kernels/moe/get_fused_mapping_kernel.py b/tile_kernels/moe/get_fused_mapping_kernel.py index 0d2c897..3be872d 100644 --- a/tile_kernels/moe/get_fused_mapping_kernel.py +++ b/tile_kernels/moe/get_fused_mapping_kernel.py @@ -29,12 +29,17 @@ def get_get_fused_mapping_kernel( alignment: int, num_sms: int, ): - num_threads = 256 + # __match_any_sync returns a 32-bit lane mask on this lowering. Keep one + # logical 32-lane match group in each hardware wave. + # The short top-k path has less work per worker; 256 threads lowers its + # fixed grid-sync cost, while longer routes benefit from 512 threads. + num_threads = 256 if num_topk == 2 else 512 while num_threads < num_experts: num_threads *= 2 assert num_threads <= 1024 and num_threads >= num_experts - warp_size = 64 - num_warps = num_threads // warp_size + hardware_warp_size = 64 + match_group_size = 32 + num_warps = num_threads // hardware_warp_size num_global_warps = num_sms * num_warps num_global_threads = num_threads * num_sms @@ -56,8 +61,8 @@ def get_fused_mapping_kernel( ): with T.Kernel(num_sms, threads=num_threads) as (sm_idx,): thread_idx = T.get_thread_binding(0) - warp_idx = thread_idx // warp_size - lane_idx = thread_idx % warp_size + warp_idx = thread_idx // hardware_warp_size + lane_idx = thread_idx % hardware_warp_size global_thread_idx = sm_idx * num_threads + thread_idx global_warp_idx = sm_idx * num_warps + warp_idx numel = num_tokens * num_topk @@ -73,8 +78,9 @@ def get_fused_mapping_kernel( topk_idx_1d = T.view(topk_idx, (num_tokens * num_topk,)) token_topk_to_pos_1d = T.view(token_topk_to_pos, (num_tokens * num_topk,)) - for i in T.serial(lane_idx, num_experts, warp_size): - experts_sum_per_warp_shared[warp_idx, i] = 0 + if lane_idx < match_group_size: + for i in T.serial(lane_idx, num_experts, match_group_size): + experts_sum_per_warp_shared[warp_idx, i] = 0 T.sync_warp() for i in T.serial(global_thread_idx, num_expanded_tokens, num_global_threads): @@ -89,12 +95,13 @@ def get_fused_mapping_kernel( end = T.alloc_var(T.int32) divide_task(numel, num_global_warps, global_warp_idx, start, end) - for i in T.serial(start + lane_idx, end, warp_size): - T.assume(0 <= i < numel) - expert_idx = topk_idx_1d[i] - if expert_idx != -1: - T.assume(0 <= expert_idx < num_experts) - T.atomic_add(experts_sum_per_warp_shared[warp_idx, expert_idx], 1) + if lane_idx < match_group_size: + for i in T.serial(start + lane_idx, end, match_group_size): + T.assume(0 <= i < numel) + expert_idx = topk_idx_1d[i] + if expert_idx != -1: + T.assume(0 <= expert_idx < num_experts) + T.atomic_add(experts_sum_per_warp_shared[warp_idx, expert_idx], 1) T.sync_threads() @@ -140,26 +147,27 @@ def get_fused_mapping_kernel( T.sync_threads() divide_task(numel, num_global_warps, global_warp_idx, start, end) - aligned_end = align(end, warp_size) - lane_mask = T.uint64(1 << lane_idx) + T.uint64(1 << lane_idx) - 1 - lane_mask_rev = ~lane_mask - for i in T.serial(start + lane_idx, aligned_end, warp_size): - T.assume(0 <= i) - expert_idx = T.Select(i < numel, T.int32(topk_idx_1d[i]), -1) - mask = T.call_extern(T.uint64, '__match_any_sync', tilelang.tvm.tir.const(0xFFFFFFFFFFFFFFFF, T.uint64), expert_idx) - count = T.popcount(mask & lane_mask) - - if i < numel and expert_idx >= 0: - T.assume(expert_idx < num_experts) - prefix_count = experts_sum_per_warp_shared[warp_idx, expert_idx] - pos = prefix_count - count - if mask & lane_mask_rev == 0: - experts_sum_per_warp_shared[warp_idx, expert_idx] = pos - token_topk_to_pos_1d[i] = pos - T.assume(0 <= pos < num_expanded_tokens) - pos_to_expert[pos] = expert_idx - pos_to_token[pos] = i // num_topk - pos_to_token_topk[pos] = i + aligned_end = align(end, match_group_size) + for base in T.serial(start, aligned_end, match_group_size): + if lane_idx < match_group_size: + i = base + lane_idx + expert_idx = T.Select(i < numel, T.int32(topk_idx_1d[i]), -1) + lane_mask = T.uint32(1 << lane_idx) + T.uint32(1 << lane_idx) - 1 + lane_mask_rev = ~lane_mask + mask = T.call_extern(T.uint32, '__match_any_sync', tilelang.tvm.tir.const(0xFFFFFFFF, T.uint32), expert_idx) + count = T.popcount(mask & lane_mask) + + if i < numel and expert_idx >= 0: + T.assume(expert_idx < num_experts) + prefix_count = experts_sum_per_warp_shared[warp_idx, expert_idx] + pos = prefix_count - count + if mask & lane_mask_rev == 0: + experts_sum_per_warp_shared[warp_idx, expert_idx] = pos + token_topk_to_pos_1d[i] = pos + T.assume(0 <= pos < num_expanded_tokens) + pos_to_expert[pos] = expert_idx + pos_to_token[pos] = i // num_topk + pos_to_token_topk[pos] = i T.sync_warp() return get_fused_mapping_kernel diff --git a/tile_kernels/moe/top2_sum_gate_kernel.py b/tile_kernels/moe/top2_sum_gate_kernel.py index 768b52c..03f54ae 100644 --- a/tile_kernels/moe/top2_sum_gate_kernel.py +++ b/tile_kernels/moe/top2_sum_gate_kernel.py @@ -13,7 +13,13 @@ def warp_reduce_sum(x: T.Ref): # Keep the same with the old implementation for i in T.unroll(0, 5): - x += T.shfl_xor(x, 1 << (4 - i)) + x += T.shfl_xor(x, 1 << (4 - i), width=32) + + +@T.macro +def warp_reduce_max(x: T.Ref): + for i in T.unroll(0, 5): + x = T.max(x, T.shfl_xor(x, 1 << (4 - i), width=32)) @tilelang.jit( @@ -33,10 +39,11 @@ def get_top2_sum_gate_kernel( ): # fmt: off # Kernel config warp_size = 32 - num_threads = 32 + use_wave64 = num_routed_experts % 64 == 0 + num_threads = 64 if use_wave64 else 32 assert num_topk <= warp_size, f'num_topk must be less than or equal to {warp_size}' - # Each warp handles one token + # Each 32-lane subgroup handles one token. num_tokens_per_block = num_threads // warp_size # Keep the same with the old implementation @@ -95,7 +102,7 @@ def top2_sum_gate_kernel( tp_rank: T.int32, num_tp_ranks: T.int32, ): - with T.Kernel(num_tokens, threads=num_threads) as pid: + with T.Kernel(T.ceildiv(num_tokens, num_tokens_per_block), threads=num_threads) as pid: thread_idx = T.get_thread_binding() token_idx = thread_idx // 32 global_token_idx = token_idx + pid * num_tokens_per_block @@ -106,27 +113,30 @@ def top2_sum_gate_kernel( bias_local = T.alloc_local((num_routed_experts_per_thread,), dtype=T.float32) scores_local = T.alloc_local((num_routed_experts_per_thread,), dtype=T.float32) + key_local = T.alloc_local((num_routed_experts_per_thread,), dtype=T.uint64) idx_local = T.alloc_local((num_routed_experts_per_thread,), dtype=T.int32) topk_group_idx_shared = T.alloc_shared((num_tokens_per_block, num_topk_groups), dtype=T.int32) - topk_scores_local = T.alloc_local(num_topk, dtype=T.float32) topk_idx_local = T.alloc_local(num_topk, dtype=T.int32) topk_group_idx_local = T.alloc_local(num_topk_groups, dtype=T.int32) logit_max_var = T.alloc_var(dtype=T.float32) logit_sum_var = T.alloc_var(dtype=T.float32) - other_idx = T.alloc_var(dtype=T.int32) + topk_key_var = T.alloc_var(dtype=T.uint64) + other_key = T.alloc_var(dtype=T.uint64) topk_score_var = T.alloc_var(dtype=T.float32) topk_idx_var = T.alloc_var(dtype=T.int64) topk_sum_var = T.alloc_var(dtype=T.float32) - - # Tokens with mask = 0 does not participate in routing - if mask_exists and not mask[global_token_idx]: - if lane_idx < num_topk and unmapped_topk_idx_exists: - unmapped_topk_idx[global_token_idx, lane_idx] = -1 - if lane_idx < num_physical_topk: - topk_idx[global_token_idx, lane_idx] = -1 - topk_weights[global_token_idx, lane_idx] = 0.0 - T.thread_return() + token_exists = T.alloc_var(dtype=T.int32, init=0) + token_active = T.alloc_var(dtype=T.int32, init=0) + fixed_route = T.alloc_var(dtype=T.int32, init=0) + + if global_token_idx < num_tokens: + token_exists = 1 + token_active = 1 + if mask_exists and not mask[global_token_idx]: + token_active = 0 + if fix_routing_mask_exists and fix_routing_mask[global_token_idx]: + fixed_route = 1 # Load and do activation functions logit_max_var = -T.infinity(T.float32) @@ -142,8 +152,12 @@ def top2_sum_gate_kernel( start_expert_idx = i * num_vectorize * warp_size + lane_idx * num_vectorize if start_expert_idx < num_routed_experts: for j in T.vectorized(num_vectorize): - scores_local[i * num_vectorize + j] = logits[global_token_idx, start_expert_idx + j] - bias_local[i * num_vectorize + j] = bias[start_expert_idx + j] + if token_active != 0: + scores_local[i * num_vectorize + j] = logits[global_token_idx, start_expert_idx + j] + bias_local[i * num_vectorize + j] = bias[start_expert_idx + j] + else: + scores_local[i * num_vectorize + j] = 0.0 + bias_local[i * num_vectorize + j] = 0.0 if scoring_type == 2: # SOFTMAX scores_shared[token_idx, start_expert_idx + j] = scores_local[i * num_vectorize + j] @@ -155,7 +169,7 @@ def top2_sum_gate_kernel( if scoring_type == 2: # SOFTMAX for i in T.unroll(num_routed_experts_per_thread): logit_max_var = T.max(logit_max_var, scores_local[i]) - logit_max_var = T.warp_reduce_max(logit_max_var) + warp_reduce_max(logit_max_var) for i in T.unroll(0, T.ceildiv(num_routed_experts, num_vectorize * warp_size)): if i * num_vectorize * warp_size + lane_idx * num_vectorize < num_routed_experts: for j in T.unroll(num_vectorize): @@ -195,20 +209,22 @@ def top2_sum_gate_kernel( # Ensure all shared memory stores are completed T.sync_warp() - if not fix_routing_mask_exists or not fix_routing_mask[global_token_idx]: + # Group ranking contains a synchronization. Run it for both logical + # 32-lane subgroups before diverging on the per-token fixed route. + if not skip_group_sort: + get_topk_group_idx( + scores_shared, + topk_group_idx_shared, + num_groups, + num_routed_experts_per_group, + num_topk_groups, + num_topk_sum, + num_vectorize_for_grouped_expert, + ) + + if fixed_route == 0: # Get `num_topk_groups` groups with the largest top2-sum if not skip_group_sort: - # Get topk group indices - get_topk_group_idx( - scores_shared, - topk_group_idx_shared, - num_groups, - num_routed_experts_per_group, - num_topk_groups, - num_topk_sum, - num_vectorize_for_grouped_expert, - ) - # Sort group indices in ascending order to ensure stable sort for i in T.vectorized(num_topk_groups): topk_group_idx_local[i] = topk_group_idx_shared[token_idx, i] @@ -228,25 +244,34 @@ def top2_sum_gate_kernel( scores_local[i] = scores_shared[token_idx, select_group_idx * num_routed_experts_per_group + lane_idx] idx_local[i] = select_group_idx * num_routed_experts_per_group + lane_idx - # Get topk via repeatly finding max + # Pack descending score and ascending index into one key so the + # warp tournament exchanges one value rather than score + index. + for i in T.unroll(num_routed_experts_per_thread): + score_bits = T.reinterpret(scores_local[i], T.uint32) + magnitude = score_bits & T.uint32(0x7FFFFFFF) + is_negative = ((score_bits & T.uint32(0x80000000)) != 0) & (magnitude != 0) + normalized_bits = T.Select(magnitude == 0, T.uint32(0), score_bits) + ordered_score = T.Select(is_negative, ~normalized_bits, normalized_bits ^ T.uint32(0x80000000)) + candidate_key = (T.uint64(ordered_score) << 32) | T.uint64(T.uint32(~idx_local[i])) + key_local[i] = T.Select(idx_local[i] >= 0, candidate_key, T.uint64(0)) + + # Get topk via repeatedly finding max. for k in T.unroll(num_topk): - # Get local max score - topk_scores_local[k] = -T.infinity(T.float32) + if k != 0: + for i in T.unroll(num_routed_experts_per_thread): + if topk_idx_local[k - 1] == idx_local[i]: + key_local[i] = 0 + + # Get local max key. + topk_key_var = T.uint64(0) for i in T.unroll(0, num_routed_experts_per_thread): - if k != 0 and topk_idx_local[k - 1] == idx_local[i]: - scores_local[i] = -T.infinity(T.float32) - # If j > i, then idx_local[j] > idx_local[i] - elif scores_local[i] > topk_scores_local[k]: - topk_scores_local[k] = scores_local[i] - topk_idx_local[k] = idx_local[i] - - # Get max score across all threads + topk_key_var = T.max(topk_key_var, key_local[i]) + + # Get max key across all threads. for i in T.unroll(5): - other_score = T.shfl_xor(topk_scores_local[k], 1 << i) - other_idx = T.shfl_xor(topk_idx_local[k], 1 << i) - if other_score > topk_scores_local[k] or (other_score == topk_scores_local[k] and other_idx < topk_idx_local[k]): - topk_scores_local[k] = other_score - topk_idx_local[k] = other_idx + other_key = T.shfl_xor(topk_key_var, 1 << i, width=32) + topk_key_var = T.max(topk_key_var, other_key) + topk_idx_local[k] = T.cast(T.uint32(~T.uint32(topk_key_var)), T.int32) topk_score_var = 0.0 if lane_idx < num_topk: @@ -262,43 +287,50 @@ def top2_sum_gate_kernel( # Get topk sum topk_sum_var = 1e-20 for i in T.unroll(num_topk): - topk_sum_var += T.shfl_sync(topk_score_var, i) + topk_sum_var += T.shfl_sync(topk_score_var, i, width=32) # Ensure one warp can handle one token T.device_assert(num_physical_topk <= warp_size) - # Normalize top-k weights - if lane_idx < num_topk: - # NOTES: If this fails, there may be some NaN values in logits input or internal error in the kernel - T.device_assert(topk_idx_var >= 0) - topk_score_var = topk_score_var / topk_sum_var * routed_scaling_factor - if unmapped_topk_idx_exists: - unmapped_topk_idx[global_token_idx, lane_idx] = topk_idx_var - elif lane_idx < num_physical_topk: - topk_score_var = 1.0 - topk_idx_var = lane_idx + (num_routed_experts - num_topk) - - # Map to physical experts - if to_physical_map_exists and lane_idx < num_physical_topk: - logical_expert_idx = topk_idx_var - num_duplicates = logical_count[logical_expert_idx] - duplicate_idx = (ep_rank + global_token_idx * large_prime_number) % num_duplicates - topk_idx_var = to_physical_map[logical_expert_idx, duplicate_idx] - - # Mask ETP idx - num_experts_per_rank = (num_routed_experts + num_extra_experts) // num_ep_ranks - num_experts_per_dp = num_experts_per_rank * num_tp_ranks - if lane_idx < num_physical_topk: - dst_ep_rank = topk_idx_var // num_experts_per_rank - if dst_ep_rank % num_tp_ranks != T.int64(tp_rank): - topk_idx_var = -1 - else: - topk_idx_var -= tp_rank * num_experts_per_rank - dst_dp_rank = topk_idx_var // num_experts_per_dp - topk_idx_var = topk_idx_var - dst_dp_rank * num_experts_per_dp + dst_dp_rank * num_experts_per_rank - topk_idx_var = T.if_then_else(topk_idx_var < 0, -1, topk_idx_var) - topk_idx[global_token_idx, lane_idx] = topk_idx_var - topk_weights[global_token_idx, lane_idx] = topk_score_var + if token_active != 0: + # Normalize top-k weights + if lane_idx < num_topk: + # NOTES: If this fails, there may be some NaN values in logits input or internal error in the kernel + T.device_assert(topk_idx_var >= 0) + topk_score_var = topk_score_var / topk_sum_var * routed_scaling_factor + if unmapped_topk_idx_exists: + unmapped_topk_idx[global_token_idx, lane_idx] = topk_idx_var + elif lane_idx < num_physical_topk: + topk_score_var = 1.0 + topk_idx_var = lane_idx + (num_routed_experts - num_topk) + + # Map to physical experts + if to_physical_map_exists and lane_idx < num_physical_topk: + logical_expert_idx = topk_idx_var + num_duplicates = logical_count[logical_expert_idx] + duplicate_idx = (ep_rank + global_token_idx * large_prime_number) % num_duplicates + topk_idx_var = to_physical_map[logical_expert_idx, duplicate_idx] + + # Mask ETP idx + num_experts_per_rank = (num_routed_experts + num_extra_experts) // num_ep_ranks + num_experts_per_dp = num_experts_per_rank * num_tp_ranks + if lane_idx < num_physical_topk: + dst_ep_rank = topk_idx_var // num_experts_per_rank + if dst_ep_rank % num_tp_ranks != T.int64(tp_rank): + topk_idx_var = -1 + else: + topk_idx_var -= tp_rank * num_experts_per_rank + dst_dp_rank = topk_idx_var // num_experts_per_dp + topk_idx_var = topk_idx_var - dst_dp_rank * num_experts_per_dp + dst_dp_rank * num_experts_per_rank + topk_idx_var = T.if_then_else(topk_idx_var < 0, -1, topk_idx_var) + topk_idx[global_token_idx, lane_idx] = topk_idx_var + topk_weights[global_token_idx, lane_idx] = topk_score_var + elif token_exists != 0: + if lane_idx < num_topk and unmapped_topk_idx_exists: + unmapped_topk_idx[global_token_idx, lane_idx] = -1 + if lane_idx < num_physical_topk: + topk_idx[global_token_idx, lane_idx] = -1 + topk_weights[global_token_idx, lane_idx] = 0.0 return top2_sum_gate_kernel diff --git a/tile_kernels/moe/topk_gate_kernel.py b/tile_kernels/moe/topk_gate_kernel.py index 1b2ce99..abaa4c2 100644 --- a/tile_kernels/moe/topk_gate_kernel.py +++ b/tile_kernels/moe/topk_gate_kernel.py @@ -4,8 +4,6 @@ import torch from tilelang import language as T -from tile_kernels.utils import align - @tilelang.jit( pass_configs={ @@ -14,43 +12,59 @@ ) def get_topk_gate_kernel(num_experts: int, num_topk: int): num_tokens = T.dynamic('num_tokens') - num_threads = 32 - num_aligned_experts = align(num_experts, num_threads) + subgroup_size = 32 + num_threads = 64 + num_tokens_per_block = num_threads // subgroup_size + num_experts_per_lane = (num_experts + subgroup_size - 1) // subgroup_size @T.prim_func def topk_gate_kernel( scores: T.Tensor[(num_tokens, num_experts), T.float32], topk_idx: T.Tensor[(num_tokens, num_topk), T.int64], ): - with T.Kernel(num_tokens, threads=num_threads) as pid: - scores_fragment = T.alloc_fragment((num_aligned_experts,), T.float32) - amax_fragment = T.alloc_fragment((1,), T.float32) - idx_fragment = T.alloc_fragment((num_aligned_experts,), T.int32) - idx_reducer = T.alloc_reducer((1,), T.int32, 'min', replication='all') - topk_idx_shared = T.alloc_shared((num_topk,), T.int32) - - for i in T.Parallel(num_aligned_experts): - if i < num_experts: - scores_fragment[i] = scores[pid, i] + with T.Kernel(T.ceildiv(num_tokens, num_tokens_per_block), threads=num_threads) as pid: + thread_idx = T.get_thread_binding() + token_idx = thread_idx // subgroup_size + lane_idx = thread_idx % subgroup_size + token = pid * num_tokens_per_block + token_idx + + key_local = T.alloc_local((num_experts_per_lane,), T.uint64) + idx_local = T.alloc_local((num_experts_per_lane,), T.int32) + topk_key_var = T.alloc_var(T.uint64) + other_key = T.alloc_var(T.uint64) + + for i in T.unroll(num_experts_per_lane): + expert_idx = lane_idx + i * subgroup_size + idx_local[i] = expert_idx + if token < num_tokens and expert_idx < num_experts: + score_bits = T.reinterpret(scores[token, expert_idx], T.uint32) + magnitude = score_bits & T.uint32(0x7FFFFFFF) + is_nan = magnitude > T.uint32(0x7F800000) + is_negative = ((score_bits & T.uint32(0x80000000)) != 0) & (magnitude != 0) + normalized_bits = T.Select(magnitude == 0, T.uint32(0), score_bits) + ordered_score = T.Select( + is_nan, + T.Select(is_negative, T.uint32(0), T.uint32(0xFFFFFFFF)), + T.Select(is_negative, ~normalized_bits, normalized_bits ^ T.uint32(0x80000000)), + ) + key_local[i] = (T.uint64(ordered_score) << 32) | T.uint64(T.uint32(~expert_idx)) else: - scores_fragment[i] = -T.infinity(T.float32) - for i in T.Parallel(num_aligned_experts): - idx_fragment[i] = i + key_local[i] = T.uint64(0) - # Get topk via repeatly finding max + # Repeated argmax with one 32-lane tournament per selected expert. for k in T.unroll(num_topk): - T.reduce_max(scores_fragment, amax_fragment) - T.fill(idx_reducer, T.max_value(T.int32)) - for i in T.Parallel(num_aligned_experts): - if scores_fragment[i] == amax_fragment[0]: - idx_reducer[0] = T.min(idx_reducer[0], idx_fragment[i]) - T.finalize_reducer(idx_reducer) - topk_idx_shared[k] = idx_reducer[0] - for i in T.Parallel(num_aligned_experts): - if idx_fragment[i] == idx_reducer[0]: - scores_fragment[i] = -T.infinity(T.float32) - - T.copy(topk_idx_shared, topk_idx[pid, 0], disable_tma=True) + topk_key_var = T.uint64(0) + for i in T.unroll(num_experts_per_lane): + topk_key_var = T.max(topk_key_var, key_local[i]) + for i in T.unroll(5): + other_key = T.shfl_xor(topk_key_var, 1 << i, width=subgroup_size) + topk_key_var = T.max(topk_key_var, other_key) + topk_idx_local = T.cast(T.uint32(~T.uint32(topk_key_var)), T.int32) + if token < num_tokens and lane_idx == 0: + topk_idx[token, k] = topk_idx_local + for i in T.unroll(num_experts_per_lane): + if idx_local[i] == topk_idx_local: + key_local[i] = T.uint64(0) return topk_gate_kernel diff --git a/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py b/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py index 77bb4f7..c44e1e8 100644 --- a/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py +++ b/tile_kernels/moe/topk_sum_and_topk_group_idx_kernel.py @@ -8,6 +8,56 @@ from tile_kernels.utils import align +@T.macro +def get_topk_group_idx_wave64( + scores_shared: T.SharedBuffer, + topk_group_idx_shared: T.SharedBuffer, + group_scores_shared: T.SharedBuffer, + num_groups: int, + num_experts_per_group: int, + num_topk_groups: int, + num_topk_sum: int, + num_vectorize_for_grouped_expert: int, +): + thread_idx = T.get_thread_binding() + token_idx = thread_idx // 32 + lane_idx = thread_idx % 32 + scores_vec_local = T.alloc_local((num_vectorize_for_grouped_expert,), dtype=T.float32) + + top1_var = T.alloc_var(dtype=T.float32, init=-T.infinity(T.float32)) + top2_var = T.alloc_var(dtype=T.float32, init=-T.infinity(T.float32)) + topk_sum_var = T.alloc_var(dtype=T.float32, init=-T.infinity(T.float32)) + count_var = T.alloc_var(dtype=T.int32, init=0) + + if lane_idx < num_groups: + num_vec_experts_per_group = num_experts_per_group // num_vectorize_for_grouped_expert + for i in T.unroll(num_vec_experts_per_group): + for j in T.vectorized(num_vectorize_for_grouped_expert): + vec_idx = (i + lane_idx) % num_vec_experts_per_group + scores_vec_local[j] = scores_shared[ + token_idx, lane_idx * num_experts_per_group + vec_idx * num_vectorize_for_grouped_expert + j + ] + if scores_vec_local[j] > top1_var: + top2_var = top1_var + top1_var = scores_vec_local[j] + elif scores_vec_local[j] > top2_var: + top2_var = scores_vec_local[j] + topk_sum_var = T.Select(num_topk_sum == 1, top1_var, top1_var + top2_var) + + group_scores_shared[token_idx, lane_idx] = topk_sum_var + T.sync_threads() + + if lane_idx < num_groups: + for i in T.unroll(num_groups): + other_topk_sum = group_scores_shared[token_idx, i] + if other_topk_sum > topk_sum_var or (other_topk_sum == topk_sum_var and i < lane_idx): + count_var += 1 + if count_var < num_topk_groups: + topk_group_idx_shared[token_idx, count_var] = lane_idx + + T.sync_threads() + + @tilelang.jit( pass_configs={ tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, @@ -21,8 +71,9 @@ def get_topk_sum_and_topk_group_idx_kernel( num_topk_groups: int, num_topk_sum: int, ): - num_threads = 32 num_experts = num_experts_per_group * num_groups + use_wave64 = num_experts % 64 == 0 + num_threads = 64 if use_wave64 else 32 num_aligned_experts = align(num_experts, num_threads) num_tokens_per_block = num_threads // 32 @@ -41,29 +92,54 @@ def topk_sum_and_topk_group_idx_kernel( scores: T.Tensor[(num_tokens, num_experts), T.float32], group_topk_idx: T.Tensor[(num_tokens, num_topk_groups), T.int64], ): - with T.Kernel(num_tokens, threads=num_threads) as pid: + with T.Kernel(T.ceildiv(num_tokens, num_tokens_per_block), threads=num_threads) as pid: scores_shared = T.alloc_shared((num_tokens_per_block, num_aligned_experts), T.float32) topk_group_idx_shared = T.alloc_shared((num_tokens_per_block, num_topk_groups), T.int32) + group_scores_shared = T.alloc_shared((num_tokens_per_block, 32), T.float32) thread_idx = T.get_thread_binding() - warp_idx = thread_idx // 32 + token_idx = thread_idx // 32 lane_idx = thread_idx % 32 - - T.copy(scores[pid * num_tokens_per_block, 0], scores_shared) - T.sync_warp() - - get_topk_group_idx( - scores_shared=scores_shared, - topk_group_idx_shared=topk_group_idx_shared, - num_groups=num_groups, - num_experts_per_group=num_experts_per_group, - num_topk_groups=num_topk_groups, - num_topk_sum=num_topk_sum, - num_vectorize_for_grouped_expert=num_vectorize_for_grouped_expert, - ) - - if lane_idx < num_topk_groups: - group_topk_idx[pid * num_tokens_per_block + warp_idx, lane_idx] = topk_group_idx_shared[warp_idx, lane_idx] + token_base = pid * num_tokens_per_block + token = token_base + token_idx + + if use_wave64: + if token_base + 1 < num_tokens: + T.copy(scores[token_base, 0], scores_shared) + elif token_base < num_tokens: + for i in T.Parallel(num_aligned_experts): + scores_shared[0, i] = T.Select(i < num_experts, scores[token_base, i], -T.infinity(T.float32)) + scores_shared[1, i] = -T.infinity(T.float32) + else: + for i in T.Parallel(num_aligned_experts): + scores_shared[token_idx, i] = -T.infinity(T.float32) + T.sync_threads() + + get_topk_group_idx_wave64( + scores_shared=scores_shared, + topk_group_idx_shared=topk_group_idx_shared, + group_scores_shared=group_scores_shared, + num_groups=num_groups, + num_experts_per_group=num_experts_per_group, + num_topk_groups=num_topk_groups, + num_topk_sum=num_topk_sum, + num_vectorize_for_grouped_expert=num_vectorize_for_grouped_expert, + ) + else: + T.copy(scores[pid, 0], scores_shared) + T.sync_warp() + get_topk_group_idx( + scores_shared=scores_shared, + topk_group_idx_shared=topk_group_idx_shared, + num_groups=num_groups, + num_experts_per_group=num_experts_per_group, + num_topk_groups=num_topk_groups, + num_topk_sum=num_topk_sum, + num_vectorize_for_grouped_expert=num_vectorize_for_grouped_expert, + ) + + if token < num_tokens and lane_idx < num_topk_groups: + group_topk_idx[token, lane_idx] = topk_group_idx_shared[token_idx, lane_idx] return topk_sum_and_topk_group_idx_kernel