Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 0 additions & 28 deletions gpt_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,41 +22,13 @@
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:
if args.yaml_cfg is not None:
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:
Expand Down
39 changes: 38 additions & 1 deletion megatron/training/argument_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading