[Refactor] Algebraic data type for nextn config + some basic refactors (#17347)

This commit is contained in:
Jerry Ji
2026-01-23 09:16:55 -08:00
committed by GitHub
parent 08fcda2f63
commit 010c17a133

View File

@@ -14,7 +14,8 @@
import concurrent.futures
import logging
from typing import Iterable, Optional, Tuple
from dataclasses import dataclass
from typing import Iterable, List, Optional, Tuple
import torch
import torch.nn as nn
@@ -66,6 +67,23 @@ logger = logging.getLogger(__name__)
NVFP4_CKPT_FP8_ATTN_QUANT_MODULES = ["q_b_proj"]
@dataclass(frozen=True)
class NextNEnabledConfig:
num_nextn_layers: int
nextn_layer_id: int
nextn_layer_prefix: str
nextn_spec_weight_names: List[str]
@dataclass(frozen=True)
class NextNDisabledConfig:
pass
"""Union type for NextN configuration, including enabled and disabled configurations."""
NextNConfig = NextNEnabledConfig | NextNDisabledConfig
class DeepseekV2WeightLoaderMixin:
"""Mixin for loading weights in DeepSeek V2/V3 models."""
@@ -76,7 +94,7 @@ class DeepseekV2WeightLoaderMixin:
num_fused_shared_experts: int
def do_load_weights(
self: nn.Module,
self,
weights: Iterable[Tuple[str, torch.Tensor]],
is_nextn: bool = False,
):
@@ -86,21 +104,10 @@ class DeepseekV2WeightLoaderMixin:
weights: Iterable of (weight_name, weight_tensor) pairs
is_nextn: Whether loading NextN speculative decoding weights
"""
if is_nextn:
if hasattr(self.config, "num_nextn_predict_layers"):
num_nextn_layers = self.config.num_nextn_predict_layers
assert num_nextn_layers == 1, "Only 1 nextn layer is supported"
# compatible with old design
nextn_layer_id = (
0
if self.config.num_hidden_layers == 1
else self.config.num_hidden_layers
)
else:
raise ValueError("num_nextn_predict_layers is not in the config")
nextn_conf = self._initialize_nextn_conf(is_nextn)
weights = self._maybe_quant_weights_to_fp8_ue8m0(
weights, NVFP4_CKPT_FP8_ATTN_QUANT_MODULES, is_nextn
weights, NVFP4_CKPT_FP8_ATTN_QUANT_MODULES, nextn_conf
)
stacked_params_mapping = [
@@ -131,15 +138,6 @@ class DeepseekV2WeightLoaderMixin:
)
cached_a_proj = {} if fuse_qkv_a_proj else None
if is_nextn:
nextn_layer_prefix = f"model.layers.{nextn_layer_id}"
nextn_spec_weight_names = [
"shared_head.norm",
"eh_proj",
"enorm",
"hnorm",
]
if self.num_fused_shared_experts > 0:
assert self.num_fused_shared_experts == 1
log_info_on_rank0(logger, "Shared experts fusion optimization enabled.")
@@ -168,37 +166,38 @@ class DeepseekV2WeightLoaderMixin:
weight_names.append(name)
if not is_nextn:
if hasattr(self.config, "num_nextn_predict_layers"):
num_nextn_layers = self.config.num_nextn_predict_layers
if num_nextn_layers > 0 and name.startswith("model.layers"):
name_list = name.split(".")
if (
len(name_list) >= 3
and int(name_list[2]) >= self.config.num_hidden_layers
):
continue
else:
if not name.startswith(nextn_layer_prefix):
continue
match nextn_conf:
case NextNEnabledConfig(
nextn_layer_prefix=layer_prefix,
nextn_spec_weight_names=spec_weight_names,
):
if not name.startswith(layer_prefix):
continue
# Use shared head and embed weights from target model
if "shared_head.head" in name or "embed_tokens" in name:
continue
# Use shared head and embed weights from target model
if "shared_head.head" in name or "embed_tokens" in name:
continue
is_decoder = True
# For nextn specific weights
for weight_name in nextn_spec_weight_names:
if weight_name in name:
name = name.replace(nextn_layer_prefix, "model")
is_decoder = False
break
# For decoder layer weights
if is_decoder:
name = name.replace(nextn_layer_prefix, "model.decoder")
# Transform name: NextN-specific → "model.*", decoder → "model.decoder.*"
if any(s in name for s in spec_weight_names):
name = name.replace(layer_prefix, "model")
else:
name = name.replace(layer_prefix, "model.decoder")
case NextNDisabledConfig():
if hasattr(self.config, "num_nextn_predict_layers"):
num_nextn_layers = self.config.num_nextn_predict_layers
if num_nextn_layers > 0 and name.startswith("model.layers"):
name_list = name.split(".")
if (
len(name_list) >= 3
and int(name_list[2])
>= self.config.num_hidden_layers
):
continue
if "rotary_emb.inv_freq" in name:
continue
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if weight_name not in name:
@@ -364,8 +363,42 @@ class DeepseekV2WeightLoaderMixin:
self.post_load_weights(is_nextn=is_nextn, weight_names=weight_names)
def _initialize_nextn_conf(self, is_nextn: bool) -> NextNConfig:
"""
Initialize the nextn configuration.
Raises:
ValueError: If num_nextn_predict_layers is not in the config.
AssertionError: If num_nextn_predict_layers is not equal to 1.
"""
if not is_nextn:
return NextNDisabledConfig()
if not hasattr(self.config, "num_nextn_predict_layers"):
raise ValueError("num_nextn_predict_layers is not in the config")
num_nextn_layers = self.config.num_nextn_predict_layers
assert num_nextn_layers == 1, "Only 1 nextn layer is supported"
# compatible with old design
nextn_layer_id = (
0 if self.config.num_hidden_layers == 1 else self.config.num_hidden_layers
)
return NextNEnabledConfig(
num_nextn_layers=num_nextn_layers,
nextn_layer_id=nextn_layer_id,
nextn_layer_prefix=f"model.layers.{nextn_layer_id}",
nextn_spec_weight_names=[
"shared_head.norm",
"eh_proj",
"enorm",
"hnorm",
],
)
def post_load_weights(
self: nn.Module,
self,
is_nextn: bool = False,
weight_names: Optional[Iterable[str]] = None,
) -> None:
@@ -577,56 +610,54 @@ class DeepseekV2WeightLoaderMixin:
self_attn.use_deep_gemm_bmm = True
def _maybe_quant_weights_to_fp8_ue8m0(
self, weights, attn_quant_modules, is_nextn=False
self,
weights,
attn_quant_modules,
nextn_conf: NextNConfig,
):
"""Optionally quantize weights to FP8 UE8M0 format for DeepSeek nvfp4 checkpoints.
Args:
weights: Iterable of (name, tensor) weight pairs
attn_quant_modules: List of attention module names to quantize
is_nextn: Whether processing NextN weights
nextn_conf: NextN configuration
Returns:
List of (name, tensor) pairs with quantized weights
"""
partial_names = []
nextn_layer_id = (
0 if self.config.num_hidden_layers == 1 else self.config.num_hidden_layers
)
weights_dict = dict(weights)
weight_block_size = [128, 128]
partial_names = []
if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get():
layer_ids = (
list(range(self.config.num_hidden_layers))
if not is_nextn
else [nextn_layer_id]
)
for layer_id in layer_ids:
for stem in attn_quant_modules:
partial_names.append(f"model.layers.{layer_id}.self_attn.{stem}")
match nextn_conf:
case NextNEnabledConfig(nextn_layer_id=layer_id):
if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get():
for stem in attn_quant_modules:
partial_names.append(
f"model.layers.{layer_id}.self_attn.{stem}"
)
if is_nextn and enable_nextn_moe_bf16_cast_to_fp8(self.quant_config):
for expert_sub_name in [
"shared_experts",
*[
f"experts.{expert_id}"
for expert_id in range(self.config.n_routed_experts)
],
]:
for stem in [
"gate_proj",
"up_proj",
"down_proj",
]:
partial_names.append(
f"model.layers.{nextn_layer_id}.mlp.{expert_sub_name}.{stem}"
)
if enable_nextn_moe_bf16_cast_to_fp8(self.quant_config):
expert_sub_names = ["shared_experts"] + [
f"experts.{i}" for i in range(self.config.n_routed_experts)
]
for expert_sub_name in expert_sub_names:
for stem in ["gate_proj", "up_proj", "down_proj"]:
partial_names.append(
f"model.layers.{layer_id}.mlp.{expert_sub_name}.{stem}"
)
if len(partial_names) > 0:
case NextNDisabledConfig():
if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get():
for layer_id in range(self.config.num_hidden_layers):
for stem in attn_quant_modules:
partial_names.append(
f"model.layers.{layer_id}.self_attn.{stem}"
)
if partial_names:
for partial_name in tqdm.tqdm(
partial_names,
desc="quant weights to fp8 ue8m0",
partial_names, desc="quant weights to fp8 ue8m0"
):
original_weight = weights_dict[f"{partial_name}.weight"]
out_w, out_s = quant_weight_ue8m0(
@@ -635,7 +666,9 @@ class DeepseekV2WeightLoaderMixin:
weights_dict[f"{partial_name}.weight"] = out_w
weights_dict[f"{partial_name}.weight_scale_inv"] = out_s
if is_nextn and enable_nextn_moe_bf16_cast_to_fp8(self.quant_config):
if isinstance(
nextn_conf, NextNEnabledConfig
) and enable_nextn_moe_bf16_cast_to_fp8(self.quant_config):
self._mark_nextn_moe_weights_as_ue8m0()
return list(weights_dict.items())