34 lines
1.3 KiB
Python
34 lines
1.3 KiB
Python
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
|
|
import logging
|
|
|
|
from megatron.training.arguments import parse_args, validate_args
|
|
from megatron.training.tokenizer.tokenizer import _vocab_size_with_padding
|
|
|
|
__all__ = ["validate_args", "parse_args", "set_default_megatron_args"]
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def set_default_megatron_args(args):
|
|
# always use zero optimizer
|
|
args.use_distributed_optimizer = True
|
|
# TODO: maybe change this after megatron has good fp8 support
|
|
args.bf16 = not args.fp16
|
|
# placeholders
|
|
args.seq_length = 4096
|
|
args.max_position_embeddings = args.seq_length
|
|
# compatible for megatron
|
|
if hasattr(args, "rope_type") and args.rope_type is None:
|
|
args.rope_type = "yarn" if args.multi_latent_attention else "rope"
|
|
|
|
if args.vocab_size and not args.padded_vocab_size:
|
|
args.padded_vocab_size = _vocab_size_with_padding(args.vocab_size, args)
|
|
|
|
if not args.tokenizer_model and not args.tokenizer_type:
|
|
logger.info("--tokenizer-model not set, use --hf-checkpoint as tokenizer model.")
|
|
args.tokenizer_model = args.hf_checkpoint
|
|
args.tokenizer_type = "HuggingFaceTokenizer"
|
|
return args
|