Files
temperature-based-fertility…/code/training_wrapper.py
T
2025-09-10 10:37:55 +02:00

454 lines
19 KiB
Python

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()