Upgrading tvm-ffi/sgl-deep-gemm/tilelang (#29554)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user