forked from Karylab-cklius/vllm
@@ -723,6 +723,7 @@ class CompilationConfig:
|
||||
"vllm::kda_attention",
|
||||
"vllm::sparse_attn_indexer",
|
||||
"vllm::rocm_aiter_sparse_attn_indexer",
|
||||
# For specialized models
|
||||
"vllm::monolithic_attn",
|
||||
]
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Monolithic DeepSeek V3.2 model optimized for SM100 (Blackwell)."""
|
||||
"""DeepSeek V3.2 model optimized for SM100 (Blackwell)."""
|
||||
|
||||
from .model import DeepseekV32ForCausalLM
|
||||
from .mtp import DeepSeekMTP
|
||||
|
||||
@@ -1,200 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Monolithic MLA attention for DeepSeek V3.2 on SM100 (Blackwell).
|
||||
|
||||
MLA forward fully inlined:
|
||||
KV cache update -> W_UK_T absorption -> sparse attn kernel -> W_UV up-proj
|
||||
MLAAttention kept only as a registration stub for KV cache / backend.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from vllm.config import CacheConfig, VllmConfig
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size
|
||||
from vllm.model_executor.layers.attention.mla_attention import MLAAttention
|
||||
from vllm.model_executor.layers.layernorm import LayerNorm
|
||||
from vllm.model_executor.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from vllm.model_executor.layers.rotary_embedding import get_rope
|
||||
from vllm.model_executor.layers.sparse_attn_indexer import SparseAttnIndexer
|
||||
from vllm.model_executor.models.deepseek_v2 import (
|
||||
DeepseekV32IndexerCache,
|
||||
yarn_get_mscale,
|
||||
)
|
||||
from vllm.v1.attention.backends.mla.indexer import get_max_prefill_buffer_size
|
||||
|
||||
|
||||
class MonolithicMLAAttention(nn.Module):
|
||||
"""
|
||||
Monolithic MLA attention for DeepSeek V3.2 targeting SM100.
|
||||
MLA forward fully inlined. MLAAttention kept only for KV cache
|
||||
registration and backend/impl initialization.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
config,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
qk_nope_head_dim: int,
|
||||
qk_rope_head_dim: int,
|
||||
v_head_dim: int,
|
||||
q_lora_rank: int,
|
||||
kv_lora_rank: int,
|
||||
max_position_embeddings: int,
|
||||
cache_config: CacheConfig,
|
||||
quant_config: QuantizationConfig | None,
|
||||
topk_indices_buffer: torch.Tensor,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.qk_rope_head_dim = qk_rope_head_dim
|
||||
self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||
self.v_head_dim = v_head_dim
|
||||
self.q_lora_rank = q_lora_rank
|
||||
self.kv_lora_rank = kv_lora_rank
|
||||
self.num_heads = num_heads
|
||||
self.num_local_heads = num_heads // get_tensor_model_parallel_world_size()
|
||||
self.scaling = self.qk_head_dim**-0.5
|
||||
self.rms_norm_eps = config.rms_norm_eps
|
||||
|
||||
# Q path
|
||||
self.q_a_layernorm_weight = nn.Parameter(
|
||||
torch.ones(q_lora_rank, dtype=torch.get_default_dtype())
|
||||
)
|
||||
self.q_b_proj = ColumnParallelLinear(
|
||||
q_lora_rank,
|
||||
num_heads * self.qk_head_dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.q_b_proj",
|
||||
)
|
||||
|
||||
# KV path
|
||||
self.kv_a_layernorm_weight = nn.Parameter(
|
||||
torch.ones(kv_lora_rank, dtype=torch.get_default_dtype())
|
||||
)
|
||||
self.kv_b_proj = ColumnParallelLinear(
|
||||
kv_lora_rank,
|
||||
num_heads * (qk_nope_head_dim + v_head_dim),
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.kv_b_proj",
|
||||
)
|
||||
|
||||
# Output projection (TP sync point)
|
||||
self.o_proj = RowParallelLinear(
|
||||
num_heads * v_head_dim,
|
||||
hidden_size,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
# RoPE
|
||||
if config.rope_parameters["rope_type"] != "default":
|
||||
config.rope_parameters["rope_type"] = (
|
||||
"deepseek_yarn"
|
||||
if config.rope_parameters.get("apply_yarn_scaling", True)
|
||||
else "deepseek_llama_scaling"
|
||||
)
|
||||
self.rotary_emb = get_rope(
|
||||
qk_rope_head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
rope_parameters=config.rope_parameters,
|
||||
is_neox_style=False,
|
||||
)
|
||||
if config.rope_parameters["rope_type"] == "deepseek_yarn":
|
||||
mscale_all_dim = config.rope_parameters.get("mscale_all_dim", False)
|
||||
scaling_factor = config.rope_parameters["factor"]
|
||||
mscale = yarn_get_mscale(scaling_factor, float(mscale_all_dim))
|
||||
self.scaling = self.scaling * mscale * mscale
|
||||
|
||||
# V3.2 Sparse Indexer (inlined)
|
||||
self.indexer_rope_emb = get_rope(
|
||||
qk_rope_head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
rope_parameters=config.rope_parameters,
|
||||
is_neox_style=not getattr(config, "indexer_rope_interleave", False),
|
||||
)
|
||||
self.topk_tokens = config.index_topk
|
||||
self.index_n_heads = config.index_n_heads
|
||||
self.index_head_dim = config.index_head_dim
|
||||
self.indexer_softmax_scale = config.index_head_dim**-0.5
|
||||
self.indexer_quant_block_size = 128
|
||||
self.topk_indices_buffer = topk_indices_buffer
|
||||
|
||||
self.indexer_wq_b = ReplicatedLinear(
|
||||
q_lora_rank,
|
||||
config.index_head_dim * config.index_n_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.indexer.wq_b",
|
||||
)
|
||||
self.indexer_wk = ReplicatedLinear(
|
||||
hidden_size,
|
||||
config.index_head_dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.indexer.wk",
|
||||
)
|
||||
self.indexer_k_norm = LayerNorm(config.index_head_dim, eps=1e-6)
|
||||
self.indexer_weights_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
config.index_n_heads,
|
||||
bias=False,
|
||||
quant_config=None,
|
||||
prefix=f"{prefix}.indexer.weights_proj",
|
||||
)
|
||||
|
||||
idx_dim = config.index_head_dim
|
||||
indexer_cache_head_dim = idx_dim + idx_dim // 128 * 4
|
||||
self.indexer_k_cache = DeepseekV32IndexerCache(
|
||||
head_dim=indexer_cache_head_dim,
|
||||
dtype=torch.uint8,
|
||||
prefix=f"{prefix}.indexer.k_cache",
|
||||
cache_config=cache_config,
|
||||
)
|
||||
self.indexer_op = SparseAttnIndexer(
|
||||
self.indexer_k_cache,
|
||||
self.indexer_quant_block_size,
|
||||
"ue8m0",
|
||||
self.topk_tokens,
|
||||
config.index_head_dim,
|
||||
vllm_config.model_config.max_model_len,
|
||||
get_max_prefill_buffer_size(vllm_config),
|
||||
self.topk_indices_buffer,
|
||||
)
|
||||
|
||||
# MLAAttention stub: only for KV cache registration + backend init.
|
||||
# We never call its forward(); we inline everything below.
|
||||
class _IndexerProxy:
|
||||
def __init__(proxy_self):
|
||||
proxy_self.topk_indices_buffer = topk_indices_buffer
|
||||
proxy_self.indexer_op = self.indexer_op
|
||||
|
||||
self._indexer_proxy = _IndexerProxy()
|
||||
self.mla_attn = MLAAttention(
|
||||
num_heads=self.num_local_heads,
|
||||
scale=self.scaling,
|
||||
qk_nope_head_dim=qk_nope_head_dim,
|
||||
qk_rope_head_dim=qk_rope_head_dim,
|
||||
v_head_dim=v_head_dim,
|
||||
q_lora_rank=q_lora_rank,
|
||||
kv_lora_rank=kv_lora_rank,
|
||||
kv_b_proj=self.kv_b_proj,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.mla_attn",
|
||||
use_sparse=True,
|
||||
indexer=self._indexer_proxy,
|
||||
)
|
||||
+238
-49
@@ -1,7 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Monolithic decoder layer for DeepSeek V3.2 on SM100 (Blackwell).
|
||||
MLA attention and decoder layer for DeepSeek V3.2 on SM100 (Blackwell).
|
||||
|
||||
MLAAttention:
|
||||
KV cache update -> W_UK_T absorption -> sparse attn kernel -> W_UV up-proj
|
||||
MLAAttention kept only as a registration stub for KV cache / backend.
|
||||
|
||||
DecoderLayer:
|
||||
Single decoder layer: norm -> attn -> norm -> MoE/MLP.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -9,20 +16,32 @@ from __future__ import annotations
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from vllm.config import VllmConfig, get_current_vllm_config
|
||||
from vllm.config import CacheConfig, VllmConfig, get_current_vllm_config
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size
|
||||
from vllm.forward_context import get_forward_context
|
||||
from vllm.model_executor.layers.layernorm import RMSNorm
|
||||
from vllm.model_executor.layers.attention.mla_attention import MLAAttention
|
||||
from vllm.model_executor.layers.layernorm import LayerNorm, RMSNorm
|
||||
from vllm.model_executor.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from vllm.model_executor.layers.rotary_embedding import get_rope
|
||||
from vllm.model_executor.layers.sparse_attn_indexer import SparseAttnIndexer
|
||||
from vllm.model_executor.models.deepseek_v2 import (
|
||||
DeepseekV32IndexerCache,
|
||||
yarn_get_mscale,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.torch_utils import direct_register_custom_op
|
||||
from vllm.v1.attention.backends.mla.indexer import get_max_prefill_buffer_size
|
||||
|
||||
from .attention import MonolithicMLAAttention
|
||||
from .kernels import fused_norm_rope, fused_q
|
||||
from .sparse_indexer import sparse_attn_indexer
|
||||
|
||||
|
||||
def monolithic_attn(
|
||||
def dsa(
|
||||
positions: torch.Tensor,
|
||||
q_c: torch.Tensor,
|
||||
kv_c: torch.Tensor,
|
||||
@@ -143,7 +162,7 @@ def monolithic_attn(
|
||||
return output
|
||||
|
||||
|
||||
def monolithic_attn_fake(
|
||||
def dsa_fake(
|
||||
positions: torch.Tensor,
|
||||
q_c: torch.Tensor,
|
||||
kv_c: torch.Tensor,
|
||||
@@ -159,14 +178,14 @@ def monolithic_attn_fake(
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="monolithic_attn",
|
||||
op_func=monolithic_attn,
|
||||
fake_impl=monolithic_attn_fake,
|
||||
op_func=dsa,
|
||||
fake_impl=dsa_fake,
|
||||
mutates_args=["output"],
|
||||
dispatch_key=current_platform.dispatch_key,
|
||||
)
|
||||
|
||||
|
||||
class MonolithicDecoderLayer(nn.Module):
|
||||
class DeepseekV32DecoderLayer(nn.Module):
|
||||
"""
|
||||
Single decoder layer: norm -> attn -> norm -> MoE/MLP.
|
||||
Norms are raw weight + direct kernel call.
|
||||
@@ -231,7 +250,7 @@ class MonolithicDecoderLayer(nn.Module):
|
||||
)
|
||||
|
||||
# MLA Attention — disable AllReduce in o_proj when using fused path
|
||||
self.attn = MonolithicMLAAttention(
|
||||
self.attn = DeepseekV32MLAAttention(
|
||||
vllm_config=vllm_config,
|
||||
config=config,
|
||||
hidden_size=config.hidden_size,
|
||||
@@ -278,45 +297,6 @@ class MonolithicDecoderLayer(nn.Module):
|
||||
prefix=f"{prefix}.mlp",
|
||||
)
|
||||
|
||||
def fuse_indexer_weights(self) -> None:
|
||||
"""Fuse Step 1 and Step 3 BF16 linears used by the monolithic path.
|
||||
|
||||
Call after model weights are loaded.
|
||||
"""
|
||||
attn = self.attn
|
||||
qkv_a = self.self_attn.fused_qkv_a_proj.weight.data # [2112, 7168]
|
||||
wk = attn.indexer_wk.weight.data # [128, 7168]
|
||||
wp = attn.indexer_weights_proj.weight.data # [64, 7168]
|
||||
if not (qkv_a.dtype == wk.dtype == wp.dtype):
|
||||
raise ValueError(
|
||||
"Cannot fuse Step 1 weights: expected matching dtypes for "
|
||||
"fused_qkv_a_proj, indexer_wk, and indexer_weights_proj."
|
||||
)
|
||||
self._fused_step1_hidden_w = nn.Parameter(
|
||||
torch.cat([qkv_a, wk, wp], dim=0), # [2304, 7168]
|
||||
requires_grad=False,
|
||||
)
|
||||
self._step1_split_sizes = [
|
||||
self.q_lora_rank,
|
||||
self.kv_lora_rank,
|
||||
self.qk_rope_head_dim,
|
||||
wk.shape[0],
|
||||
wp.shape[0],
|
||||
]
|
||||
|
||||
wq_b = attn.indexer_wq_b.weight.data
|
||||
q_b = attn.q_b_proj.weight.data
|
||||
if wq_b.dtype != q_b.dtype:
|
||||
raise ValueError(
|
||||
"Cannot fuse Step 3 weights: expected matching dtypes for "
|
||||
"indexer_wq_b and q_b_proj."
|
||||
)
|
||||
self._fused_step3_q_w = nn.Parameter(
|
||||
torch.cat([wq_b, q_b], dim=0),
|
||||
requires_grad=False,
|
||||
)
|
||||
self._step3_index_q_dim = wq_b.shape[0]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
@@ -361,6 +341,45 @@ class MonolithicDecoderLayer(nn.Module):
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
return hidden_states, residual
|
||||
|
||||
def fuse_indexer_weights(self) -> None:
|
||||
"""Fuse Step 1 and Step 3 BF16 linears used by the inlined path.
|
||||
|
||||
Call after model weights are loaded.
|
||||
"""
|
||||
attn = self.attn
|
||||
qkv_a = self.self_attn.fused_qkv_a_proj.weight.data # [2112, 7168]
|
||||
wk = attn.indexer_wk.weight.data # [128, 7168]
|
||||
wp = attn.indexer_weights_proj.weight.data # [64, 7168]
|
||||
if not (qkv_a.dtype == wk.dtype == wp.dtype):
|
||||
raise ValueError(
|
||||
"Cannot fuse Step 1 weights: expected matching dtypes for "
|
||||
"fused_qkv_a_proj, indexer_wk, and indexer_weights_proj."
|
||||
)
|
||||
self._fused_step1_hidden_w = nn.Parameter(
|
||||
torch.cat([qkv_a, wk, wp], dim=0), # [2304, 7168]
|
||||
requires_grad=False,
|
||||
)
|
||||
self._step1_split_sizes = [
|
||||
self.q_lora_rank,
|
||||
self.kv_lora_rank,
|
||||
self.qk_rope_head_dim,
|
||||
wk.shape[0],
|
||||
wp.shape[0],
|
||||
]
|
||||
|
||||
wq_b = attn.indexer_wq_b.weight.data
|
||||
q_b = attn.q_b_proj.weight.data
|
||||
if wq_b.dtype != q_b.dtype:
|
||||
raise ValueError(
|
||||
"Cannot fuse Step 3 weights: expected matching dtypes for "
|
||||
"indexer_wq_b and q_b_proj."
|
||||
)
|
||||
self._fused_step3_q_w = nn.Parameter(
|
||||
torch.cat([wq_b, q_b], dim=0),
|
||||
requires_grad=False,
|
||||
)
|
||||
self._step3_index_q_dim = wq_b.shape[0]
|
||||
|
||||
def fuse_shared_expert_act_quant(self) -> None:
|
||||
"""Fuse SiLU-and-Mul + NVFP4 quantize in the shared expert MLP.
|
||||
|
||||
@@ -403,3 +422,173 @@ class MonolithicDecoderLayer(nn.Module):
|
||||
return out
|
||||
|
||||
shared_experts.forward = _fused_forward
|
||||
|
||||
|
||||
class DeepseekV32MLAAttention(nn.Module):
|
||||
"""
|
||||
MLA attention for DeepSeek V3.2 targeting SM100.
|
||||
MLA forward fully inlined. MLAAttention kept only for KV cache
|
||||
registration and backend/impl initialization.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
config,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
qk_nope_head_dim: int,
|
||||
qk_rope_head_dim: int,
|
||||
v_head_dim: int,
|
||||
q_lora_rank: int,
|
||||
kv_lora_rank: int,
|
||||
max_position_embeddings: int,
|
||||
cache_config: CacheConfig,
|
||||
quant_config: QuantizationConfig | None,
|
||||
topk_indices_buffer: torch.Tensor,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.qk_rope_head_dim = qk_rope_head_dim
|
||||
self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||
self.v_head_dim = v_head_dim
|
||||
self.q_lora_rank = q_lora_rank
|
||||
self.kv_lora_rank = kv_lora_rank
|
||||
self.num_heads = num_heads
|
||||
self.num_local_heads = num_heads // get_tensor_model_parallel_world_size()
|
||||
self.scaling = self.qk_head_dim**-0.5
|
||||
self.rms_norm_eps = config.rms_norm_eps
|
||||
|
||||
# Q path
|
||||
self.q_a_layernorm_weight = nn.Parameter(
|
||||
torch.ones(q_lora_rank, dtype=torch.get_default_dtype())
|
||||
)
|
||||
self.q_b_proj = ColumnParallelLinear(
|
||||
q_lora_rank,
|
||||
num_heads * self.qk_head_dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.q_b_proj",
|
||||
)
|
||||
|
||||
# KV path
|
||||
self.kv_a_layernorm_weight = nn.Parameter(
|
||||
torch.ones(kv_lora_rank, dtype=torch.get_default_dtype())
|
||||
)
|
||||
self.kv_b_proj = ColumnParallelLinear(
|
||||
kv_lora_rank,
|
||||
num_heads * (qk_nope_head_dim + v_head_dim),
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.kv_b_proj",
|
||||
)
|
||||
|
||||
# Output projection (TP sync point)
|
||||
self.o_proj = RowParallelLinear(
|
||||
num_heads * v_head_dim,
|
||||
hidden_size,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
# RoPE
|
||||
if config.rope_parameters["rope_type"] != "default":
|
||||
config.rope_parameters["rope_type"] = (
|
||||
"deepseek_yarn"
|
||||
if config.rope_parameters.get("apply_yarn_scaling", True)
|
||||
else "deepseek_llama_scaling"
|
||||
)
|
||||
self.rotary_emb = get_rope(
|
||||
qk_rope_head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
rope_parameters=config.rope_parameters,
|
||||
is_neox_style=False,
|
||||
)
|
||||
if config.rope_parameters["rope_type"] == "deepseek_yarn":
|
||||
mscale_all_dim = config.rope_parameters.get("mscale_all_dim", False)
|
||||
scaling_factor = config.rope_parameters["factor"]
|
||||
mscale = yarn_get_mscale(scaling_factor, float(mscale_all_dim))
|
||||
self.scaling = self.scaling * mscale * mscale
|
||||
|
||||
# V3.2 Sparse Indexer (inlined)
|
||||
self.indexer_rope_emb = get_rope(
|
||||
qk_rope_head_dim,
|
||||
max_position=max_position_embeddings,
|
||||
rope_parameters=config.rope_parameters,
|
||||
is_neox_style=not getattr(config, "indexer_rope_interleave", False),
|
||||
)
|
||||
self.topk_tokens = config.index_topk
|
||||
self.index_n_heads = config.index_n_heads
|
||||
self.index_head_dim = config.index_head_dim
|
||||
self.indexer_softmax_scale = config.index_head_dim**-0.5
|
||||
self.indexer_quant_block_size = 128
|
||||
self.topk_indices_buffer = topk_indices_buffer
|
||||
|
||||
self.indexer_wq_b = ReplicatedLinear(
|
||||
q_lora_rank,
|
||||
config.index_head_dim * config.index_n_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.indexer.wq_b",
|
||||
)
|
||||
self.indexer_wk = ReplicatedLinear(
|
||||
hidden_size,
|
||||
config.index_head_dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.indexer.wk",
|
||||
)
|
||||
self.indexer_k_norm = LayerNorm(config.index_head_dim, eps=1e-6)
|
||||
self.indexer_weights_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
config.index_n_heads,
|
||||
bias=False,
|
||||
quant_config=None,
|
||||
prefix=f"{prefix}.indexer.weights_proj",
|
||||
)
|
||||
|
||||
idx_dim = config.index_head_dim
|
||||
indexer_cache_head_dim = idx_dim + idx_dim // 128 * 4
|
||||
self.indexer_k_cache = DeepseekV32IndexerCache(
|
||||
head_dim=indexer_cache_head_dim,
|
||||
dtype=torch.uint8,
|
||||
prefix=f"{prefix}.indexer.k_cache",
|
||||
cache_config=cache_config,
|
||||
)
|
||||
self.indexer_op = SparseAttnIndexer(
|
||||
self.indexer_k_cache,
|
||||
self.indexer_quant_block_size,
|
||||
"ue8m0",
|
||||
self.topk_tokens,
|
||||
config.index_head_dim,
|
||||
vllm_config.model_config.max_model_len,
|
||||
get_max_prefill_buffer_size(vllm_config),
|
||||
self.topk_indices_buffer,
|
||||
)
|
||||
|
||||
# MLAAttention stub: only for KV cache registration + backend init.
|
||||
# We never call its forward(); we inline everything below.
|
||||
class _IndexerProxy:
|
||||
def __init__(proxy_self):
|
||||
proxy_self.topk_indices_buffer = topk_indices_buffer
|
||||
proxy_self.indexer_op = self.indexer_op
|
||||
|
||||
self._indexer_proxy = _IndexerProxy()
|
||||
self.mla_attn = MLAAttention(
|
||||
num_heads=self.num_local_heads,
|
||||
scale=self.scaling,
|
||||
qk_nope_head_dim=qk_nope_head_dim,
|
||||
qk_rope_head_dim=qk_rope_head_dim,
|
||||
v_head_dim=v_head_dim,
|
||||
q_lora_rank=q_lora_rank,
|
||||
kv_lora_rank=kv_lora_rank,
|
||||
kv_b_proj=self.kv_b_proj,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.mla_attn",
|
||||
use_sparse=True,
|
||||
indexer=self._indexer_proxy,
|
||||
)
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Monolithic DeepSeek V3.2 model for SM100 (Blackwell).
|
||||
DeepSeek V3.2 model for SM100 (Blackwell).
|
||||
No PP, TP only, same checkpoint compatibility.
|
||||
"""
|
||||
|
||||
@@ -22,13 +22,13 @@ from vllm.model_executor.layers.vocab_parallel_embedding import (
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
from .decoder_layer import MonolithicDecoderLayer
|
||||
from .layer import DeepseekV32DecoderLayer
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@support_torch_compile
|
||||
class MonolithicDeepseekV32Model(nn.Module):
|
||||
class DeepseekV32Model(nn.Module):
|
||||
"""Transformer backbone."""
|
||||
|
||||
fall_back_to_pt_during_load = False
|
||||
@@ -57,7 +57,7 @@ class MonolithicDeepseekV32Model(nn.Module):
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
MonolithicDecoderLayer(
|
||||
DeepseekV32DecoderLayer(
|
||||
vllm_config=vllm_config,
|
||||
config=config,
|
||||
layer_idx=i,
|
||||
@@ -89,7 +89,7 @@ class MonolithicDeepseekV32Model(nn.Module):
|
||||
|
||||
class DeepseekV32ForCausalLM(nn.Module):
|
||||
"""
|
||||
Monolithic DeepSeek V3.2 CausalLM for SM100.
|
||||
DeepSeek V3.2 CausalLM for SM100.
|
||||
"""
|
||||
|
||||
packed_modules_mapping = {
|
||||
@@ -105,7 +105,7 @@ class DeepseekV32ForCausalLM(nn.Module):
|
||||
self.quant_config = quant_config
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
|
||||
self.model = MonolithicDeepseekV32Model(
|
||||
self.model = DeepseekV32Model(
|
||||
vllm_config=vllm_config,
|
||||
prefix=f"{prefix}.model" if prefix else "model",
|
||||
)
|
||||
@@ -165,11 +165,11 @@ class DeepseekV32ForCausalLM(nn.Module):
|
||||
# Only remap layernorms and indexer (underscore prefix).
|
||||
# Everything else (fused_qkv_a_proj, experts, gate, etc.) uses the
|
||||
# same module paths as the original model.
|
||||
return remap_monolithic_weight_name(name)
|
||||
return remap_weight_name(name)
|
||||
|
||||
|
||||
def remap_monolithic_weight_name(name: str) -> str:
|
||||
"""Remap checkpoint names that differ from the monolithic module layout."""
|
||||
def remap_weight_name(name: str) -> str:
|
||||
"""Remap checkpoint names that differ from the module layout."""
|
||||
replacements = [
|
||||
(
|
||||
"self_attn.q_a_layernorm.weight",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Monolithic DeepSeek V3.2 MTP model."""
|
||||
"""DeepSeek V3.2 MTP model for SM100 (Blackwell)."""
|
||||
|
||||
from collections.abc import Iterable
|
||||
|
||||
@@ -28,8 +28,8 @@ from vllm.model_executor.models.utils import maybe_prefix
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.sequence import IntermediateTensors
|
||||
|
||||
from .decoder_layer import MonolithicDecoderLayer
|
||||
from .model import remap_monolithic_weight_name
|
||||
from .layer import DeepseekV32DecoderLayer
|
||||
from .model import remap_weight_name
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -61,7 +61,7 @@ class DeepSeekMultiTokenPredictorLayer(DeepSeekMultiTokenPredictorLayerBase):
|
||||
self.shared_head = SharedHead(
|
||||
config=config, prefix=prefix, quant_config=quant_config
|
||||
)
|
||||
self.mtp_block = MonolithicDecoderLayer(
|
||||
self.mtp_block = DeepseekV32DecoderLayer(
|
||||
vllm_config=vllm_config,
|
||||
config=config,
|
||||
layer_idx=int(prefix.rsplit(".", 1)[-1]),
|
||||
@@ -124,7 +124,7 @@ class DeepSeekMTP(DeepSeekMTPBase):
|
||||
example_moe = None
|
||||
for layer in self.model.layers.values():
|
||||
layer = layer.mtp_block
|
||||
assert isinstance(layer, MonolithicDecoderLayer)
|
||||
assert isinstance(layer, DeepseekV32DecoderLayer)
|
||||
if isinstance(layer.mlp, DeepseekV2MoE):
|
||||
example_moe = layer.mlp
|
||||
self.moe_mlp_layers.append(layer.mlp)
|
||||
@@ -154,7 +154,7 @@ class DeepSeekMTP(DeepSeekMTPBase):
|
||||
|
||||
def _rewrite_spec_layer_name(self, spec_layer: int, name: str) -> str:
|
||||
name = super()._rewrite_spec_layer_name(spec_layer, name)
|
||||
return remap_monolithic_weight_name(name)
|
||||
return remap_weight_name(name)
|
||||
|
||||
|
||||
@torch.compile
|
||||
|
||||
Reference in New Issue
Block a user