[DeepSeek-V4] Add mhc_fused_post_pre kernel (#25976)
Co-authored-by: Qichao Li <liqichao@baidu.com>
This commit is contained in:
@@ -657,6 +657,7 @@ class Envs:
|
|||||||
SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True)
|
SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True)
|
||||||
SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True)
|
SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True)
|
||||||
SGLANG_OPT_USE_TRITON_FUSED_MHC = EnvBool(True)
|
SGLANG_OPT_USE_TRITON_FUSED_MHC = EnvBool(True)
|
||||||
|
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
|
||||||
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
||||||
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
|
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
|
||||||
SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True)
|
SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True)
|
||||||
|
|||||||
@@ -896,3 +896,500 @@ def mhc_post(
|
|||||||
residual.shape[-1],
|
residual.shape[-1],
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@tilelang.jit(
|
||||||
|
pass_configs={
|
||||||
|
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||||
|
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||||
|
tilelang.PassConfigKey.TL_PTXAS_REGISTER_USAGE_LEVEL: 10,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
def mhc_fused_post_pre_fma_tilelang(
|
||||||
|
prev_comb_mix,
|
||||||
|
prev_residual,
|
||||||
|
prev_post_mix,
|
||||||
|
hidden_in,
|
||||||
|
pre_fn,
|
||||||
|
mixes_partial_out,
|
||||||
|
sqrsum_partial_out,
|
||||||
|
cur_residual_out,
|
||||||
|
hc: int,
|
||||||
|
hidden_size: int,
|
||||||
|
num_mix_outputs: int,
|
||||||
|
n_thr: int = 256,
|
||||||
|
tile_mix_outputs: int = 1,
|
||||||
|
split_k: int = 1,
|
||||||
|
) -> tilelang.JITKernel:
|
||||||
|
num_tokens = T.dynamic("num_tokens")
|
||||||
|
split_k = T.dynamic("split_k")
|
||||||
|
|
||||||
|
hidden_per_split = (hidden_size + split_k - 1) // split_k
|
||||||
|
num_mix_output_tiles = (num_mix_outputs + tile_mix_outputs - 1) // tile_mix_outputs
|
||||||
|
|
||||||
|
prev_comb_mix: T.Tensor((num_tokens, hc, hc), T.float32)
|
||||||
|
prev_residual: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16)
|
||||||
|
prev_post_mix: T.Tensor((num_tokens, hc), T.float32)
|
||||||
|
hidden_in: T.Tensor((num_tokens, hidden_size), T.bfloat16)
|
||||||
|
pre_fn: T.Tensor((num_mix_outputs, hc, hidden_size), T.float32)
|
||||||
|
|
||||||
|
mixes_partial_out: T.Tensor((split_k, num_tokens, num_mix_outputs), T.float32)
|
||||||
|
sqrsum_partial_out: T.Tensor((split_k, num_tokens), T.float32)
|
||||||
|
cur_residual_out: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16)
|
||||||
|
|
||||||
|
hidden_iters_per_thread = (hidden_per_split + n_thr - 1) // n_thr
|
||||||
|
num_warps = n_thr // 32
|
||||||
|
|
||||||
|
ENABLE_PDL = is_arch_support_pdl()
|
||||||
|
|
||||||
|
# CTA assignment:
|
||||||
|
# token_idx : this CTA handles one token.
|
||||||
|
# mix_output_tile_idx : this CTA handles a small tile of mix output columns.
|
||||||
|
# For HC=4, num_mix_outputs = 24:
|
||||||
|
# [0:4] -> pre logits
|
||||||
|
# [4:8] -> post logits
|
||||||
|
# [8:24] -> comb logits
|
||||||
|
# hidden_split_idx : this CTA handles one split of the hidden dimension.
|
||||||
|
#
|
||||||
|
# Thread assignment inside one CTA:
|
||||||
|
# Each thread owns several hidden positions in this hidden split:
|
||||||
|
# hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx
|
||||||
|
#
|
||||||
|
# For each owned hidden_idx, the thread computes:
|
||||||
|
# 1. post result: cur_residual[token, :, hidden_idx]
|
||||||
|
# 2. sqrsum partial for pre RMS
|
||||||
|
# 3. GEMM partial for several mix output columns
|
||||||
|
with T.Kernel(
|
||||||
|
num_tokens,
|
||||||
|
num_mix_output_tiles,
|
||||||
|
split_k,
|
||||||
|
threads=n_thr,
|
||||||
|
) as (token_idx, mix_output_tile_idx, hidden_split_idx):
|
||||||
|
thread_idx = T.get_thread_binding()
|
||||||
|
warp_idx = T.get_warp_idx()
|
||||||
|
lane_idx = T.get_lane_idx()
|
||||||
|
|
||||||
|
warp_partials = T.alloc_shared((num_warps, tile_mix_outputs + 1), T.float32)
|
||||||
|
post_mix_smem = T.alloc_shared((hc,), T.float32)
|
||||||
|
comb_mix_smem = T.alloc_shared((hc, hc), T.float32)
|
||||||
|
|
||||||
|
post_mix_for_token = T.alloc_local((hc,), T.float32)
|
||||||
|
comb_mix_for_token = T.alloc_local((hc, hc), T.float32)
|
||||||
|
|
||||||
|
mix_acc = T.alloc_local((tile_mix_outputs,), T.float32)
|
||||||
|
sqrsum_acc = T.alloc_local((1,), T.float32)
|
||||||
|
cur_residual_values = T.alloc_local((hc,), T.float32)
|
||||||
|
|
||||||
|
T.clear(mix_acc)
|
||||||
|
T.clear(sqrsum_acc)
|
||||||
|
|
||||||
|
hidden_split_start = hidden_split_idx * hidden_per_split
|
||||||
|
|
||||||
|
if ENABLE_PDL:
|
||||||
|
T.pdl_sync()
|
||||||
|
|
||||||
|
# Load post/comb coefficients for this token.
|
||||||
|
#
|
||||||
|
# PyTorch equivalent:
|
||||||
|
# post = prev_post_mix[token_idx] # [HC]
|
||||||
|
# comb = prev_comb_mix[token_idx] # [HC, HC]
|
||||||
|
T.copy(prev_post_mix[token_idx, 0], post_mix_smem)
|
||||||
|
T.copy(prev_comb_mix[token_idx, 0, 0], comb_mix_smem)
|
||||||
|
|
||||||
|
for route_idx in T.unroll(hc):
|
||||||
|
post_mix_for_token[route_idx] = post_mix_smem[route_idx]
|
||||||
|
|
||||||
|
for old_route_idx in T.unroll(hc):
|
||||||
|
for new_route_idx in T.unroll(hc):
|
||||||
|
comb_mix_for_token[old_route_idx, new_route_idx] = comb_mix_smem[
|
||||||
|
old_route_idx, new_route_idx
|
||||||
|
]
|
||||||
|
|
||||||
|
for hidden_iter in T.serial(hidden_iters_per_thread):
|
||||||
|
hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx
|
||||||
|
|
||||||
|
if hidden_idx < hidden_size:
|
||||||
|
# Step A: fused post.
|
||||||
|
#
|
||||||
|
# PyTorch equivalent:
|
||||||
|
# cur_residual =
|
||||||
|
# post.unsqueeze(-1) * hidden_in.unsqueeze(1)
|
||||||
|
# + (
|
||||||
|
# comb.unsqueeze(-1)
|
||||||
|
# * prev_residual.unsqueeze(2)
|
||||||
|
# ).sum(dim=1)
|
||||||
|
#
|
||||||
|
# Scalar form for this token and this hidden position:
|
||||||
|
# cur_residual[j, h]
|
||||||
|
# = post[j] * hidden_in[h]
|
||||||
|
# + sum_k comb[k, j] * prev_residual[k, h]
|
||||||
|
for new_route_idx in T.unroll(hc):
|
||||||
|
cur_residual_values[new_route_idx] = (
|
||||||
|
post_mix_for_token[new_route_idx]
|
||||||
|
* hidden_in[token_idx, hidden_idx]
|
||||||
|
)
|
||||||
|
|
||||||
|
for old_route_idx in T.unroll(hc):
|
||||||
|
cur_residual_values[new_route_idx] += (
|
||||||
|
comb_mix_for_token[old_route_idx, new_route_idx]
|
||||||
|
* prev_residual[token_idx, old_route_idx, hidden_idx]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Match the unfused path:
|
||||||
|
# mhc_post writes bf16 residual,
|
||||||
|
# then mhc_pre reads bf16 residual.
|
||||||
|
for route_idx in T.unroll(hc):
|
||||||
|
cur_residual_values[route_idx] = T.bfloat16(
|
||||||
|
cur_residual_values[route_idx]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Step B1: pre sqrsum partial.
|
||||||
|
#
|
||||||
|
# PyTorch equivalent:
|
||||||
|
# x_flat = cur_residual.reshape(T, HC * H).float()
|
||||||
|
# sqrsum = (x_flat * x_flat).sum(dim=-1)
|
||||||
|
#
|
||||||
|
# Only mix_output_tile_idx == 0 writes cur_residual and sqrsum,
|
||||||
|
# otherwise different output-column CTAs would duplicate this work.
|
||||||
|
if mix_output_tile_idx == 0:
|
||||||
|
for route_idx in T.unroll(hc):
|
||||||
|
cur_residual_out[token_idx, route_idx, hidden_idx] = (
|
||||||
|
cur_residual_values[route_idx]
|
||||||
|
)
|
||||||
|
sqrsum_acc[0] += (
|
||||||
|
cur_residual_values[route_idx]
|
||||||
|
* cur_residual_values[route_idx]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Step B2: pre GEMM partial.
|
||||||
|
#
|
||||||
|
# PyTorch equivalent:
|
||||||
|
# mixes = F.linear(x_flat, fn)
|
||||||
|
#
|
||||||
|
# Scalar form:
|
||||||
|
# mixes[token, o] +=
|
||||||
|
# pre_fn[o, route, hidden] * cur_residual[route, hidden]
|
||||||
|
#
|
||||||
|
# This CTA computes only tile_mix_outputs columns of mixes.
|
||||||
|
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||||
|
mix_output_idx = (
|
||||||
|
mix_output_tile_idx * tile_mix_outputs + tile_col_idx
|
||||||
|
)
|
||||||
|
|
||||||
|
if mix_output_idx < num_mix_outputs:
|
||||||
|
for route_idx in T.unroll(hc):
|
||||||
|
mix_acc[tile_col_idx] += (
|
||||||
|
pre_fn[mix_output_idx, route_idx, hidden_idx]
|
||||||
|
* cur_residual_values[route_idx]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Reduce thread partials inside each warp.
|
||||||
|
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||||
|
mix_acc[tile_col_idx] = T.warp_reduce_sum(mix_acc[tile_col_idx])
|
||||||
|
|
||||||
|
if mix_output_tile_idx == 0:
|
||||||
|
sqrsum_acc[0] = T.warp_reduce_sum(sqrsum_acc[0])
|
||||||
|
|
||||||
|
# One lane per warp writes warp-level partials to shared memory.
|
||||||
|
if lane_idx == 0:
|
||||||
|
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||||
|
warp_partials[warp_idx, tile_col_idx] = mix_acc[tile_col_idx]
|
||||||
|
|
||||||
|
if mix_output_tile_idx == 0:
|
||||||
|
warp_partials[warp_idx, tile_mix_outputs] = sqrsum_acc[0]
|
||||||
|
|
||||||
|
T.sync_threads()
|
||||||
|
|
||||||
|
# Reduce across warps and write split partials.
|
||||||
|
#
|
||||||
|
# The full PyTorch result would be:
|
||||||
|
# mixes = F.linear(cur_residual.reshape(T, HC * H), fn)
|
||||||
|
# sqrsum = (cur_residual.float() ** 2).sum(dim=(1, 2))
|
||||||
|
#
|
||||||
|
# This kernel is split along hidden, so each CTA writes only:
|
||||||
|
# mixes_partial_out[hidden_split_idx, token, o]
|
||||||
|
# sqrsum_partial_out[hidden_split_idx, token]
|
||||||
|
#
|
||||||
|
# Later mhc_pre_big_fuse does:
|
||||||
|
# mixes = mixes_partial_out.sum(dim=0)
|
||||||
|
# sqrsum = sqrsum_partial_out.sum(dim=0)
|
||||||
|
# rms = rsqrt(sqrsum / (HC * H) + eps)
|
||||||
|
# mixes *= rms
|
||||||
|
# mixes -> pre/post/comb
|
||||||
|
# layer_input = sum_j pre[j] * cur_residual[j]
|
||||||
|
if warp_idx == 0:
|
||||||
|
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||||
|
mix_output_idx = mix_output_tile_idx * tile_mix_outputs + tile_col_idx
|
||||||
|
|
||||||
|
if mix_output_idx < num_mix_outputs and lane_idx == tile_col_idx:
|
||||||
|
mix_output_partial = T.alloc_var(T.float32, init=0.0)
|
||||||
|
|
||||||
|
for reduce_warp_idx in T.unroll(num_warps):
|
||||||
|
mix_output_partial += warp_partials[
|
||||||
|
reduce_warp_idx, tile_col_idx
|
||||||
|
]
|
||||||
|
|
||||||
|
mixes_partial_out[hidden_split_idx, token_idx, mix_output_idx] = (
|
||||||
|
mix_output_partial
|
||||||
|
)
|
||||||
|
|
||||||
|
if mix_output_tile_idx == 0 and lane_idx == 0:
|
||||||
|
sqrsum_partial = T.alloc_var(T.float32, init=0.0)
|
||||||
|
|
||||||
|
for reduce_warp_idx in T.unroll(num_warps):
|
||||||
|
sqrsum_partial += warp_partials[reduce_warp_idx, tile_mix_outputs]
|
||||||
|
|
||||||
|
sqrsum_partial_out[hidden_split_idx, token_idx] = sqrsum_partial
|
||||||
|
|
||||||
|
if ENABLE_PDL:
|
||||||
|
T.pdl_trigger()
|
||||||
|
|
||||||
|
|
||||||
|
def mhc_fused_post_pre(
|
||||||
|
x: torch.Tensor,
|
||||||
|
residual: torch.Tensor,
|
||||||
|
post_layer_mix: torch.Tensor,
|
||||||
|
comb_res_mix: torch.Tensor,
|
||||||
|
fn: torch.Tensor,
|
||||||
|
hc_scale: torch.Tensor,
|
||||||
|
hc_base: torch.Tensor,
|
||||||
|
rms_eps: float,
|
||||||
|
hc_pre_eps: float,
|
||||||
|
hc_sinkhorn_eps: float,
|
||||||
|
hc_post_mult_value: float,
|
||||||
|
sinkhorn_repeat: int,
|
||||||
|
n_splits: int = 1,
|
||||||
|
tile_n: int = 1,
|
||||||
|
*,
|
||||||
|
norm_weight: torch.Tensor | None = None,
|
||||||
|
norm_eps: float | None = None,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
"""Fuse the boundary between one mHC post step and the next mHC pre step.
|
||||||
|
|
||||||
|
The unfused sequence is ``mhc_post -> pre-norm GEMM -> mhc_pre big_fuse``.
|
||||||
|
This wrapper keeps the numerically sensitive ``mhc_pre_big_fuse`` stage,
|
||||||
|
including optional RMSNorm, but removes the separate post/pre boundary.
|
||||||
|
Small token batches use the FMA kernel above to combine ``mhc_post`` and the
|
||||||
|
pre-norm GEMM in one launch; larger batches keep DeepGEMM for throughput and
|
||||||
|
only fuse the Python/model-level scheduling boundary.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
residual_cur: post-mapped residual, shape (..., hc_mult, hidden_size)
|
||||||
|
post_mix_cur: shape (..., hc_mult, 1)
|
||||||
|
comb_mix_cur: shape (..., hc_mult, hc_mult)
|
||||||
|
layer_input_cur: shape (..., hidden_size)
|
||||||
|
"""
|
||||||
|
|
||||||
|
assert residual.dtype == torch.bfloat16
|
||||||
|
assert x.dtype == torch.bfloat16
|
||||||
|
assert post_layer_mix.dtype == torch.float32
|
||||||
|
assert comb_res_mix.dtype == torch.float32
|
||||||
|
assert fn.dtype == torch.float32
|
||||||
|
assert hc_scale.dtype == torch.float32
|
||||||
|
assert hc_base.dtype == torch.float32
|
||||||
|
|
||||||
|
hc_mult = residual.shape[-2]
|
||||||
|
hidden_size = residual.shape[-1]
|
||||||
|
hc_mult2 = hc_mult * hc_mult
|
||||||
|
hc_mult3 = hc_mult * 2 + hc_mult2
|
||||||
|
hc_hidden_size = hc_mult * hidden_size
|
||||||
|
outer_shape = residual.shape[:-2]
|
||||||
|
|
||||||
|
assert x.shape == (*outer_shape, hidden_size)
|
||||||
|
assert post_layer_mix.shape in (
|
||||||
|
(*outer_shape, hc_mult, 1),
|
||||||
|
(*outer_shape, hc_mult),
|
||||||
|
)
|
||||||
|
assert comb_res_mix.shape == (*outer_shape, hc_mult, hc_mult)
|
||||||
|
assert fn.shape == (hc_mult3, hc_hidden_size)
|
||||||
|
assert hc_scale.shape == (3,)
|
||||||
|
assert hc_base.shape == (hc_mult3,)
|
||||||
|
|
||||||
|
residual_flat = residual.view(-1, hc_mult, hidden_size)
|
||||||
|
num_tokens = residual_flat.shape[0]
|
||||||
|
if num_tokens == 0:
|
||||||
|
# Some DP/EP ranks can receive no tokens; return correctly typed empty
|
||||||
|
# tensors so later fused layers keep the same contracts as mhc_pre/hc_post.
|
||||||
|
return (
|
||||||
|
torch.empty_like(residual),
|
||||||
|
torch.empty(
|
||||||
|
(*outer_shape, hc_mult, 1), dtype=torch.float32, device=residual.device
|
||||||
|
),
|
||||||
|
torch.empty(
|
||||||
|
(*outer_shape, hc_mult, hc_mult),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=residual.device,
|
||||||
|
),
|
||||||
|
torch.empty(
|
||||||
|
(*outer_shape, hidden_size),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=residual.device,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
x_flat = x.view(num_tokens, hidden_size)
|
||||||
|
|
||||||
|
# The scalar-FMA kernel wins only for small batches where launch
|
||||||
|
# overhead dominates; beyond the threshold DeepGEMM's tensor-core path wins.
|
||||||
|
fma_token_threshold = 32
|
||||||
|
if num_tokens <= fma_token_threshold:
|
||||||
|
tile_n = 2 if num_tokens < 8 else 3
|
||||||
|
n_splits = 8 if (num_tokens < 8 and hidden_size <= 4096) else 4
|
||||||
|
else:
|
||||||
|
n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size)
|
||||||
|
|
||||||
|
gemm_out_mul = torch.empty(
|
||||||
|
n_splits,
|
||||||
|
num_tokens,
|
||||||
|
hc_mult3,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=residual.device,
|
||||||
|
)
|
||||||
|
gemm_out_sqrsum = torch.empty(
|
||||||
|
n_splits,
|
||||||
|
num_tokens,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=residual.device,
|
||||||
|
)
|
||||||
|
residual_cur = torch.empty_like(residual_flat)
|
||||||
|
|
||||||
|
if num_tokens <= fma_token_threshold:
|
||||||
|
# Small-batch path: one TileLang launch computes hc_post, the bf16
|
||||||
|
# residual write, GEMM partials, and the RMS square-sum partials.
|
||||||
|
mhc_fused_post_pre_fma_tilelang(
|
||||||
|
comb_res_mix.view(num_tokens, hc_mult, hc_mult),
|
||||||
|
residual_flat,
|
||||||
|
post_layer_mix.view(num_tokens, hc_mult),
|
||||||
|
x_flat,
|
||||||
|
fn.view(hc_mult3, hc_mult, hidden_size),
|
||||||
|
gemm_out_mul,
|
||||||
|
gemm_out_sqrsum,
|
||||||
|
residual_cur,
|
||||||
|
hc_mult,
|
||||||
|
hidden_size,
|
||||||
|
hc_mult3,
|
||||||
|
tile_mix_outputs=tile_n,
|
||||||
|
split_k=n_splits,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Large-batch path: keep the existing high-throughput TileLang hc_post +
|
||||||
|
# DeepGEMM pre-norm GEMM decomposition instead of replacing tensor cores.
|
||||||
|
mhc_post_tilelang(
|
||||||
|
comb_res_mix.view(num_tokens, hc_mult, hc_mult),
|
||||||
|
residual_flat,
|
||||||
|
post_layer_mix.view(num_tokens, hc_mult),
|
||||||
|
x_flat,
|
||||||
|
residual_cur,
|
||||||
|
hc_mult,
|
||||||
|
hidden_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get():
|
||||||
|
import deep_gemm
|
||||||
|
|
||||||
|
deep_gemm.tf32_hc_prenorm_gemm(
|
||||||
|
residual_cur.view(num_tokens, hc_hidden_size),
|
||||||
|
fn,
|
||||||
|
gemm_out_mul,
|
||||||
|
gemm_out_sqrsum,
|
||||||
|
num_splits=n_splits,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Fallback mirrors mhc_pre when DeepGEMM prenorm is disabled.
|
||||||
|
n_splits = 1
|
||||||
|
gemm_out_mul_2d = torch.empty(
|
||||||
|
num_tokens, hc_mult3, dtype=torch.float32, device=residual.device
|
||||||
|
)
|
||||||
|
gemm_out_sqrsum_1d = torch.empty(
|
||||||
|
num_tokens, dtype=torch.float32, device=residual.device
|
||||||
|
)
|
||||||
|
mhc_pre_gemm_sqrsum_tilelang(
|
||||||
|
residual_cur.view(num_tokens, hc_hidden_size),
|
||||||
|
fn,
|
||||||
|
gemm_out_mul_2d,
|
||||||
|
gemm_out_sqrsum_1d,
|
||||||
|
hc_mult3,
|
||||||
|
hc_hidden_size,
|
||||||
|
)
|
||||||
|
gemm_out_mul = gemm_out_mul_2d.unsqueeze(0)
|
||||||
|
gemm_out_sqrsum = gemm_out_sqrsum_1d.unsqueeze(0)
|
||||||
|
|
||||||
|
post_mix_cur = torch.empty(
|
||||||
|
num_tokens,
|
||||||
|
hc_mult,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=residual.device,
|
||||||
|
)
|
||||||
|
comb_mix_cur = torch.empty(
|
||||||
|
num_tokens,
|
||||||
|
hc_mult2,
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=residual.device,
|
||||||
|
)
|
||||||
|
layer_input_cur = torch.empty(
|
||||||
|
num_tokens,
|
||||||
|
hidden_size,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=residual.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
if norm_weight is not None:
|
||||||
|
# Final mhc_pre stage: convert GEMM partials into post/comb/layer_input
|
||||||
|
# and fuse the following RMSNorm when the model passed a norm weight.
|
||||||
|
assert norm_eps is not None
|
||||||
|
assert norm_weight.shape == (hidden_size,)
|
||||||
|
norm_weight_bf = (
|
||||||
|
norm_weight.bfloat16()
|
||||||
|
if norm_weight.dtype != torch.bfloat16
|
||||||
|
else norm_weight
|
||||||
|
)
|
||||||
|
if not norm_weight_bf.is_contiguous():
|
||||||
|
norm_weight_bf = norm_weight_bf.contiguous()
|
||||||
|
mhc_pre_big_fuse_with_norm_tilelang(
|
||||||
|
gemm_out_mul,
|
||||||
|
gemm_out_sqrsum,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
residual_cur,
|
||||||
|
post_mix_cur,
|
||||||
|
comb_mix_cur,
|
||||||
|
layer_input_cur,
|
||||||
|
norm_weight_bf,
|
||||||
|
hidden_size,
|
||||||
|
rms_eps,
|
||||||
|
hc_pre_eps,
|
||||||
|
hc_sinkhorn_eps,
|
||||||
|
hc_post_mult_value,
|
||||||
|
sinkhorn_repeat,
|
||||||
|
norm_eps,
|
||||||
|
n_splits,
|
||||||
|
hc_mult,
|
||||||
|
hc_mult3,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Same mhc_pre finalization without the model-layer RMSNorm.
|
||||||
|
mhc_pre_big_fuse_tilelang(
|
||||||
|
gemm_out_mul,
|
||||||
|
gemm_out_sqrsum,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
residual_cur,
|
||||||
|
post_mix_cur,
|
||||||
|
comb_mix_cur,
|
||||||
|
layer_input_cur,
|
||||||
|
hidden_size,
|
||||||
|
rms_eps,
|
||||||
|
hc_pre_eps,
|
||||||
|
hc_sinkhorn_eps,
|
||||||
|
hc_post_mult_value,
|
||||||
|
sinkhorn_repeat,
|
||||||
|
n_splits,
|
||||||
|
hc_mult,
|
||||||
|
hc_mult3,
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
residual_cur.view(*outer_shape, hc_mult, hidden_size),
|
||||||
|
post_mix_cur.view(*outer_shape, hc_mult, 1),
|
||||||
|
comb_mix_cur.view(*outer_shape, hc_mult, hc_mult),
|
||||||
|
layer_input_cur.view(*outer_shape, hidden_size),
|
||||||
|
)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
@@ -61,6 +62,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||||
|
from sglang.srt.layers.mhc import mhc_fused_post_pre
|
||||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
||||||
@@ -110,6 +112,18 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get()
|
_FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get()
|
||||||
|
_MHC_POST_MULT_VALUE = 2.0
|
||||||
|
|
||||||
|
|
||||||
|
def _is_fused_mhc_post_pre_enabled() -> bool:
|
||||||
|
# The fused path directly reuses TileLang mhc_post/mhc_pre kernels and their
|
||||||
|
# tensor layout assumptions, so keep it disabled when either dependency is off.
|
||||||
|
return (
|
||||||
|
envs.SGLANG_OPT_FUSE_MHC_POST_PRE.get()
|
||||||
|
and envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get()
|
||||||
|
and envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
_is_gfx95_supported = is_gfx95_supported()
|
_is_gfx95_supported = is_gfx95_supported()
|
||||||
@@ -976,6 +990,133 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32))
|
self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32))
|
||||||
self.rms_norm_eps = config.rms_norm_eps
|
self.rms_norm_eps = config.rms_norm_eps
|
||||||
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
||||||
|
self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled()
|
||||||
|
self._input_layernorm_weight_bf16 = None
|
||||||
|
self._post_attention_layernorm_weight_bf16 = None
|
||||||
|
|
||||||
|
def refresh_mhc_norm_weight_cache(self):
|
||||||
|
# Cache bf16 norm weights so the fused path does not allocate/cast per forward.
|
||||||
|
self._input_layernorm_weight_bf16 = (
|
||||||
|
self.input_layernorm.weight.data.bfloat16().contiguous()
|
||||||
|
)
|
||||||
|
self._post_attention_layernorm_weight_bf16 = (
|
||||||
|
self.post_attention_layernorm.weight.data.bfloat16().contiguous()
|
||||||
|
)
|
||||||
|
|
||||||
|
def prewarm_mhc_token_counts(
|
||||||
|
self, token_counts: Tuple[int, ...], device: torch.device
|
||||||
|
) -> None:
|
||||||
|
paths = (
|
||||||
|
(
|
||||||
|
"attn",
|
||||||
|
self.hc_attn_fn,
|
||||||
|
self.hc_attn_scale,
|
||||||
|
self.hc_attn_base,
|
||||||
|
self.input_layernorm,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"ffn",
|
||||||
|
self.hc_ffn_fn,
|
||||||
|
self.hc_ffn_scale,
|
||||||
|
self.hc_ffn_base,
|
||||||
|
self.post_attention_layernorm,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
for num_tokens in token_counts:
|
||||||
|
for path_name, hc_fn, hc_scale, hc_base, norm in paths:
|
||||||
|
tic = time.perf_counter()
|
||||||
|
residual = torch.empty(
|
||||||
|
(num_tokens, self.hc_mult, self.hidden_size),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
y, post, comb, _ = self.hc_pre(
|
||||||
|
residual,
|
||||||
|
hc_fn,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
norm=norm,
|
||||||
|
)
|
||||||
|
del residual, y, post, comb
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
logger.info(
|
||||||
|
"DeepSeek V4 MHC prewarm path=%s num_tokens=%s completed in %.3fs",
|
||||||
|
path_name,
|
||||||
|
num_tokens,
|
||||||
|
time.perf_counter() - tic,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_fused_mhc_post_pre:
|
||||||
|
for num_tokens in token_counts:
|
||||||
|
for path_name, hc_fn, hc_scale, hc_base, norm in paths:
|
||||||
|
tic = time.perf_counter()
|
||||||
|
# Dummy inputs matching the fused kernel's expected shapes.
|
||||||
|
x = torch.empty(
|
||||||
|
(num_tokens, self.hidden_size),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
residual = torch.empty(
|
||||||
|
(num_tokens, self.hc_mult, self.hidden_size),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
post_mix = torch.empty(
|
||||||
|
(num_tokens, self.hc_mult, 1),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
comb_mix = torch.empty(
|
||||||
|
(num_tokens, self.hc_mult, self.hc_mult),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
norm_weight = norm.weight.data.bfloat16().contiguous()
|
||||||
|
mhc_fused_post_pre(
|
||||||
|
x,
|
||||||
|
residual,
|
||||||
|
post_mix,
|
||||||
|
comb_mix,
|
||||||
|
hc_fn,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
self.rms_norm_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
_MHC_POST_MULT_VALUE,
|
||||||
|
self.hc_sinkhorn_iters,
|
||||||
|
norm_weight=norm_weight,
|
||||||
|
norm_eps=norm.variance_epsilon,
|
||||||
|
)
|
||||||
|
del x, residual, post_mix, comb_mix, norm_weight
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
logger.info(
|
||||||
|
"DeepSeek V4 MHC fused prewarm path=%s num_tokens=%s completed in %.3fs",
|
||||||
|
path_name,
|
||||||
|
num_tokens,
|
||||||
|
time.perf_counter() - tic,
|
||||||
|
)
|
||||||
|
|
||||||
|
def prewarm_mhc_token_count_buckets(
|
||||||
|
self, max_num_tokens: int, device: torch.device
|
||||||
|
) -> Tuple[int, ...]:
|
||||||
|
from sglang.srt.layers.mhc import get_mhc_pre_token_count_representatives
|
||||||
|
|
||||||
|
token_counts = get_mhc_pre_token_count_representatives(
|
||||||
|
max_num_tokens, self.hc_mult * self.hidden_size
|
||||||
|
)
|
||||||
|
if not token_counts:
|
||||||
|
return token_counts
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"DeepSeek V4 MHC prewarm max_num_tokens=%s representative token counts: %s",
|
||||||
|
max_num_tokens,
|
||||||
|
token_counts,
|
||||||
|
)
|
||||||
|
self.prewarm_mhc_token_counts(token_counts, device)
|
||||||
|
return token_counts
|
||||||
|
|
||||||
def hc_pre(
|
def hc_pre(
|
||||||
self,
|
self,
|
||||||
@@ -1001,9 +1142,9 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
|
|
||||||
if x.shape[0] == 0:
|
if x.shape[0] == 0:
|
||||||
y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device)
|
y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device)
|
||||||
post = torch.empty((0, self.hc_mult), dtype=dtype, device=x.device)
|
post = torch.empty((0, self.hc_mult), dtype=torch.float32, device=x.device)
|
||||||
comb = torch.empty(
|
comb = torch.empty(
|
||||||
(0, self.hc_mult, self.hc_mult), dtype=dtype, device=x.device
|
(0, self.hc_mult, self.hc_mult), dtype=torch.float32, device=x.device
|
||||||
)
|
)
|
||||||
return y, post, comb, False
|
return y, post, comb, False
|
||||||
|
|
||||||
@@ -1023,7 +1164,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
rms_eps=self.rms_norm_eps,
|
rms_eps=self.rms_norm_eps,
|
||||||
hc_pre_eps=self.hc_eps,
|
hc_pre_eps=self.hc_eps,
|
||||||
hc_sinkhorn_eps=self.hc_eps,
|
hc_sinkhorn_eps=self.hc_eps,
|
||||||
hc_post_mult_value=2.0,
|
hc_post_mult_value=_MHC_POST_MULT_VALUE,
|
||||||
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
||||||
**norm_kwargs,
|
**norm_kwargs,
|
||||||
)
|
)
|
||||||
@@ -1040,7 +1181,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
rms_eps=self.rms_norm_eps,
|
rms_eps=self.rms_norm_eps,
|
||||||
hc_pre_eps=self.hc_eps,
|
hc_pre_eps=self.hc_eps,
|
||||||
hc_sinkhorn_eps=self.hc_eps,
|
hc_sinkhorn_eps=self.hc_eps,
|
||||||
hc_post_mult_value=2.0,
|
hc_post_mult_value=_MHC_POST_MULT_VALUE,
|
||||||
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
||||||
)
|
)
|
||||||
return y, post.squeeze(-1), comb, False
|
return y, post.squeeze(-1), comb, False
|
||||||
@@ -1122,27 +1263,60 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
input_ids_global: torch.Tensor,
|
input_ids_global: torch.Tensor,
|
||||||
) -> torch.Tensor:
|
prev_residual: Optional[torch.Tensor] = None,
|
||||||
residual = hidden_states
|
prev_post: Optional[torch.Tensor] = None,
|
||||||
hidden_states, post, comb, norm_fused = self.hc_pre(
|
prev_comb: Optional[torch.Tensor] = None,
|
||||||
hidden_states,
|
) -> Tuple[
|
||||||
self.hc_attn_fn,
|
torch.Tensor,
|
||||||
self.hc_attn_scale,
|
Optional[torch.Tensor],
|
||||||
self.hc_attn_base,
|
Optional[torch.Tensor],
|
||||||
norm=self.input_layernorm,
|
Optional[torch.Tensor],
|
||||||
)
|
]:
|
||||||
if not norm_fused:
|
use_fused = self.use_fused_mhc_post_pre
|
||||||
if _use_aiter and _is_gfx95_supported:
|
|
||||||
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
|
if prev_residual is not None and use_fused:
|
||||||
hidden_states,
|
residual, post, comb, hidden_states = mhc_fused_post_pre(
|
||||||
self.input_layernorm.weight,
|
hidden_states,
|
||||||
self.rms_norm_eps,
|
prev_residual,
|
||||||
)
|
prev_post,
|
||||||
else:
|
prev_comb,
|
||||||
hidden_states = self.input_layernorm(hidden_states)
|
self.hc_attn_fn,
|
||||||
x_quant = None
|
self.hc_attn_scale,
|
||||||
else:
|
self.hc_attn_base,
|
||||||
|
self.rms_norm_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
_MHC_POST_MULT_VALUE,
|
||||||
|
self.hc_sinkhorn_iters,
|
||||||
|
norm_weight=(
|
||||||
|
self._input_layernorm_weight_bf16
|
||||||
|
if self._input_layernorm_weight_bf16 is not None
|
||||||
|
else self.input_layernorm.weight.data
|
||||||
|
),
|
||||||
|
norm_eps=self.input_layernorm.variance_epsilon,
|
||||||
|
)
|
||||||
x_quant = None
|
x_quant = None
|
||||||
|
else:
|
||||||
|
residual = hidden_states
|
||||||
|
hidden_states, post, comb, norm_fused = self.hc_pre(
|
||||||
|
hidden_states,
|
||||||
|
self.hc_attn_fn,
|
||||||
|
self.hc_attn_scale,
|
||||||
|
self.hc_attn_base,
|
||||||
|
norm=self.input_layernorm,
|
||||||
|
)
|
||||||
|
if not norm_fused:
|
||||||
|
if _use_aiter and _is_gfx95_supported:
|
||||||
|
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
|
||||||
|
hidden_states,
|
||||||
|
self.input_layernorm.weight,
|
||||||
|
self.rms_norm_eps,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states = self.input_layernorm(hidden_states)
|
||||||
|
x_quant = None
|
||||||
|
else:
|
||||||
|
x_quant = None
|
||||||
|
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
x=hidden_states,
|
x=hidden_states,
|
||||||
@@ -1151,35 +1325,58 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
x_quant=x_quant,
|
x_quant=x_quant,
|
||||||
)
|
)
|
||||||
|
|
||||||
fused_mhc = try_fused_hc_post_pre(
|
if use_fused:
|
||||||
hidden_states,
|
fused_mhc = try_fused_hc_post_pre(
|
||||||
residual,
|
hidden_states,
|
||||||
post,
|
residual,
|
||||||
comb,
|
post,
|
||||||
self.hc_ffn_fn.T,
|
comb,
|
||||||
self.hc_ffn_scale,
|
self.hc_ffn_fn.T,
|
||||||
self.hc_ffn_base,
|
self.hc_ffn_scale,
|
||||||
self.hc_mult,
|
self.hc_ffn_base,
|
||||||
self.rms_norm_eps,
|
self.hc_mult,
|
||||||
self.hc_eps,
|
self.rms_norm_eps,
|
||||||
2.0,
|
self.hc_eps,
|
||||||
self.hc_sinkhorn_iters,
|
_MHC_POST_MULT_VALUE,
|
||||||
_is_gfx95_supported,
|
self.hc_sinkhorn_iters,
|
||||||
)
|
_is_gfx95_supported,
|
||||||
if fused_mhc is not None:
|
)
|
||||||
residual, hidden_states, post, comb, norm_fused = fused_mhc
|
if fused_mhc is not None:
|
||||||
|
residual, hidden_states, post, comb, norm_fused = fused_mhc
|
||||||
|
else:
|
||||||
|
residual, post, comb, hidden_states = mhc_fused_post_pre(
|
||||||
|
hidden_states,
|
||||||
|
residual,
|
||||||
|
post.unsqueeze(-1) if post.ndim == 2 else post,
|
||||||
|
comb,
|
||||||
|
self.hc_ffn_fn,
|
||||||
|
self.hc_ffn_scale,
|
||||||
|
self.hc_ffn_base,
|
||||||
|
self.rms_norm_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
self.hc_eps,
|
||||||
|
_MHC_POST_MULT_VALUE,
|
||||||
|
self.hc_sinkhorn_iters,
|
||||||
|
norm_weight=(
|
||||||
|
self._post_attention_layernorm_weight_bf16
|
||||||
|
if self._post_attention_layernorm_weight_bf16 is not None
|
||||||
|
else self.post_attention_layernorm.weight.data
|
||||||
|
),
|
||||||
|
norm_eps=self.post_attention_layernorm.variance_epsilon,
|
||||||
|
)
|
||||||
|
norm_fused = True
|
||||||
else:
|
else:
|
||||||
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
||||||
residual = hidden_states # [n, hc, d]
|
residual = hidden_states
|
||||||
hidden_states, post, comb, norm_fused = self.hc_pre(
|
hidden_states, post, comb, norm_fused = self.hc_pre(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
self.hc_ffn_fn,
|
self.hc_ffn_fn,
|
||||||
self.hc_ffn_scale,
|
self.hc_ffn_scale,
|
||||||
self.hc_ffn_base,
|
self.hc_ffn_base,
|
||||||
norm=self.post_attention_layernorm,
|
norm=self.post_attention_layernorm,
|
||||||
) # -> [n, d]
|
)
|
||||||
if not norm_fused:
|
if not norm_fused:
|
||||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||||
|
|
||||||
_use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch)
|
_use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch)
|
||||||
_use_tp_moe_gather = (
|
_use_tp_moe_gather = (
|
||||||
@@ -1233,9 +1430,13 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
attn_tp_all_gather(gathered, hidden_states.contiguous())
|
attn_tp_all_gather(gathered, hidden_states.contiguous())
|
||||||
hidden_states = torch.cat(gathered)
|
hidden_states = torch.cat(gathered)
|
||||||
|
|
||||||
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
if not use_fused:
|
||||||
|
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
||||||
|
return hidden_states, None, None, None
|
||||||
|
|
||||||
return hidden_states
|
# Return the deferred FFN hc_post state; the next layer consumes it with
|
||||||
|
# cross-layer fusion, and the final layer is completed in DeepseekV4Model.
|
||||||
|
return hidden_states, residual, post, comb
|
||||||
|
|
||||||
|
|
||||||
class DeepseekV4Model(nn.Module):
|
class DeepseekV4Model(nn.Module):
|
||||||
@@ -1302,6 +1503,7 @@ class DeepseekV4Model(nn.Module):
|
|||||||
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
|
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
|
||||||
|
|
||||||
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
||||||
|
self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled()
|
||||||
if self.dsa_enable_prefill_cp:
|
if self.dsa_enable_prefill_cp:
|
||||||
self.cp_size = get_attention_cp_size()
|
self.cp_size = get_attention_cp_size()
|
||||||
|
|
||||||
@@ -1375,21 +1577,32 @@ class DeepseekV4Model(nn.Module):
|
|||||||
# forks alt-streams; later per-layer calls become no-ops.
|
# forks alt-streams; later per-layer calls become no-ops.
|
||||||
get_attn_backend()._maybe_upgrade_forward_metadata()
|
get_attn_backend()._maybe_upgrade_forward_metadata()
|
||||||
|
|
||||||
|
use_fused = self.use_fused_mhc_post_pre
|
||||||
|
prev_residual, prev_post, prev_comb = None, None, None
|
||||||
|
last_layer = None
|
||||||
for i in range(self.start_layer, self.end_layer):
|
for i in range(self.start_layer, self.end_layer):
|
||||||
layer = self.layers[i]
|
layer = self.layers[i]
|
||||||
|
last_layer = layer
|
||||||
ctx = (
|
ctx = (
|
||||||
nullcontext()
|
nullcontext()
|
||||||
if not get_global_server_args().disable_piecewise_cuda_graph
|
if not get_global_server_args().disable_piecewise_cuda_graph
|
||||||
else get_global_expert_distribution_recorder().with_current_layer(i)
|
else get_global_expert_distribution_recorder().with_current_layer(i)
|
||||||
)
|
)
|
||||||
with ctx:
|
with ctx:
|
||||||
hidden_states = layer(
|
hidden_states, prev_residual, prev_post, prev_comb = layer(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
input_ids_global=input_ids_global,
|
input_ids_global=input_ids_global,
|
||||||
|
prev_residual=prev_residual,
|
||||||
|
prev_post=prev_post,
|
||||||
|
prev_comb=prev_comb,
|
||||||
)
|
)
|
||||||
|
if use_fused and last_layer is not None:
|
||||||
|
hidden_states = last_layer.hc_post(
|
||||||
|
hidden_states, prev_residual, prev_post, prev_comb
|
||||||
|
)
|
||||||
|
|
||||||
# CP all-gather only on the last PP rank; PP IPC carries CP-split tensors.
|
# CP all-gather only on the last PP rank; PP IPC carries CP-split tensors.
|
||||||
if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch):
|
if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch):
|
||||||
@@ -1589,6 +1802,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
and not self_attn.indexer.compressor.ape_converted
|
and not self_attn.indexer.compressor.ape_converted
|
||||||
):
|
):
|
||||||
self_attn.indexer.compressor.apply_ape_hotfix()
|
self_attn.indexer.compressor.apply_ape_hotfix()
|
||||||
|
layer.refresh_mhc_norm_weight_cache()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def remap_weight_name_to_dpsk_hf_format(
|
def remap_weight_name_to_dpsk_hf_format(
|
||||||
|
|||||||
@@ -170,13 +170,17 @@ class DeepseekV4ModelNextN(nn.Module):
|
|||||||
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
||||||
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
||||||
|
|
||||||
hidden_states = self.decoder(
|
hidden_states, residual, post, comb = self.decoder(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
input_ids_global=input_ids_global,
|
input_ids_global=input_ids_global,
|
||||||
)
|
)
|
||||||
|
if residual is not None:
|
||||||
|
# NextN has a single decoder layer, so no later layer can consume a
|
||||||
|
# deferred fused hc_post state.
|
||||||
|
hidden_states = self.decoder.hc_post(hidden_states, residual, post, comb)
|
||||||
|
|
||||||
if dsa_use_prefill_cp(forward_batch):
|
if dsa_use_prefill_cp(forward_batch):
|
||||||
hidden_states = cp_all_gather_rerange_output(
|
hidden_states = cp_all_gather_rerange_output(
|
||||||
|
|||||||
@@ -0,0 +1,111 @@
|
|||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
import sglang.srt.layers.mhc as mhc
|
||||||
|
from sglang.srt.layers.mhc import mhc_fused_post_pre, mhc_post, mhc_pre
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("hidden_size", [4096, 7168])
|
||||||
|
@pytest.mark.parametrize("num_tokens", [0, 1, 8, 17, 32, 64])
|
||||||
|
@pytest.mark.parametrize("use_norm", [False, True])
|
||||||
|
def test_mhc_fused_post_pre_matches_unfused(
|
||||||
|
monkeypatch, hidden_size, num_tokens, use_norm
|
||||||
|
):
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
pytest.skip("CUDA is required for TileLang mHC kernels")
|
||||||
|
|
||||||
|
monkeypatch.setattr(mhc, "is_dsa_prefill_cp_round_robin_split", lambda: False)
|
||||||
|
torch.manual_seed(0)
|
||||||
|
device = torch.device("cuda")
|
||||||
|
hc_mult = 4
|
||||||
|
hc_mult3 = hc_mult * 2 + hc_mult * hc_mult
|
||||||
|
hc_hidden_size = hc_mult * hidden_size
|
||||||
|
|
||||||
|
x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) * 0.1
|
||||||
|
residual = (
|
||||||
|
torch.randn(
|
||||||
|
num_tokens, hc_mult, hidden_size, device=device, dtype=torch.bfloat16
|
||||||
|
)
|
||||||
|
* 0.1
|
||||||
|
)
|
||||||
|
post_prev = torch.rand(num_tokens, hc_mult, 1, device=device, dtype=torch.float32)
|
||||||
|
comb_prev = (
|
||||||
|
torch.rand(num_tokens, hc_mult, hc_mult, device=device, dtype=torch.float32)
|
||||||
|
* 0.25
|
||||||
|
)
|
||||||
|
fn = (
|
||||||
|
torch.randn(hc_mult3, hc_hidden_size, device=device, dtype=torch.float32) * 0.01
|
||||||
|
)
|
||||||
|
hc_scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32)
|
||||||
|
hc_base = torch.zeros(hc_mult3, device=device, dtype=torch.float32)
|
||||||
|
norm_weight = (
|
||||||
|
torch.ones(hidden_size, device=device, dtype=torch.bfloat16)
|
||||||
|
if use_norm
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
norm_eps = 1e-6 if use_norm else None
|
||||||
|
|
||||||
|
rms_eps = 1e-6
|
||||||
|
hc_eps = 1e-6
|
||||||
|
sinkhorn_repeat = 2
|
||||||
|
|
||||||
|
residual_ref = post_ref = comb_ref = layer_ref = None
|
||||||
|
if num_tokens > 0:
|
||||||
|
residual_ref = mhc_post(x, residual, post_prev, comb_prev)
|
||||||
|
post_ref, comb_ref, layer_ref = mhc_pre(
|
||||||
|
residual_ref,
|
||||||
|
fn,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
rms_eps,
|
||||||
|
hc_eps,
|
||||||
|
hc_eps,
|
||||||
|
2.0,
|
||||||
|
sinkhorn_repeat,
|
||||||
|
norm_weight=norm_weight,
|
||||||
|
norm_eps=norm_eps,
|
||||||
|
)
|
||||||
|
residual_out, post_out, comb_out, layer_out = mhc_fused_post_pre(
|
||||||
|
x,
|
||||||
|
residual,
|
||||||
|
post_prev,
|
||||||
|
comb_prev,
|
||||||
|
fn,
|
||||||
|
hc_scale,
|
||||||
|
hc_base,
|
||||||
|
rms_eps,
|
||||||
|
hc_eps,
|
||||||
|
hc_eps,
|
||||||
|
2.0,
|
||||||
|
sinkhorn_repeat,
|
||||||
|
norm_weight=norm_weight,
|
||||||
|
norm_eps=norm_eps,
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
if num_tokens == 0:
|
||||||
|
assert residual_out.shape == residual.shape
|
||||||
|
assert post_out.shape == (0, hc_mult, 1)
|
||||||
|
assert comb_out.shape == (0, hc_mult, hc_mult)
|
||||||
|
assert layer_out.shape == (0, hidden_size)
|
||||||
|
assert residual_out.dtype == torch.bfloat16
|
||||||
|
assert post_out.dtype == torch.float32
|
||||||
|
assert comb_out.dtype == torch.float32
|
||||||
|
assert layer_out.dtype == torch.bfloat16
|
||||||
|
return
|
||||||
|
|
||||||
|
assert residual_ref is not None
|
||||||
|
assert post_ref is not None
|
||||||
|
assert comb_ref is not None
|
||||||
|
assert layer_ref is not None
|
||||||
|
assert residual_out.shape == residual_ref.shape
|
||||||
|
assert post_out.shape == post_ref.shape
|
||||||
|
assert comb_out.shape == comb_ref.shape
|
||||||
|
assert layer_out.shape == layer_ref.shape
|
||||||
|
|
||||||
|
torch.testing.assert_close(residual_out, residual_ref, atol=0, rtol=0)
|
||||||
|
torch.testing.assert_close(post_out, post_ref, atol=1e-3, rtol=1e-3)
|
||||||
|
torch.testing.assert_close(comb_out, comb_ref, atol=1e-3, rtol=1e-3)
|
||||||
|
layer_atol = 2e-2 if use_norm else 2e-3
|
||||||
|
layer_rtol = 2e-2 if use_norm else 2e-3
|
||||||
|
torch.testing.assert_close(layer_out, layer_ref, atol=layer_atol, rtol=layer_rtol)
|
||||||
Reference in New Issue
Block a user