diff --git a/examples/post_training/modelopt/convert_model.py b/examples/post_training/modelopt/convert_model.py index 90a1859dc61..9790d73fc4c 100644 --- a/examples/post_training/modelopt/convert_model.py +++ b/examples/post_training/modelopt/convert_model.py @@ -18,12 +18,13 @@ 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import report_current_memory_info, to_empty_if_meta from megatron.training import get_args, get_tokenizer from megatron.training.checkpointing import save_checkpoint from megatron.training.initialize import initialize_megatron from megatron.training.utils import print_rank_0, unwrap_model +from model_provider import model_provider ALGO_TO_CONFIG = { "eagle1": mtsp.config.EAGLE1_DEFAULT_CFG, @@ -120,13 +121,13 @@ def check_arguments(): UserWarning, ) - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) report_current_memory_info() unwrapped_model = unwrap_model(model)[0] if args.pretrained_model_path is not None: - import_dtype = torch.float16 if args.fp16 else torch.bfloat16 + import_dtype = torch.float16 if args.fp16 else torch.bfloat16 unwrapped_model = unwrap_model(model)[0] workspace_dir = os.environ.get("MLM_WORK_DIR", "/tmp") print_rank_0("Import model from Hugging Face checkpoint in dtype {}.".format(str(import_dtype))) @@ -168,10 +169,10 @@ def check_arguments(): tokenizer = get_tokenizer() for i in range(unwrapped_model.eagle_config.parallel_draft_step - 1): mask_token = "[MASK_{}]".format(i) - tokenizer._tokenizer.add_tokens([mask_token], special_tokens=True) + tokenizer._tokenizer.add_tokens([mask_token], special_tokens=True) token_id = tokenizer._tokenizer.convert_tokens_to_ids(mask_token) setattr(unwrapped_model, "mask_token_{}".format(i), torch.tensor(token_id)) - + elif args.algorithm == "medusa": config = {"medusa_num_heads": args.export_num_medusa_heads, "medusa_num_layers": 1} unwrapped_model = mtsp.convert(unwrapped_model, [("medusa", config)]) diff --git a/examples/post_training/modelopt/export.py b/examples/post_training/modelopt/export.py index 4509af0b200..8794c4c738c 100644 --- a/examples/post_training/modelopt/export.py +++ b/examples/post_training/modelopt/export.py @@ -13,10 +13,11 @@ from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.training import get_args, get_model from megatron.training.initialize import initialize_megatron from megatron.training.utils import unwrap_model +from model_provider import model_provider warnings.filterwarnings('ignore') @@ -64,11 +65,11 @@ def add_modelopt_export_args(parser): UserWarning, ) - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) - # Materialize the model from meta device to cpu before loading the checkpoint. + # Materialize the model from meta device to cpu before loading the checkpoint. unwrapped_model = unwrap_model(model)[0] - unwrapped_model.to_empty(device="cpu") + unwrapped_model.to_empty(device="cpu") if args.load is not None: _ = load_modelopt_checkpoint(model) diff --git a/examples/post_training/modelopt/finetune.py b/examples/post_training/modelopt/finetune.py index ef3f83f773b..bd0569bb513 100755 --- a/examples/post_training/modelopt/finetune.py +++ b/examples/post_training/modelopt/finetune.py @@ -2,12 +2,12 @@ """Supervised Finetuning GPT.""" import itertools +import json import os import sys from functools import partial from typing import Any, Dict, Optional -import json import jsonlines sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../"))) @@ -20,17 +20,16 @@ from megatron.core.enums import ModelType from megatron.core.models.gpt import GPTModel from megatron.post_training.arguments import add_modelopt_args -from megatron.post_training.model_provider import model_provider from megatron.post_training.loss_func import loss_func +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.non_loss_data_func import report_draft_acceptance_length from megatron.training import get_args, get_timers, get_tokenizer, pretrain from megatron.training.utils import ( - average_losses_across_data_parallel_group, get_batch_on_this_cp_rank, get_ltor_masks_and_position_ids, print_rank_0, - unwrap_model, ) +from model_provider import model_provider REMOVE_THINK_CHAT_TEMPLATE = ( "{% if '' in content %}{% set content = content.split('')[-1] %}{% endif %}" @@ -48,7 +47,7 @@ def add_finetune_args(parser): def get_eos_id(): """Return the eos token id. - + We insert eos_token between two samples during packing. However, if the eos_token is used in message or after turns, we need to replace it with some other special tokens that do not appear in message.""" tokenizer = get_tokenizer() @@ -79,7 +78,7 @@ def __init__(self, data_dir: str, num_samples): def __len__(self): return self.num_samples - + def __getitem__(self, idx): idx = idx % len(self.file_paths) file_path = self.file_paths[idx] @@ -387,7 +386,7 @@ def train_valid_test_sft_datasets_provider(train_val_test_num_samples): def get_batch(data_iterator): """Generate a batch. - + For OfflineDataset, the aux_hidden_states and final hidden_states from the base model are loaded for offline speculative model training.""" # TODO: this is pretty hacky, find a better way @@ -495,7 +494,7 @@ def forward_step(data_iterator, model: GPTModel): if __name__ == "__main__": pretrain( train_valid_test_sft_datasets_provider, - model_provider, + partial(model_provider, modelopt_gpt_mamba_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 f923863f0a7..a773ea89f00 100644 --- a/examples/post_training/modelopt/generate.py +++ b/examples/post_training/modelopt/generate.py @@ -14,10 +14,11 @@ 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import report_current_memory_info, to_empty_if_meta 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 warnings.filterwarnings('once') @@ -96,7 +97,7 @@ def get_conversations(example): UserWarning, ) - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) report_current_memory_info() unwrapped_model = unwrap_model(model)[0] diff --git a/examples/post_training/modelopt/generation_server.py b/examples/post_training/modelopt/generation_server.py index fcd02669354..b32cca0d73f 100644 --- a/examples/post_training/modelopt/generation_server.py +++ b/examples/post_training/modelopt/generation_server.py @@ -4,6 +4,7 @@ import os import sys import warnings +from functools import partial sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../"))) import os @@ -11,6 +12,10 @@ from argparse import Namespace from contextlib import nullcontext +import torch + +from megatron.core import mpu +from megatron.core.inference.engines import AbstractEngine, StaticInferenceEngine from megatron.core.inference.engines.abstract_engine import AbstractEngine from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import ( InferenceWrapperConfig, @@ -19,13 +24,6 @@ from megatron.core.inference.text_generation_controllers.text_generation_controller import ( TextGenerationController, ) -import torch - -from megatron.core.inference.engines import AbstractEngine, StaticInferenceEngine -from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import ( - InferenceWrapperConfig, -) -from megatron.training import get_model from megatron.core.transformer.module import MegatronModule from megatron.inference.text_generation import beam_search_and_post_process from megatron.inference.text_generation.mcore_engine_server import ( @@ -33,17 +31,12 @@ run_mcore_engine, ) from megatron.inference.text_generation_server import MegatronServer -from megatron.training import print_rank_0 - -from megatron.core import mpu -from megatron.training import get_args, get_model, get_tokenizer +from megatron.post_training.arguments import add_modelopt_args +from megatron.training import get_args, get_model, get_tokenizer, print_rank_0 from megatron.training.checkpointing import load_checkpoint from megatron.training.initialize import initialize_megatron -from megatron.post_training.arguments import add_modelopt_args - - def get_inference_engine(args: Namespace, model: MegatronModule) -> AbstractEngine: """Get the relevant backend for running inference @@ -144,9 +137,10 @@ def main(model_provider: str = "gpt"): load_context = fp8_model_init() with load_context: - from megatron.post_training.model_provider import model_provider as modelopt_model_provider + from megatron.post_training.model_builder import modelopt_gpt_mamba_builder + from model_provider import model_provider as root_model_provider if model_provider == "gpt": - model = get_model(modelopt_model_provider, wrap_with_ddp=False) + model = get_model(partial(root_model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) elif model_provider == "mamba": pass else: diff --git a/examples/post_training/modelopt/mmlu.py b/examples/post_training/modelopt/mmlu.py index dbe150b6dca..1446afc8392 100644 --- a/examples/post_training/modelopt/mmlu.py +++ b/examples/post_training/modelopt/mmlu.py @@ -14,10 +14,11 @@ 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import report_current_memory_info 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 warnings.filterwarnings('ignore') @@ -148,7 +149,7 @@ def generate_prompt(test_example, dev_examples, few_shots=0, no_subject_prompt=F UserWarning, ) - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) report_current_memory_info() disable_tqdm = args.disable_tqdm or torch.distributed.get_rank() > 0 diff --git a/examples/post_training/modelopt/offline_feature_extract.py b/examples/post_training/modelopt/offline_feature_extract.py index e33d68b0420..80207faf2b2 100644 --- a/examples/post_training/modelopt/offline_feature_extract.py +++ b/examples/post_training/modelopt/offline_feature_extract.py @@ -4,19 +4,21 @@ import functools import os import sys + import torch sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../"))) +from examples.post_training.modelopt.finetune import SFTDataset 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_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 -from examples.post_training.modelopt.finetune import SFTDataset def add_extract_args(parser): """Add additional arguments for feature extraction.""" @@ -51,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, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) load_modelopt_checkpoint(model, strict=not args.untie_embeddings_and_output_weights) print_rank_0("Done loading checkpoint") @@ -68,7 +70,7 @@ def extract_feature(dataset, model, output_dir, idx_start, idx_end): "shard_index": mpu.get_expert_data_parallel_rank(), } sft_dataset = SFTDataset(args.num_samples, None, **kwargs) - + extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "train"), 0, int(args.num_samples * 0.98)) extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "valid"), int(args.num_samples * 0.98), int(args.num_samples * 0.99)) extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "test"), int(args.num_samples * 0.99), args.num_samples) diff --git a/examples/post_training/modelopt/prune.py b/examples/post_training/modelopt/prune.py index 7b91370ed58..7819b2ed2af 100644 --- a/examples/post_training/modelopt/prune.py +++ b/examples/post_training/modelopt/prune.py @@ -23,11 +23,12 @@ 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import report_current_memory_info from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron from megatron.training.checkpointing import save_checkpoint from megatron.training.utils import print_rank_0, unwrap_model +from model_provider import model_provider warnings.filterwarnings("ignore") @@ -138,7 +139,7 @@ def get_calib_dataloader(calib_size=1024, max_sequence_length=512): check_arguments(args) tokenizer = get_tokenizer()._tokenizer - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) unwrapped_model = unwrap_model(model)[0] report_current_memory_info() diff --git a/examples/post_training/modelopt/quantize.py b/examples/post_training/modelopt/quantize.py index fd6b685960f..737aed68b6a 100644 --- a/examples/post_training/modelopt/quantize.py +++ b/examples/post_training/modelopt/quantize.py @@ -20,11 +20,12 @@ 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_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import report_current_memory_info from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron from megatron.training.checkpointing import save_checkpoint from megatron.training.utils import print_rank_0, unwrap_model +from model_provider import model_provider warnings.filterwarnings("ignore") @@ -157,7 +158,7 @@ def get_calib_dataloader(calib_size=512, max_sequence_length=512): args = get_args() tokenizer = get_tokenizer()._tokenizer - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) report_current_memory_info() diff --git a/examples/post_training/modelopt/validate.py b/examples/post_training/modelopt/validate.py index 7e0b119991a..ee8bf64eccb 100644 --- a/examples/post_training/modelopt/validate.py +++ b/examples/post_training/modelopt/validate.py @@ -2,29 +2,23 @@ """Sample Generate GPT.""" import functools +import json import os import sys import warnings -import json sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../"))) -import modelopt -from modelopt.torch.speculative.plugins.megatron_eagle import MegatronARValidation import torch -from datasets import load_dataset -from tqdm import tqdm +from modelopt.torch.speculative.plugins.megatron_eagle import MegatronARValidation -from megatron.core import mpu -from megatron.core.inference.communication_utils import broadcast_from_last_pipeline_stage -from megatron.core.pipeline_parallel import get_forward_backward_func -from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.checkpointing import load_modelopt_checkpoint -from megatron.post_training.model_provider import model_provider +from megatron.post_training.model_builder import modelopt_gpt_mamba_builder from megatron.post_training.utils import get_mtbench_chat_data from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron -from megatron.training.utils import get_ltor_masks_and_position_ids, print_rank_0, unwrap_model +from megatron.training.utils import print_rank_0, unwrap_model +from model_provider import model_provider warnings.filterwarnings('ignore') @@ -121,7 +115,7 @@ def report_current_memory_info(): ground_truth = [None for _ in range(len(prompts))] tokenizer = get_tokenizer()._tokenizer - model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False) + model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False) report_current_memory_info() diff --git a/megatron/post_training/loss_func.py b/megatron/post_training/loss_func.py index a8f664be7d7..eb8dbca1c6a 100644 --- a/megatron/post_training/loss_func.py +++ b/megatron/post_training/loss_func.py @@ -2,14 +2,12 @@ """Pretrain GPT loss function(s).""" -import os - import torch from megatron.core import parallel_state from megatron.core.models.gpt import GPTModel from megatron.training import get_args -from megatron.training.utils import average_losses_across_data_parallel_group, unwrap_model +from megatron.training.utils import unwrap_model def _mask_loss(output_tensor, loss_mask): @@ -38,24 +36,6 @@ def _mask_loss(output_tensor, loss_mask): return loss -def _allreduce_losses(losses): - """Reduce losses across all GPUs.""" - args = get_args() - - # Check individual rank losses are not NaN prior to DP all-reduce. - if args.check_for_nan_in_loss_and_grad: - global_rank = torch.distributed.get_rank() - for loss in losses: - assert not loss.isnan(), ( - f'Rank {global_rank}: found NaN in local forward loss calculation. ' - f'Device: {torch.cuda.current_device()}, node: {os.uname()[1]}' - ) - - # Reduce loss for logging. - # TODO(aanoosheh): This should ideally be done with num_tokens separately reduced and averaged. - return average_losses_across_data_parallel_group(losses) - - def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: GPTModel): """Loss function (with KD Loss support). diff --git a/megatron/post_training/model_provider.py b/megatron/post_training/model_builder.py similarity index 95% rename from megatron/post_training/model_provider.py rename to megatron/post_training/model_builder.py index d56e8b47377..34daa279651 100644 --- a/megatron/post_training/model_provider.py +++ b/megatron/post_training/model_builder.py @@ -2,13 +2,11 @@ """ModelOpt GPT model provider.""" -import json import os from argparse import Namespace from typing import Any, Dict import modelopt.torch.distill as mtd -import modelopt.torch.opt as mto import yaml from megatron.core.models.gpt import GPTModel as MCoreGPTModel @@ -34,6 +32,7 @@ def count_parameters_in_layer(model, layer_name): print_rank_0(f" - {name}: {param.numel()}") return num_params + def _add_load_convert_hooks(model: MCoreGPTModel): """Register some load_state_dict prehooks to handle some known state_dict key mismatch. """ @@ -130,22 +129,19 @@ def _teacher_provider(config: Namespace, model_kwargs: Dict[str, Any]) -> MCoreG return teacher -def model_provider(pre_process=True, post_process=True, parallel_output=True) -> MCoreGPTModel: +def modelopt_gpt_mamba_builder(args, pre_process, post_process, vp_stage=None, config=None) -> MCoreGPTModel | MCoreMambaModel: """Builds the model. - If you set the use_legacy_models to True, it will return the legacy GPT model and if not the core GPT model. - Args: + args (Namespace): The arguments namespace. pre_process (bool, optional): Set to true if you need to compute embedings. Defaults to True. post_process (bool, optional): Set to true if you need to want to compute output logits/loss. Defaults to True. - parallel_output (bool): whether to allgather the output logits? This must be - True if `model_provider` is called in text_generation_server. + vp_stage (int, optional): The virtual pipeline stage. + config (TransformerConfig, optional): The configuration object. Returns: - MCoreGPTModel: The returned model + MCoreGPTModel | MCoreMambaModel: The returned model """ - args = get_args() - print_rank_0("building GPT model ...") # ModelOpt by default assumes none homogenous layers. This affect the storage format of the sharded checkpoint. @@ -166,6 +162,8 @@ def model_provider(pre_process=True, post_process=True, parallel_output=True) -> config.yarn_mscale_all_dim = 0.0 config.yarn_correction_range_round_to_int = False + if vp_stage is not None: + raise ValueError("ModelOpt integration does not currently support virtual pipeline parallel.") if args.use_legacy_models: raise ValueError( "ModelOpt integration only support MCore models. Use --use-mcore-modules instead." @@ -174,8 +172,8 @@ def model_provider(pre_process=True, post_process=True, parallel_output=True) -> raise ValueError("ModelOpt integration does not support custom args.spec.") # Llama-4 Scout/Maverick support - config.qk_l2_norm = args.export_qk_l2_norm - config.moe_apply_probs_on_input = args.export_moe_apply_probs_on_input + config.qk_l2_norm = args.export_qk_l2_norm + config.moe_apply_probs_on_input = args.export_moe_apply_probs_on_input if args.export_model_type == "GPTModel": if args.export_offline_model: @@ -216,7 +214,7 @@ def model_provider(pre_process=True, post_process=True, parallel_output=True) -> "pre_process": pre_process, "post_process": post_process, "fp16_lm_cross_entropy": args.fp16_lm_cross_entropy, - "parallel_output": parallel_output, + "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, @@ -257,7 +255,7 @@ def model_provider(pre_process=True, post_process=True, parallel_output=True) -> raise ValueError("ModelOpt does not support model type {}".format(args.export_model_type)) # [IMPORTANT] Load modelopt_state immediately before returning the model back to `get_model()`. - # + # # ModelOpt can create additional trainable parameters (e.g. for online speculative # decoding training or PEFT). Hence resuming modelopt_state during checkpoint loading is already # too late since Megatron created the optimizer right after calling model_provider before loading diff --git a/model_provider.py b/model_provider.py index 4d8b0daac71..2d43f94eb53 100644 --- a/model_provider.py +++ b/model_provider.py @@ -11,8 +11,7 @@ from megatron.training import get_args, print_rank_0 try: - from megatron.post_training.model_provider import model_provider as model_provider_modelopt - + from megatron.post_training.model_builder import modelopt_gpt_mamba_builder has_nvidia_modelopt = True except ImportError: has_nvidia_modelopt = False @@ -34,15 +33,11 @@ def model_provider( 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 """ args = get_args() - if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt] - return model_provider_modelopt(pre_process, post_process) - if args.record_memory_history: torch.cuda.memory._record_memory_history( True, @@ -65,6 +60,10 @@ def oom_observer(device, alloc, device_alloc, device_free): torch._C._cuda_attach_out_of_memory_observer(oom_observer) + 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 + return model_builder(args, pre_process, post_process, vp_stage) diff --git a/pretrain_gpt.py b/pretrain_gpt.py index 6316aef03bf..69f26f3271a 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -2,29 +2,29 @@ """Pretrain and SFT GPT.""" -import torch - from functools import partial from typing import List, Optional, Tuple + +import torch + +from gpt_builders import gpt_builder from megatron.core import parallel_state -from megatron.training import inprocess_restart 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.models.gpt import GPTModel from megatron.core.rerun_state_machine import get_rerun_state_machine -from megatron.core.utils import get_attr_wrapped_model, StragglerDetector from megatron.core.tokenizers.text.utils.build_tokenizer import build_tokenizer -from megatron.training import get_args, get_timers, get_tokenizer, pretrain, print_rank_0 +from megatron.core.utils import StragglerDetector, get_attr_wrapped_model +from megatron.training import get_args, get_timers, get_tokenizer, inprocess_restart, pretrain, print_rank_0 +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 megatron.training.datasets.sft_dataset import SFTDataset from model_provider import model_provider -from gpt_builders import gpt_builder try: from megatron.post_training.arguments import add_modelopt_args @@ -75,11 +75,14 @@ def loss_func( args = get_args() if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt] - return loss_func_modelopt(loss_mask, output_tensor, model=model) + 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) - 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() @@ -112,10 +115,7 @@ def loss_func( fatal=False, ) - num_tokens = loss_mask.sum().clone().detach().to(torch.int) - reporting_loss = torch.cat([loss.clone().detach().view(1), num_tokens.view(1)]) - - return (loss, num_tokens, {'lm loss': reporting_loss}) + return loss, num_tokens, report def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = False): diff --git a/pretrain_mamba.py b/pretrain_mamba.py index ba084c73478..45b646a6cc0 100644 --- a/pretrain_mamba.py +++ b/pretrain_mamba.py @@ -1,38 +1,30 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. """Pretrain and SFT Mamba.""" -import os -import torch from functools import partial from typing import List, Optional, Tuple -from model_provider import model_provider -from mamba_builders import mamba_builder +import torch -from megatron.training import get_args -from megatron.training import get_tokenizer -from megatron.training import inprocess_restart -from megatron.training import print_rank_0 -from megatron.training import get_timers +from mamba_builders import mamba_builder from megatron.core import mpu -from megatron.core.enums import ModelType -from megatron.core.tokenizers.text.utils.build_tokenizer import build_tokenizer from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder -from megatron.core.datasets.gpt_dataset import GPTDatasetConfig -from megatron.core.datasets.gpt_dataset import MockGPTDataset, GPTDataset -from megatron.core.rerun_state_machine import get_rerun_state_machine +from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset +from megatron.core.enums import ModelType from megatron.core.models.mamba import MambaModel -from megatron.training import pretrain -from megatron.core.utils import get_attr_wrapped_model, StragglerDetector +from megatron.core.rerun_state_machine import get_rerun_state_machine +from megatron.core.tokenizers.text.utils.build_tokenizer import build_tokenizer +from megatron.core.utils import StragglerDetector, get_attr_wrapped_model +from megatron.training import get_args, get_timers, get_tokenizer, inprocess_restart, pretrain, print_rank_0 +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 megatron.training.datasets.sft_dataset import SFTDataset +from model_provider import model_provider -# modelopt distillation try: from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.loss_func import loss_func as loss_func_modelopt @@ -42,6 +34,7 @@ stimer = StragglerDetector() + def get_batch(data_iterator, vp_stage=None): """Generate a batch.""" @@ -61,7 +54,6 @@ def get_batch(data_iterator, vp_stage=None): # 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. @@ -77,11 +69,14 @@ def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: Optio """ args = get_args() if has_nvidia_modelopt and getattr(args, 'modelopt_enabled', False): # [ModelOpt] - return loss_func_modelopt(loss_mask, output_tensor, model=model) + 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) - 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() @@ -114,10 +109,7 @@ def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: Optio fatal=False, ) - num_tokens = loss_mask.sum().clone().detach().to(torch.int) - reporting_loss = torch.cat([loss.clone().detach().view(1), num_tokens.view(1)]) - - return (loss, num_tokens, {'lm loss': reporting_loss}) + return loss, num_tokens, report def forward_step(data_iterator, model: MambaModel): @@ -127,7 +119,6 @@ def forward_step(data_iterator, model: MambaModel): data_iterator : Input data iterator model (MambaModel): The GPT Model """ - args = get_args() timers = get_timers() # Get the batch. @@ -188,7 +179,6 @@ def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None 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) if args.sft: