[NemotronH] PP support (#16172)

Signed-off-by: Roi Koren <roik@nvidia.com>
This commit is contained in:
roikoren755
2025-12-31 05:16:15 +02:00
committed by GitHub
parent c0fc7a89e7
commit 47a660d5b9
2 changed files with 94 additions and 35 deletions

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.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