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"]:

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…