Prewarm DSV4 MHC post kernel at model load (#30741)
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
This commit is contained in:
@@ -2752,8 +2752,8 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
|
||||
return name
|
||||
|
||||
def _prewarm_mhc_pre_kernels(self) -> None:
|
||||
"""One-shot mhc_pre() JIT prewarm at load time, synced across ranks.
|
||||
def _prewarm_mhc_kernels(self) -> None:
|
||||
"""One-shot MHC JIT prewarm at load time, synced across ranks.
|
||||
|
||||
Runs before any forward so the compile burst stays off the serving
|
||||
path; the barrier keeps ranks from proceeding while a peer is still
|
||||
@@ -2774,16 +2774,17 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
if layer is None:
|
||||
return
|
||||
|
||||
from sglang.kernels.ops.layernorm.mhc import prewarm_mhc_pre
|
||||
from sglang.kernels.ops.layernorm.mhc import mhc_post, prewarm_mhc_pre
|
||||
|
||||
tic = time.perf_counter()
|
||||
residual = torch.zeros(
|
||||
(1, layer.hc_mult, layer.hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=layer.hc_attn_fn.device,
|
||||
)
|
||||
prewarm_mhc_pre(
|
||||
# Template carrying dtype/device; buckets allocate their own sizes.
|
||||
residual=torch.zeros(
|
||||
(1, layer.hc_mult, layer.hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=layer.hc_attn_fn.device,
|
||||
),
|
||||
residual=residual,
|
||||
fn=layer.hc_attn_fn,
|
||||
hc_scale=layer.hc_attn_scale,
|
||||
hc_base=layer.hc_attn_base,
|
||||
@@ -2797,13 +2798,27 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
norm_weight=layer.input_layernorm.weight.data,
|
||||
norm_eps=layer.input_layernorm.variance_epsilon,
|
||||
)
|
||||
mhc_post(
|
||||
x=residual.new_zeros((1, layer.hidden_size)),
|
||||
residual=residual,
|
||||
post_layer_mix=torch.zeros(
|
||||
(1, layer.hc_mult, 1),
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
),
|
||||
comb_res_mix=torch.zeros(
|
||||
(1, layer.hc_mult, layer.hc_mult),
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
),
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
compile_secs = time.perf_counter() - tic
|
||||
# Runs before init_memory_pool(); don't let transients skew pool sizing.
|
||||
torch.cuda.empty_cache()
|
||||
get_tp_group().barrier()
|
||||
logger.info(
|
||||
"DeepSeek V4 MHC prenorm prewarm at load: compile %.1fs, rank sync +%.1fs",
|
||||
"DeepSeek V4 MHC prewarm at load: compile %.1fs, rank sync +%.1fs",
|
||||
compile_secs,
|
||||
time.perf_counter() - tic - compile_secs,
|
||||
)
|
||||
@@ -3153,7 +3168,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
self.post_load_weights(is_nextn=is_nextn, weight_names=weight_names)
|
||||
|
||||
if not is_nextn:
|
||||
self._prewarm_mhc_pre_kernels()
|
||||
self._prewarm_mhc_kernels()
|
||||
|
||||
def get_embed_and_head(self):
|
||||
return self.model.embed_tokens.weight, self.lm_head.weight
|
||||
|
||||
Reference in New Issue
Block a user