Upgrading tvm-ffi/sgl-deep-gemm/tilelang (#29554)

This commit is contained in:
Baizhou Zhang
2026-07-01 12:32:18 -07:00
committed by GitHub
parent 779ea4a9b5
commit c312cdd3a7
5 changed files with 45 additions and 31 deletions
+3 -3
View File
@@ -18,7 +18,7 @@ classifiers = [
dependencies = [ dependencies = [
"aiohttp", "aiohttp",
"anthropic>=0.20.0", "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')", "av ; sys_platform == 'linux' and (platform_machine == 'aarch64' or platform_machine == 'arm64' or platform_machine == 'armv7l')",
"blobfile==3.0.0", "blobfile==3.0.0",
"build", "build",
@@ -65,12 +65,12 @@ dependencies = [
"scipy", "scipy",
"sentencepiece", "sentencepiece",
"setproctitle", "setproctitle",
"sgl-deep-gemm==0.1.3", "sgl-deep-gemm==0.1.4",
"sglang-kernel==0.4.4", "sglang-kernel==0.4.4",
"smg-grpc-servicer>=0.5.0", "smg-grpc-servicer>=0.5.0",
"soundfile==0.13.1", "soundfile==0.13.1",
"tiktoken", "tiktoken",
"tilelang==0.1.8", "tilelang==0.1.11",
"timm==1.0.16", "timm==1.0.16",
"tokenspeed_mla==0.1.7", "tokenspeed_mla==0.1.7",
"torch==2.11.0", "torch==2.11.0",
@@ -577,22 +577,14 @@ def sparse_attention_fwd_kernel_v2(
acc_s[h_i, bi_i] = T.if_then_else( acc_s[h_i, bi_i] = T.if_then_else(
is_kv_valid_0[bi_i], 0, -T.infinity(acc_s.dtype) is_kv_valid_0[bi_i], 0, -T.infinity(acc_s.dtype)
) )
T.gemm( T.gemm(Q_shared_l, KV_shared_0_l, acc_s, transpose_B=True)
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)
)
T.gemm(
Q_shared_r, KV_shared_0_r, acc_s, transpose_B=True, wg_wait=-1
)
T.gemm( T.gemm(
Q_tail_shared, Q_tail_shared,
K_tail_shared_0, K_tail_shared_0,
acc_s, acc_s,
transpose_B=True, transpose_B=True,
wg_wait=-1,
) )
T.wait_wgmma(0)
if i_i != 0: if i_i != 0:
T.barrier_arrive(bar_sScale_and_sS_free) T.barrier_arrive(bar_sScale_and_sS_free)
T.barrier_wait(bar_sScale_and_sS_free, ((i_i * 2) & 1) ^ 1) 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( acc_s[h_i, bi_i] = T.if_then_else(
is_kv_valid_1[bi_i], 0, -T.infinity(acc_s.dtype) is_kv_valid_1[bi_i], 0, -T.infinity(acc_s.dtype)
) )
T.gemm( T.gemm(Q_shared_l, KV_shared_1_l, acc_s, transpose_B=True)
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)
)
T.gemm(
Q_shared_r, KV_shared_1_r, acc_s, transpose_B=True, wg_wait=-1
)
T.gemm( T.gemm(
Q_tail_shared, Q_tail_shared,
K_tail_shared_1, K_tail_shared_1,
acc_s, acc_s,
transpose_B=True, transpose_B=True,
wg_wait=-1,
) )
T.wait_wgmma(0)
T.barrier_arrive(bar_sScale_and_sS_free) T.barrier_arrive(bar_sScale_and_sS_free)
T.barrier_wait(bar_sScale_and_sS_free, ((i_i * 2 + 1) & 1) ^ 1) T.barrier_wait(bar_sScale_and_sS_free, ((i_i * 2 + 1) & 1) ^ 1)
@@ -280,7 +280,7 @@ def _empty_token_fp8(size):
*dims, k = size *dims, k = size
return ( return (
torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn), torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn),
torch.empty( torch.ones(
(*dims, ceil_div(k, _BLOCK_SIZE)), device="cuda", dtype=torch.float32 (*dims, ceil_div(k, _BLOCK_SIZE)), device="cuda", dtype=torch.float32
), ),
) )
@@ -290,7 +290,7 @@ def _empty_block_fp8(size):
*dims, n, k = size *dims, n, k = size
return ( return (
torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn), torch.empty(size, device="cuda", dtype=torch.float8_e4m3fn),
torch.empty( torch.ones(
(*dims, ceil_div(n, _BLOCK_SIZE), ceil_div(k, _BLOCK_SIZE)), (*dims, ceil_div(n, _BLOCK_SIZE), ceil_div(k, _BLOCK_SIZE)),
device="cuda", device="cuda",
dtype=torch.float32, dtype=torch.float32,
-2
View File
@@ -345,7 +345,6 @@ def mhc_pre_gemm_sqrsum_tilelang(
out_frag, out_frag,
transpose_A=False, transpose_A=False,
transpose_B=True, transpose_B=True,
wg_wait=0,
clear_accum=False, clear_accum=False,
) )
sqrsum_l = T.alloc_fragment(token_block, T.float32) sqrsum_l = T.alloc_fragment(token_block, T.float32)
@@ -426,7 +425,6 @@ def mhc_pre_gemm_sqrsum_splitk_kernel(
out_frag, out_frag,
transpose_A=False, transpose_A=False,
transpose_B=True, transpose_B=True,
wg_wait=0,
clear_accum=False, clear_accum=False,
) )
+36 -4
View File
@@ -267,12 +267,42 @@ def _run_mega_routed(
return y 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: def build_mega_moe_experts_weights(experts) -> None:
from deep_gemm import ( from deep_gemm import (
transform_sf_into_required_layout, transform_sf_into_required_layout,
transform_weights_for_mega_moe, 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): if getattr(experts, "_mega_moe_weights_built", False):
return 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 # the deep-ep path consumes the non-transposed interleaved scale and a
# swizzle-aware activation kernel. L2 weight is untouched by the mega # swizzle-aware activation kernel. L2 weight is untouched by the mega
# transform, so the existing `w2_weight.data` is shared directly. # transform, so the existing `w2_weight.data` is shared directly.
w13_interleaved, w13_sf_interleaved = _interleave_l1_weights((w13, w13_sf)) w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights(
w13_sf_utccp = _transpose_sf_for_utccp(w13_sf_interleaved) (w13, w13_sf)
w2_sf_utccp = _transpose_sf_for_utccp(w2_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.data = w13_interleaved
experts.w13_weight_scale_inv.data = w13_sf_interleaved experts.w13_weight_scale_inv.data = w13_sf_interleaved