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
52 changes: 49 additions & 3 deletions megatron/training/argument_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
StragglerDetectionConfig,
RerunStateMachineConfig, CheckpointConfig, ProfilingConfig
)
from megatron.training.models.hybrid import HybridModelConfig
from megatron.training.models import HybridModelConfig, GPTModelConfig
# TODO: support arg renames

class TypeInferenceError(Exception):
Expand Down Expand Up @@ -388,6 +388,50 @@ def _default_config_from_args(cls: type, args: Namespace, return_instance: bool
return kwargs


def gpt_config_from_args(args: Namespace, config: TransformerConfig | None=None) -> Any:
"""Create a GPTModelConfig from the appropriate values in the `args` Namespace."""

kwargs = {}
if config is None:
if args.yaml_cfg is not None:
from megatron.training.yaml_arguments import core_transformer_config_from_yaml

transformer_cfg = core_transformer_config_from_yaml(args, "language_model")
else:
transformer_cfg = core_transformer_config_from_args(args)
else:
transformer_cfg = config
kwargs["transformer"] = transformer_cfg

if args.spec is not None:
kwargs["transformer_layer_spec"] = import_module(args.spec)


kwargs["fp16_lm_cross_entropy"] = args.fp16_lm_cross_entropy
kwargs["position_embedding_type"] = args.position_embedding_type
kwargs["rotary_percent"] = args.rotary_percent
kwargs["rotary_base"] = args.rotary_base
kwargs["make_vocab_size_divisible_by"] = args.make_vocab_size_divisible_by
kwargs["rope_scaling"] = args.use_rope_scaling

kwargs["seq_len_interpolation_factor"] = args.rotary_seq_len_interpolation_factor
kwargs["seq_length"] = args.max_position_embeddings
kwargs["share_embeddings_and_output_weights"] = not args.untie_embeddings_and_output_weights

# GPTModelConfig supports either automatically padding vocab size or using exact provided
# vocab size via "should_pad_vocab" to support loading third-party checkpoints. Here,
# that is just mapped to settings in args appropriately.
if args.padded_vocab_size is not None:
kwargs["vocab_size"] = args.padded_vocab_size
kwargs["should_pad_vocab"] = False
else:
assert args.vocab_size is not None, "Either --padded-vocab-size or --vocab-size must be specified."
kwargs["vocab_size"] = args.vocab_size
kwargs["should_pad_vocab"] = True

return GPTModelConfig(**kwargs)


def hybrid_config_from_args(args: Namespace, config: TransformerConfig | None=None) -> Any:
"""Create a HybridModelConfig from the appropriate values in the `args` Namespace."""

Expand Down Expand Up @@ -417,11 +461,13 @@ def hybrid_config_from_args(args: Namespace, config: TransformerConfig | None=No
kwargs["seq_length"] = args.max_position_embeddings
kwargs["share_embeddings_and_output_weights"] = not args.untie_embeddings_and_output_weights

# HybridModelConfig supports either automatically padding vocab size or using exact provided
# vocab size via "should_pad_vocab" to support loading third-party checkpoints. Here,
# that is just mapped to settings in args appropriately.
if args.padded_vocab_size is not None:
kwargs["vocab_size"] = args.padded_vocab_size
kwargs["should_pad_vocab"] = False
else:
# Megatron-Bridge uses an explicit setting "should_pad_vocab" so that
# when converting model configs from HF, we can set a vocab size and disable padding.
assert args.vocab_size is not None, "Either --padded-vocab-size or --vocab-size must be specified."
kwargs["vocab_size"] = args.vocab_size
kwargs["should_pad_vocab"] = True
Expand Down
4 changes: 2 additions & 2 deletions megatron/training/config/container.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
)
from megatron.training.config.utils import sanitize_dataclass_config
from megatron.training.config.yaml_utils import safe_yaml_representers
from megatron.training.models import Serializable, HybridModelConfig
from megatron.training.models import GPTModelConfig, Serializable, HybridModelConfig

T = TypeVar("T", bound="ConfigContainerBase")

Expand Down Expand Up @@ -233,7 +233,7 @@ class PretrainConfigContainer(ConfigContainerBase):

train: TrainingConfig
validation: ValidationConfig = field(default_factory=ValidationConfig)
model: HybridModelConfig # TODO (@maanug): add support for GPTModelConfig
model: HybridModelConfig | GPTModelConfig
optimizer: OptimizerConfig
scheduler: SchedulerConfig
# dataset: GPTDatasetConfig # TODO (@maanug): add support
Expand Down
3 changes: 3 additions & 0 deletions megatron/training/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
unimodal_build_distributed_models,
)
from megatron.training.models.hybrid import HybridModelBuilder, HybridModelConfig
from megatron.training.models.gpt import GPTModelBuilder, GPTModelConfig

MambaModelConfig = HybridModelConfig
MambaModelBuilder = HybridModelBuilder
Expand All @@ -21,4 +22,6 @@
"HybridModelBuilder",
"MambaModelConfig",
"MambaModelBuilder",
"GPTModelConfig",
"GPTModelBuilder",
]
Loading
Loading