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