diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 85a0c24b3..81b92e2fb 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -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.quantization import QuantizationConfig 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 ( DEFAULT_VOCAB_PADDING_SIZE, ParallelLMHead, @@ -65,7 +66,7 @@ from sglang.srt.utils import ( add_prefix, get_current_device_stream_fast, is_cuda, - make_layers_non_pp, + make_layers, ) from sglang.utils import logger @@ -526,21 +527,32 @@ class NemotronHModel(nn.Module): ) self.vocab_size = config.vocab_size + lora_vocab self.org_vocab_size = config.vocab_size + self.pp_group = get_pp_group() - self.embed_tokens = VocabParallelEmbedding( - self.vocab_size, - config.hidden_size, - org_num_embeddings=config.vocab_size, - ) + if self.pp_group.is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + self.vocab_size, + config.hidden_size, + org_num_embeddings=config.vocab_size, + ) + else: + self.embed_tokens = PPMissingLayer() def get_layer(idx: int, prefix: str): layer_class = ALL_DECODER_LAYER_TYPES[config.hybrid_override_pattern[idx]] return layer_class(config, idx, quant_config=quant_config, prefix=prefix) - self.layers = make_layers_non_pp( - len(config.hybrid_override_pattern), get_layer, prefix=f"{prefix}.layers" + self.layers, self.start_layer, self.end_layer = make_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( self, @@ -550,7 +562,7 @@ class NemotronHModel(nn.Module): pp_proxy_tensors: Optional[PPProxyTensors] = None, inputs_embeds: Optional[torch.Tensor] = None, ) -> Union[torch.Tensor, PPProxyTensors]: - if get_pp_group().is_first_rank: + if self.pp_group.is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds else: @@ -561,8 +573,8 @@ class NemotronHModel(nn.Module): hidden_states = pp_proxy_tensors["hidden_states"] residual = pp_proxy_tensors["residual"] - residual = None - for layer in self.layers: + for i in range(self.start_layer, self.end_layer): + layer = self.layers[i] if not isinstance(layer, Layers): raise ValueError(f"Unknown layer type: {type(layer)}") hidden_states, residual = layer.forward( @@ -571,7 +583,7 @@ class NemotronHModel(nn.Module): forward_batch=forward_batch, ) - if not get_pp_group().is_last_rank: + if not self.pp_group.is_last_rank: return PPProxyTensors( {"hidden_states": hidden_states, "residual": residual} ) @@ -606,26 +618,45 @@ class NemotronHForCausalLM(nn.Module): self.model = self._init_model( config=config, quant_config=quant_config, prefix=prefix ) - if self.config.tie_word_embeddings: - self.lm_head = self.model.embed_tokens + self.pp_group = get_pp_group() + + 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: - 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), - ) + self.lm_head = PPMissingLayer() + + if self.pp_group.world_size > 1 and self.config.tie_word_embeddings: + if self.pp_group.is_first_rank: + self.pp_group.send( + self.model.embed_tokens.weight, dst=self.pp_group.last_rank + ) + elif self.pp_group.is_last_rank: + emb_token_weight = self.pp_group.recv( + size=self.lm_head.weight.shape, + dtype=next(self.model.parameters()).dtype, + src=self.pp_group.first_rank, + ) + self.lm_head.weight.copy_(emb_token_weight) + self.logits_processor = LogitsProcessor(config) def _init_model( @@ -653,9 +684,12 @@ class NemotronHForCausalLM(nn.Module): hidden_states = self.model.forward( input_ids, positions, forward_batch, pp_proxy_tensors, input_embeds ) - return self.logits_processor( - input_ids, hidden_states, self.lm_head, forward_batch - ) + if self.pp_group.is_last_rank: + 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): return self.mamba_cache.copy_inputs_before_cuda_graphs(input_buffers, **kwargs) @@ -689,6 +723,25 @@ class NemotronHForCausalLM(nn.Module): if name is None: 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: if weight_name not in name: continue diff --git a/test/srt/models/test_nvidia_nemotron_nano_v2.py b/test/srt/models/test_nvidia_nemotron_nano_v2.py index 336db3718..656eef838 100644 --- a/test/srt/models/test_nvidia_nemotron_nano_v2.py +++ b/test/srt/models/test_nvidia_nemotron_nano_v2.py @@ -11,6 +11,12 @@ class TestNvidiaNemotronNanoV2BF16(GSM8KMixin, DefaultServerBase): 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): gsm8k_accuracy_thres = 0.87 model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2-FP8"