Compare commits

...
Author SHA1 Message Date
Jee Jee LiandGitHub 065e6819e1 Merge branch 'main' into k3-dspark-ar-fusion 2026-07-29 18:41:01 +08:00
Jee Jee Li 0c97e6e102 init
Signed-off-by: Jee Jee Li <jeejeelee@inferact.ai>
2026-07-29 10:37:47 +00:00
4 changed files with 23 additions and 7 deletions
+2
View File
@@ -2,8 +2,10 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Ops shared across model implementations."""
from .fused_allreduce_rms_norm import fused_allreduce_rms_norm
from .fused_qk_rmsnorm import fused_q_kv_rmsnorm
__all__ = [
"fused_allreduce_rms_norm",
"fused_q_kv_rmsnorm",
]
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Fused ops for deepseek_v32 (eager / breakable-cudagraph path).
"""Fused all-reduce + residual-add + RMSNorm for eager model paths.
These recover fusions that vLLM's torch.compile passes would normally do but
that don't fire when running eager under the breakable CUDA graph.
This recovers a fusion that vLLM's torch.compile passes would normally do but
that doesn't fire for models running eager (or under a breakable CUDA graph).
"""
import torch
+1 -1
View File
@@ -38,10 +38,10 @@ from vllm.model_executor.models.utils import (
make_layers,
sequence_parallel_chunk,
)
from vllm.models.common.ops import fused_allreduce_rms_norm
from vllm.sequence import IntermediateTensors
from .attention import DeepseekV32Attention
from .fused_ops import fused_allreduce_rms_norm
def _all_gather_sp_states(
+17 -3
View File
@@ -20,6 +20,7 @@ from vllm.model_executor.models.utils import (
get_draft_quant_config,
maybe_prefix,
)
from vllm.models.common.ops import fused_allreduce_rms_norm
from vllm.models.kimi_k3.nvidia.mla import MultiHeadLatentAttention
from vllm.models.kimi_k3.nvidia.model import KimiMLP
from vllm.utils.torch_utils import is_quantized_kv_cache
@@ -74,11 +75,15 @@ class K3DSparkDecoderLayer(nn.Module):
use_rope=True,
non_causal_multi_token_decode=True,
)
# Both row-parallel outputs stay un-reduced; their all-reduces are fused
# into the RMSNorm that follows via fused_allreduce_rms_norm.
self.self_attn.o_proj.reduce_results = False
self.mlp = KimiMLP(
hidden_size=config.hidden_size,
intermediate_size=config.intermediate_size,
hidden_act=config.hidden_act,
quant_config=quant_config,
reduce_results=False,
prefix=maybe_prefix(prefix, f"layers.{start_layer_id + layer_idx}.mlp"),
)
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
@@ -93,16 +98,23 @@ class K3DSparkDecoderLayer(nn.Module):
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
if residual is None:
# First layer: hidden_states is the (already reduced) embedding.
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(hidden_states, residual)
hidden_states, residual = fused_allreduce_rms_norm(
hidden_states, residual, self.input_layernorm
)
hidden_states = self.self_attn(
positions=positions,
hidden_states=hidden_states,
)
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
hidden_states, residual = fused_allreduce_rms_norm(
hidden_states, residual, self.post_attention_layernorm
)
# The MLP output is reduced by the next layer's input_layernorm (or by
# the model's final_norm).
hidden_states = self.mlp(hidden_states)
return hidden_states, residual
@@ -430,7 +442,9 @@ class K3DSparkModel(nn.Module):
hidden_states=hidden_states,
residual=residual,
)
hidden_states, _ = self.final_norm(hidden_states, residual)
hidden_states, _ = fused_allreduce_rms_norm(
hidden_states, residual, self.final_norm
)
return hidden_states