Feat/nemotron nano v3 support (#12690)
This commit is contained in:
@@ -22,8 +22,13 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.configs import NemotronHConfig
|
||||
from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA, MLP
|
||||
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size
|
||||
from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA, MLP, MOE
|
||||
from sglang.srt.distributed import (
|
||||
get_moe_ep_group,
|
||||
get_pp_group,
|
||||
get_tensor_model_parallel_world_size,
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.srt.layers.activation import ReLU2
|
||||
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
||||
HybridLinearAttnBackend,
|
||||
@@ -34,9 +39,13 @@ from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
ReplicatedLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||
from sglang.srt.layers.moe.topk import TopK
|
||||
from sglang.srt.layers.quantization import QuantizationConfig
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
@@ -51,31 +60,30 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
replace_prefix,
|
||||
replace_substrings,
|
||||
)
|
||||
from sglang.srt.utils import add_prefix, make_layers_non_pp
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
get_current_device_stream_fast,
|
||||
is_cuda,
|
||||
make_layers_non_pp,
|
||||
)
|
||||
from sglang.utils import logger
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
|
||||
|
||||
class NemotronHMLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: NemotronHConfig,
|
||||
layer_idx: int,
|
||||
intermediate_size: int,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
bias: bool = False,
|
||||
reduce_results: bool = True,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hybrid_override_pattern = config.hybrid_override_pattern
|
||||
mlp_index = hybrid_override_pattern[: layer_idx + 1].count("-") - 1
|
||||
if isinstance(config.intermediate_size, list):
|
||||
if len(config.intermediate_size) == 1:
|
||||
intermediate_size = config.intermediate_size[0]
|
||||
else:
|
||||
intermediate_size = config.intermediate_size[mlp_index]
|
||||
else:
|
||||
intermediate_size = config.intermediate_size
|
||||
|
||||
self.up_proj = ColumnParallelLinear(
|
||||
input_size=config.hidden_size,
|
||||
output_size=intermediate_size,
|
||||
@@ -88,6 +96,7 @@ class NemotronHMLP(nn.Module):
|
||||
output_size=config.hidden_size,
|
||||
bias=bias,
|
||||
quant_config=quant_config,
|
||||
reduce_results=reduce_results,
|
||||
prefix=f"{prefix}.down_proj",
|
||||
)
|
||||
self.act_fn = ReLU2()
|
||||
@@ -99,6 +108,148 @@ class NemotronHMLP(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
_alt_stream = None
|
||||
|
||||
|
||||
def _get_or_create_alt_stream(device_module):
|
||||
global _alt_stream
|
||||
if _alt_stream is None:
|
||||
_alt_stream = device_module.Stream()
|
||||
return _alt_stream
|
||||
|
||||
|
||||
class NemotronHMoE(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: NemotronHConfig,
|
||||
layer_idx: int,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.routed_scaling_factor = config.routed_scaling_factor
|
||||
self.device_module = torch.get_device_module()
|
||||
|
||||
self.ep_group = get_moe_ep_group().device_group
|
||||
self.ep_rank = self.ep_group.rank()
|
||||
self.ep_size = self.ep_group.size()
|
||||
self.n_routed_experts = config.n_routed_experts
|
||||
self.n_shared_experts = config.n_shared_experts
|
||||
|
||||
self.gate = ReplicatedLinear(
|
||||
config.hidden_size,
|
||||
config.n_routed_experts,
|
||||
bias=False,
|
||||
params_dtype=torch.float32,
|
||||
quant_config=None,
|
||||
prefix=f"{prefix}.gate",
|
||||
)
|
||||
self.gate.e_score_correction_bias = nn.Parameter(
|
||||
torch.empty(config.n_routed_experts, dtype=torch.float32)
|
||||
)
|
||||
|
||||
self.topk = TopK(
|
||||
top_k=config.num_experts_per_tok,
|
||||
use_grouped_topk=True,
|
||||
topk_group=config.topk_group,
|
||||
num_expert_group=config.n_group,
|
||||
renormalize=config.norm_topk_prob,
|
||||
scoring_func="sigmoid",
|
||||
correction_bias=self.gate.e_score_correction_bias,
|
||||
routed_scaling_factor=1.0,
|
||||
)
|
||||
self.experts = get_moe_impl_class(quant_config)(
|
||||
num_experts=config.n_routed_experts
|
||||
+ get_global_server_args().ep_num_redundant_experts,
|
||||
top_k=config.num_experts_per_tok,
|
||||
hidden_size=config.hidden_size,
|
||||
intermediate_size=config.moe_intermediate_size,
|
||||
reduce_results=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.experts",
|
||||
activation=config.mlp_hidden_act,
|
||||
layer_id=layer_idx,
|
||||
is_gated=False,
|
||||
)
|
||||
if config.n_shared_experts:
|
||||
self.shared_experts = NemotronHMLP(
|
||||
config,
|
||||
intermediate_size=config.moe_shared_expert_intermediate_size
|
||||
* config.n_shared_experts,
|
||||
quant_config=quant_config,
|
||||
reduce_results=False,
|
||||
prefix=f"{prefix}.shared_experts",
|
||||
)
|
||||
else:
|
||||
self.shared_experts = None
|
||||
|
||||
def _forward_core(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
if _is_cuda:
|
||||
return self._forward_core_shared_routed_overlap(hidden_states)
|
||||
else:
|
||||
return self._forward_core_normal(hidden_states)
|
||||
|
||||
def _forward_core_normal(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
# router_scores: [num_tokens, num_experts]
|
||||
router_logits, _ = self.gate(hidden_states.to(dtype=torch.float32))
|
||||
if self.shared_experts is not None:
|
||||
shared_output = self.shared_experts(hidden_states)
|
||||
else:
|
||||
shared_output = None
|
||||
topk_output = self.topk(hidden_states, router_logits)
|
||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||
return final_hidden_states, shared_output
|
||||
|
||||
def _forward_core_shared_routed_overlap(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
alt_stream = _get_or_create_alt_stream(self.device_module)
|
||||
|
||||
alt_stream.wait_stream(get_current_device_stream_fast())
|
||||
|
||||
if self.shared_experts is not None:
|
||||
shared_output = self.shared_experts(hidden_states)
|
||||
else:
|
||||
shared_output = None
|
||||
|
||||
with self.device_module.stream(alt_stream):
|
||||
# router_scores: [num_tokens, num_experts]
|
||||
router_logits, _ = self.gate(hidden_states.to(dtype=torch.float32))
|
||||
topk_output = self.topk(hidden_states, router_logits)
|
||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||
get_current_device_stream_fast().wait_stream(alt_stream)
|
||||
|
||||
return final_hidden_states, shared_output
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
num_tokens, hidden_dim = hidden_states.shape
|
||||
final_hidden_states, shared_output = self._forward_core(hidden_states)
|
||||
|
||||
# Fix FP16 overflow
|
||||
if hidden_states.dtype != torch.float16:
|
||||
final_hidden_states *= self.routed_scaling_factor
|
||||
elif self.shared_experts is not None:
|
||||
assert shared_output is not None
|
||||
shared_output *= 1.0 / self.routed_scaling_factor
|
||||
|
||||
if shared_output is not None:
|
||||
final_hidden_states += shared_output
|
||||
|
||||
if self.tp_size > 1:
|
||||
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
|
||||
|
||||
return final_hidden_states.view(num_tokens, hidden_dim)
|
||||
|
||||
|
||||
class NemotronHMLPDecoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -110,15 +261,61 @@ class NemotronHMLPDecoderLayer(nn.Module):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
|
||||
hybrid_override_pattern = config.hybrid_override_pattern
|
||||
mlp_index = hybrid_override_pattern[: layer_idx + 1].count("-") - 1
|
||||
if isinstance(config.intermediate_size, list):
|
||||
if len(config.intermediate_size) == 1:
|
||||
intermediate_size = config.intermediate_size[0]
|
||||
else:
|
||||
intermediate_size = config.intermediate_size[mlp_index]
|
||||
else:
|
||||
intermediate_size = config.intermediate_size
|
||||
|
||||
self.mixer = NemotronHMLP(
|
||||
config,
|
||||
intermediate_size=intermediate_size,
|
||||
quant_config=quant_config,
|
||||
bias=config.mlp_bias,
|
||||
prefix=f"{prefix}.mixer",
|
||||
layer_idx=layer_idx,
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
*,
|
||||
hidden_states: torch.Tensor,
|
||||
residual: Optional[torch.Tensor],
|
||||
forward_batch: ForwardBatch,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
hidden_states = self.norm(hidden_states)
|
||||
else:
|
||||
hidden_states, residual = self.norm(hidden_states, residual)
|
||||
|
||||
hidden_states = self.mixer.forward(hidden_states)
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class NemotronHMoEDecoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: NemotronHConfig,
|
||||
layer_idx: int,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.mixer = NemotronHMoE(
|
||||
config,
|
||||
layer_idx=layer_idx,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.mixer",
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -154,13 +351,13 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
||||
use_conv_bias=config.use_conv_bias,
|
||||
use_bias=config.use_bias,
|
||||
n_groups=config.mamba_n_groups,
|
||||
rms_norm_eps=config.rms_norm_eps,
|
||||
rms_norm_eps=config.layer_norm_epsilon,
|
||||
activation=config.mamba_hidden_act,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.mixer",
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -275,7 +472,7 @@ class NemotronHAttentionDecoderLayer(nn.Module):
|
||||
prefix=f"{prefix}.mixer",
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -300,11 +497,13 @@ Layers = (
|
||||
NemotronHAttentionDecoderLayer
|
||||
| NemotronHMLPDecoderLayer
|
||||
| NemotronHMambaDecoderLayer
|
||||
| NemotronHMoEDecoderLayer
|
||||
)
|
||||
ALL_DECODER_LAYER_TYPES: dict[str, type[Layers]] = {
|
||||
ATTENTION: NemotronHAttentionDecoderLayer,
|
||||
MLP: NemotronHMLPDecoderLayer,
|
||||
MAMBA: NemotronHMambaDecoderLayer,
|
||||
MOE: NemotronHMoEDecoderLayer,
|
||||
}
|
||||
|
||||
|
||||
@@ -341,7 +540,7 @@ class NemotronHModel(nn.Module):
|
||||
self.layers = make_layers_non_pp(
|
||||
len(config.hybrid_override_pattern), get_layer, prefix=f"{prefix}.layers"
|
||||
)
|
||||
self.norm_f = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
self.norm_f = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.embed_tokens(input_ids)
|
||||
@@ -473,6 +672,18 @@ class NemotronHForCausalLM(nn.Module):
|
||||
name = replace_prefix(name, self.remap_prefix)
|
||||
name = replace_substrings(name, self.remap_substr)
|
||||
updated_weights.append((name, loaded_weight))
|
||||
|
||||
# - FusedMoe.w1 (aka gate_proj) should be up_proj since that's
|
||||
# what the activation is applied to
|
||||
# - FusedMoe.w3 (aka up_proj) should be ignored since we're
|
||||
# using non-gated MoE
|
||||
expert_params_mapping = FusedMoE.make_expert_params_mapping(
|
||||
ckpt_gate_proj_name="up_proj",
|
||||
ckpt_down_proj_name="down_proj",
|
||||
ckpt_up_proj_name="",
|
||||
num_experts=self.config.n_routed_experts,
|
||||
)
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
|
||||
for name, loaded_weight in updated_weights:
|
||||
@@ -495,17 +706,37 @@ class NemotronHForCausalLM(nn.Module):
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if name in params_dict.keys():
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(
|
||||
param, "weight_loader", default_weight_loader
|
||||
is_expert_weight = False
|
||||
for mapping in expert_params_mapping:
|
||||
param_name, weight_name, expert_id, shard_id = mapping
|
||||
if weight_name not in name:
|
||||
continue
|
||||
is_expert_weight = True
|
||||
name_mapped = name.replace(weight_name, param_name)
|
||||
param = params_dict[name_mapped]
|
||||
param.weight_loader(
|
||||
param,
|
||||
loaded_weight,
|
||||
name_mapped,
|
||||
shard_id=shard_id,
|
||||
expert_id=expert_id,
|
||||
)
|
||||
weight_loader(param, loaded_weight)
|
||||
name = name_mapped
|
||||
break
|
||||
else:
|
||||
logger.warning(f"Parameter {name} not found in params_dict")
|
||||
if is_expert_weight:
|
||||
continue
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if name in params_dict.keys():
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(
|
||||
param, "weight_loader", default_weight_loader
|
||||
)
|
||||
weight_loader(param, loaded_weight)
|
||||
else:
|
||||
logger.warning(f"Parameter {name} not found in params_dict")
|
||||
|
||||
|
||||
EntryClass = [NemotronHForCausalLM]
|
||||
|
||||
Reference in New Issue
Block a user