Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
065e6819e1 | ||
|
|
0c97e6e102 |
@@ -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",
|
||||
]
|
||||
|
||||
+3
-3
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user