added code
This commit is contained in:
+164
-60
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user