diff --git a/examples/mamba/run_text_gen_server_8b.sh b/examples/mamba/run_text_gen_server_8b.sh index d228e0c0edb..f183dea4ad1 100755 --- a/examples/mamba/run_text_gen_server_8b.sh +++ b/examples/mamba/run_text_gen_server_8b.sh @@ -22,7 +22,7 @@ export NCCL_IB_QPS_PER_CONNECTION=4 export TRITON_CACHE_DIR="./triton-cache/" export TRITON_CACHE_MANAGER="megatron.core.ssm.triton_cache_manager:ParallelFileCacheManager" -torchrun $DISTRIBUTED_ARGS ../../tools/run_mamba_text_generation_server.py \ +torchrun $DISTRIBUTED_ARGS ../../tools/run_hybrid_text_generation_server.py \ --tensor-model-parallel-size 1 \ --pipeline-model-parallel-size 1 \ --untie-embeddings-and-output-weights \ @@ -46,5 +46,5 @@ torchrun $DISTRIBUTED_ARGS ../../tools/run_mamba_text_generation_server.py \ --bf16 \ --micro-batch-size 1 \ --use-mcore-models \ - --spec megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec \ + --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ --seed 42 diff --git a/examples/mamba/train.sh b/examples/mamba/train.sh index ba83f0d4e33..f971242ff0b 100755 --- a/examples/mamba/train.sh +++ b/examples/mamba/train.sh @@ -96,8 +96,8 @@ options=" \ --eval-iters 32 \ --bf16 \ --use-mcore-models \ - --spec megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec \ + --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ --no-create-attention-mask-in-dataloader \ --tensorboard-dir ${TENSORBOARD_DIR}" -torchrun --nproc_per_node 8 ../../pretrain_mamba.py ${options} +torchrun --nproc_per_node 8 ../../pretrain_hybrid.py ${options} diff --git a/examples/multimodal/layer_specs.py b/examples/multimodal/layer_specs.py index ad24850b631..acced15eeb6 100644 --- a/examples/multimodal/layer_specs.py +++ b/examples/multimodal/layer_specs.py @@ -1,8 +1,8 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. import torch from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add -from megatron.core.ssm.mamba_block import MambaStack, MambaStackSubmodules +from megatron.core.models.hybrid.hybrid_block import HybridStack, HybridStackSubmodules from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules from megatron.core.ssm.mlp_layer import MLPLayer @@ -125,15 +125,15 @@ def get_layer_spec_te(is_vit=False, padding=False) -> ModuleSpec: ) -def get_mamba_layer_spec_te(padding=False) -> ModuleSpec: +def get_hybrid_layer_spec_te(padding=False) -> ModuleSpec: attn_mask_type = AttnMaskType.causal # Padding mask is needed for e.g. Context Parallel. if padding: attn_mask_type = AttnMaskType.padding_causal return ModuleSpec( - module=MambaStack, - submodules=MambaStackSubmodules( + module=HybridStack, + submodules=HybridStackSubmodules( mamba_layer=ModuleSpec( module=MambaLayer, submodules=MambaLayerSubmodules( diff --git a/examples/multimodal/model.py b/examples/multimodal/model.py index 494a854099e..a2d83428338 100644 --- a/examples/multimodal/model.py +++ b/examples/multimodal/model.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. import warnings import logging from copy import deepcopy @@ -6,7 +6,7 @@ import torch from config import get_language_model_config, get_vision_model_config, get_vision_projection_config from layer_specs import (get_layer_spec, get_layer_spec_te, get_mlp_module_spec, get_norm_mlp_module_spec_te, - get_mamba_layer_spec_te) + get_hybrid_layer_spec_te) from megatron.core.models.multimodal.llava_model import IMAGE_TOKEN, LLaVAModel from megatron.core.models.vision.clip_vit_model import get_num_image_embeddings @@ -99,7 +99,7 @@ def model_provider( # Padding mask needed for SP/CP. padding = args.context_parallel_size > 1 and args.sequence_parallel if args.language_model_type.startswith('nemotron5-hybrid'): - language_transformer_layer_spec = get_mamba_layer_spec_te(padding=padding) + language_transformer_layer_spec = get_hybrid_layer_spec_te(padding=padding) else: language_transformer_layer_spec = get_layer_spec_te( is_vit=False, padding=padding diff --git a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16.sh b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16.sh index 1fa00889e99..805302498fc 100644 --- a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16.sh +++ b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16.sh @@ -51,5 +51,5 @@ MODEL_ARGS=" \ --bf16 \ --seq-length 8192 \ --max-position-embeddings 8192 \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " diff --git a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16.sh b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16.sh index 977be033df0..b9da9429eb5 100644 --- a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16.sh +++ b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-BF16.sh @@ -58,5 +58,5 @@ MODEL_ARGS=" \ --bf16 \ --seq-length 8192 \ --max-position-embeddings 8192 \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " diff --git a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-Nano-9B-v2.sh b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-Nano-9B-v2.sh index 83867430a97..51aff10a22a 100644 --- a/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-Nano-9B-v2.sh +++ b/examples/post_training/modelopt/conf/nvidia/NVIDIA-Nemotron-Nano-9B-v2.sh @@ -35,6 +35,6 @@ MODEL_ARGS=" \ --tokenizer-type HuggingFaceTokenizer \ --make-vocab-size-divisible-by 1 \ --use-mcore-models \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ --padded-vocab-size 131072 \ " diff --git a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-47B-Reasoning-128K.sh b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-47B-Reasoning-128K.sh index 901e607f298..e2da6a3c33d 100644 --- a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-47B-Reasoning-128K.sh +++ b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-47B-Reasoning-128K.sh @@ -33,5 +33,5 @@ MODEL_ARGS=" \ --max-position-embeddings 8192 \ --tokenizer-type HuggingFaceTokenizer \ --use-mcore-models \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " diff --git a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-4B-Instruct.sh b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-4B-Instruct.sh index 084db49e0eb..523f7d521b0 100644 --- a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-4B-Instruct.sh +++ b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-4B-Instruct.sh @@ -38,5 +38,5 @@ MODEL_ARGS=" \ --make-vocab-size-divisible-by 1 \ --use-mcore-models \ --rotary-base 10000 \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " diff --git a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-56B-Base-8K.sh b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-56B-Base-8K.sh index 645a159d075..be80d8a9a19 100644 --- a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-56B-Base-8K.sh +++ b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-56B-Base-8K.sh @@ -35,5 +35,5 @@ MODEL_ARGS=" \ --max-position-embeddings 8192 \ --tokenizer-type HuggingFaceTokenizer \ --bf16 \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " diff --git a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-8B-Base-8K.sh b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-8B-Base-8K.sh index 66f3ad368b4..36b242e36dd 100644 --- a/examples/post_training/modelopt/conf/nvidia/Nemotron-H-8B-Base-8K.sh +++ b/examples/post_training/modelopt/conf/nvidia/Nemotron-H-8B-Base-8K.sh @@ -37,6 +37,6 @@ MODEL_ARGS=" \ --use-mcore-models \ --rotary-percent 0.5 \ --rotary-base 500000 \ - --export-model-type MambaModel \ + --export-model-type HybridModel \ " # --rotary-base 10000 \ diff --git a/examples/post_training/modelopt/convert_model.py b/examples/post_training/modelopt/convert_model.py index eaec9789e1e..136b0273724 100644 --- a/examples/post_training/modelopt/convert_model.py +++ b/examples/post_training/modelopt/convert_model.py @@ -19,7 +19,7 @@ from megatron.core.parallel_state import destroy_model_parallel from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import ( report_current_memory_info, to_empty_if_meta, @@ -129,7 +129,7 @@ def check_arguments(): ) model = get_model( - functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False + functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False ) report_current_memory_info() diff --git a/examples/post_training/modelopt/distillation.md b/examples/post_training/modelopt/distillation.md index 49f73c4edde..9946723364e 100644 --- a/examples/post_training/modelopt/distillation.md +++ b/examples/post_training/modelopt/distillation.md @@ -53,7 +53,7 @@ Without this configuration file, the default logits-only distillation with scale ### Training -Distillation is triggered by calling `pretrain_gpt.py` or `pretrain_mamba.py` with the following arguments: +Distillation is triggered by calling `pretrain_gpt.py` or `pretrain_hybrid.py` with the following arguments: ```bash --export-kd-teacher-load diff --git a/examples/post_training/modelopt/export.py b/examples/post_training/modelopt/export.py index 5e3b2a1716e..31c72d95eab 100755 --- a/examples/post_training/modelopt/export.py +++ b/examples/post_training/modelopt/export.py @@ -15,7 +15,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.training import get_args, get_model from megatron.training.initialize import initialize_megatron from megatron.training.utils import unwrap_model @@ -74,7 +74,7 @@ def add_modelopt_export_args(parser): ) model = get_model( - functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False + functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False ) # Materialize the model from meta device to cpu before loading the checkpoint. diff --git a/examples/post_training/modelopt/finetune.py b/examples/post_training/modelopt/finetune.py index f7f7c24f970..2efd3cde6a4 100755 --- a/examples/post_training/modelopt/finetune.py +++ b/examples/post_training/modelopt/finetune.py @@ -19,7 +19,7 @@ from megatron.core.models.gpt import GPTModel from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.loss_func import loss_func -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.non_loss_data_func import report_draft_acceptance_length from megatron.training import get_args, get_timers, pretrain from megatron.training.utils import ( @@ -486,7 +486,7 @@ def forward_step(data_iterator, model: GPTModel): if __name__ == "__main__": pretrain( train_valid_test_sft_datasets_provider, - partial(model_provider, modelopt_gpt_mamba_builder), + partial(model_provider, modelopt_gpt_hybrid_builder), ModelType.encoder_or_decoder, forward_step, extra_args_provider=add_finetune_args, diff --git a/examples/post_training/modelopt/generate.py b/examples/post_training/modelopt/generate.py index 3d3f6571b34..cc4c4e37a80 100644 --- a/examples/post_training/modelopt/generate.py +++ b/examples/post_training/modelopt/generate.py @@ -14,7 +14,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint from megatron.post_training.generate import simple_generate -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import report_current_memory_info, to_empty_if_meta from megatron.training import get_args, get_model, initialize_megatron from utils import get_hf_tokenizer @@ -100,7 +100,7 @@ def get_conversations(example): UserWarning, ) - model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False) report_current_memory_info() unwrapped_model = unwrap_model(model)[0] diff --git a/examples/post_training/modelopt/mmlu.py b/examples/post_training/modelopt/mmlu.py index 5aa5d1c24c7..466d5052b50 100644 --- a/examples/post_training/modelopt/mmlu.py +++ b/examples/post_training/modelopt/mmlu.py @@ -17,7 +17,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint from megatron.post_training.generate import simple_generate -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import report_current_memory_info from megatron.training import get_args, get_model, initialize_megatron from utils import get_hf_tokenizer @@ -158,7 +158,7 @@ def generate_prompt(test_example, dev_examples, few_shots=0, no_subject_prompt=F UserWarning, ) - model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False) report_current_memory_info() # Materialize the model from meta device to gpu before loading the checkpoint. diff --git a/examples/post_training/modelopt/offline_feature_extract.py b/examples/post_training/modelopt/offline_feature_extract.py index 80207faf2b2..92500b2950e 100644 --- a/examples/post_training/modelopt/offline_feature_extract.py +++ b/examples/post_training/modelopt/offline_feature_extract.py @@ -14,7 +14,7 @@ from megatron.core import mpu from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron from megatron.training.utils import print_rank_0, unwrap_model from model_provider import model_provider @@ -53,7 +53,7 @@ def extract_feature(dataset, model, output_dir, idx_start, idx_end): args = get_args() tokenizer = get_tokenizer() - model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False) load_modelopt_checkpoint(model, strict=not args.untie_embeddings_and_output_weights) print_rank_0("Done loading checkpoint") diff --git a/examples/post_training/modelopt/prune.py b/examples/post_training/modelopt/prune.py index 56bbffa0cd0..99e351a6198 100644 --- a/examples/post_training/modelopt/prune.py +++ b/examples/post_training/modelopt/prune.py @@ -28,7 +28,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint from megatron.post_training.generate import simple_generate -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import ( report_current_memory_info, ) @@ -163,7 +163,7 @@ def get_params(model): tokenizer = get_hf_tokenizer() model = get_model( - functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False + functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False ) unwrapped_model = unwrap_model(model)[0] diff --git a/examples/post_training/modelopt/quantize.py b/examples/post_training/modelopt/quantize.py index dc4947038e5..0c10696df84 100644 --- a/examples/post_training/modelopt/quantize.py +++ b/examples/post_training/modelopt/quantize.py @@ -39,7 +39,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint from megatron.post_training.generate import simple_generate -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import ( print_distributed_quant_summary, report_current_memory_info, @@ -362,7 +362,7 @@ def get_calib_dataloader( tokenizer = get_hf_tokenizer() model = get_model( - functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False + functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False ) report_current_memory_info() diff --git a/examples/post_training/modelopt/train.sh b/examples/post_training/modelopt/train.sh index 1ebb8bf3d76..3afcd4f5be7 100755 --- a/examples/post_training/modelopt/train.sh +++ b/examples/post_training/modelopt/train.sh @@ -69,8 +69,8 @@ fi export HF_TOKEN=${HF_TOKEN} -if [[ ${MODEL_ARGS} == *"MambaModel"* ]]; then - PRETRAIN_EXE=${SCRIPT_DIR}/../../../pretrain_mamba.py +if [[ ${MODEL_ARGS} == *"HybridModel"* ]] || [[ ${MODEL_ARGS} == *"MambaModel"* ]]; then + PRETRAIN_EXE=${SCRIPT_DIR}/../../../pretrain_hybrid.py else PRETRAIN_EXE=${SCRIPT_DIR}/../../../pretrain_gpt.py fi diff --git a/examples/post_training/modelopt/validate.py b/examples/post_training/modelopt/validate.py index 8b8f1ffc9dd..4d9757da00c 100644 --- a/examples/post_training/modelopt/validate.py +++ b/examples/post_training/modelopt/validate.py @@ -14,7 +14,7 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_builder import modelopt_gpt_mamba_builder +from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder from megatron.post_training.utils import get_mtbench_chat_data from megatron.training import get_args, get_model, initialize_megatron from utils import get_hf_tokenizer @@ -116,7 +116,7 @@ def report_current_memory_info(): ground_truth = [None for _ in range(len(prompts))] tokenizer = get_hf_tokenizer() - model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_hybrid_builder), wrap_with_ddp=False) report_current_memory_info() diff --git a/examples/rl/model_configs/nemotron5_56b.sh b/examples/rl/model_configs/nemotron5_56b.sh index 23b9f99a72a..b4fcee17a8e 100644 --- a/examples/rl/model_configs/nemotron5_56b.sh +++ b/examples/rl/model_configs/nemotron5_56b.sh @@ -69,7 +69,7 @@ MODEL_OPTIONS="\ \ --fp8-recipe tensorwise \ --hybrid-layer-pattern M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- \ - --spec megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec \ + --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ --mamba-state-dim 256 \ --per-split-data-args-path ${BLEND_PATH} \ --tiktoken-pattern v2 \ diff --git a/examples/rl/model_configs/nemotron5_8b.sh b/examples/rl/model_configs/nemotron5_8b.sh index c18149f03d6..198efd2a163 100644 --- a/examples/rl/model_configs/nemotron5_8b.sh +++ b/examples/rl/model_configs/nemotron5_8b.sh @@ -61,7 +61,7 @@ MODEL_OPTIONS="\ --inference-max-requests $MAX_INFERENCE_BS \ --pretrained-checkpoint $CHECKPOINT \ --hybrid-layer-pattern M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- \ - --spec megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec \ + --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ --tiktoken-pattern v2 \ --distributed-timeout-minutes 60 \ --use-mcore-models \ diff --git a/examples/rl/model_configs/nemotron5p5_12b_H.sh b/examples/rl/model_configs/nemotron5p5_12b_H.sh index 1826d57e913..bfb4c7e4727 100644 --- a/examples/rl/model_configs/nemotron5p5_12b_H.sh +++ b/examples/rl/model_configs/nemotron5p5_12b_H.sh @@ -76,7 +76,7 @@ MODEL_OPTIONS="\ --disable-gloo-process-groups \ --mamba-head-dim 80 \ --hybrid-layer-pattern M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M- \ - --spec megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec \ + --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ --tiktoken-pattern v2 \ --distributed-timeout-minutes 10 \ --use-mcore-models \ diff --git a/hybrid_builders.py b/hybrid_builders.py new file mode 100644 index 00000000000..36a87a3940b --- /dev/null +++ b/hybrid_builders.py @@ -0,0 +1,54 @@ +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. + +from model_provider import count_parameters_in_layer +from megatron.core.models.hybrid.hybrid_model import HybridModel +from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.spec_utils import import_module +from megatron.training import print_rank_0 +from megatron.training.arguments import core_transformer_config_from_args +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_inference_stack_spec + + +def hybrid_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None): + print_rank_0('building Hybrid model ...') + if config is None: + config = core_transformer_config_from_args(args, TransformerConfig) + assert args.use_legacy_models is False, "Hybrid model only supported in Mcore!" + + if config.transformer_impl == "inference_optimized": + hybrid_stack_spec = hybrid_inference_stack_spec + assert ( + not config.inference_fuse_tp_communication + ), "inference_fuse_tp_communication is not supported for HybridModel" + elif args.spec is not None: + hybrid_stack_spec = import_module(args.spec) + else: + raise ValueError("You must provide a valid hybrid layer spec via --spec") + + model = HybridModel( + config=config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=args.padded_vocab_size, + max_sequence_length=args.max_position_embeddings, + hybrid_layer_pattern=args.hybrid_layer_pattern, + 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, + pg_collection=pg_collection, + vp_stage=vp_stage, + ) + + for l in range(model.decoder.num_layers_per_pipeline_rank): + layer_params = count_parameters_in_layer(model, f'decoder.layers.{l}.') + print_rank_0(f" == params layer {l}: {layer_params}") + + return model + + +# Backward-compatible alias +mamba_builder = hybrid_builder diff --git a/mamba_builders.py b/mamba_builders.py index 650ea4a719f..f824fce9be3 100644 --- a/mamba_builders.py +++ b/mamba_builders.py @@ -1,50 +1,15 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. +"""Backward-compatible re-export of hybrid_builders. -from model_provider import count_parameters_in_layer -from megatron.core.models.mamba import MambaModel -from megatron.core.transformer import TransformerConfig -from megatron.core.transformer.spec_utils import import_module -from megatron.training import print_rank_0 -from megatron.training.arguments import core_transformer_config_from_args -from megatron.core.models.mamba.mamba_layer_specs import mamba_inference_stack_spec +Deprecated. Use hybrid_builders instead. +""" +import warnings +warnings.warn( + "mamba_builders has been deprecated. Use hybrid_builders instead.", + DeprecationWarning, + stacklevel=2, +) -def mamba_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None): - print_rank_0('building MAMBA model ...') - if config is None: - config = core_transformer_config_from_args(args, TransformerConfig) - assert args.use_legacy_models is False, "Mamba only supported in Mcore!" - - if config.transformer_impl == "inference_optimized": - mamba_stack_spec = mamba_inference_stack_spec - assert ( - not config.inference_fuse_tp_communication - ), "inference_fuse_tp_communication is not supported for Mamba" - elif args.spec is not None: - mamba_stack_spec = import_module(args.spec) - else: - raise ValueError("You must provide a valid Mamba layer spec via --spec") - - model = MambaModel( - config=config, - mamba_stack_spec=mamba_stack_spec, - vocab_size=args.padded_vocab_size, - max_sequence_length=args.max_position_embeddings, - hybrid_layer_pattern=args.hybrid_layer_pattern, - 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, - pg_collection=pg_collection, - vp_stage=vp_stage, - ) - - for l in range(model.decoder.num_layers_per_pipeline_rank): - layer_params = count_parameters_in_layer(model, f'decoder.layers.{l}.') - print_rank_0(f" == params layer {l}: {layer_params}") - - return model +from hybrid_builders import * # noqa: F401,F403 +from hybrid_builders import hybrid_builder as mamba_builder # noqa: F401 diff --git a/megatron/inference/utils.py b/megatron/inference/utils.py index a1204db487a..bb65c754ab1 100644 --- a/megatron/inference/utils.py +++ b/megatron/inference/utils.py @@ -7,7 +7,7 @@ import torch from gpt_builders import gpt_builder -from mamba_builders import mamba_builder +from hybrid_builders import hybrid_builder from megatron.core.inference.config import ( InferenceConfig, KVCacheManagementMode, @@ -43,8 +43,16 @@ def get_model_for_inference() -> MegatronModule: if args.model_provider == "gpt": model_builder = gpt_builder - elif args.model_provider == "mamba": - model_builder = mamba_builder + elif args.model_provider in ("hybrid", "mamba"): + if args.model_provider == "mamba": + import warnings + + warnings.warn( + '--model-provider "mamba" is deprecated. Use --model-provider "hybrid" instead.', + DeprecationWarning, + stacklevel=2, + ) + model_builder = hybrid_builder else: raise ValueError(f"Invalid model provider {args.model_provider}") @@ -158,7 +166,11 @@ def add_inference_args(parser: ArgumentParser) -> ArgumentParser: "total number of requests. Set to -1 to add all requests together.", ) group.add_argument( - "--model-provider", choices=["mamba", "gpt"], default="gpt", help="Model provider" + "--model-provider", + choices=["hybrid", "mamba", "gpt"], + default="gpt", + help='Model provider. Use "hybrid" for HybridModel (formerly MambaModel). ' + '"mamba" is accepted for backward compatibility but deprecated.', ) group.add_argument( "--skip-prompt-log-probs", action='store_true', default=False, help='Skip prompt log probs.' diff --git a/megatron/post_training/arguments.py b/megatron/post_training/arguments.py index dc98c6d28e4..47c667b4d0a 100644 --- a/megatron/post_training/arguments.py +++ b/megatron/post_training/arguments.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. def add_modelopt_args(parser): @@ -10,8 +10,9 @@ def add_modelopt_args(parser): "--export-model-type", type=str, default="GPTModel", - choices=["GPTModel", "MambaModel"], - help="Model type to use in model_provider.", + choices=["GPTModel", "HybridModel", "MambaModel"], + help='Model type to use in model_provider. Use "HybridModel" for hybrid models ' + '(formerly MambaModel). "MambaModel" is accepted for backward compatibility but deprecated.', ) group.add_argument( "--export-legacy-megatron", diff --git a/megatron/post_training/model_builder.py b/megatron/post_training/model_builder.py index 085d188e811..383ae6ec8aa 100644 --- a/megatron/post_training/model_builder.py +++ b/megatron/post_training/model_builder.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. """ModelOpt GPT model provider.""" @@ -16,7 +16,7 @@ from megatron.core.models.gpt.heterogeneous.heterogeneous_layer_specs import ( get_gpt_heterogeneous_layer_spec, ) -from megatron.core.models.mamba import MambaModel as MCoreMambaModel +from megatron.core.models.hybrid.hybrid_model import HybridModel as MCoreHybridModel from megatron.core.post_training.modelopt.gpt.model_specs import get_gpt_modelopt_spec from megatron.core.post_training.modelopt.gpt.state_dict_hooks import ( mcore_gpt_load_te_state_dict_pre_hook, @@ -124,7 +124,7 @@ def _load_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, A # _load_teacher_model_config, so config_raw.hybrid_layer_pattern is always set here. model_kwargs["hybrid_layer_pattern"] = config_raw.hybrid_layer_pattern - teacher = MCoreMambaModel(config=config, **model_kwargs) + teacher = MCoreHybridModel(config=config, **model_kwargs) else: # GPT layer spec needs re-creation since it depends on number of model layers. if config.heterogeneous_block_specs: @@ -158,14 +158,14 @@ def _load_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, A return teacher -def modelopt_gpt_mamba_builder( +def modelopt_gpt_hybrid_builder( args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None, -) -> MCoreGPTModel | MCoreMambaModel: +) -> MCoreGPTModel | MCoreHybridModel: """Builds the model. Args: @@ -179,7 +179,7 @@ def modelopt_gpt_mamba_builder( attached to the returned model for downstream routing/resharding utilities. Returns: - MCoreGPTModel | MCoreMambaModel: The returned model + MCoreGPTModel | MCoreHybridModel: The returned model """ print_rank_0("building GPT model ...") @@ -259,8 +259,17 @@ def modelopt_gpt_mamba_builder( "pg_collection": pg_collection, } model = MCoreGPTModel(config=config, **model_kwargs) - elif args.export_model_type == "MambaModel" or getattr(args, 'hybrid_layer_pattern', None) is not None: - from megatron.core.post_training.modelopt.mamba.model_specs import get_mamba_stack_modelopt_spec + elif args.export_model_type in ("HybridModel", "MambaModel") or getattr(args, 'hybrid_layer_pattern', None) is not None: + if args.export_model_type == "MambaModel": + import warnings + + warnings.warn( + '--export-model-type "MambaModel" is deprecated. ' + 'Use --export-model-type "HybridModel" instead.', + DeprecationWarning, + stacklevel=2, + ) + from megatron.core.post_training.modelopt.hybrid.model_specs import get_hybrid_stack_modelopt_spec if args.export_default_te_spec and args.export_te_mcore_model: logging.getLogger(__name__).warning( @@ -269,12 +278,12 @@ def modelopt_gpt_mamba_builder( ) args.export_te_mcore_model = False - mamba_stack_spec = get_mamba_stack_modelopt_spec( + hybrid_stack_spec = get_hybrid_stack_modelopt_spec( remap_te_layernorm=args.export_te_mcore_model, use_default_te_spec=args.export_default_te_spec, ) model_kwargs = { - "mamba_stack_spec": mamba_stack_spec, + "hybrid_stack_spec": hybrid_stack_spec, "vocab_size": args.padded_vocab_size, "max_sequence_length": args.max_position_embeddings, "hybrid_layer_pattern": args.hybrid_layer_pattern, @@ -289,7 +298,7 @@ def modelopt_gpt_mamba_builder( "pg_collection": pg_collection, } - model = MCoreMambaModel(config=config, **model_kwargs) + model = MCoreHybridModel(config=config, **model_kwargs) for l in range(model.decoder.num_layers_per_pipeline_rank): layer_params = count_parameters_in_layer(model, f'decoder.layers.{l}.') @@ -352,3 +361,7 @@ def modelopt_gpt_mamba_builder( print_distributed_quant_summary(model) return model + + +# Backward-compatible alias +modelopt_gpt_mamba_builder = modelopt_gpt_hybrid_builder diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 684a8ace6fb..eaa5b15ee53 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -657,7 +657,7 @@ def validate_args(args, defaults={}): args.rank, ) - from megatron.core.ssm.mamba_hybrid_layer_allocation import ( + from megatron.core.models.hybrid.hybrid_layer_allocation import ( Symbols, parse_hybrid_pattern, get_hybrid_total_layer_count, get_hybrid_total_pipeline_segment_count, ) @@ -1770,7 +1770,7 @@ def core_transformer_config_from_args(args, config_class=None): kw_args['cp_comm_type'] = args.cp_comm_type[0] if args.hybrid_layer_pattern is not None: kw_args['is_hybrid_model'] = True - from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols + from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols if Symbols.DS_ATTENTION in args.hybrid_layer_pattern: kw_args['experimental_attention_variant'] = 'dsa' diff --git a/megatron/training/training.py b/megatron/training/training.py index 8fed9afd570..8172c6a0cd5 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -677,7 +677,7 @@ def transformer_flops(): # Calculate the number of each type of layer. from operator import itemgetter - from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols, get_hybrid_layer_counts + from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols, get_hybrid_layer_counts num_mamba_layers, num_gdn_layers, num_attn_layers, num_mlp_layers, num_moe_layers = ( itemgetter(Symbols.MAMBA, Symbols.GDN, Symbols.ATTENTION, Symbols.MLP, Symbols.MOE)( get_hybrid_layer_counts(args.hybrid_layer_pattern) @@ -2169,7 +2169,7 @@ def training_log( if is_hybrid_model(args): from operator import itemgetter - from megatron.core.ssm.mamba_hybrid_layer_allocation import ( + from megatron.core.models.hybrid.hybrid_layer_allocation import ( Symbols, get_hybrid_layer_counts, ) layers = itemgetter(Symbols.MOE)(get_hybrid_layer_counts(args.hybrid_layer_pattern)) diff --git a/model_provider.py b/model_provider.py index 0c80c54dfdb..3e61343bb0d 100644 --- a/model_provider.py +++ b/model_provider.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. """Common functions used in train_*.py and pretrain_*.py scripts.""" @@ -7,11 +7,11 @@ import torch from megatron.core.models.gpt import GPTModel -from megatron.core.models.mamba import MambaModel +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.training import get_args, print_rank_0 try: - from megatron.post_training.model_builder import modelopt_gpt_mamba_builder + from megatron.post_training.model_builder import modelopt_gpt_hybrid_builder has_nvidia_modelopt = True except ImportError: has_nvidia_modelopt = False @@ -23,18 +23,18 @@ def model_provider( model_builder: Callable, pre_process=True, post_process=True, vp_stage: Optional[int] = None, config=None, pg_collection=None, -) -> Union[GPTModel, megatron.legacy.model.GPTModel, MambaModel]: +) -> Union[GPTModel, megatron.legacy.model.GPTModel, HybridModel]: """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: - model_builder: A callable that builds the actual model, its signature is the same as model_provider's with an exception of the first argument which is a builder itself. In addition might take a config passed from outside to skip its own config loading. See gpt_builder or mamba_builder for an example, see _gpt_model_builder in train_rl.py to see how to augment a default gpt builder and pass the config from outside + model_builder: A callable that builds the actual model, its signature is the same as model_provider's with an exception of the first argument which is a builder itself. In addition might take a config passed from outside to skip its own config loading. See gpt_builder or hybrid_builder for an example, see _gpt_model_builder in train_rl.py to see how to augment a default gpt builder and pass the config from outside 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 compute output logits/loss. Defaults to True. Returns: - Union[GPTModel, megatron.legacy.model.GPTModel, MambaModel]: The returned model + Union[GPTModel, megatron.legacy.model.GPTModel, HybridModel]: The returned model """ args = get_args() @@ -58,7 +58,7 @@ def oom_observer(device, alloc, device_alloc, device_free): if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt]: Use custom builder + spec when modelopt is enabled - model_builder = modelopt_gpt_mamba_builder + model_builder = modelopt_gpt_hybrid_builder return model_builder(args, pre_process, post_process, vp_stage, config=config, pg_collection=pg_collection) diff --git a/pretrain_gpt.py b/pretrain_gpt.py index 929a9d0f866..02ca5dd72c0 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -90,7 +90,7 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): - MTP ranks (``mtp_on_this_rank``) also receive the full batch, regardless of pipeline stage. - Difference from ``pretrain_mamba.py``: + Difference from ``pretrain_hybrid.py``: - Return format: GPT returns a 6-tuple ``(tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params)`` where ``packed_seq_params`` is a diff --git a/pretrain_hybrid.py b/pretrain_hybrid.py new file mode 100644 index 00000000000..f073e8e9ab3 --- /dev/null +++ b/pretrain_hybrid.py @@ -0,0 +1,369 @@ +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. +"""Pretrain and SFT Hybrid.""" + +# Capture the true program start time BEFORE any heavy imports. +import time +_PROGRAM_START_TIME = time.time() + +import json + +# Suppress warnings on all ranks but rank 0. +import os +import warnings +rank = int(os.environ.get('RANK', 0)) +if rank != 0: + warnings.filterwarnings("ignore", category=UserWarning) + warnings.filterwarnings("ignore", category=FutureWarning) + +from functools import partial +from typing import List, Optional, Tuple + +import torch + +from hybrid_builders import hybrid_builder +from megatron.core import mpu +from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder +from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset +from megatron.core.enums import ModelType +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.parallel_state import ( + get_context_parallel_rank, + get_context_parallel_world_size, +) +from megatron.core.models.hybrid.hybrid_model import HybridModel +from megatron.core.rerun_state_machine import get_rerun_state_machine +from megatron.core.tokenizers.utils.build_tokenizer import build_tokenizer +from megatron.core.utils import get_attr_wrapped_model, is_te_min_version, StragglerDetector +from megatron.training import ( + get_args, + get_timers, + inprocess_restart, + pretrain, + print_rank_0, + set_startup_timestamps, +) +from megatron.training.arguments import parse_and_validate_args +from megatron.training.datasets.sft_dataset import SFTDataset +from megatron.training.utils import ( + get_batch_on_this_cp_rank, + get_batch_on_this_tp_rank, + get_blend_and_blend_per_split, + is_first_or_last_pipeline_stage, +) +from model_provider import model_provider + +try: + from megatron.post_training.arguments import add_modelopt_args + from megatron.post_training.loss_func import loss_func as loss_func_modelopt + has_nvidia_modelopt = True +except ImportError: + has_nvidia_modelopt = False + +try: + # Register the TE CUDA kernels + import transformer_engine # pylint: disable=unused-import + + # Alias the PyTorch wrapper so we can call tex.* APIs + import transformer_engine_torch as tex +except ImportError: + # TE isn’t installed or the torch wrapper is missing + tex = None + +stimer = StragglerDetector() + + +def get_batch(data_iterator, vp_stage=None): + """Generate a batch.""" + + empty_batch = { + 'tokens': None, + 'labels': None, + 'loss_mask': None, + 'attention_mask': None, + 'position_ids': None, + 'cu_seqlens': None, + 'max_seqlen': None, + } + + # TODO(duncan): Is there a more efficient way to access is_packed_sequence here? + is_packed_sequence = get_args().sft # SFT always uses packed sequence + if not is_first_or_last_pipeline_stage(vp_stage) and not is_packed_sequence: + return empty_batch.values() + + batch = get_batch_on_this_tp_rank(data_iterator) + + cu_seqlens = batch['cu_seqlens'] + # Unused at the moment + cu_seqlens_padded = batch.pop('cu_seqlens_padded', None) + # Support for Hybrid Context Parallel (Unused in this script) + local_cp_size = batch.pop('local_cp_size', None) + + if cu_seqlens is not None: + assert ( + cu_seqlens.dim() == 2 and cu_seqlens.shape[0] == 1 + ), "micro-batch-size must be 1 for packing" + cu_seqlens = cu_seqlens[0] + batch['cu_seqlens'] = cu_seqlens + + max_seqlen = batch['max_seqlen'] + assert max_seqlen.dim() == 1 + # TODO(duncan): can this be kept as a 0-D tensor? + batch['max_seqlen'] = int(max_seqlen[0].item()) + + if mpu.is_pipeline_first_stage(ignore_virtual=(vp_stage is None), vp_stage=vp_stage): + total_tokens = batch['tokens'].size(1) + elif mpu.is_pipeline_last_stage(ignore_virtual=(vp_stage is None), vp_stage=vp_stage): + total_tokens = batch['labels'].size(1) + else: # packed sequence + empty_batch['cu_seqlens'] = cu_seqlens + empty_batch['max_seqlen'] = max_seqlen + return empty_batch.values() + + if cu_seqlens is None: + # slice batch along sequence dimension for context parallelism + batch = get_batch_on_this_cp_rank(batch) # The implementation of this function is in MCore + else: # Packed THD format + cp_size = get_context_parallel_world_size() + if cp_size > 1: # slice batch along sequence dimension for context parallelism + assert tex is not None and is_te_min_version("1.10.0"), ( + "Please update Transformer Engine to >= 1.10 to use " + "Context Parallel with THD format data" + ) + cp_rank = get_context_parallel_rank() + index = tex.thd_get_partitioned_indices( + cu_seqlens, + total_tokens, + cp_size, + cp_rank, + ) + for key, data in batch.items(): + if key in {'attention_mask', 'cu_seqlens', 'max_seqlen'}: + continue + if data is not None: + # On first PP rank, labels and loss_mask can be None. + # On last PP rank, tokens and position_ids can be None. + batch[key] = data.index_select(1, index) + + return batch.values() + + +# define spiky loss as a loss that's 10x the max loss observed +SPIKY_LOSS_FACTOR = 10 + +def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: Optional[HybridModel] = None): + """Loss function. + + Args: + loss_mask (torch.Tensor): Used to mask out some portions of the loss + output_tensor (torch.Tensor): The tensor with the losses + + Returns: + the loss scalar for this micro-batch + the number of non-padded tokens in this microbatch + a dict containing reporting metrics on the loss and number of tokens across + the data parallel ranks + """ + args = get_args() + if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt] + loss, num_tokens, report = loss_func_modelopt(loss_mask, output_tensor, model=model) + else: + losses = output_tensor.view(-1).float() + loss_mask = loss_mask.view(-1).float() + loss = torch.sum(losses * loss_mask) + + num_tokens = loss_mask.sum().clone().detach().to(torch.int) + report = {'lm loss': torch.cat([loss.clone().detach().view(1), num_tokens.view(1)])} + + # Check individual rank losses are not NaN prior to DP all-reduce. + rerun_state_machine = get_rerun_state_machine() + if args.check_for_nan_in_loss_and_grad: + rerun_state_machine.validate_result( + result=loss, + rejection_func=torch.isnan, + message="found NaN in local forward loss calculation", + tolerance=0.0, # forward pass calculations are deterministic + fatal=True, + ) + rerun_state_machine.validate_result( + result=loss, + rejection_func=torch.isinf, + message="found Inf in local forward loss calculation", + tolerance=0.0, # forward pass calculations are deterministic + fatal=True, + ) + # Check for spiky loss + if args.check_for_spiky_loss: + rerun_state_machine.validate_result( + result=loss, + rejection_func=partial( + rerun_state_machine.is_unexpectedly_large, + threshold=SPIKY_LOSS_FACTOR, + context="loss", + ), + message="Spiky loss", + tolerance=0.0, # forward pass calculations are deterministic + fatal=False, + ) + + return loss, num_tokens, report + + +def forward_step(data_iterator, model: HybridModel): + """Forward training step. + + Args: + data_iterator : Input data iterator + model (HybridModel): The Model + """ + timers = get_timers() + + # Get the batch. + timers('batch-generator', log_level=2).start() + + global stimer + + with stimer(bdata=True): + vp_stage = get_attr_wrapped_model(model, "vp_stage") + ( + tokens, + labels, + loss_mask, + attention_mask, + position_ids, + cu_seqlens, + max_seqlen, + ) = get_batch(data_iterator, vp_stage) + + if cu_seqlens is None: + packed_seq_params = None + else: + total_tokens = tokens.size(1) if tokens is not None else labels.size(1) + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=None, + cu_seqlens_kv_padded=None, + max_seqlen_q=max_seqlen, + max_seqlen_kv=max_seqlen, + total_tokens=total_tokens, + ) + + timers('batch-generator').stop() + + with stimer: + output_tensor = model( + tokens, + position_ids, + attention_mask, + labels=labels, + packed_seq_params=packed_seq_params, + loss_mask=loss_mask + ) + + # [ModelOpt]: model is needed to access ModelOpt distillation losses + return output_tensor, partial(loss_func, loss_mask, model=model) + + +def is_dataset_built_on_rank(vp_stage=None, is_packed_sequence=False): + if mpu.get_tensor_model_parallel_rank() != 0: + return False + elif is_packed_sequence: + return True + else: + return is_first_or_last_pipeline_stage(vp_stage) + + +def core_gpt_dataset_config_from_args(args): + tokenizer = build_tokenizer(args) + + # Sometimes --data-path is too long, instead we parse it from a file. + blend: Optional[Tuple[List[str], Optional[List[float]]]] + blend_per_split: Optional[List[Optional[Tuple[List[str], Optional[List[float]]]]]] + blend, blend_per_split = get_blend_and_blend_per_split(args) + + sequences_per_dataset = None + if args.per_dataset_sequences_path is not None: + with open(args.per_dataset_sequences_path, "r") as f: + sequences_per_dataset = json.load(f) + + return GPTDatasetConfig( + random_seed=args.seed, + sequence_length=args.seq_length, + blend=blend, + blend_per_split=blend_per_split, + split=args.split, + num_dataset_builder_threads=args.num_dataset_builder_threads, + path_to_cache=args.data_cache_path, + mmap_bin_files=args.mmap_bin_files, + tokenizer=tokenizer, + reset_position_ids=args.reset_position_ids, + reset_attention_mask=args.reset_attention_mask, + eod_mask_loss=args.eod_mask_loss, + create_attention_mask=args.create_attention_mask_in_dataloader, + object_storage_cache_path=args.object_storage_cache_path, + mid_level_dataset_surplus=args.mid_level_dataset_surplus, + allow_ambiguous_pad_tokens=args.allow_ambiguous_pad_tokens, + fast_cache_load=args.dataloader_fast_cache_load, + sequences_per_dataset=sequences_per_dataset, + defer_npy_index_mmap=args.dataloader_defer_npy_index_mmap, + context_parallel_size=args.context_parallel_size, + ) + + +def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None): + """Build the train test and validation datasets. + + Args: + train_val_test_num_samples : A list containing the number of samples in train test and validation. + """ + args = get_args() + config = core_gpt_dataset_config_from_args(args) + + is_packed_sequence = False + if args.sft: + dataset_type = SFTDataset + is_packed_sequence = True # SFT always uses packed sequence + else: + if args.mock_data: + dataset_type = MockGPTDataset + else: + dataset_type = GPTDataset + + print_rank_0("> building train, validation, and test datasets for GPT ...") + + train_ds, valid_ds, test_ds = BlendedMegatronDatasetBuilder( + dataset_type, + train_val_test_num_samples, + partial(is_dataset_built_on_rank, vp_stage=vp_stage, is_packed_sequence=is_packed_sequence), + config + ).build() + + print_rank_0("> finished creating GPT datasets ...") + + return train_ds, valid_ds, test_ds + + +if __name__ == "__main__": + # Timestamp right after entering __main__ block (after all imports/library setup) + _MAIN_ENTRY_TIME = time.time() + + # Register startup timestamps for timing report in pretrain() + set_startup_timestamps(program_start=_PROGRAM_START_TIME, main_entry=_MAIN_ENTRY_TIME) + + # Temporary for transition to core datasets + train_valid_test_datasets_provider.is_distributed = True + + # Optionally enable inprocess restart on pretrain + pretrain, store = inprocess_restart.maybe_wrap_for_inprocess_restart(pretrain) + + args = parse_and_validate_args( + extra_args_provider=add_modelopt_args if has_nvidia_modelopt else None, + args_defaults={'tokenizer_type': 'GPT2BPETokenizer'}, + ) + pretrain(train_valid_test_datasets_provider, + partial(model_provider, hybrid_builder), + ModelType.encoder_or_decoder, + forward_step, + store=store, + ) diff --git a/pretrain_mamba.py b/pretrain_mamba.py index 590eb92ab28..7eb7f461cab 100644 --- a/pretrain_mamba.py +++ b/pretrain_mamba.py @@ -1,369 +1,19 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. -"""Pretrain and SFT Mamba.""" +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. +"""Backward-compatible wrapper for pretrain_hybrid.py. -# Capture the true program start time BEFORE any heavy imports. -import time -_PROGRAM_START_TIME = time.time() - -import json - -# Suppress warnings on all ranks but rank 0. +Deprecated. Use pretrain_hybrid.py instead. +""" import os +import runpy import warnings -rank = int(os.environ.get('RANK', 0)) -if rank != 0: - warnings.filterwarnings("ignore", category=UserWarning) - warnings.filterwarnings("ignore", category=FutureWarning) - -from functools import partial -from typing import List, Optional, Tuple - -import torch -from mamba_builders import mamba_builder -from megatron.core import mpu -from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder -from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset -from megatron.core.enums import ModelType -from megatron.core.packed_seq_params import PackedSeqParams -from megatron.core.parallel_state import ( - get_context_parallel_rank, - get_context_parallel_world_size, +warnings.warn( + "pretrain_mamba.py has been deprecated. Use pretrain_hybrid.py instead.", + DeprecationWarning, + stacklevel=2, ) -from megatron.core.models.mamba import MambaModel -from megatron.core.rerun_state_machine import get_rerun_state_machine -from megatron.core.tokenizers.utils.build_tokenizer import build_tokenizer -from megatron.core.utils import get_attr_wrapped_model, is_te_min_version, StragglerDetector -from megatron.training import ( - get_args, - get_timers, - inprocess_restart, - pretrain, - print_rank_0, - set_startup_timestamps, -) -from megatron.training.arguments import parse_and_validate_args -from megatron.training.datasets.sft_dataset import SFTDataset -from megatron.training.utils import ( - get_batch_on_this_cp_rank, - get_batch_on_this_tp_rank, - get_blend_and_blend_per_split, - is_first_or_last_pipeline_stage, -) -from model_provider import model_provider - -try: - from megatron.post_training.arguments import add_modelopt_args - from megatron.post_training.loss_func import loss_func as loss_func_modelopt - has_nvidia_modelopt = True -except ImportError: - has_nvidia_modelopt = False - -try: - # Register the TE CUDA kernels - import transformer_engine # pylint: disable=unused-import - - # Alias the PyTorch wrapper so we can call tex.* APIs - import transformer_engine_torch as tex -except ImportError: - # TE isn’t installed or the torch wrapper is missing - tex = None - -stimer = StragglerDetector() - - -def get_batch(data_iterator, vp_stage=None): - """Generate a batch.""" - - empty_batch = { - 'tokens': None, - 'labels': None, - 'loss_mask': None, - 'attention_mask': None, - 'position_ids': None, - 'cu_seqlens': None, - 'max_seqlen': None, - } - - # TODO(duncan): Is there a more efficient way to access is_packed_sequence here? - is_packed_sequence = get_args().sft # SFT always uses packed sequence - if not is_first_or_last_pipeline_stage(vp_stage) and not is_packed_sequence: - return empty_batch.values() - - batch = get_batch_on_this_tp_rank(data_iterator) - - cu_seqlens = batch['cu_seqlens'] - # Unused at the moment - cu_seqlens_padded = batch.pop('cu_seqlens_padded', None) - # Support for Hybrid Context Parallel (Unused in this script) - local_cp_size = batch.pop('local_cp_size', None) - - if cu_seqlens is not None: - assert ( - cu_seqlens.dim() == 2 and cu_seqlens.shape[0] == 1 - ), "micro-batch-size must be 1 for packing" - cu_seqlens = cu_seqlens[0] - batch['cu_seqlens'] = cu_seqlens - - max_seqlen = batch['max_seqlen'] - assert max_seqlen.dim() == 1 - # TODO(duncan): can this be kept as a 0-D tensor? - batch['max_seqlen'] = int(max_seqlen[0].item()) - - if mpu.is_pipeline_first_stage(ignore_virtual=(vp_stage is None), vp_stage=vp_stage): - total_tokens = batch['tokens'].size(1) - elif mpu.is_pipeline_last_stage(ignore_virtual=(vp_stage is None), vp_stage=vp_stage): - total_tokens = batch['labels'].size(1) - else: # packed sequence - empty_batch['cu_seqlens'] = cu_seqlens - empty_batch['max_seqlen'] = max_seqlen - return empty_batch.values() - - if cu_seqlens is None: - # slice batch along sequence dimension for context parallelism - batch = get_batch_on_this_cp_rank(batch) # The implementation of this function is in MCore - else: # Packed THD format - cp_size = get_context_parallel_world_size() - if cp_size > 1: # slice batch along sequence dimension for context parallelism - assert tex is not None and is_te_min_version("1.10.0"), ( - "Please update Transformer Engine to >= 1.10 to use " - "Context Parallel with THD format data" - ) - cp_rank = get_context_parallel_rank() - index = tex.thd_get_partitioned_indices( - cu_seqlens, - total_tokens, - cp_size, - cp_rank, - ) - for key, data in batch.items(): - if key in {'attention_mask', 'cu_seqlens', 'max_seqlen'}: - continue - if data is not None: - # On first PP rank, labels and loss_mask can be None. - # On last PP rank, tokens and position_ids can be None. - batch[key] = data.index_select(1, index) - - return batch.values() - - -# define spiky loss as a loss that's 10x the max loss observed -SPIKY_LOSS_FACTOR = 10 - -def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: Optional[MambaModel] = None): - """Loss function. - - Args: - loss_mask (torch.Tensor): Used to mask out some portions of the loss - output_tensor (torch.Tensor): The tensor with the losses - - Returns: - the loss scalar for this micro-batch - the number of non-padded tokens in this microbatch - a dict containing reporting metrics on the loss and number of tokens across - the data parallel ranks - """ - args = get_args() - if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt] - loss, num_tokens, report = loss_func_modelopt(loss_mask, output_tensor, model=model) - else: - losses = output_tensor.view(-1).float() - loss_mask = loss_mask.view(-1).float() - loss = torch.sum(losses * loss_mask) - - num_tokens = loss_mask.sum().clone().detach().to(torch.int) - report = {'lm loss': torch.cat([loss.clone().detach().view(1), num_tokens.view(1)])} - - # Check individual rank losses are not NaN prior to DP all-reduce. - rerun_state_machine = get_rerun_state_machine() - if args.check_for_nan_in_loss_and_grad: - rerun_state_machine.validate_result( - result=loss, - rejection_func=torch.isnan, - message="found NaN in local forward loss calculation", - tolerance=0.0, # forward pass calculations are deterministic - fatal=True, - ) - rerun_state_machine.validate_result( - result=loss, - rejection_func=torch.isinf, - message="found Inf in local forward loss calculation", - tolerance=0.0, # forward pass calculations are deterministic - fatal=True, - ) - # Check for spiky loss - if args.check_for_spiky_loss: - rerun_state_machine.validate_result( - result=loss, - rejection_func=partial( - rerun_state_machine.is_unexpectedly_large, - threshold=SPIKY_LOSS_FACTOR, - context="loss", - ), - message="Spiky loss", - tolerance=0.0, # forward pass calculations are deterministic - fatal=False, - ) - - return loss, num_tokens, report - - -def forward_step(data_iterator, model: MambaModel): - """Forward training step. - - Args: - data_iterator : Input data iterator - model (MambaModel): The GPT Model - """ - timers = get_timers() - - # Get the batch. - timers('batch-generator', log_level=2).start() - - global stimer - - with stimer(bdata=True): - vp_stage = get_attr_wrapped_model(model, "vp_stage") - ( - tokens, - labels, - loss_mask, - attention_mask, - position_ids, - cu_seqlens, - max_seqlen, - ) = get_batch(data_iterator, vp_stage) - - if cu_seqlens is None: - packed_seq_params = None - else: - total_tokens = tokens.size(1) if tokens is not None else labels.size(1) - packed_seq_params = PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu_seqlens, - cu_seqlens_kv=cu_seqlens, - cu_seqlens_q_padded=None, - cu_seqlens_kv_padded=None, - max_seqlen_q=max_seqlen, - max_seqlen_kv=max_seqlen, - total_tokens=total_tokens, - ) - - timers('batch-generator').stop() - - with stimer: - output_tensor = model( - tokens, - position_ids, - attention_mask, - labels=labels, - packed_seq_params=packed_seq_params, - loss_mask=loss_mask - ) - - # [ModelOpt]: model is needed to access ModelOpt distillation losses - return output_tensor, partial(loss_func, loss_mask, model=model) - - -def is_dataset_built_on_rank(vp_stage=None, is_packed_sequence=False): - if mpu.get_tensor_model_parallel_rank() != 0: - return False - elif is_packed_sequence: - return True - else: - return is_first_or_last_pipeline_stage(vp_stage) - - -def core_gpt_dataset_config_from_args(args): - tokenizer = build_tokenizer(args) - - # Sometimes --data-path is too long, instead we parse it from a file. - blend: Optional[Tuple[List[str], Optional[List[float]]]] - blend_per_split: Optional[List[Optional[Tuple[List[str], Optional[List[float]]]]]] - blend, blend_per_split = get_blend_and_blend_per_split(args) - - sequences_per_dataset = None - if args.per_dataset_sequences_path is not None: - with open(args.per_dataset_sequences_path, "r") as f: - sequences_per_dataset = json.load(f) - - return GPTDatasetConfig( - random_seed=args.seed, - sequence_length=args.seq_length, - blend=blend, - blend_per_split=blend_per_split, - split=args.split, - num_dataset_builder_threads=args.num_dataset_builder_threads, - path_to_cache=args.data_cache_path, - mmap_bin_files=args.mmap_bin_files, - tokenizer=tokenizer, - reset_position_ids=args.reset_position_ids, - reset_attention_mask=args.reset_attention_mask, - eod_mask_loss=args.eod_mask_loss, - create_attention_mask=args.create_attention_mask_in_dataloader, - object_storage_cache_path=args.object_storage_cache_path, - mid_level_dataset_surplus=args.mid_level_dataset_surplus, - allow_ambiguous_pad_tokens=args.allow_ambiguous_pad_tokens, - fast_cache_load=args.dataloader_fast_cache_load, - sequences_per_dataset=sequences_per_dataset, - defer_npy_index_mmap=args.dataloader_defer_npy_index_mmap, - context_parallel_size=args.context_parallel_size, - ) - - -def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None): - """Build the train test and validation datasets. - - Args: - train_val_test_num_samples : A list containing the number of samples in train test and validation. - """ - args = get_args() - config = core_gpt_dataset_config_from_args(args) - - is_packed_sequence = False - if args.sft: - dataset_type = SFTDataset - is_packed_sequence = True # SFT always uses packed sequence - else: - if args.mock_data: - dataset_type = MockGPTDataset - else: - dataset_type = GPTDataset - - print_rank_0("> building train, validation, and test datasets for GPT ...") - - train_ds, valid_ds, test_ds = BlendedMegatronDatasetBuilder( - dataset_type, - train_val_test_num_samples, - partial(is_dataset_built_on_rank, vp_stage=vp_stage, is_packed_sequence=is_packed_sequence), - config - ).build() - - print_rank_0("> finished creating GPT datasets ...") - - return train_ds, valid_ds, test_ds - if __name__ == "__main__": - # Timestamp right after entering __main__ block (after all imports/library setup) - _MAIN_ENTRY_TIME = time.time() - - # Register startup timestamps for timing report in pretrain() - set_startup_timestamps(program_start=_PROGRAM_START_TIME, main_entry=_MAIN_ENTRY_TIME) - - # Temporary for transition to core datasets - train_valid_test_datasets_provider.is_distributed = True - - # Optionally enable inprocess restart on pretrain - pretrain, store = inprocess_restart.maybe_wrap_for_inprocess_restart(pretrain) - - args = parse_and_validate_args( - extra_args_provider=add_modelopt_args if has_nvidia_modelopt else None, - args_defaults={'tokenizer_type': 'GPT2BPETokenizer'}, - ) - pretrain(train_valid_test_datasets_provider, - partial(model_provider, mamba_builder), - ModelType.encoder_or_decoder, - forward_step, - store=store, - ) + # Execute pretrain_hybrid.py as if it were invoked directly. + _this_dir = os.path.dirname(os.path.abspath(__file__)) + runpy.run_path(os.path.join(_this_dir, "pretrain_hybrid.py"), run_name="__main__") diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml index f5de6eaac72..4b258afe0d6 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml @@ -24,7 +24,7 @@ MODEL_ARGS: --pipeline-model-parallel-size: 1 --expert-model-parallel-size: 1 --use-mcore-models: true - --model-provider: mamba + --model-provider: hybrid --init-method-std: 0.0198 --untie-embeddings-and-output-weights: true --disable-bias-linear: true @@ -35,7 +35,7 @@ MODEL_ARGS: --num-attention-heads: 16 --kv-channels: 128 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- - --spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --normalization: RMSNorm --swiglu: true --attention-dropout: 0.0 diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml index b10698d521f..bd86d2faa44 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml @@ -24,7 +24,7 @@ MODEL_ARGS: --pipeline-model-parallel-size: 1 --expert-model-parallel-size: 1 --use-mcore-models: true - --model-provider: mamba + --model-provider: hybrid --init-method-std: 0.0198 --untie-embeddings-and-output-weights: true --disable-bias-linear: true @@ -35,7 +35,7 @@ MODEL_ARGS: --num-attention-heads: 16 --kv-channels: 128 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- - --spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --normalization: RMSNorm --swiglu: true --attention-dropout: 0.0 diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml index 6d40098499d..9add53f8a49 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml @@ -9,7 +9,7 @@ MODEL_ARGS: --group-query-attention: true --num-query-groups: 8 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M- - --spec: "[megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec]" + --spec: "[megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec]" --log-params-norm: true --log-num-zeros-in-grad: true --log-validation-ppl-to-tensorboard: true diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml index 51492f98c6e..25df6aa0359 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml @@ -9,7 +9,7 @@ MODEL_ARGS: --group-query-attention: true --num-query-groups: 8 --hybrid-layer-pattern: M-M-M-M*-M-|M-M-M*-M-M-|M-M*-M-M-M-|M*-M-M-M-M- - --spec: "[megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec]" + --spec: "[megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec]" --log-params-norm: true --log-num-zeros-in-grad: true --log-validation-ppl-to-tensorboard: true diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml index 6eff846884a..fe4f9e63714 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml @@ -9,7 +9,7 @@ MODEL_ARGS: --group-query-attention: true --num-query-groups: 8 --hybrid-layer-pattern: M-M-M-M*-M-|M-M-M*-M-M-|M-M*-M-M-M-|M*-M-M-M-M- - --spec: "[megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec]" + --spec: "[megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec]" --log-params-norm: true --log-num-zeros-in-grad: true --log-validation-ppl-to-tensorboard: true diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml index 8c655bc135c..2339f7a7ce9 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml @@ -9,7 +9,7 @@ MODEL_ARGS: --group-query-attention: true --num-query-groups: 8 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M- - --spec: "[megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec]" + --spec: "[megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec]" --log-params-norm: true --log-num-zeros-in-grad: true --log-validation-ppl-to-tensorboard: true diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml index 44b588ee140..3efc155949f 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml @@ -9,7 +9,7 @@ MODEL_ARGS: --group-query-attention: true --num-query-groups: 8 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M*-M-M-M-M- - --spec: "[megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec]" + --spec: "[megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec]" --log-params-norm: true --log-num-zeros-in-grad: true --log-validation-ppl-to-tensorboard: true diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml index 26708b32a60..02c5cc3055c 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml @@ -22,7 +22,7 @@ MODEL_ARGS: --pipeline-model-parallel-size: 1 --expert-model-parallel-size: 1 --use-mcore-models: true - --model-provider: mamba + --model-provider: hybrid --init-method-std: 0.0198 --untie-embeddings-and-output-weights: true --disable-bias-linear: true @@ -33,7 +33,7 @@ MODEL_ARGS: --num-attention-heads: 16 --kv-channels: 128 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- - --spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --normalization: RMSNorm --swiglu: true --attention-dropout: 0.0 diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml index 3964bcb8ecb..2543f59e668 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml @@ -22,7 +22,7 @@ MODEL_ARGS: --pipeline-model-parallel-size: 1 --expert-model-parallel-size: 1 --use-mcore-models: true - --model-provider: mamba + --model-provider: hybrid --init-method-std: 0.0198 --untie-embeddings-and-output-weights: true --disable-bias-linear: true @@ -33,7 +33,7 @@ MODEL_ARGS: --num-attention-heads: 16 --kv-channels: 128 --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- - --spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --normalization: RMSNorm --swiglu: true --attention-dropout: 0.0 diff --git a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_g200/model_config.yaml b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_g200/model_config.yaml index 1147dda6118..9c5f1807c2d 100644 --- a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_g200/model_config.yaml +++ b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_g200/model_config.yaml @@ -42,7 +42,7 @@ MODEL_ARGS: # Network size args --use-mcore-models: true - --spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --is-hybrid-model: true --mamba-num-heads: 128 --num-layers: 88 @@ -90,7 +90,7 @@ MODEL_ARGS: --moe-shared-expert-compute-before-router: true # MTP args - --mtp-spec: megatron.core.models.mamba.mamba_layer_specs mamba_stack_spec + --mtp-spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec --mtp-num-layers: 2 --mtp-hybrid-override-pattern: \"*E\" --calculate-per-token-loss: true diff --git a/tests/test_utils/recipes/h100/mamba.yaml b/tests/test_utils/recipes/h100/mamba.yaml index 703fb53160f..72b44495617 100644 --- a/tests/test_utils/recipes/h100/mamba.yaml +++ b/tests/test_utils/recipes/h100/mamba.yaml @@ -44,7 +44,7 @@ spec: "TENSORBOARD_PATH={assets_dir}/tensorboard" "CHECKPOINT_SAVE_PATH={artifacts_dir}/checkpoints" "CHECKPOINT_LOAD_PATH=/mnt/artifacts/model/{name}" - "TRAINING_SCRIPT_PATH=pretrain_mamba.py" + "TRAINING_SCRIPT_PATH=pretrain_hybrid.py" "TRAINING_PARAMS_PATH=./tests/functional_tests/test_cases/{model}/{test_case}/model_config.yaml" "GOLDEN_VALUES_PATH=./tests/functional_tests/test_cases/{model}/{test_case}/golden_values_{environment}_{platforms}.json" "N_REPEAT={n_repeat}" diff --git a/tests/unit_tests/inference/contexts/test_dynamic_context.py b/tests/unit_tests/inference/contexts/test_dynamic_context.py index 06acdcfec9f..721e69212e3 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_context.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_context.py @@ -16,7 +16,7 @@ ) from megatron.core.inference.inference_request import DynamicInferenceRequest from megatron.core.inference.sampling_params import SamplingParams -from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols +from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.transformer_config import TransformerConfig from tests.unit_tests.test_utilities import Utils diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index fe2b8fc5802..b23e9562242 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import asyncio import gc @@ -46,8 +46,8 @@ get_gpt_mtp_block_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel -from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec -from megatron.core.models.mamba.mamba_model import MambaModel +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.ssm.mamba_mixer import _check_mamba_sequence_packing_support from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.cuda_graphs import CudaGraphManager, _CudagraphGlobalRecord @@ -65,7 +65,7 @@ def skip_if_mamba_sequence_packing_not_available(model_provider: str): - if model_provider == "mamba": + if model_provider == "hybrid": sequence_packing_available, reason_for_no_sequence_packing = ( _check_mamba_sequence_packing_support() ) @@ -368,7 +368,7 @@ def _build_test_env(cls, test_config): mtp_block_spec=mtp_block_spec, position_embedding_type=test_config.position_embedding_type, ).cuda() - elif test_config.model_provider == "mamba": + elif test_config.model_provider == "hybrid": pp_size = test_config.pipeline_model_parallel_size # Transformer config. transformer_config = TransformerConfig( @@ -407,7 +407,7 @@ def _build_test_env(cls, test_config): is_hybrid_model=True, # Needs to be set for correct out_proj init ) - # Mamba model. + # Hybrid model. # When speculative tokens are configured, append MTP depth sections # to the hybrid layer pattern so the model creates MTP blocks. mtp_suffix = "/M" * test_config.num_speculative_tokens @@ -415,9 +415,9 @@ def _build_test_env(cls, test_config): mamba_pattern = "M*-" + mtp_suffix else: mamba_pattern = "M*-|M*-" + mtp_suffix - model = MambaModel( + model = HybridModel( config=transformer_config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=test_config.vocab_size, max_sequence_length=test_config.max_sequence_length, parallel_output=True, @@ -574,7 +574,7 @@ def teardown_class(cls): @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) @pytest.mark.parametrize("num_cuda_graphs", [None, 1, 4, -1]) @pytest.mark.parametrize("cuda_graph_scope", [[], [CudaGraphScope.full_iteration_inference]]) def test_simple(self, model_provider, num_cuda_graphs, cuda_graph_scope) -> None: @@ -632,7 +632,7 @@ def test_simple(self, model_provider, num_cuda_graphs, cuda_graph_scope) -> None if model_provider == "gpt": expected_generated_tokens_list = gpt_expected_generated_tokens - elif model_provider == "mamba": + elif model_provider == "hybrid": expected_generated_tokens_list = mamba_expected_generated_tokens else: raise ValueError(f"Invalid model_provider {model_provider}") @@ -693,7 +693,7 @@ def test_token_overflow_nontransient(self) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) def test_block_overflow(self, model_provider: str) -> None: """Test block overflow.""" skip_if_mamba_sequence_packing_not_available(model_provider) @@ -739,7 +739,7 @@ def test_block_overflow_insufficient_kv_cache(self) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) def test_multi_add(self, model_provider: str) -> None: """Test adding multiple requests simultaneously.""" skip_if_mamba_sequence_packing_not_available(model_provider) @@ -749,7 +749,7 @@ def test_multi_add(self, model_provider: str) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) def test_fixed_output_lengths(self, model_provider: str) -> None: """Test generating a fixed number of output tokens.""" skip_if_mamba_sequence_packing_not_available(model_provider) @@ -792,7 +792,7 @@ def test_cuda_graph_token_counts(self) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) @torch.inference_mode() def test_generate_function(self, model_provider: str) -> None: """Test the generate function that processes multiple prompts at once.""" @@ -886,7 +886,7 @@ async def test_run_engine(self): not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) @pytest.mark.skipif(not is_te_min_version("2.2.0"), reason="TE 2.2.0 is required") - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) def test_fp8_inference(self, model_provider: str): skip_if_mamba_sequence_packing_not_available(model_provider) @@ -1092,7 +1092,7 @@ def test_log_probs_token_correspondence(self): @pytest.mark.parametrize("ep_size", [1, 2]) @pytest.mark.parametrize("pp_size", [1, 2]) @pytest.mark.parametrize("tp_size", [1, 2]) - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) @pytest.mark.parametrize("transformer_impl", ["local", "inference_optimized"]) @torch.inference_mode() def test_parallel_inference( @@ -1131,7 +1131,7 @@ def test_parallel_inference( "when tp_size > 1." ) ) - if model_provider == "mamba": + if model_provider == "hybrid": pytest.skip( reason="Mamba model is not supported with the inference optimized transformer." ) @@ -1299,11 +1299,11 @@ def test_mamba_chunked_prefill(self): """ Test chunked prefill with a Mamba model. """ - skip_if_mamba_sequence_packing_not_available("mamba") + skip_if_mamba_sequence_packing_not_available("hybrid") # Context max tokens = 50. test_config = DynamicEngineTestConfig( - model_provider="mamba", + model_provider="hybrid", num_requests=0, num_tokens_to_generate=None, num_tokens_total=200, @@ -4319,7 +4319,7 @@ def test_speculative_decoding_mamba_hybrid(self, rejection_mode): Two requests run simultaneously to exercise batched rewind indexing where mamba_metadata.request_to_mamba_state_idx differs per request. """ - skip_if_mamba_sequence_packing_not_available("mamba") + skip_if_mamba_sequence_packing_not_available("hybrid") num_tokens_to_generate = 8 test_config = DynamicEngineTestConfig( @@ -4329,7 +4329,7 @@ def test_speculative_decoding_mamba_hybrid(self, rejection_mode): num_tokens_to_generate=num_tokens_to_generate, num_speculative_tokens=2, materialize_only_last_token_logits=False, - model_provider="mamba", + model_provider="hybrid", ) env = self._build_test_env(test_config) @@ -4460,7 +4460,7 @@ def _create_model(self, model_provider, num_cuda_graphs): pre_process=parallel_state.is_pipeline_first_stage(), post_process=parallel_state.is_pipeline_last_stage(), ).cuda() - elif model_provider == "mamba": + elif model_provider == "hybrid": config = TransformerConfig( params_dtype=torch.bfloat16, num_layers=3, @@ -4476,9 +4476,9 @@ def _create_model(self, model_provider, num_cuda_graphs): add_bias_linear=True, is_hybrid_model=True, ) - model = MambaModel( + model = HybridModel( config=config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=CHUNKED_CG_VOCAB_SIZE, max_sequence_length=CHUNKED_CG_MAX_SEQ_LEN, parallel_output=True, @@ -4564,7 +4564,7 @@ def _run_to_completion(self, engine, prompts, num_tokens_to_generate): return finished, step_count - @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @pytest.mark.parametrize("model_provider", ["gpt", "hybrid"]) @pytest.mark.parametrize("chunked_prefill", [False, True]) @pytest.mark.parametrize("num_cuda_graphs", [None, 2]) @torch.inference_mode() diff --git a/tests/unit_tests/inference/engines/test_mamba_prefix_caching_e2e.py b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py similarity index 99% rename from tests/unit_tests/inference/engines/test_mamba_prefix_caching_e2e.py rename to tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py index ce21c775b73..303cf76d122 100644 --- a/tests/unit_tests/inference/engines/test_mamba_prefix_caching_e2e.py +++ b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py @@ -54,8 +54,8 @@ from megatron.core.inference.text_generation_controllers.text_generation_controller import ( TextGenerationController, ) -from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec -from megatron.core.models.mamba.mamba_model import MambaModel +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.ssm.mamba_mixer import _check_mamba_sequence_packing_support from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.cuda_graphs import CudaGraphManager, _CudagraphGlobalRecord @@ -131,9 +131,9 @@ def _create_model(self, num_cuda_graphs=None): add_bias_linear=True, is_hybrid_model=True, ) - model = MambaModel( + model = HybridModel( config=transformer_config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=VOCAB_SIZE, max_sequence_length=MAX_SEQ_LEN, parallel_output=True, diff --git a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py index 52a05f7f80f..26a81c5baef 100644 --- a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py +++ b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py @@ -37,8 +37,8 @@ ) from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec from megatron.core.models.gpt.gpt_model import GPTModel -from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec -from megatron.core.models.mamba.mamba_model import MambaModel +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.ssm.mamba_mixer import _check_mamba_sequence_packing_support from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.cuda_graphs import CudaGraphManager, _CudagraphGlobalRecord @@ -121,9 +121,9 @@ def _create_model(self, model_type, num_cuda_graphs=None): add_bias_linear=True, is_hybrid_model=True, ) - model = MambaModel( + model = HybridModel( config=config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=VOCAB_SIZE, max_sequence_length=MAX_SEQ_LEN, parallel_output=True, @@ -343,9 +343,9 @@ def _create_hybrid_model(self, num_cuda_graphs=None): add_bias_linear=True, is_hybrid_model=True, ) - model = MambaModel( + model = HybridModel( config=config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=VOCAB_SIZE, max_sequence_length=MAX_SEQ_LEN, parallel_output=True, diff --git a/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py b/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py index 96b782fad85..229af268a79 100644 --- a/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py +++ b/tests/unit_tests/models/test_dsa_gpt_mamba_equivalence.py @@ -1,6 +1,6 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. """ -Equivalence tests: GPTModel with DSA vs MambaModel with DSA pattern. +Equivalence tests: GPTModel with DSA vs HybridModel with DSA pattern. A small DeepSeek-V3.2 proxy model (4 GPT layers / 8 Mamba layers) is built, weights are remapped GPT→Mamba, and logprobs are compared to verify they are @@ -9,8 +9,8 @@ Architecture equivalence ------------------------ GPTModel layer N (combined attention + MLP in one TransformerLayer) - ≡ MambaModel layer 2N (D, DSA TransformerLayer: input_layernorm + MLASelfAttention) - + MambaModel layer 2N+1 (-, MLPLayer: fused-norm MLP) + ≡ HybridModel layer 2N (D, DSA TransformerLayer: input_layernorm + MLASelfAttention) + + HybridModel layer 2N+1 (-, MLPLayer: fused-norm MLP) Run with:: @@ -35,10 +35,10 @@ get_transformer_block_with_experimental_attention_variant_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel -from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec -from megatron.core.models.mamba.mamba_model import MambaModel +from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.ssm.mamba_hybrid_layer_allocation import validate_segment_layers from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.rl.rl_utils import selective_log_softmax @@ -193,15 +193,15 @@ def _build_mamba_model( layer_pattern: str, pre_process: bool = True, post_process: bool = True, -) -> MambaModel: - """Build a MambaModel with the given hybrid layer pattern.""" +) -> HybridModel: + """Build a HybridModel with the given hybrid layer pattern.""" layer_type_list = validate_segment_layers(layer_pattern) mamba_config = copy.deepcopy(config) mamba_config.num_layers = len(layer_type_list) assert mamba_config.num_layers == _NUM_GPT_LAYERS * 2 - model = MambaModel( + model = HybridModel( config=mamba_config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=_VOCAB_SIZE, max_sequence_length=_MAX_SEQ_LEN, pre_process=pre_process, @@ -221,14 +221,14 @@ def _build_mamba_model( def _remap_gpt_to_mamba_state_dict( gpt_sd: Dict[str, torch.Tensor], num_local_gpt_layers: int ) -> Dict[str, torch.Tensor]: - """Remap a GPTModel state_dict to a MambaModel state_dict. + """Remap a GPTModel state_dict to a HybridModel state_dict. GPTModel layer N (combined attention + MLP) maps to: - * MambaModel layer 2N – DSA attention (input_layernorm + self_attention) - * MambaModel layer 2N+1 – MLP (mlp.*) + * HybridModel layer 2N – DSA attention (input_layernorm + self_attention) + * HybridModel layer 2N+1 – MLP (mlp.*) Additionally, ``decoder.final_layernorm.*`` (TransformerBlock naming) is - remapped to ``decoder.final_norm.*`` (MambaStack naming). + remapped to ``decoder.final_norm.*`` (HybridStack naming). All other keys (embedding, output_layer, rotary_pos_emb, …) are unchanged. @@ -238,7 +238,7 @@ def _remap_gpt_to_mamba_state_dict( pipeline stage (i.e. ``len(gpt_model.decoder.layers)``). Returns: - Remapped state dict ready for MambaModel.load_state_dict(strict=True). + Remapped state dict ready for HybridModel.load_state_dict(strict=True). """ mamba_sd: Dict[str, torch.Tensor] = {} layer_prefix = "decoder.layers." @@ -380,12 +380,12 @@ def _compare_against_golden_values( @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") @pytest.mark.parametrize("tp,pp", [(1, 1), (2, 1), (1, 2)]) class TestDSAGPTMambaEquivalence: - """Verify logprob equivalence between GPTModel+DSA and MambaModel+DSA. + """Verify logprob equivalence between GPTModel+DSA and HybridModel+DSA. For each distributed configuration (TP, PP), the test: 1. Builds a GPTModel with 4 DSA layers. - 2. Builds a MambaModel with pattern "D-D-D-D-" (8 layers). - 3. Remaps and loads GPT weights into MambaModel (strict=True). + 2. Builds a HybridModel with pattern "D-D-D-D-" (8 layers). + 3. Remaps and loads GPT weights into HybridModel (strict=True). 4. Runs the same random tokens through both models. 5. Asserts logprob tensors are numerically close. """ @@ -416,7 +416,7 @@ def test_dsa_logprobs_match(self, tp: int, pp: int) -> None: num_local_gpt_layers = len(gpt_model.decoder.layers) gpt_sd = gpt_model.state_dict() - # ---- Build MambaModel ---- + # ---- Build HybridModel ---- mamba_model = _build_mamba_model( gpt_config, _MAMBA_PATTERN, pre_process=pre_process, post_process=post_process ) @@ -481,7 +481,7 @@ def test_weight_loading_strict(self, tp: int, pp: int) -> None: assert not unexpected, f"Unexpected keys: {unexpected}" def test_record_and_compare_golden_values(self, tp: int, pp: int) -> None: - """Record GPTModel logprobs as golden values, then compare MambaModel against them. + """Record GPTModel logprobs as golden values, then compare HybridModel against them. Golden values are written to the functional test directory so they can be committed and used by the CI inference golden-value tests. @@ -508,7 +508,7 @@ def test_record_and_compare_golden_values(self, tp: int, pp: int) -> None: gpt_logprobs = _forward_logprobs_pp1(gpt_model, tokens) mamba_logprobs = _forward_logprobs_pp1(mamba_model, tokens) - # Verify MambaModel matches golden values + # Verify HybridModel matches golden values _compare_against_golden_values(mamba_logprobs, gpt_logprobs, abs_tol=1e-3) @@ -520,7 +520,7 @@ def test_record_and_compare_golden_values(self, tp: int, pp: int) -> None: @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") @pytest.mark.parametrize("tp,pp", [(1, 1), (2, 1), (1, 2)]) class TestDSAMoEGPTMambaEquivalence: - """Verify logprob equivalence between GPTModel+DSA+MoE and MambaModel+DSA+MoE. + """Verify logprob equivalence between GPTModel+DSA+MoE and HybridModel+DSA+MoE. Architecture: 4 GPT layers with moe_layer_freq=[0,0,1,1] (first 2 dense, last 2 MoE) maps to 8 Mamba layers with pattern "D-D-DEDE": @@ -556,7 +556,7 @@ def test_dsa_moe_logprobs_match(self, tp: int, pp: int) -> None: num_local_gpt_layers = len(gpt_model.decoder.layers) gpt_sd = gpt_model.state_dict() - # ---- Build MambaModel with MoE pattern ---- + # ---- Build HybridModel with MoE pattern ---- mamba_model = _build_mamba_model( gpt_config, _MOE_MAMBA_PATTERN, pre_process=pre_process, post_process=post_process ) @@ -618,7 +618,7 @@ def test_moe_weight_loading_strict(self, tp: int, pp: int) -> None: assert not unexpected, f"Unexpected keys: {unexpected}" def test_moe_record_and_compare_golden_values(self, tp: int, pp: int) -> None: - """Record GPTModel+MoE logprobs as golden values, then compare MambaModel+MoE.""" + """Record GPTModel+MoE logprobs as golden values, then compare HybridModel+MoE.""" self._skip_if_insufficient_gpus(tp, pp) if tp != 1 or pp != 1: pytest.skip("Golden-value recording only runs for tp=1, pp=1") @@ -640,5 +640,5 @@ def test_moe_record_and_compare_golden_values(self, tp: int, pp: int) -> None: gpt_logprobs = _forward_logprobs_pp1(gpt_model, tokens) mamba_logprobs = _forward_logprobs_pp1(mamba_model, tokens) - # Verify MambaModel matches golden values + # Verify HybridModel matches golden values _compare_against_golden_values(mamba_logprobs, gpt_logprobs, abs_tol=1e-3) diff --git a/tests/unit_tests/post_training/test_modelopt_model_builder.py b/tests/unit_tests/post_training/test_modelopt_model_builder.py index b489d659ec4..2ab8ebfe947 100644 --- a/tests/unit_tests/post_training/test_modelopt_model_builder.py +++ b/tests/unit_tests/post_training/test_modelopt_model_builder.py @@ -39,7 +39,7 @@ def test_model_provider_switches_to_modelopt_builder(monkeypatch): monkeypatch.setattr(mp, "has_nvidia_modelopt", True) monkeypatch.setattr(mp, "get_args", lambda: args) monkeypatch.setattr( - mp, "modelopt_gpt_mamba_builder", _sentinel_builder(modelopt_result, modelopt_calls) + mp, "modelopt_gpt_hybrid_builder", _sentinel_builder(modelopt_result, modelopt_calls) ) # original_builder should be ignored when ModelOpt is enabled. diff --git a/tests/unit_tests/resharding/test_model_swap.py b/tests/unit_tests/resharding/test_model_swap.py index 70d81d97829..e2d6a2bd096 100644 --- a/tests/unit_tests/resharding/test_model_swap.py +++ b/tests/unit_tests/resharding/test_model_swap.py @@ -1,4 +1,4 @@ -# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. import copy import gc import os @@ -37,8 +37,8 @@ try: import mamba_ssm # noqa: F401 - from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec - from megatron.core.models.mamba.mamba_model import MambaModel + from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec + from megatron.core.models.hybrid.hybrid_model import HybridModel has_mamba_deps = True except Exception: @@ -203,9 +203,9 @@ def _build_mamba( parallel_output: bool = True, ): pre_process, post_process = _pp_flags(pg_collection) - model = MambaModel( + model = HybridModel( config=config, - mamba_stack_spec=mamba_stack_spec, + hybrid_stack_spec=hybrid_stack_spec, vocab_size=vocab_size, max_sequence_length=seq_len, hybrid_layer_pattern=hybrid_layer_pattern, diff --git a/tools/checkpoint/remap_gpt_dsa_to_mamba.py b/tools/checkpoint/remap_gpt_dsa_to_mamba.py index 8a6888d1dc7..3d11c981c25 100644 --- a/tools/checkpoint/remap_gpt_dsa_to_mamba.py +++ b/tools/checkpoint/remap_gpt_dsa_to_mamba.py @@ -1,16 +1,16 @@ #!/usr/bin/env python3 # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. -"""Convert a GPTModel DSA checkpoint to a MambaModel-compatible checkpoint. +"""Convert a GPTModel DSA checkpoint to a HybridModel-compatible checkpoint. A GPTModel with ``--experimental-attention-variant dsa`` uses one combined -TransformerLayer per model layer (attention + MLP). The equivalent MambaModel +TransformerLayer per model layer (attention + MLP). The equivalent HybridModel with pattern ``D-D-...`` stores them as two separate layers: * Layer 2N – DSA attention (TransformerLayer: input_layernorm + MLASelfAttention) * Layer 2N+1 – MLP (MLPLayer: fused-norm MLP) This script loads a GPTModel Distributed Checkpoint (DCP), remaps the state-dict -keys, and saves a new DCP that can be loaded by MambaModel. +keys, and saves a new DCP that can be loaded by HybridModel. Usage ----- @@ -43,14 +43,14 @@ def _remap_key(key: str, num_gpt_layers: int) -> str: - """Return the MambaModel state-dict key corresponding to *key* from GPTModel. + """Return the HybridModel state-dict key corresponding to *key* from GPTModel. Args: key: A key from the GPTModel state dict. num_gpt_layers: Total number of GPT decoder layers (across all PP stages). Returns: - The remapped key for MambaModel. + The remapped key for HybridModel. Raises: ValueError: If an unexpected sub-key is encountered in a decoder layer. @@ -58,7 +58,7 @@ def _remap_key(key: str, num_gpt_layers: int) -> str: layer_prefix = "decoder.layers." final_ln_prefix = "decoder.final_layernorm." - # Final layernorm name differs between TransformerBlock and MambaStack + # Final layernorm name differs between TransformerBlock and HybridStack if key.startswith(final_ln_prefix): return "decoder.final_norm." + key[len(final_ln_prefix):] @@ -96,11 +96,11 @@ def _remap_state_dict( def convert(input_path: Path, output_path: Path, num_gpt_layers: int) -> None: - """Load a GPTModel DCP checkpoint, remap keys, and save as MambaModel DCP. + """Load a GPTModel DCP checkpoint, remap keys, and save as HybridModel DCP. Args: input_path: Path to the GPTModel DCP checkpoint directory. - output_path: Destination directory for the MambaModel DCP checkpoint. + output_path: Destination directory for the HybridModel DCP checkpoint. num_gpt_layers: Number of GPT decoder layers in the original model. """ try: @@ -139,7 +139,7 @@ def convert(input_path: Path, output_path: Path, num_gpt_layers: int) -> None: output_path.mkdir(parents=True, exist_ok=True) torch_save_to_dcp(str(tmp_mamba), str(output_path)) - print(f"MambaModel DCP checkpoint saved to: {output_path}") + print(f"HybridModel DCP checkpoint saved to: {output_path}") finally: for tmp in (tmp_flat, output_path.parent / "_tmp_mamba_flat.pt"): @@ -149,7 +149,7 @@ def convert(input_path: Path, output_path: Path, num_gpt_layers: int) -> None: def main() -> None: parser = argparse.ArgumentParser( - description="Convert GPTModel DSA checkpoint to MambaModel-compatible format." + description="Convert GPTModel DSA checkpoint to HybridModel-compatible format." ) parser.add_argument( "--input", required=True, type=Path, @@ -157,7 +157,7 @@ def main() -> None: ) parser.add_argument( "--output", required=True, type=Path, - help="Destination path for the MambaModel DCP checkpoint.", + help="Destination path for the HybridModel DCP checkpoint.", ) parser.add_argument( "--num-gpt-layers", required=True, type=int, diff --git a/tools/run_mamba_text_generation_server.py b/tools/run_hybrid_text_generation_server.py similarity index 89% rename from tools/run_mamba_text_generation_server.py rename to tools/run_hybrid_text_generation_server.py index 33465f1bb4a..e70e5389e88 100644 --- a/tools/run_mamba_text_generation_server.py +++ b/tools/run_hybrid_text_generation_server.py @@ -8,4 +8,4 @@ from run_text_generation_server import main if __name__ == "__main__": - main(model_type="mamba") + main(model_type="hybrid") diff --git a/tools/run_mamba_text_generation_server_completions.py b/tools/run_hybrid_text_generation_server_completions.py similarity index 89% rename from tools/run_mamba_text_generation_server_completions.py rename to tools/run_hybrid_text_generation_server_completions.py index 33465f1bb4a..e70e5389e88 100644 --- a/tools/run_mamba_text_generation_server_completions.py +++ b/tools/run_hybrid_text_generation_server_completions.py @@ -8,4 +8,4 @@ from run_text_generation_server import main if __name__ == "__main__": - main(model_type="mamba") + main(model_type="hybrid") diff --git a/tools/run_inference_performance_test.py b/tools/run_inference_performance_test.py index ac9e92d3639..d42453c62ed 100644 --- a/tools/run_inference_performance_test.py +++ b/tools/run_inference_performance_test.py @@ -9,7 +9,7 @@ import torch from gpt_builders import gpt_builder -from mamba_builders import mamba_builder +from hybrid_builders import hybrid_builder from megatron.core.inference.contexts import StaticInferenceContext from megatron.core.inference.engines import DynamicInferenceEngine, StaticInferenceEngine from megatron.core.inference.engines.abstract_engine import AbstractEngine diff --git a/tools/run_text_generation_server.py b/tools/run_text_generation_server.py index 5a2940f1a4c..e871214e739 100644 --- a/tools/run_text_generation_server.py +++ b/tools/run_text_generation_server.py @@ -15,7 +15,7 @@ import torch from gpt_builders import gpt_builder -from mamba_builders import mamba_builder +from hybrid_builders import hybrid_builder from megatron.core.inference.contexts import StaticInferenceContext from megatron.core.inference.engines import AbstractEngine, StaticInferenceEngine from megatron.core.inference.engines.abstract_engine import AbstractEngine @@ -140,8 +140,16 @@ def main(model_type: str = "gpt"): # Set up model and load checkpoint if model_type == "gpt": model_builder = gpt_builder - elif model_type == "mamba": - model_builder = mamba_builder + elif model_type in ("hybrid", "mamba"): + if model_type == "mamba": + import warnings + + warnings.warn( + 'model_type="mamba" is deprecated. Use model_type="hybrid" instead.', + DeprecationWarning, + stacklevel=2, + ) + model_builder = hybrid_builder else: raise ValueError(f"Invalid model provider {model_type}") model = get_model(partial(model_provider, model_builder), wrap_with_ddp=False) diff --git a/train_rl.py b/train_rl.py index 8bcee5f096d..3e4ccdf4f39 100644 --- a/train_rl.py +++ b/train_rl.py @@ -8,7 +8,7 @@ import torch from gpt_builders import gpt_builder -from mamba_builders import mamba_builder +from hybrid_builders import hybrid_builder from megatron.core import mpu from megatron.core.enums import ModelType from megatron.core.models.gpt import GPTModel @@ -392,7 +392,7 @@ def _model_builder( args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None ): if is_hybrid_model(args): - return mamba_builder( + return hybrid_builder( args, pre_process, post_process,