INNER CODE UNIT · Python
load_megatron_dataset
Guitaricet/relora · torchrun_main.py:276
def load_megatron_dataset(args, world_size, start_iteration):
logger.info(f"Loading Megatron dataset arguments from {args.megatron_dataset_config}")
with open(args.megatron_dataset_config) as f:
dataset_config_yaml = yaml.safe_load(f)
dataset_config_yaml["global_num_gpus"] = world_size
dataset_config_yaml["train_micro_batch_size_per_gpu"] = args.batch_size
dataset_config_yaml["gradient_accumulation_steps"] = args.gradient_accumulation
dataset_config_yaml["train_batch_size"] = args.total_batch_size
dataset_config_yaml["num_workers"] = args.workers
if args.max_length != dataset_config_yaml["seq_length"]:
logger.warning(f"rags.max_length ({args.max_length}) does not match "
f"seq_length ({dataset_config_yaml['seq_length']}) in the dataset config")
logger.warning(f"Overwriting max_length with seq_length")
args.max_length = dataset_config_yaml["seq_length"]
if args.num_training_steps > dataset_config_yaml["train_iters"]: