# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 # Adapt from https://github.com/NVIDIA/Megatron-LM/blob/b1efb3c7126ef7615e8c333432d76e08038e17ff/pretrain_gpt.py import argparse import inspect from contextlib import nullcontext from typing import Literal import torch from megatron.core import tensor_parallel from megatron.core.models.gpt import GPTModel from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_decoder_block_spec, get_gpt_layer_local_spec, get_gpt_layer_with_transformer_engine_spec, ) from megatron.core.transformer.spec_utils import import_module from megatron.core.transformer.transformer_config import TransformerConfig from megatron.training.arguments import core_transformer_config_from_args # Adapt from https://github.com/volcengine/verl/blob/c3b20575d2bc815fcccd84bddb4c0401fc4b632b/verl/models/llama/megatron/layers/parallel_linear.py#L82 class LinearForLastLayer(torch.nn.Linear): def __init__( self, input_size: int, output_size: int, *, config: TransformerConfig, bias: bool = True, ) -> None: super().__init__(in_features=input_size, out_features=output_size, bias=bias) self.sequence_parallel = config.sequence_parallel if self.sequence_parallel: self.weight.sequence_parallel = True self.weight.data.normal_(mean=0.0, std=0.02) if bias: self.bias.data.zero_() def forward( self, input_: torch.Tensor, weight: torch.Tensor | None = None, runtime_gather_output: bool | None = None, ) -> tuple[torch.Tensor, None]: logits = super().forward(input_) logits = logits.float() if self.sequence_parallel: logits = tensor_parallel.gather_from_sequence_parallel_region(logits, tensor_parallel_output_grad=False) return logits, None def get_model_provider_func( args: argparse.Namespace, role: Literal["actor", "critic"] = "actor", ): if args.megatron_to_hf_mode == "bridge": from megatron.bridge import AutoBridge bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True) provider = bridge.to_megatron_provider(load_weights=False) # TODO: we should not manually set this... provider.tensor_model_parallel_size = args.tensor_model_parallel_size provider.pipeline_model_parallel_size = args.pipeline_model_parallel_size provider.expert_model_parallel_size = args.expert_model_parallel_size provider.expert_tensor_parallel_size = args.expert_tensor_parallel_size provider.sequence_parallel = args.sequence_parallel provider.finalize() return provider.provide def model_provider(pre_process: bool = True, post_process: bool = True, vp_stage: int | None = None) -> GPTModel: """Builds the model. If you set the use_legacy_models to True, it will return the legacy GPT model and if not the mcore GPT model. Args: pre_process (bool, optional): Set to true if you need to compute embedings. Defaults to True. post_process (bool, optional): Set to true if you need to want to compute output logits/loss. Defaults to True. Returns: Union[GPTModel, megatron.legacy.model.GPTModel]: The returned model """ use_te = args.transformer_impl == "transformer_engine" # Experimental loading arguments from yaml config: TransformerConfig = core_transformer_config_from_args(args) if args.spec is not None: transformer_layer_spec = import_module(args.spec) # Allow the spec to be a function so that user can use customized Megatron easier. if callable(transformer_layer_spec): transformer_layer_spec = transformer_layer_spec(args, config, vp_stage) else: if args.num_experts: # Define the decoder block spec kwargs = { "use_transformer_engine": use_te, } if vp_stage is not None: kwargs["vp_stage"] = vp_stage transformer_layer_spec = get_gpt_decoder_block_spec(config, **kwargs) else: # Define the decoder layer spec if use_te: transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( num_experts=args.num_experts, moe_grouped_gemm=args.moe_grouped_gemm, qk_layernorm=args.qk_layernorm, multi_latent_attention=args.multi_latent_attention, moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm, ) else: transformer_layer_spec = get_gpt_layer_local_spec( num_experts=args.num_experts, moe_grouped_gemm=args.moe_grouped_gemm, qk_layernorm=args.qk_layernorm, multi_latent_attention=args.multi_latent_attention, moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm, ) build_model_context = nullcontext build_model_context_args = {} if args.fp8_param_gather: try: from transformer_engine.pytorch import fp8_model_init build_model_context = fp8_model_init build_model_context_args["enabled"] = True # Check if fp8_model_init supports preserve_high_precision_init_val if "preserve_high_precision_init_val" in inspect.signature(fp8_model_init).parameters: build_model_context_args["preserve_high_precision_init_val"] = True except Exception as e: raise RuntimeError( "--fp8-param-gather requires `fp8_model_init` from TransformerEngine, but not found." ) from e kwargs = { "config": config, "transformer_layer_spec": transformer_layer_spec, "vocab_size": args.padded_vocab_size, "max_sequence_length": args.max_position_embeddings, "pre_process": pre_process, "post_process": post_process, "fp16_lm_cross_entropy": args.fp16_lm_cross_entropy, "parallel_output": True, "share_embeddings_and_output_weights": not args.untie_embeddings_and_output_weights, "position_embedding_type": args.position_embedding_type, "rotary_percent": args.rotary_percent, "rotary_base": args.rotary_base, "rope_scaling": args.use_rope_scaling, } if vp_stage is not None: kwargs["vp_stage"] = vp_stage if args.mtp_num_layers: from megatron.core.models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec mtp_kwargs = { "use_transformer_engine": use_te, } if vp_stage is not None: mtp_kwargs["vp_stage"] = vp_stage mtp_block_spec = get_gpt_mtp_block_spec(config, transformer_layer_spec, **mtp_kwargs) kwargs["mtp_block_spec"] = mtp_block_spec with build_model_context(**build_model_context_args): model = GPTModel(**kwargs) if post_process and role == "critic": model.output_layer = LinearForLastLayer(input_size=config.hidden_size, output_size=1, config=config) return model return model_provider