From 7862f099df884b11d8c5001b93089c549e8e9626 Mon Sep 17 00:00:00 2001 From: Antoni-Joan Solergibert Date: Sun, 28 Jun 2026 09:56:27 +0200 Subject: [PATCH] Fix gl ci Signed-off-by: Antoni-Joan Solergibert --- gpt_builders.py | 28 --------------------- megatron/training/argument_utils.py | 39 ++++++++++++++++++++++++++++- 2 files changed, 38 insertions(+), 29 deletions(-) diff --git a/gpt_builders.py b/gpt_builders.py index 2f3a8c3aff7..3512918efe6 100644 --- a/gpt_builders.py +++ b/gpt_builders.py @@ -22,33 +22,6 @@ from megatron.training.yaml_arguments import core_transformer_config_from_yaml -def _apply_yarn_config_from_args(config, args) -> None: - """Populate YaRN fields on config from args when not already set. - - Preserves values already present on ``config`` (e.g. from YAML or a caller- - supplied config). YaRN-specific hyperparameters must be supplied via CLI - when ``position_embedding_type == 'yarn'`` (see functional test configs). - """ - if args.position_embedding_type != 'yarn': - return - - def _set_if_missing(attr: str, value) -> None: - if value is None: - return - if not hasattr(config, attr): - setattr(config, attr, value) - - _set_if_missing('yarn_rotary_scaling_factor', args.rotary_scaling_factor) - _set_if_missing( - 'yarn_original_max_position_embeddings', args.yarn_original_max_position_embeddings - ) - _set_if_missing('yarn_beta_fast', args.yarn_beta_fast) - _set_if_missing('yarn_beta_slow', args.yarn_beta_slow) - _set_if_missing('yarn_mscale', args.mscale) - _set_if_missing('yarn_mscale_all_dim', args.mscale_all_dim) - _set_if_missing('yarn_correction_range_round_to_int', args.yarn_correction_range_round_to_int) - - def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_collection=None): print_rank_0('building GPT model ...') if config is None: @@ -56,7 +29,6 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ config = core_transformer_config_from_yaml(args, "language_model") else: config = core_transformer_config_from_args(args) - _apply_yarn_config_from_args(config, args) if args.spec is not None: transformer_layer_spec = import_module(args.spec) else: diff --git a/megatron/training/argument_utils.py b/megatron/training/argument_utils.py index abe437e2ee7..d8a757ddfc6 100644 --- a/megatron/training/argument_utils.py +++ b/megatron/training/argument_utils.py @@ -369,8 +369,45 @@ def core_transformer_config_from_args(args, config_class=None): if hasattr(args, "kitchen_attention_backend"): kw_args['kitchen_attention_backend'] = args.kitchen_attention_backend + # Build config. + config = config_class(**kw_args) + + _apply_yarn_config_from_args(config, args) + # Return config. - return config_class(**kw_args) + return config + + +def _apply_yarn_config_from_args(config, args) -> None: + """Populate ``config.yarn_*`` attributes from args for non-MLA YaRN models. + + GPTModel's ``yarn`` branch and ``yarn_rotary_pos_embedding`` read these as + dynamic attributes off the config (``getattr(config, "yarn_rotary_scaling_factor")`` + etc.) with no default, so the attributes must exist whenever + ``position_embedding_type == 'yarn'``. The CLI exposes some of these without a + ``yarn_`` prefix (``--rotary-scaling-factor``, ``--mscale``, ``--mscale-all-dim``), + so the mapping is explicit. Pre-existing values on ``config`` (e.g. from YAML or a + ModelOpt GPT-OSS builder) are preserved. Defaults mirror ``YarnRotaryEmbedding``. + """ + if getattr(args, 'position_embedding_type', None) != 'yarn': + return + if getattr(args, 'multi_latent_attention', False): + # MLATransformerConfig declares the unprefixed YaRN fields and its + # attention path consumes them directly; do not shadow them here. + return + + def _set(attr: str, value, default) -> None: + if hasattr(config, attr): + return + setattr(config, attr, value if value is not None else default) + + _set('yarn_rotary_scaling_factor', args.rotary_scaling_factor, 1.0) + _set('yarn_original_max_position_embeddings', args.yarn_original_max_position_embeddings, 4096) + _set('yarn_beta_fast', args.yarn_beta_fast, 32.0) + _set('yarn_beta_slow', args.yarn_beta_slow, 1.0) + _set('yarn_mscale', args.mscale, 1.0) + _set('yarn_mscale_all_dim', args.mscale_all_dim, 0.0) + _set('yarn_correction_range_round_to_int', args.yarn_correction_range_round_to_int, True) def _default_config_from_args(cls: type, args: Namespace, return_instance: bool = True) -> Any: