From 487cec0701269faa82e56a1223cbc6e1a886552c Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Mon, 20 Apr 2026 11:25:08 -0700 Subject: [PATCH 1/4] Move KD teacher loading to after Float16Module Signed-off-by: Asha Anoosheh --- megatron/post_training/model_builder.py | 24 ++++++------------------ megatron/training/checkpointing.py | 25 ++++++++++++++++++++----- 2 files changed, 26 insertions(+), 23 deletions(-) diff --git a/megatron/post_training/model_builder.py b/megatron/post_training/model_builder.py index 085d188e811..df7cf87a6d6 100644 --- a/megatron/post_training/model_builder.py +++ b/megatron/post_training/model_builder.py @@ -21,12 +21,11 @@ from megatron.core.post_training.modelopt.gpt.state_dict_hooks import ( mcore_gpt_load_te_state_dict_pre_hook, ) -from megatron.post_training.checkpointing import load_modelopt_checkpoint, load_modelopt_state +from megatron.post_training.checkpointing import load_modelopt_state +from megatron.post_training.utils import print_distributed_quant_summary from megatron.training import get_args, print_rank_0 from megatron.training.arguments import core_transformer_config_from_args -from megatron.post_training.utils import print_distributed_quant_summary - def count_parameters_in_layer(model, layer_name): num_params = 0 @@ -114,7 +113,7 @@ def _load_teacher_model_config(checkpoint_path: str) -> Namespace: return Namespace(**args_dict) -def _load_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, Any]) -> MCoreGPTModel: +def _build_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, Any]) -> MCoreGPTModel: """Teacher model creator.""" args = get_args() @@ -141,19 +140,8 @@ def _load_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, A use_arbitrary_attention_mask=False, ) teacher = MCoreGPTModel(config=config, **model_kwargs) - _add_load_convert_hooks(teacher) - print_rank_0(f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ...") - # [WAR]: load checkpoint will check checkpoint's saved args and rng state if not finetune. - # To avoid error out on loading teacher's checkpoint, we temporarily set args.finetune to - # True while loading the teacher checkpoint. - original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format - args.finetune = True - if args.export_kd_teacher_ckpt_format is not None: - args.ckpt_format = args.export_kd_teacher_ckpt_format - load_modelopt_checkpoint([teacher], load_arg='export_kd_teacher_load') - args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format - print_rank_0("...teacher loaded successfully.") + _add_load_convert_hooks(teacher) return teacher @@ -338,7 +326,7 @@ def modelopt_gpt_mamba_builder( args.export_kd_cfg, student_cfg=config, teacher_cfg=teacher_config ) kd_config = { - "teacher_model": _load_teacher_model(teacher_config, teacher_config_raw, model_kwargs), + "teacher_model": _build_teacher_model(teacher_config, teacher_config_raw, model_kwargs), "criterion": distill_cfg.criterion, "loss_balancer": distill_cfg.loss_balancer, } @@ -349,6 +337,6 @@ def modelopt_gpt_mamba_builder( mtd_mcore.adjust_distillation_model_for_mcore(model, distill_cfg) # Also remove KD mode state to prevent issues with re-conversion after restore. mto.ModeloptStateManager(model).state_dict().pop() # TODO(aanoosheh): remove once fixed in ModelOpt - + print_distributed_quant_summary(model) return model diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 1441a71518d..28324e01ac3 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -38,7 +38,7 @@ from megatron.core.num_microbatches_calculator import update_num_microbatches from megatron.core.optimizer import DistributedOptimizer from megatron.core.rerun_state_machine import get_rerun_state_machine -from megatron.core.utils import get_pg_rank, get_pg_size, get_torch_version, is_torch_min_version +from megatron.core.utils import get_attr_wrapped_model, get_pg_rank, get_pg_size from ..core.dist_checkpointing.utils import _clean_metadata_for_serialization from . import ft_integration, wandb_utils @@ -2037,10 +2037,7 @@ def load_model_state_dict(module, state_dict, strict: bool): f'[ t {mpu.get_tensor_model_parallel_rank() + 1}/{mpu.get_tensor_model_parallel_world_size()}, ' f'p {mpu.get_pipeline_model_parallel_rank() + 1}/{mpu.get_pipeline_model_parallel_world_size()} ] ' f'at iteration {iteration}') - - if has_nvidia_modelopt: - print_distributed_quant_summary(model, msg="After loading checkpoint") - + # Additional callback for wandb (last rank) if not torch.distributed.is_initialized() \ or is_last_rank(): @@ -2063,6 +2060,24 @@ def load_model_state_dict(module, state_dict, strict: bool): print_rank_0(">>> Inserting 'default_config' field into optimizer.param_groups...") log_printed = True + if has_nvidia_modelopt: + print_distributed_quant_summary(model, msg="After loading checkpoint") + if args.export_kd_teacher_load: + from megatron.post_training.checkpointing import load_modelopt_checkpoint + + teacher = get_attr_wrapped_model(model[0], 'teacher_model') + print_rank_0(f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ...") + # [WAR]: load checkpoint will check checkpoint's saved args and rng state if not finetune. + # To avoid error out on loading teacher's checkpoint, we temporarily set args.finetune to + # True while loading the teacher checkpoint. + original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format + args.finetune = True + if args.export_kd_teacher_ckpt_format is not None: + args.ckpt_format = args.export_kd_teacher_ckpt_format + load_modelopt_checkpoint([teacher], load_arg='export_kd_teacher_load') + args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format + print_rank_0("...teacher loaded successfully.") + return iteration, num_floating_point_operations_so_far From 8f1fa3243627c67d74ff99c033e20aadae6c34fa Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Mon, 20 Apr 2026 12:37:54 -0700 Subject: [PATCH 2/4] Prevent recursion Signed-off-by: Asha Anoosheh --- megatron/training/checkpointing.py | 30 +++++++++++++++++------------- 1 file changed, 17 insertions(+), 13 deletions(-) diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 28324e01ac3..8873beba26c 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -38,7 +38,7 @@ from megatron.core.num_microbatches_calculator import update_num_microbatches from megatron.core.optimizer import DistributedOptimizer from megatron.core.rerun_state_machine import get_rerun_state_machine -from megatron.core.utils import get_attr_wrapped_model, get_pg_rank, get_pg_size +from megatron.core.utils import get_pg_rank, get_pg_size from ..core.dist_checkpointing.utils import _clean_metadata_for_serialization from . import ft_integration, wandb_utils @@ -2062,21 +2062,25 @@ def load_model_state_dict(module, state_dict, strict: bool): if has_nvidia_modelopt: print_distributed_quant_summary(model, msg="After loading checkpoint") + + # Load teacher model in Distillation mode. if args.export_kd_teacher_load: from megatron.post_training.checkpointing import load_modelopt_checkpoint - teacher = get_attr_wrapped_model(model[0], 'teacher_model') - print_rank_0(f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ...") - # [WAR]: load checkpoint will check checkpoint's saved args and rng state if not finetune. - # To avoid error out on loading teacher's checkpoint, we temporarily set args.finetune to - # True while loading the teacher checkpoint. - original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format - args.finetune = True - if args.export_kd_teacher_ckpt_format is not None: - args.ckpt_format = args.export_kd_teacher_ckpt_format - load_modelopt_checkpoint([teacher], load_arg='export_kd_teacher_load') - args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format - print_rank_0("...teacher loaded successfully.") + unwrapped_model = unwrap_model(model)[0] + # Note: load_modelopt_checkpoint may call this function so we prevent infinite recursion. + if hasattr(unwrapped_model, 'teacher_model'): + teacher = unwrapped_model.teacher_model + print_rank_0(f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ...") + # [WAR]: To avoid error out on loading teacher's checkpoint, we temporarily + # set args.finetune to True while loading the teacher checkpoint. + original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format + args.finetune = True + if args.export_kd_teacher_ckpt_format is not None: + args.ckpt_format = args.export_kd_teacher_ckpt_format + load_modelopt_checkpoint([teacher], load_arg='export_kd_teacher_load') + args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format + print_rank_0("... teacher loaded successfully.") return iteration, num_floating_point_operations_so_far From 38a53301ff14ce081bf9ce4d35e10c9e7d784126 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Tue, 28 Apr 2026 21:14:48 +0200 Subject: [PATCH 3/4] Add comment Signed-off-by: Asha Anoosheh --- megatron/post_training/model_builder.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/megatron/post_training/model_builder.py b/megatron/post_training/model_builder.py index df7cf87a6d6..eb79e0756bf 100644 --- a/megatron/post_training/model_builder.py +++ b/megatron/post_training/model_builder.py @@ -143,6 +143,8 @@ def _build_teacher_model(config, config_raw: Namespace, model_kwargs: Dict[str, _add_load_convert_hooks(teacher) + # NOTE: Checkpoint loading now handled in `megatron/training/checkpointing.py`. + return teacher From a06147788a07105b20244e15c8dd16e6da471e26 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 30 Apr 2026 18:07:11 +0200 Subject: [PATCH 4/4] Safeguard args Signed-off-by: Asha Anoosheh --- megatron/training/checkpointing.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 8873beba26c..2daaff9ce3b 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -2064,7 +2064,7 @@ def load_model_state_dict(module, state_dict, strict: bool): print_distributed_quant_summary(model, msg="After loading checkpoint") # Load teacher model in Distillation mode. - if args.export_kd_teacher_load: + if getattr(args, "export_kd_teacher_load", None): from megatron.post_training.checkpointing import load_modelopt_checkpoint unwrapped_model = unwrap_model(model)[0]