[NemotronH] PP support (#16172)
Signed-off-by: Roi Koren <roik@nvidia.com>
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user