From c312cdd3a7dba789a6603abfc06d4514466daec6 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Wed, 1 Jul 2026 12:32:18 -0700 Subject: [PATCH] Upgrading tvm-ffi/sgl-deep-gemm/tilelang (#29554) --- python/pyproject.toml | 6 +-- .../layers/attention/dsa/tilelang_kernel.py | 24 ++--------- .../layers/deep_gemm_wrapper/compile_utils.py | 4 +- python/sglang/srt/layers/mhc.py | 2 - python/sglang/srt/layers/moe/mega_moe.py | 40 +++++++++++++++++-- 5 files changed, 45 insertions(+), 31 deletions(-) diff --git a/python/pyproject.toml b/python/pyproject.toml index 4a35a1e65..c196aedcd 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -18,7 +18,7 @@ classifiers = [ dependencies = [ "aiohttp", "anthropic>=0.20.0", - "apache-tvm-ffi==0.1.9", + "apache-tvm-ffi==0.1.11", "av ; sys_platform == 'linux' and (platform_machine == 'aarch64' or platform_machine == 'arm64' or platform_machine == 'armv7l')", "blobfile==3.0.0", "build", @@ -65,12 +65,12 @@ dependencies = [ "scipy", "sentencepiece", "setproctitle", - "sgl-deep-gemm==0.1.3", + "sgl-deep-gemm==0.1.4", "sglang-kernel==0.4.4", "smg-grpc-servicer>=0.5.0", "soundfile==0.13.1", "tiktoken", - "tilelang==0.1.8", + "tilelang==0.1.11", "timm==1.0.16", "tokenspeed_mla==0.1.7", "torch==2.11.0", diff --git a/python/sglang/srt/layers/attention/dsa/tilelang_kernel.py b/python/sglang/srt/layers/attention/dsa/tilelang_kernel.py index 62509c308..a648a8de8 100644 --- a/python/sglang/srt/layers/attention/dsa/tilelang_kernel.py +++ b/python/sglang/srt/layers/attention/dsa/tilelang_kernel.py @@ -577,22 +577,14 @@ def sparse_attention_fwd_kernel_v2( acc_s[h_i, bi_i] = T.if_then_else( is_kv_valid_0[bi_i], 0, -T.infinity(acc_s.dtype) ) - T.gemm( - Q_shared_l, KV_shared_0_l, acc_s, transpose_B=True, wg_wait=-1 - ) - T.gemm( - Q_shared_r, KV_shared_0_r, acc_s, transpose_B=True, wg_wait=-1 - ) + T.gemm(Q_shared_l, KV_shared_0_l, acc_s, transpose_B=True) + T.gemm(Q_shared_r, KV_shared_0_r, acc_s, transpose_B=True) T.gemm( Q_tail_shared, K_tail_shared_0, acc_s, transpose_B=True, - wg_wait=-1, ) - - T.wait_wgmma(0) - if i_i != 0: T.barrier_arrive(bar_sScale_and_sS_free) T.barrier_wait(bar_sScale_and_sS_free, ((i_i * 2) & 1) ^ 1) @@ -629,22 +621,14 @@ def sparse_attention_fwd_kernel_v2( acc_s[h_i, bi_i] = T.if_then_else( is_kv_valid_1[bi_i], 0, -T.infinity(acc_s.dtype) ) - T.gemm( - Q_shared_l, KV_shared_1_l, acc_s, transpose_B=True, wg_wait=-1 - ) - T.gemm( - Q_shared_r, KV_shared_1_r, acc_s, transpose_B=True, wg_wait=-1 - ) + T.gemm(Q_shared_l, KV_shared_1_l, acc_s, transpose_B=True) + T.gemm(Q_shared_r, KV_shared_1_r, acc_s, transpose_B=True) T.gemm( Q_tail_shared, K_tail_shared_1, acc_s, transpose_B=True, - wg_wait=-1, ) - - T.wait_wgmma(0) - T.barrier_arrive(bar_sScale_and_sS_free) T.barrier_wait(bar_sScale_and_sS_free, ((i_i * 2 + 1) & 1) ^ 1) diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index 2b424bef9..8ff9ac066 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -280,7 +280,7 @@ def _empty_token_fp8(size): *dims, k = size return ( torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn), - torch.empty( + torch.ones( (*dims, ceil_div(k, _BLOCK_SIZE)), device="cuda", dtype=torch.float32 ), ) @@ -290,7 +290,7 @@ def _empty_block_fp8(size): *dims, n, k = size return ( torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn), - torch.empty( + torch.ones( (*dims, ceil_div(n, _BLOCK_SIZE), ceil_div(k, _BLOCK_SIZE)), device="cuda", dtype=torch.float32, diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/srt/layers/mhc.py index 71d4d36df..11d26cae4 100644 --- a/python/sglang/srt/layers/mhc.py +++ b/python/sglang/srt/layers/mhc.py @@ -345,7 +345,6 @@ def mhc_pre_gemm_sqrsum_tilelang( out_frag, transpose_A=False, transpose_B=True, - wg_wait=0, clear_accum=False, ) sqrsum_l = T.alloc_fragment(token_block, T.float32) @@ -426,7 +425,6 @@ def mhc_pre_gemm_sqrsum_splitk_kernel( out_frag, transpose_A=False, transpose_B=True, - wg_wait=0, clear_accum=False, ) diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 8e8474f16..cad7ea92b 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -267,12 +267,42 @@ def _run_mega_routed( return y +def _interleave_mega_moe_gate_up(t: torch.Tensor, gran: int = 8) -> torch.Tensor: + # Match DeepGEMM's L1 gate/up layout: + # [gate: 0..7, up: 0..7, gate: 8..15, up: 8..15, ...]. + num_groups, n, *rest = t.shape + half = n // 2 + gate = t[:, :half].reshape(num_groups, half // gran, gran, *rest) + up = t[:, half:].reshape(num_groups, half // gran, gran, *rest) + result = torch.stack([gate, up], dim=2).reshape(num_groups, n, *rest) + return torch.empty_like(t).copy_(result) + + +def _interleave_mega_moe_l1_weights( + l1_weights: tuple[torch.Tensor, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + return ( + _interleave_mega_moe_gate_up(l1_weights[0]), + _interleave_mega_moe_gate_up(l1_weights[1]), + ) + + +def _transpose_mega_moe_sf_for_utccp(sf: torch.Tensor) -> torch.Tensor: + num_groups, mn, packed_sf_k = sf.shape + assert sf.dtype == torch.int and mn % 128 == 0 + result = ( + sf.reshape(num_groups, -1, 4, 32, packed_sf_k) + .transpose(2, 3) + .reshape(num_groups, mn, packed_sf_k) + ) + return torch.empty_like(sf).copy_(result) + + def build_mega_moe_experts_weights(experts) -> None: from deep_gemm import ( transform_sf_into_required_layout, transform_weights_for_mega_moe, ) - from deep_gemm.mega import _interleave_l1_weights, _transpose_sf_for_utccp if getattr(experts, "_mega_moe_weights_built", False): return @@ -311,9 +341,11 @@ def build_mega_moe_experts_weights(experts) -> None: # the deep-ep path consumes the non-transposed interleaved scale and a # swizzle-aware activation kernel. L2 weight is untouched by the mega # transform, so the existing `w2_weight.data` is shared directly. - w13_interleaved, w13_sf_interleaved = _interleave_l1_weights((w13, w13_sf)) - w13_sf_utccp = _transpose_sf_for_utccp(w13_sf_interleaved) - w2_sf_utccp = _transpose_sf_for_utccp(w2_sf) + w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights( + (w13, w13_sf) + ) + w13_sf_utccp = _transpose_mega_moe_sf_for_utccp(w13_sf_interleaved) + w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf) experts.w13_weight.data = w13_interleaved experts.w13_weight_scale_inv.data = w13_sf_interleaved