import json import argparse import json import math import os import socket import time from datetime import datetime import dotenv import lmdb import torch import torch.distributed as dist from utils.data_utils import LMDBIterableDataset, get_collated_batch_for_key from utils.evaluation import evaluate_model from utils.infrastructure_utils import custom_barrier_with_timeout from utils.model_utils import get_model_config from utils.training import train_model from utils.training_utils import * from utils.utils import get_variable_from_module, get_logger, convert_for_json print(f"PID: {os.getpid()} on host: {socket.gethostname()}") dotenv.load_dotenv() logger = None local_rank = None world_size = None training_complete_file_name = "training_completed.flag" run_complete_file_name = "run_completed.flag" def get_input_parameters(model_configuration: dict) -> dict: keys = [ "input_window_length", "output_window_length", "output_window_offset", "preprocessing", ] return {key: model_configuration[key] for key in keys if key in model_configuration} def prepare_run(results_dir: str, lmdb_root_dir: str, base_model_configuration: dict, base_training_configuration: dict): # create model configuration from base logger.info(f"Rank {local_rank}: Creating model configuration") model_configuration = get_model_config(base_model_configuration, results_dir, lmdb_root_dir) feature_config = model_configuration["feature_config"] dataset_dir = f"{lmdb_root_dir}/{feature_config['feature_set_name']}" logger.info(f"Rank {local_rank}: Model configuration name: {model_configuration['id']}") logger.info(f"Rank {local_rank}: Model configuration directory: {model_configuration['model_dir']}") # create training configuration logger.info(f"Rank {local_rank}: Creating training configuration") training_configuration = get_training_config(base_training_configuration, model_configuration) logger.info(f"Rank {local_rank}: Training configuration: {training_configuration['id']}") logger.info(f"Rank {local_rank}: Training directory: {training_configuration['training_dir']}") logger.info(f"Rank {local_rank}: Fetching data ids, limit: {item_limit}") train_ids, val_ids, test_ids = get_data_ids(model_configuration, training_configuration, dataset_dir, item_limit) logger.info(f"Rank {local_rank}: Train ids: {len(train_ids)}, Val ids: {len(val_ids)}, Test ids: {len(test_ids)}") return model_configuration, training_configuration, train_ids, val_ids, test_ids, dataset_dir def train( model_configuration: dict, training_configuration: dict, train_ids: list, val_ids: list, dataset_dir: str, log_dir: str, num_dataloader_workers: int = 1) -> None: # instantiate model logger.info(f"Rank {local_rank}: Creating model {model_configuration['model_name']}") logger.info(f"Rank {local_rank}: Creating model with parameters: {model_configuration['model_parameters']}") logger.info(f"Rank {local_rank}: Input_parameters: {get_input_parameters(model_configuration)}") model, training_configuration = load_model_for_usage( model_configuration, training_configuration, train_ids[0] ) batch_size = training_configuration["batch_size"] device = "cpu" if not torch.distributed.is_initialized() else f"cuda:{local_rank}" model.to(device) logger.info( f"Rank {local_rank}: Estimated batch size: {batch_size} for GPU {torch.cuda.get_device_name(local_rank)}") # create training and validation loaders logger.info(f"Rank {local_rank}: Creating training loaders") train_loader = LMDBIterableDataset(dataset_dir, train_ids, model_configuration=model_configuration, batch_size=training_configuration["batch_size"]) logger.info(f"Rank {local_rank}: Creating validation loaders") val_loader = LMDBIterableDataset(dataset_dir, val_ids, model_configuration=model_configuration, batch_size=training_configuration["batch_size"]) logger.info(f"Rank {local_rank}: Started training for model {model_configuration['id']}") train_model( model=model, model_configuration=model_configuration, training_configuration=training_configuration, train_dataset=train_loader, val_dataset=val_loader, log_dir=log_dir, logger=logger, num_dataloader_workers=num_dataloader_workers, ) def evaluate(model_configuration: dict, training_configuration: dict, evaluation_configuration: dict, test_ids: list, log_dir: str) -> dict: logger.info(f"Evaluating model {model_configuration['id']}") if "eval_functions_getter" in evaluation_configuration: eval_functions = evaluation_configuration["eval_functions_getter"](model_configuration) else: try: eval_functions = evaluation_configuration["eval_functions"] except KeyError as e: logger.error(f"Evaluation configuration does not contain 'eval_functions' or 'eval_functions_getter': {e}") raise results = evaluate_model(model_configuration, training_configuration, test_ids, eval_functions, log_dir=log_dir, logger=logger, print_progress=True) # save results to file results_dir = training_configuration["training_dir"] if not os.path.exists(results_dir): os.makedirs(results_dir) results_file = os.path.join(results_dir, "evaluation_results.json") with open(results_file, "w") as f: json.dump(convert_for_json(results), f, indent=4) return results if __name__ == "__main__": # set up distributed training dist.init_process_group(backend="nccl", init_method="env://") # sleeping for a bit to allow all processes to initialize time.sleep(5) local_rank = torch.distributed.get_rank() world_size = torch.distributed.get_world_size() torch.cuda.set_device(local_rank) print(f"Rank {local_rank} initialized with world size {world_size}") parser = argparse.ArgumentParser(description="Training wrapper for model training") parser.add_argument("run_configuration_module", type=str, help="Path to the run configuration module and config in module separated by / or .") parser.add_argument("--results_dir", type=str, required=False, default=None, help="Directory to save the results") parser.add_argument("--lmdb_root_dir", type=str, required=False, default=None, help="Path to the lmdb directory") parser.add_argument("--log_dir", type=str, required=False, default=None, help="Directory to save the logs") parser.add_argument("--item_limit", type=int, required=False, default=None, help="Limit the number of items to process, default is None (no limit)") parser.add_argument("--num_dataloader_workers", type=int, required=False, default=1, help="Number of workers for the dataloader, default is 1") parser.add_argument("--evaluate_only", action="store_true", help="If set to True, only evaluation will be performed, default is False") args = parser.parse_args() # get run config module and name from input input_raw = args.run_configuration_module.replace(".py", "").replace("/", ".") # last slash / dot separates module from run config in module run_configuration_module_name = ".".join(input_raw.split(".")[:-1]) run_config_name = input_raw.split(".")[-1] print(run_configuration_module_name) print(run_config_name) # load run configuration from module run_configuration = get_variable_from_module( # make sure to replace / with . and remove .py to get proper module tree module_path=run_configuration_module_name, variable_name=run_config_name) # check for variables in run_configuration if not args.results_dir: if "base_results_dir" not in run_configuration: results_dir = os.getenv("RESULTS_ROOT_DIR") else: results_dir = run_configuration["base_results_dir"] else: results_dir = args.results_dir if results_dir is None: raise ValueError( "No results directory specified. Please set the RESULTS_ROOT_DIR environment variable or provide a results_dir argument.") if not args.lmdb_root_dir: if "base_lmdb_root_dir" not in run_configuration: lmdb_root_dir = os.getenv("LMDB_ROOT_DIR") else: lmdb_root_dir = run_configuration["base_lmdb_root_dir"] else: lmdb_root_dir = args.lmdb_root_dir if lmdb_root_dir is None: raise ValueError( "No LMDB root directory specified. Please set the LMDB_ROOT_DIR environment variable or provide a lmdb_root_dir argument.") if not args.log_dir: if "base_log_dir" not in run_configuration: log_dir = os.getenv("LOG_DIR") else: log_dir = run_configuration["log_dir"] else: log_dir = args.log_dir if log_dir is None: raise ValueError( "No log directory specified. Please set the LOG_DIR environment variable or provide a log_dir argument.") try: item_limit = args.item_limit if args.item_limit else run_configuration["item_limit"] if item_limit == -1: item_limit = None except: item_limit = None run_name = run_configuration["name"] runs = run_configuration["runs"] evaluate_only = args.evaluate_only if evaluate_only: print(f"Rank {local_rank}: Running in evaluation only mode for run {run_name}") # set up run_id on rank 0 if local_rank == 0: run_id = f"{run_name}" if not os.path.exists(log_dir): os.makedirs(log_dir) if not os.path.exists(results_dir): os.makedirs(results_dir) else: run_id = None if torch.distributed.is_initialized(): torch.distributed.barrier() # broadcast run_id to all processes if dist.is_initialized(): print(f"Rank {local_rank}: Broadcasting run_id {run_id}") run_id_list = [run_id] torch.distributed.broadcast_object_list(run_id_list, src=0) run_id = run_id_list[0] # make sure run_id is a string run_id = str(run_id) print(f"Rank {local_rank}: fetched run_id {run_id}") # append run_id to results_dir and log_dir results_dir = os.path.join(results_dir, run_id) log_dir = os.path.join(log_dir, run_id) print(f"Rank {local_rank}: Results directory: {results_dir}") print(f"Rank {local_rank}: Log directory: {log_dir}") # create directories if they do not exist, only on the main process if torch.distributed.get_rank() == 0 or not dist.is_initialized(): print(f"Rank {local_rank}: Creating directories for run {run_name}") if not os.path.exists(results_dir): os.makedirs(results_dir) if not os.path.exists(log_dir): os.makedirs(log_dir) # sync all processes to make sure the directories are created if dist.is_initialized(): print(f"Rank {local_rank}: Waiting for all processes to create directories") dist.barrier() # set up logger for all processes logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, f"main_{local_rank}.log")) logger.info(f"Rank {local_rank}: Starting run {run_name}") logger.info(f"Rank {local_rank}: Run ID: {run_id}") logger.info(f"Rank {local_rank}: Using {args.num_dataloader_workers} dataloader workers") # print parameter values logger.info(f"Rank {local_rank}: Results directory: {results_dir}") logger.info(f"Rank {local_rank}: LMDB root directory: {lmdb_root_dir}") logger.info(f"Rank {local_rank}: Log directory: {log_dir}") logger.info(f"Rank {local_rank}: Item limit: {item_limit}") if not runs or len(runs) == 0: raise ValueError("No runs specified in the run configuration. Please provide a list of runs to train.") logger.info(f"Rank {local_rank}: Number of trainings: {len(runs)}") try: for run in runs: run_step_name = run["name"] run_description = run["description"] run_model_configuration = run["model_configuration"] run_training_configuration = run["training_configuration"] run_evaluation_configuration = run_configuration["evaluation_configuration"] logger.info(f"Rank {local_rank}: Running: {run_step_name}") logger.info(f"Rank {local_rank}: Description: {run_description}") # prepare run, make sure rank 0 is the first to avoid race conditions if local_rank == 0: logger.info(f"Rank {local_rank}: Preparing run {run_step_name}") run_model_configuration, run_training_configuration, train_ids, val_ids, test_ids, dataset_dir = prepare_run( results_dir=results_dir, lmdb_root_dir=lmdb_root_dir, base_model_configuration=run_model_configuration, base_training_configuration=run_training_configuration ) logger.info(f"Rank {local_rank}: Finished preparing run {run_step_name}") # all ranks sync here if dist.is_initialized(): if local_rank != 0: logger.info(f"Rank {local_rank}: Waiting for rank 0 to finish preparing run {run_step_name}") dist.barrier() if local_rank != 0: logger.info( f"Rank {local_rank}: Finished waiting for rank 0 to finish preparing run {run_step_name}") if local_rank != 0: # after the barrier, rank 0 will have prepared the run run_model_configuration, run_training_configuration, train_ids, val_ids, test_ids, dataset_dir = prepare_run( results_dir=results_dir, lmdb_root_dir=lmdb_root_dir, base_model_configuration=run_model_configuration, base_training_configuration=run_training_configuration ) # check, if run was already done training_dir = run_training_configuration["training_dir"] run_complete_file_path = os.path.join(training_dir, run_complete_file_name) if os.path.exists(run_complete_file_path) and not evaluate_only: logger.info(f"Rank {local_rank}: Training for {run_step_name} completed.") continue # check, if training was done training_complete_file_path = os.path.join(training_dir, training_complete_file_name) training_completed = os.path.exists(training_complete_file_path) if not training_completed and not evaluate_only: # check, if training was already done training_dir = run_training_configuration["training_dir"] if os.path.exists(training_dir): run_complete_file_path = os.path.join(training_dir, run_complete_file_name) if os.path.exists(run_complete_file_path): logger.info(f"Rank {local_rank}: Training for {run_step_name} already done, skipping.") continue else: logger.info(f"Rank {local_rank}: Training for {run_step_name} not done yet.") train( model_configuration=run_model_configuration, training_configuration=run_training_configuration, train_ids=train_ids, val_ids=val_ids, dataset_dir=dataset_dir, log_dir=log_dir, num_dataloader_workers=args.num_dataloader_workers, ) # sync after training if dist.is_initialized(): dist.barrier() logger.info(f"Rank {torch.distributed.get_rank()} finished training {run_step_name}") # create flag file to indicate that the training was completed if local_rank == 0: flag_file = os.path.join(training_dir, training_complete_file_name) with open(flag_file, "w") as f: f.write(f"Training for {run_step_name} completed at {datetime.now().isoformat()}\n") logger.info(f"Rank {local_rank}: Created flag file {flag_file} for run {run_step_name}") elif not training_completed and evaluate_only: logger.warning( f"Rank {local_rank}: Training for {run_step_name} not completed, but evaluation scheduled, skipping...") continue else: logger.info(f"Rank {local_rank}: Training for {run_step_name} already done.") # if no evaluation configuration is provided, skip evaluation if run_evaluation_configuration is not None: # run evaluation, only on rank 0 local_rank = torch.distributed.get_rank() if local_rank == 0: logger.info(f"Rank {local_rank} starting evaluation for {run_step_name}") results = evaluate( model_configuration=run_model_configuration, training_configuration=run_training_configuration, evaluation_configuration=run_evaluation_configuration, test_ids=test_ids, log_dir=log_dir, ) logger.info( f"Rank {local_rank} finished evaluation for {run_step_name} with model {run_model_configuration['id']}") logger.info(f"Results: {results}") else: logger.info( f"Rank {local_rank}: No evaluation configuration provided for {run_step_name}, skipping evaluation.") # create flag file to indicate that the run was completed if local_rank == 0: flag_file = os.path.join(training_dir, run_complete_file_name) with open(flag_file, "w") as f: f.write(f"Run {run_step_name} completed at {datetime.now().isoformat()}\n") logger.info(f"Rank {local_rank}: Created flag file {flag_file} for run {run_step_name}") # sync after evaluation if dist.is_initialized(): # make sure there is a longer timeout, since evaluation on rank 0 might take longer custom_barrier_with_timeout() finally: # clean up if dist.is_initialized(): dist.destroy_process_group()