[NemotronH] PP support (#16172)

Signed-off-by: Roi Koren <roik@nvidia.com>
This commit is contained in:
roikoren755
2025-12-31 11:16:15 +08:00
committed by GitHub
parent c0fc7a89e7
commit 47a660d5b9
2 changed files with 94 additions and 35 deletions
+88 -35
View File
@@ -48,6 +48,7 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.topk import TopK from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.layers.vocab_parallel_embedding import (
DEFAULT_VOCAB_PADDING_SIZE, DEFAULT_VOCAB_PADDING_SIZE,
ParallelLMHead, ParallelLMHead,
@@ -65,7 +66,7 @@ from sglang.srt.utils import (
add_prefix, add_prefix,
get_current_device_stream_fast, get_current_device_stream_fast,
is_cuda, is_cuda,
make_layers_non_pp, make_layers,
) )
from sglang.utils import logger from sglang.utils import logger
@@ -526,21 +527,32 @@ class NemotronHModel(nn.Module):
) )
self.vocab_size = config.vocab_size + lora_vocab self.vocab_size = config.vocab_size + lora_vocab
self.org_vocab_size = config.vocab_size self.org_vocab_size = config.vocab_size
self.pp_group = get_pp_group()
self.embed_tokens = VocabParallelEmbedding( if self.pp_group.is_first_rank:
self.vocab_size, self.embed_tokens = VocabParallelEmbedding(
config.hidden_size, self.vocab_size,
org_num_embeddings=config.vocab_size, config.hidden_size,
) org_num_embeddings=config.vocab_size,
)
else:
self.embed_tokens = PPMissingLayer()
def get_layer(idx: int, prefix: str): def get_layer(idx: int, prefix: str):
layer_class = ALL_DECODER_LAYER_TYPES[config.hybrid_override_pattern[idx]] layer_class = ALL_DECODER_LAYER_TYPES[config.hybrid_override_pattern[idx]]
return layer_class(config, idx, quant_config=quant_config, prefix=prefix) return layer_class(config, idx, quant_config=quant_config, prefix=prefix)
self.layers = make_layers_non_pp( self.layers, self.start_layer, self.end_layer = make_layers(
len(config.hybrid_override_pattern), get_layer, prefix=f"{prefix}.layers" len(config.hybrid_override_pattern),
get_layer,
pp_rank=self.pp_group.rank_in_group,
pp_size=self.pp_group.world_size,
prefix=f"{prefix}.layers",
) )
self.norm_f = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon) if self.pp_group.is_last_rank:
self.norm_f = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
else:
self.norm_f = PPMissingLayer(return_tuple=True)
def forward( def forward(
self, self,
@@ -550,7 +562,7 @@ class NemotronHModel(nn.Module):
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
inputs_embeds: Optional[torch.Tensor] = None, inputs_embeds: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, PPProxyTensors]: ) -> Union[torch.Tensor, PPProxyTensors]:
if get_pp_group().is_first_rank: if self.pp_group.is_first_rank:
if inputs_embeds is not None: if inputs_embeds is not None:
hidden_states = inputs_embeds hidden_states = inputs_embeds
else: else:
@@ -561,8 +573,8 @@ class NemotronHModel(nn.Module):
hidden_states = pp_proxy_tensors["hidden_states"] hidden_states = pp_proxy_tensors["hidden_states"]
residual = pp_proxy_tensors["residual"] residual = pp_proxy_tensors["residual"]
residual = None for i in range(self.start_layer, self.end_layer):
for layer in self.layers: layer = self.layers[i]
if not isinstance(layer, Layers): if not isinstance(layer, Layers):
raise ValueError(f"Unknown layer type: {type(layer)}") raise ValueError(f"Unknown layer type: {type(layer)}")
hidden_states, residual = layer.forward( hidden_states, residual = layer.forward(
@@ -571,7 +583,7 @@ class NemotronHModel(nn.Module):
forward_batch=forward_batch, forward_batch=forward_batch,
) )
if not get_pp_group().is_last_rank: if not self.pp_group.is_last_rank:
return PPProxyTensors( return PPProxyTensors(
{"hidden_states": hidden_states, "residual": residual} {"hidden_states": hidden_states, "residual": residual}
) )
@@ -606,26 +618,45 @@ class NemotronHForCausalLM(nn.Module):
self.model = self._init_model( self.model = self._init_model(
config=config, quant_config=quant_config, prefix=prefix config=config, quant_config=quant_config, prefix=prefix
) )
if self.config.tie_word_embeddings: self.pp_group = get_pp_group()
self.lm_head = self.model.embed_tokens
if self.pp_group.is_last_rank:
if self.pp_group.world_size == 1 and self.config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
else:
self.unpadded_vocab_size = config.vocab_size
if lora_config:
self.unpadded_vocab_size += lora_config.lora_extra_vocab_size
self.lm_head = ParallelLMHead(
self.unpadded_vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
padding_size=(
DEFAULT_VOCAB_PADDING_SIZE
# We need bigger padding if using lora for kernel
# compatibility
if not lora_config
else lora_config.lora_vocab_padding_size
),
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
else: else:
self.unpadded_vocab_size = config.vocab_size self.lm_head = PPMissingLayer()
if lora_config:
self.unpadded_vocab_size += lora_config.lora_extra_vocab_size if self.pp_group.world_size > 1 and self.config.tie_word_embeddings:
self.lm_head = ParallelLMHead( if self.pp_group.is_first_rank:
self.unpadded_vocab_size, self.pp_group.send(
config.hidden_size, self.model.embed_tokens.weight, dst=self.pp_group.last_rank
org_num_embeddings=config.vocab_size, )
padding_size=( elif self.pp_group.is_last_rank:
DEFAULT_VOCAB_PADDING_SIZE emb_token_weight = self.pp_group.recv(
# We need bigger padding if using lora for kernel size=self.lm_head.weight.shape,
# compatibility dtype=next(self.model.parameters()).dtype,
if not lora_config src=self.pp_group.first_rank,
else lora_config.lora_vocab_padding_size )
), self.lm_head.weight.copy_(emb_token_weight)
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
self.logits_processor = LogitsProcessor(config) self.logits_processor = LogitsProcessor(config)
def _init_model( def _init_model(
@@ -653,9 +684,12 @@ class NemotronHForCausalLM(nn.Module):
hidden_states = self.model.forward( hidden_states = self.model.forward(
input_ids, positions, forward_batch, pp_proxy_tensors, input_embeds input_ids, positions, forward_batch, pp_proxy_tensors, input_embeds
) )
return self.logits_processor( if self.pp_group.is_last_rank:
input_ids, hidden_states, self.lm_head, forward_batch return self.logits_processor(
) input_ids, hidden_states, self.lm_head, forward_batch
)
else:
return hidden_states
def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs): def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
return self.mamba_cache.copy_inputs_before_cuda_graphs(input_buffers, **kwargs) return self.mamba_cache.copy_inputs_before_cuda_graphs(input_buffers, **kwargs)
@@ -689,6 +723,25 @@ class NemotronHForCausalLM(nn.Module):
if name is None: if name is None:
continue continue
layer_id = get_layer_id(name)
if (
layer_id is not None
and hasattr(self.model, "start_layer")
and (
layer_id < self.model.start_layer
or layer_id >= self.model.end_layer
)
):
continue
if "embed_tokens" in name and not self.pp_group.is_first_rank:
continue
if (
"norm_f" in name or "lm_head" in name
) and not self.pp_group.is_last_rank:
continue
for param_name, weight_name, shard_id in self.stacked_params_mapping: for param_name, weight_name, shard_id in self.stacked_params_mapping:
if weight_name not in name: if weight_name not in name:
continue continue
@@ -11,6 +11,12 @@ class TestNvidiaNemotronNanoV2BF16(GSM8KMixin, DefaultServerBase):
other_args = ["--max-mamba-cache-size", "256"] other_args = ["--max-mamba-cache-size", "256"]
class TestNvidiaNemotronNanoV2BF16PP(GSM8KMixin, DefaultServerBase):
model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2"
gsm8k_accuracy_thres = 0.87
other_args = ["--max-mamba-cache-size", "256", "--pp-size", "2"]
class TestNvidiaNemotronNanoV2FP8(GSM8KMixin, DefaultServerBase): class TestNvidiaNemotronNanoV2FP8(GSM8KMixin, DefaultServerBase):
gsm8k_accuracy_thres = 0.87 gsm8k_accuracy_thres = 0.87
model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2-FP8" model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2-FP8"