code update

This commit is contained in:
Alex Blank
2025-05-19 22:54:53 +02:00
parent 58305effdb
commit cb7896900b
7 changed files with 238 additions and 51 deletions
+1 -1
View File
@@ -43,7 +43,7 @@ run_configuration = {
"take_every_nth": take_every_nth, "take_every_nth": take_every_nth,
"shift_in_hours": shift_in_hours, "shift_in_hours": shift_in_hours,
"output_window_offset": output_window_offset, "output_window_offset": output_window_offset,
"batch_size": 256, "batch_size": 128,
"max_lr": 1e-5, "max_lr": 1e-5,
"num_epochs": 10, "num_epochs": 10,
"patience": 3, "patience": 3,
+5
View File
@@ -55,6 +55,11 @@ def simple_x_y_collate(batch):
collated_x.append(np.concatenate(current_x, axis=1)) collated_x.append(np.concatenate(current_x, axis=1))
collated_y.append(np.concatenate(current_y, axis=1)) collated_y.append(np.concatenate(current_y, axis=1))
# convert to numpy arrays
collated_x = np.array(collated_x)
collated_y = np.array(collated_y)
# convert to torch tensors
return torch.tensor(collated_x, dtype=torch.float32), torch.tensor(collated_y, dtype=torch.float32) return torch.tensor(collated_x, dtype=torch.float32), torch.tensor(collated_y, dtype=torch.float32)
+58
View File
@@ -0,0 +1,58 @@
#!/bin/bash
# Parameter check
if [[ $# -ne 5 ]]; then
echo "Usage: $0 <run_config_module> <partition> <gpu_type> <num_gpus> <num_cpus_per_gpu>"
exit 1
fi
# Assign parameters
run_config=$1
partition=$2
gpu_type=$3
num_gpus=$4
num_cpus_per_gpu=$5
# Derive run name
run_name=$(basename "$run_config")
run_name="${run_name%.*}"
# Confirm inputs
echo "Starting run with configuration: $run_config"
echo "Partition: $partition"
echo "GPU type: $gpu_type"
echo "Number of GPUs: $num_gpus"
echo "Number of CPUs per GPU: $num_cpus_per_gpu"
# Set paths
current_dir_path=$(dirname "$(realpath "$0")")
venv_path=$(realpath "$current_dir_path/../../venv/bin/activate")
# Create log directory
log_dir="/work/rr41qemu-MA/logs/${run_name}_$(date +%Y-%m-%d_%H-%M-%S)"
mkdir -p "$log_dir"
echo "Logging to $log_dir"
# Submit job
sbatch <<EOF
#!/bin/bash
#SBATCH --job-name=$run_name
#SBATCH --output=$log_dir/log.out
#SBATCH --error=$log_dir/log.err
#SBATCH --time=48:00:00
#SBATCH --ntasks=1
#SBATCH --cpus-per-task=$(($num_gpus * $num_cpus_per_gpu))
#SBATCH --mem=32G
#SBATCH --partition=$partition
#SBATCH --gpus=$gpu_type:$num_gpus
echo "Loading python virtual environment..."
source $venv_path
echo "Loading python 3.10..."
module load Python/3.10.4-GCCcore-11.3.0
cd $current_dir_path
torchrun --nproc_per_node=$num_gpus training_wrapper.py $run_config --item_limit -1
EOF
+57 -25
View File
@@ -1,8 +1,10 @@
import json import json
import socket
import sys import sys
import os import os
import argparse import argparse
import logging import logging
import time
from datetime import datetime from datetime import datetime
import lmdb import lmdb
@@ -19,15 +21,12 @@ from utils.data_utils import LMDBIterableDataset
from utils.utils import get_variable_from_module, get_logger, convert_for_json from utils.utils import get_variable_from_module, get_logger, convert_for_json
from utils.training import train_model from utils.training import train_model
print(f"PID: {os.getpid()} on host: {socket.gethostname()}")
dotenv.load_dotenv() dotenv.load_dotenv()
logger = None 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, def prepare_run(results_dir: str,
lmdb_root_dir: str, lmdb_root_dir: str,
@@ -59,7 +58,8 @@ def train(
train_ids: list, train_ids: list,
val_ids: list, val_ids: list,
dataset_dir: str, dataset_dir: str,
log_dir: str) -> None: log_dir: str,
num_dataloader_workers: int = 1) -> None:
# create training and validation loaders # create training and validation loaders
logger.info("Creating training loaders") logger.info("Creating training loaders")
train_loader = LMDBIterableDataset(dataset_dir, train_loader = LMDBIterableDataset(dataset_dir,
@@ -88,14 +88,14 @@ def train(
train_dataset=train_loader, train_dataset=train_loader,
val_dataset=val_loader, val_dataset=val_loader,
log_dir=log_dir, log_dir=log_dir,
logger=logger logger=logger,
num_dataloader_workers=num_dataloader_workers,
) )
def evaluate(model_configuration: dict, def evaluate(model_configuration: dict,
training_configuration: dict, training_configuration: dict,
test_ids: list, test_ids: list) -> dict:
device) -> dict:
logger.info(f"Evaluating model {model_configuration['id']}") logger.info(f"Evaluating model {model_configuration['id']}")
eval_functions = get_eval_functions(model_configuration) eval_functions = get_eval_functions(model_configuration)
results = evaluate_model(model_configuration, results = evaluate_model(model_configuration,
@@ -114,6 +114,18 @@ def evaluate(model_configuration: dict,
if __name__ == "__main__": 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 = argparse.ArgumentParser(description="Training wrapper for model training")
parser.add_argument("run_configuration_module", parser.add_argument("run_configuration_module",
type=str, type=str,
@@ -144,11 +156,11 @@ if __name__ == "__main__":
required=False, required=False,
default=None, default=None,
help="Limit the number of items to process, default is None (no limit)") help="Limit the number of items to process, default is None (no limit)")
parser.add_argument("--device", parser.add_argument("--num_dataloader_workers",
type=str, type=int,
required=False, required=False,
default="cuda", default=1,
help="Device to use for training, default is cuda") help="Number of workers for the dataloader, default is 1")
args = parser.parse_args() args = parser.parse_args()
# load run configuration from module # load run configuration from module
@@ -201,8 +213,11 @@ if __name__ == "__main__":
if not os.path.exists(log_dir): if not os.path.exists(log_dir):
os.makedirs(log_dir) os.makedirs(log_dir)
item_limit = args.item_limit if args.item_limit else run_configuration["item_limit"] try:
if item_limit == -1: 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 item_limit = None
run_name = run_configuration["name"] run_name = run_configuration["name"]
@@ -214,29 +229,42 @@ if __name__ == "__main__":
else: else:
run_id = None run_id = None
if torch.distributed.is_initialized():
torch.distributed.barrier()
# broadcast run_id to all processes # broadcast run_id to all processes
if dist.is_initialized(): if dist.is_initialized():
print(f"Rank {local_rank}: Broadcasting run_id {run_id}")
run_id_list = [run_id] run_id_list = [run_id]
torch.distributed.broadcast_object_list(run_id_list, src=0) torch.distributed.broadcast_object_list(run_id_list, src=0)
run_id = run_id_list[0] run_id = run_id_list[0]
# make sure run_id is a string # make sure run_id is a string
run_id = str(run_id) run_id = str(run_id)
print(f"Rank {local_rank}: fetched run_id {run_id}")
# append run_id to results_dir and log_dir # append run_id to results_dir and log_dir
results_dir = os.path.join(results_dir, run_id) results_dir = os.path.join(results_dir, run_id)
log_dir = os.path.join(log_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 # create directories if they do not exist, only on the main process
if torch.distributed.get_rank() == 0 or not dist.is_initialized(): 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): if not os.path.exists(results_dir):
os.makedirs(results_dir) os.makedirs(results_dir)
if not os.path.exists(log_dir): if not os.path.exists(log_dir):
os.makedirs(log_dir) os.makedirs(log_dir)
# set up logger # sync all processes to make sure the directories are created
logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, "main.log")) 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}: Starting run {run_name}")
logger.info(f"Rank {local_rank}: Run ID: {run_id}") logger.info(f"Rank {local_rank}: Run ID: {run_id}")
@@ -270,14 +298,18 @@ if __name__ == "__main__":
base_training_configuration=run_training_configuration base_training_configuration=run_training_configuration
) )
logger.info(f"Rank {local_rank}: Finished preparing run {run_step_name}") logger.info(f"Rank {local_rank}: Finished preparing run {run_step_name}")
# sync after preparing run
if dist.is_initialized(): # all ranks sync here
dist.barrier() if dist.is_initialized():
else: if local_rank != 0:
# 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}") 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( run_model_configuration, run_training_configuration, train_ids, val_ids, test_ids, dataset_dir = prepare_run(
results_dir=results_dir, results_dir=results_dir,
lmdb_root_dir=lmdb_root_dir, lmdb_root_dir=lmdb_root_dir,
+8 -3
View File
@@ -479,6 +479,10 @@ class LMDBIterableDataset(IterableDataset):
self.key_subset = key_subset self.key_subset = key_subset
# reset length # reset length
self.len = None self.len = None
# reset lmdb env
if self.lmdb_env is not None:
self.lmdb_env.close()
self.lmdb_env = None
def get_length_of_data_subset(self, key_set: list[str]): def get_length_of_data_subset(self, key_set: list[str]):
""" """
@@ -522,12 +526,13 @@ class LMDBIterableDataset(IterableDataset):
return num_steps return num_steps
def __iter__(self): def __iter__(self):
self.init_lmdb_env()
# if no key subset is set, use all keys # if no key subset is set, use all keys
if self.key_subset is None: if self.key_subset is None:
self.key_subset = self.lmdb_keys self.key_subset = self.lmdb_keys
self.init_lmdb_env()
random.shuffle(self.key_subset) random.shuffle(self.key_subset)
batch = list() batch = list()
@@ -563,12 +568,12 @@ class LMDBIterableDataset(IterableDataset):
meminit=False) meminit=False)
def __len__(self): def __len__(self):
self.init_lmdb_env()
# if no key subset is set, use all keys # if no key subset is set, use all keys
if self.key_subset is None: if self.key_subset is None:
self.key_subset = self.lmdb_keys self.key_subset = self.lmdb_keys
self.init_lmdb_env()
if self.len is None: if self.len is None:
num_steps = self.get_length_of_data_subset(self.key_subset) num_steps = self.get_length_of_data_subset(self.key_subset)
self.len = num_steps self.len = num_steps
+2
View File
@@ -130,6 +130,8 @@ def evaluate_model(model_configuration: dict,
if eval_fn is not None: if eval_fn is not None:
eval_fn_name = eval_fn["name"] eval_fn_name = eval_fn["name"]
accumulation_fn = eval_fn["accumulation_fn"] accumulation_fn = eval_fn["accumulation_fn"]
if eval_fn_name not in errors:
continue
for key in errors[eval_fn_name]: for key in errors[eval_fn_name]:
if len(errors[eval_fn_name][key]) == 0: if len(errors[eval_fn_name][key]) == 0:
errors[eval_fn_name][key] = np.nan errors[eval_fn_name][key] = np.nan
+107 -22
View File
@@ -104,7 +104,8 @@ def train_model(model: nn.Module,
train_dataset: LMDBIterableDataset, train_dataset: LMDBIterableDataset,
val_dataset: LMDBIterableDataset, val_dataset: LMDBIterableDataset,
log_dir: str = "./logs", log_dir: str = "./logs",
logger=None) -> torch.nn.Module: num_dataloader_workers: int = 4,
logger=None) -> None:
if logger is None: if logger is None:
logger = get_logger(__name__, f"{log_dir}/{model_configuration['id']}_{training_configuration['id']}.log") logger = get_logger(__name__, f"{log_dir}/{model_configuration['id']}_{training_configuration['id']}.log")
@@ -117,7 +118,6 @@ def train_model(model: nn.Module,
# get computation rank # get computation rank
if torch.distributed.is_initialized(): if torch.distributed.is_initialized():
local_rank = torch.distributed.get_rank() local_rank = torch.distributed.get_rank()
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}") device = torch.device(f"cuda:{local_rank}")
world_size = torch.distributed.get_world_size() world_size = torch.distributed.get_world_size()
else: else:
@@ -132,10 +132,46 @@ def train_model(model: nn.Module,
all_train_ids = train_dataset.lmdb_keys all_train_ids = train_dataset.lmdb_keys
train_subsets = [get_ranked_ids(all_train_ids, i, local_rank, world_size) for i in range(num_epochs)] train_subsets = [get_ranked_ids(all_train_ids, i, local_rank, world_size) for i in range(num_epochs)]
# calc total number of steps for gpu, as it is dependent on subsets # calc total number of steps for gpu, as it is dependent on subsets
total_train_steps = sum([train_dataset.get_length_of_data_subset(subset) for subset in train_subsets]) local_steps = [train_dataset.get_length_of_data_subset(subset) for subset in train_subsets]
total_train_steps = sum(local_steps)
all_val_ids = val_dataset.lmdb_keys all_val_ids = val_dataset.lmdb_keys
val_subsets = [get_ranked_ids(all_val_ids, i, local_rank, world_size) for i in range(num_epochs)] val_subsets = [get_ranked_ids(all_val_ids, i, local_rank, world_size) for i in range(num_epochs)]
val_subset_lengths = [len(subset) for subset in val_subsets]
def get_synced_values(local_values):
"""
Sync values across all processes.
Args:
local_values: list of local values
Returns:
list of synced values
"""
if not isinstance(local_values, torch.Tensor):
local_values = torch.tensor(local_values, device=device, dtype=torch.float32)
else:
local_values = local_values.to(device)
gathered_values = [torch.zeros_like(local_values) for _ in range(world_size)]
torch.distributed.all_gather(gathered_values, local_values)
return gathered_values
# sync lengths of subsets and adjust to minimum length for equal sized training lengths
global_lengths = get_synced_values(local_steps)
global_lengths = [x.cpu().numpy() for x in global_lengths]
train_epoch_lengths = list()
for i in range(len(global_lengths[0])):
current_epoch_lengths = [x[i] for x in global_lengths]
train_epoch_lengths.append(int(min(current_epoch_lengths)))
print(f"Rank {local_rank}: Global lengths: {global_lengths} cut to {train_epoch_lengths}")
# also sync the val subsets
global_val_lengths = get_synced_values(val_subset_lengths)
global_val_lengths = [x.cpu().numpy() for x in global_val_lengths]
val_epoch_lengths = list()
for i in range(len(global_val_lengths[0])):
current_epoch_lengths = [x[i] for x in global_val_lengths]
val_epoch_lengths.append(int(min(current_epoch_lengths)))
print(f"Rank {local_rank}: Global val lengths: {global_val_lengths} cut to {val_epoch_lengths}")
else: else:
train_subsets = [train_dataset.lmdb_keys] * num_epochs train_subsets = [train_dataset.lmdb_keys] * num_epochs
val_subsets = [val_dataset.lmdb_keys] * num_epochs val_subsets = [val_dataset.lmdb_keys] * num_epochs
@@ -156,12 +192,12 @@ def train_model(model: nn.Module,
train_dataloader = DataLoader( train_dataloader = DataLoader(
train_dataset, train_dataset,
batch_size=None, batch_size=None,
num_workers=4, num_workers=num_dataloader_workers,
) )
val_dataloader = DataLoader( val_dataloader = DataLoader(
val_dataset, val_dataset,
batch_size=None, batch_size=None,
num_workers=4, num_workers=num_dataloader_workers,
) )
logger.info(f"Rank {local_rank}: Training {training_id} with {num_epochs} epochs") logger.info(f"Rank {local_rank}: Training {training_id} with {num_epochs} epochs")
@@ -171,6 +207,12 @@ def train_model(model: nn.Module,
model.to(device) model.to(device)
# wrap model in DDP if distributed training
if torch.distributed.is_initialized():
model = torch.nn.parallel.DistributedDataParallel(model,
device_ids=[local_rank],
output_device=local_rank)
# load training state from training configuration, if available # load training state from training configuration, if available
optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"]) optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"])
scheduler = OneCycleLR(optimizer, scheduler = OneCycleLR(optimizer,
@@ -204,21 +246,37 @@ def train_model(model: nn.Module,
train_dataloader = DataLoader( train_dataloader = DataLoader(
train_dataset, train_dataset,
batch_size=None, batch_size=None,
num_workers=4, num_workers=num_dataloader_workers,
# set multiprocessing start method to spawn
# multiprocessing_context="forkserver",
) )
train_length = train_epoch_lengths[epoch - 1]
val_dataloader = DataLoader( val_dataloader = DataLoader(
val_dataset, val_dataset,
batch_size=None, batch_size=None,
num_workers=4, num_workers=num_dataloader_workers,
# set multiprocessing start method to spawn
# multiprocessing_context="forkserver",
) )
val_length = val_epoch_lengths[epoch - 1]
iterator = iter(train_dataloader) iterator = iter(train_dataloader)
for step in tqdm(range(len(train_dataloader)), total=len(train_dataloader)): for step in tqdm(range(train_length)):
loss = batch_loss_fn(model, try:
iterator, loss = batch_loss_fn(model,
loss_functions, iterator,
device, loss_functions,
model_configuration) device,
model_configuration)
except StopIteration:
# if the iterator is exhausted, reset it
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
iterator = iter(train_dataloader)
loss = batch_loss_fn(model,
iterator,
loss_functions,
device,
model_configuration)
optimizer.zero_grad() optimizer.zero_grad()
loss.backward() loss.backward()
@@ -237,20 +295,35 @@ def train_model(model: nn.Module,
writer.flush() writer.flush()
avg_train_loss = total_train_loss / len(train_dataloader) avg_train_loss = total_train_loss / len(train_dataloader)
logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Train loss: {avg_train_loss:.4f}") logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Local Train loss: {avg_train_loss:.4f}")
# sync before validation
if torch.distributed.is_initialized():
torch.distributed.barrier()
# Validation # Validation
# if local_rank == 0 or not torch.distributed.is_initialized():
model.eval() model.eval()
total_val_loss = 0 total_val_loss = 0
logger.info(f"Rank {local_rank}: Validation") logger.info(f"Rank {local_rank}: Validation")
with torch.no_grad(): with torch.no_grad():
val_iter = iter(val_dataloader) val_iter = iter(val_dataloader)
for step in tqdm(range(len(val_dataloader))): for step in tqdm(range(val_length)):
loss = batch_loss_fn(model, try:
val_iter, loss = batch_loss_fn(model,
loss_functions, val_iter,
device, loss_functions,
model_configuration) device,
model_configuration)
except StopIteration:
# if the iterator is exhausted, reset it
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
val_iter = iter(val_dataloader)
loss = batch_loss_fn(model,
val_iter,
loss_functions,
device,
model_configuration)
total_val_loss += loss.item() total_val_loss += loss.item()
@@ -259,11 +332,20 @@ def train_model(model: nn.Module,
avg_val_loss = None avg_val_loss = None
else: else:
avg_val_loss = total_val_loss / len(val_dataloader) avg_val_loss = total_val_loss / len(val_dataloader)
# else:
# avg_val_loss = None
# logger.info(f"Rank {local_rank}: Validation skipped, using 0 as validation loss.")
# sync before logging
if torch.distributed.is_initialized():
torch.distributed.barrier()
# publish validation loss and wait for other gpus # publish validation loss and wait for other gpus
if torch.distributed.is_initialized(): if torch.distributed.is_initialized():
if avg_val_loss is not None: if avg_val_loss is not None:
avg_val_loss_global = torch.tensor(avg_val_loss, device=device, dtype=torch.float32) avg_val_loss_global = torch.tensor(avg_val_loss, device=device, dtype=torch.float32)
else:
avg_val_loss_global = torch.tensor(0.0, device=device, dtype=torch.float32)
torch.distributed.all_reduce(avg_val_loss_global) torch.distributed.all_reduce(avg_val_loss_global)
avg_val_loss_global /= torch.distributed.get_world_size() avg_val_loss_global /= torch.distributed.get_world_size()
@@ -276,7 +358,7 @@ def train_model(model: nn.Module,
# only rank 0 checks for early stopping # only rank 0 checks for early stopping
if local_rank == 0: if local_rank == 0:
logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Val loss: {avg_val_loss:.4f}") logger.info(f"Rank {local_rank}: Overall val loss: {avg_val_loss_global:.4f}")
writer.add_scalar("Loss/Train_Epoch", avg_train_loss_global, epoch) writer.add_scalar("Loss/Train_Epoch", avg_train_loss_global, epoch)
writer.add_scalar("Loss/Val_Epoch", avg_val_loss_global, epoch) writer.add_scalar("Loss/Val_Epoch", avg_val_loss_global, epoch)
writer.flush() writer.flush()
@@ -290,7 +372,10 @@ def train_model(model: nn.Module,
epochs_no_improve = 0 epochs_no_improve = 0
# torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt")) # torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt"))
save_fn = model_configuration["model_save_fn"] save_fn = model_configuration["model_save_fn"]
save_fn(model, training_configuration) if torch.distributed.is_initialized():
save_fn(model.module, training_configuration)
else:
save_fn(model, training_configuration)
else: else:
epochs_no_improve += 1 epochs_no_improve += 1
logger.info( logger.info(