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

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