fixes
This commit is contained in:
@@ -0,0 +1,321 @@
|
||||
import json
|
||||
import sys
|
||||
import os
|
||||
import argparse
|
||||
import logging
|
||||
from datetime import datetime
|
||||
|
||||
import lmdb
|
||||
import dotenv
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
import torch.distributed as dist
|
||||
|
||||
from experiment_setup import get_eval_functions
|
||||
from utils.evaluation import evaluate_model
|
||||
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
|
||||
|
||||
dotenv.load_dotenv()
|
||||
|
||||
logger = None
|
||||
|
||||
# set up distributed training
|
||||
dist.init_process_group(backend="nccl", init_method="env://")
|
||||
local_rank = torch.distributed.get_rank()
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
|
||||
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")
|
||||
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']}")
|
||||
|
||||
# create training configuration
|
||||
logger.info("Creating training configuration")
|
||||
training_configuration = get_training_config(base_training_configuration, model_configuration)
|
||||
logger.info(f"Training configuration: {training_configuration['id']}")
|
||||
|
||||
logger.info(f"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)}")
|
||||
|
||||
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) -> None:
|
||||
# create training and validation loaders
|
||||
logger.info("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")
|
||||
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']}")
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
def evaluate(model_configuration: dict,
|
||||
training_configuration: dict,
|
||||
test_ids: list,
|
||||
device) -> dict:
|
||||
logger.info(f"Evaluating model {model_configuration['id']}")
|
||||
eval_functions = get_eval_functions(model_configuration)
|
||||
results = evaluate_model(model_configuration,
|
||||
training_configuration,
|
||||
test_ids,
|
||||
eval_functions)
|
||||
|
||||
# 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__":
|
||||
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")
|
||||
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("--device",
|
||||
type=str,
|
||||
required=False,
|
||||
default="cuda",
|
||||
help="Device to use for training, default is cuda")
|
||||
|
||||
args = parser.parse_args()
|
||||
# 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)
|
||||
|
||||
# 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 os.path.exists(results_dir):
|
||||
os.makedirs(results_dir)
|
||||
|
||||
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 os.path.exists(lmdb_root_dir):
|
||||
os.makedirs(lmdb_root_dir)
|
||||
|
||||
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.")
|
||||
if not os.path.exists(log_dir):
|
||||
os.makedirs(log_dir)
|
||||
|
||||
item_limit = args.item_limit if args.item_limit else run_configuration["item_limit"]
|
||||
if item_limit == -1:
|
||||
item_limit = None
|
||||
|
||||
run_name = run_configuration["name"]
|
||||
runs = run_configuration["runs"]
|
||||
|
||||
# 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')}"
|
||||
else:
|
||||
run_id = None
|
||||
|
||||
# broadcast run_id to all processes
|
||||
if dist.is_initialized():
|
||||
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)
|
||||
|
||||
|
||||
# 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)
|
||||
|
||||
# create directories if they do not exist, only on the main process
|
||||
if torch.distributed.get_rank() == 0 or not dist.is_initialized():
|
||||
if not os.path.exists(results_dir):
|
||||
os.makedirs(results_dir)
|
||||
if not os.path.exists(log_dir):
|
||||
os.makedirs(log_dir)
|
||||
|
||||
# set up logger
|
||||
logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, "main.log"))
|
||||
|
||||
logger.info(f"Rank {local_rank}: Starting run {run_name}")
|
||||
logger.info(f"Rank {local_rank}: Run ID: {run_id}")
|
||||
|
||||
# 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.")
|
||||
|
||||
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"]
|
||||
|
||||
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}")
|
||||
# sync after preparing run
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
else:
|
||||
# wait for rank 0 to finish preparing run
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
logger.info(f"Rank {local_rank}: Waiting for rank 0 to finish 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
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
logger.info(f"Rank {torch.distributed.get_rank()} finished training {run_step_name}")
|
||||
|
||||
# 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,
|
||||
test_ids=test_ids,
|
||||
)
|
||||
logger.info(
|
||||
f"Rank {local_rank} finished evaluation for {run_step_name} with model {run_model_configuration['id']}")
|
||||
logger.info(f"Results: {results}")
|
||||
|
||||
# sync after evaluation
|
||||
if dist.is_initialized():
|
||||
dist.barrier()
|
||||
finally:
|
||||
# clean up
|
||||
if dist.is_initialized():
|
||||
dist.destroy_process_group()
|
||||
Reference in New Issue
Block a user