added code

This commit is contained in:
2025-09-10 10:37:55 +02:00
parent 36901c736d
commit c78a68de80
199 changed files with 3561 additions and 22579 deletions
+164 -60
View File
@@ -1,25 +1,24 @@
import json
import socket
import sys
import os
import argparse
import logging
import json
import math
import os
import socket
import time
from datetime import datetime
import lmdb
import dotenv
import lmdb
import torch
from torch.utils.data import DataLoader
import torch.distributed as dist
from experiment_setup import get_eval_functions
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_utils import get_training_config, get_data_ids
from utils.data_utils import LMDBIterableDataset
from utils.utils import get_variable_from_module, get_logger, convert_for_json
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()}")
@@ -27,27 +26,45 @@ 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("Creating model configuration")
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"Model configuration name: {model_configuration['id']}")
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("Creating training configuration")
logger.info(f"Rank {local_rank}: Creating training configuration")
training_configuration = get_training_config(base_training_configuration, model_configuration)
logger.info(f"Training configuration: {training_configuration['id']}")
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"Fetching data ids, limit: {item_limit}")
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"Train ids: {len(train_ids)}, Val ids: {len(val_ids)}, Test ids: {len(test_ids)}")
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
@@ -60,27 +77,35 @@ def train(
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("Creating training 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("Creating validation loaders")
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"])
# instantiate model
logger.info("Creating model")
lmdb_env = lmdb.open(dataset_dir, readonly=True)
model_creation_fn = model_configuration["model_creation_fn"]
model = model_creation_fn(model_configuration,
train_ids[0],
lmdb_env=lmdb_env)
logger.info(f"Started training for model {model_configuration['id']}")
logger.info(f"Rank {local_rank}: Started training for model {model_configuration['id']}")
train_model(
model=model,
model_configuration=model_configuration,
@@ -95,13 +120,25 @@ def train(
def evaluate(model_configuration: dict,
training_configuration: dict,
test_ids: list) -> dict:
evaluation_configuration: dict,
test_ids: list,
log_dir: str) -> dict:
logger.info(f"Evaluating model {model_configuration['id']}")
eval_functions = get_eval_functions(model_configuration)
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)
eval_functions,
log_dir=log_dir,
logger=logger,
print_progress=True)
# save results to file
results_dir = training_configuration["training_dir"]
@@ -129,12 +166,7 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Training wrapper for model training")
parser.add_argument("run_configuration_module",
type=str,
help="Path to the run configuration module")
parser.add_argument("--run_configuration_variable",
type=str,
required=False,
default="run_configuration",
help="Name of the run configuration variable in the module")
help="Path to the run configuration module and config in module separated by / or .")
parser.add_argument("--results_dir",
type=str,
required=False,
@@ -161,13 +193,25 @@ if __name__ == "__main__":
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=args.run_configuration_module.replace(".py", "").replace("/", "."),
variable_name=args.run_configuration_variable)
module_path=run_configuration_module_name,
variable_name=run_config_name)
# check for variables in run_configuration
if not args.results_dir:
@@ -215,9 +259,13 @@ if __name__ == "__main__":
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}_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
run_id = f"{run_name}"
if not os.path.exists(log_dir):
os.makedirs(log_dir)
if not os.path.exists(results_dir):
@@ -264,6 +312,7 @@ if __name__ == "__main__":
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}")
@@ -274,12 +323,15 @@ if __name__ == "__main__":
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}")
@@ -313,36 +365,88 @@ if __name__ == "__main__":
base_training_configuration=run_training_configuration
)
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,
)
# sync after training
if dist.is_initialized():
dist.barrier()
# 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
logger.info(f"Rank {torch.distributed.get_rank()} finished training {run_step_name}")
# 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)
# 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(
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,
test_ids=test_ids,
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} finished evaluation for {run_step_name} with model {run_model_configuration['id']}")
logger.info(f"Results: {results}")
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():
dist.barrier()
# 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():