[Refactor] Algebraic data type for nextn config + some basic refactors (#17347)
This commit is contained in:
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user