diff --git a/code/configs/combined.py b/code/configs/combined.py index f0670a3..7a8545b 100644 --- a/code/configs/combined.py +++ b/code/configs/combined.py @@ -43,7 +43,7 @@ run_configuration = { "take_every_nth": take_every_nth, "shift_in_hours": shift_in_hours, "output_window_offset": output_window_offset, - "batch_size": 256, + "batch_size": 128, "max_lr": 1e-5, "num_epochs": 10, "patience": 3, diff --git a/code/models/collation.py b/code/models/collation.py index 77e3e24..3f1e523 100644 --- a/code/models/collation.py +++ b/code/models/collation.py @@ -55,6 +55,11 @@ def simple_x_y_collate(batch): collated_x.append(np.concatenate(current_x, 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) diff --git a/code/slurm_start.sh b/code/slurm_start.sh new file mode 100755 index 0000000..7124805 --- /dev/null +++ b/code/slurm_start.sh @@ -0,0 +1,58 @@ +#!/bin/bash + +# Parameter check +if [[ $# -ne 5 ]]; then + echo "Usage: $0 " + 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 < None: + log_dir: str, + num_dataloader_workers: int = 1) -> None: # create training and validation loaders logger.info("Creating training loaders") train_loader = LMDBIterableDataset(dataset_dir, @@ -88,14 +88,14 @@ def train( train_dataset=train_loader, val_dataset=val_loader, log_dir=log_dir, - logger=logger + logger=logger, + num_dataloader_workers=num_dataloader_workers, ) def evaluate(model_configuration: dict, training_configuration: dict, - test_ids: list, - device) -> dict: + test_ids: list) -> dict: logger.info(f"Evaluating model {model_configuration['id']}") eval_functions = get_eval_functions(model_configuration) results = evaluate_model(model_configuration, @@ -114,6 +114,18 @@ def evaluate(model_configuration: dict, 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, @@ -144,11 +156,11 @@ if __name__ == "__main__": required=False, default=None, help="Limit the number of items to process, default is None (no limit)") - parser.add_argument("--device", - type=str, + parser.add_argument("--num_dataloader_workers", + type=int, required=False, - default="cuda", - help="Device to use for training, default is cuda") + default=1, + help="Number of workers for the dataloader, default is 1") args = parser.parse_args() # load run configuration from module @@ -201,8 +213,11 @@ if __name__ == "__main__": 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: + 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"] @@ -214,29 +229,42 @@ if __name__ == "__main__": 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) - # set up logger - logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, "main.log")) + # 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}") @@ -270,14 +298,18 @@ if __name__ == "__main__": 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() + + # 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, diff --git a/code/utils/data_utils.py b/code/utils/data_utils.py index e0bba83..4d66a01 100644 --- a/code/utils/data_utils.py +++ b/code/utils/data_utils.py @@ -479,6 +479,10 @@ class LMDBIterableDataset(IterableDataset): self.key_subset = key_subset # reset length 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]): """ @@ -522,12 +526,13 @@ class LMDBIterableDataset(IterableDataset): return num_steps def __iter__(self): - self.init_lmdb_env() # if no key subset is set, use all keys if self.key_subset is None: self.key_subset = self.lmdb_keys + self.init_lmdb_env() + random.shuffle(self.key_subset) batch = list() @@ -563,12 +568,12 @@ class LMDBIterableDataset(IterableDataset): meminit=False) def __len__(self): - self.init_lmdb_env() - # if no key subset is set, use all keys if self.key_subset is None: self.key_subset = self.lmdb_keys + self.init_lmdb_env() + if self.len is None: num_steps = self.get_length_of_data_subset(self.key_subset) self.len = num_steps diff --git a/code/utils/evaluation.py b/code/utils/evaluation.py index 54f68bc..c412c2c 100644 --- a/code/utils/evaluation.py +++ b/code/utils/evaluation.py @@ -130,6 +130,8 @@ def evaluate_model(model_configuration: dict, if eval_fn is not None: eval_fn_name = eval_fn["name"] accumulation_fn = eval_fn["accumulation_fn"] + if eval_fn_name not in errors: + continue for key in errors[eval_fn_name]: if len(errors[eval_fn_name][key]) == 0: errors[eval_fn_name][key] = np.nan diff --git a/code/utils/training.py b/code/utils/training.py index 20ff69d..d6d4b6e 100644 --- a/code/utils/training.py +++ b/code/utils/training.py @@ -104,7 +104,8 @@ def train_model(model: nn.Module, train_dataset: LMDBIterableDataset, val_dataset: LMDBIterableDataset, log_dir: str = "./logs", - logger=None) -> torch.nn.Module: + num_dataloader_workers: int = 4, + logger=None) -> None: if logger is None: 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 if torch.distributed.is_initialized(): local_rank = torch.distributed.get_rank() - torch.cuda.set_device(local_rank) device = torch.device(f"cuda:{local_rank}") world_size = torch.distributed.get_world_size() else: @@ -132,10 +132,46 @@ def train_model(model: nn.Module, 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)] # 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 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: train_subsets = [train_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_dataset, batch_size=None, - num_workers=4, + num_workers=num_dataloader_workers, ) val_dataloader = DataLoader( val_dataset, batch_size=None, - num_workers=4, + num_workers=num_dataloader_workers, ) 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) + # 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 optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"]) scheduler = OneCycleLR(optimizer, @@ -204,21 +246,37 @@ def train_model(model: nn.Module, train_dataloader = DataLoader( train_dataset, 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_dataset, 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) - for step in tqdm(range(len(train_dataloader)), total=len(train_dataloader)): - loss = batch_loss_fn(model, - iterator, - loss_functions, - device, - model_configuration) + for step in tqdm(range(train_length)): + try: + loss = batch_loss_fn(model, + iterator, + loss_functions, + 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() loss.backward() @@ -237,20 +295,35 @@ def train_model(model: nn.Module, writer.flush() 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 + # if local_rank == 0 or not torch.distributed.is_initialized(): model.eval() total_val_loss = 0 logger.info(f"Rank {local_rank}: Validation") with torch.no_grad(): val_iter = iter(val_dataloader) - for step in tqdm(range(len(val_dataloader))): - loss = batch_loss_fn(model, - val_iter, - loss_functions, - device, - model_configuration) + for step in tqdm(range(val_length)): + try: + loss = batch_loss_fn(model, + val_iter, + loss_functions, + 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() @@ -259,11 +332,20 @@ def train_model(model: nn.Module, avg_val_loss = None else: 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 if torch.distributed.is_initialized(): if avg_val_loss is not None: 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) 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 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/Val_Epoch", avg_val_loss_global, epoch) writer.flush() @@ -290,7 +372,10 @@ def train_model(model: nn.Module, epochs_no_improve = 0 # torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt")) 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: epochs_no_improve += 1 logger.info(