Feat/nemotron nano v3 support (#12690)

This commit is contained in:
roikoren755
2025-11-21 13:53:05 -08:00
committed by GitHub
parent a24aefe562
commit 1b48e1b974
13 changed files with 776 additions and 68 deletions
+25 -6
View File
@@ -26,6 +26,7 @@ logger = logging.get_logger(__name__)
MAMBA = "M"
ATTENTION = "*"
MLP = "-"
MOE = "E"
class NemotronHConfig(PretrainedConfig):
@@ -189,6 +190,15 @@ class NemotronHConfig(PretrainedConfig):
mamba_proj_bias=False,
mamba_chunk_size=256,
rescale_prenorm_residual=True,
n_routed_experts=8,
n_shared_experts=1,
moe_intermediate_size=7688,
moe_shared_expert_intermediate_size=7688,
num_experts_per_tok=2,
routed_scaling_factor=1.0,
n_group=1,
topk_group=1,
norm_topk_prob=True,
**kwargs,
):
self.vocab_size = vocab_size
@@ -206,12 +216,12 @@ class NemotronHConfig(PretrainedConfig):
# Validate hybrid_override_pattern
# M: Mamba2, *: Attention, -: MLP
assert len(self.hybrid_override_pattern) == self.num_hidden_layers, (
"hybrid_override_pattern must have same length as " "num_hidden_layers"
)
assert re.match(r"^[*-M]+$", self.hybrid_override_pattern), (
"hybrid_override_pattern must only contain characters " "'M', '*', or '-'"
)
assert (
len(self.hybrid_override_pattern) == self.num_hidden_layers
), "hybrid_override_pattern must have same length as num_hidden_layers"
assert re.match(
r"^[*\-ME]+$", self.hybrid_override_pattern
), "hybrid_override_pattern must only contain characters 'M', '*', '-' or 'E'"
# for backward compatibility
if num_key_value_heads is None:
@@ -245,6 +255,15 @@ class NemotronHConfig(PretrainedConfig):
self.mamba_proj_bias = mamba_proj_bias
self.mamba_chunk_size = mamba_chunk_size
self.rescale_prenorm_residual = rescale_prenorm_residual
self.n_routed_experts = n_routed_experts
self.n_shared_experts = n_shared_experts
self.moe_intermediate_size = moe_intermediate_size
self.moe_shared_expert_intermediate_size = moe_shared_expert_intermediate_size
self.num_experts_per_tok = num_experts_per_tok
self.routed_scaling_factor = routed_scaling_factor
self.n_group = n_group
self.topk_group = topk_group
self.norm_topk_prob = norm_topk_prob
super().__init__(
pad_token_id=pad_token_id,