[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.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()
|
||||||
|
|
||||||
|
if self.pp_group.is_first_rank:
|
||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
self.vocab_size,
|
self.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
org_num_embeddings=config.vocab_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",
|
||||||
)
|
)
|
||||||
|
if self.pp_group.is_last_rank:
|
||||||
self.norm_f = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
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,7 +618,10 @@ 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()
|
||||||
|
|
||||||
|
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
|
self.lm_head = self.model.embed_tokens
|
||||||
else:
|
else:
|
||||||
self.unpadded_vocab_size = config.vocab_size
|
self.unpadded_vocab_size = config.vocab_size
|
||||||
@@ -626,6 +641,22 @@ class NemotronHForCausalLM(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("lm_head", prefix),
|
prefix=add_prefix("lm_head", prefix),
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
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)
|
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
|
||||||
)
|
)
|
||||||
|
if self.pp_group.is_last_rank:
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids, hidden_states, self.lm_head, forward_batch
|
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"
|
||||||
|
|||||||
Reference in New Issue
Block a user