import math import torch from torch import nn from torch.optim import AdamW from torch.optim.lr_scheduler import OneCycleLR from torch.utils.data import IterableDataset, DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm from utils.data_utils import LMDBIterableDataset from utils.utils import get_logger def process_tft_batch(model: nn.Module, data_iterator: IterableDataset, loss_functions: list, device: torch.device, model_configuration: dict) -> torch.Tensor: batch = next(data_iterator) batch = {k: v.to(device) for k, v in batch.items() if v is not None} input_window_length = model_configuration["model_parameters"]["encoder_length"] preds = model(batch).cpu() # [B, decoder_len, Q] target = batch["target"][:, input_window_length:, :].cpu() # match decoder segment loss = get_x_y_loss(preds, target, loss_functions) return loss def get_x_y_loss(pred: torch.Tensor, target: torch.Tensor, loss_functions: list, *args, **kwargs) -> torch.Tensor: if len(loss_functions) > 1: losses = list() for dim in range(target.shape[-1]): if len(pred.shape) > 2: current_preds = pred[:, :, dim].ravel() else: current_preds = pred[:, dim] if len(target.shape) > 2: current_target = target[:, :, dim].ravel() else: current_target = target[:, dim].ravel() # skip dimension, if it contains only NaN values, as loss cens nan_indices = torch.isnan(current_target) if torch.all(nan_indices): continue current_target = current_target[~nan_indices] current_preds = current_preds[~nan_indices] if len(current_target) == 0: continue loss = loss_functions[dim](current_preds, current_target) losses.append(loss) loss = torch.stack(losses).mean() else: nan_indices = torch.isnan(target) if torch.all(nan_indices): return torch.tensor(0.0) current_target = target[~nan_indices] current_preds = pred[~nan_indices] if len(current_target) == 0: return torch.tensor(0.0) loss = loss_functions[0](current_preds, current_target) return loss def get_model_loss(model: nn.Module, data_iterator: IterableDataset, loss_functions: list, device: str, *args, **kwargs) -> torch.Tensor: batch_x, batch_y = next(data_iterator) batch_x = batch_x.to(device).float() target = batch_y.to(device).float() pred = model(batch_x) loss = get_x_y_loss(pred, target, loss_functions) return loss def get_ranked_ids(all_ids, epoch, rank, world_size, base_seed=42): """ Get ranked ids for distributed training. Args: all_ids: list of all available ids epoch: current epoch rank: rank of the current process world_size: number of processes base_seed: base seed for random number generator Returns: list of ids for the current process """ g = torch.Generator() g.manual_seed(base_seed + epoch) permuted = torch.randperm(len(all_ids), generator=g).tolist() return [all_ids[i] for i in permuted[rank::world_size]] def train_model(model: nn.Module, model_configuration: dict, training_configuration: dict, train_dataset: LMDBIterableDataset, val_dataset: LMDBIterableDataset, log_dir: str = "./logs", 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") learning_parameters = training_configuration["learning_parameters"] num_epochs = learning_parameters["epochs"] patience = learning_parameters["patience"] training_id = training_configuration["id"] # get computation rank if torch.distributed.is_initialized(): local_rank = torch.distributed.get_rank() device = torch.device(f"cuda:{local_rank}") world_size = torch.distributed.get_world_size() else: local_rank = 0 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") world_size = 1 logger.info(f"Rank {local_rank}: Using device: {device}, world size: {world_size}") # get subsets for distributed training if torch.distributed.is_initialized(): 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 local_train_steps = [train_dataset.get_length_of_data_subset(subset) for subset in train_subsets] total_train_steps = sum(local_train_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)] local_val_steps = [val_dataset.get_length_of_data_subset(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_train_lengths = get_synced_values(local_train_steps) global_train_lengths = [x.cpu().numpy() for x in global_train_lengths] train_epoch_lengths = list() for i in range(len(global_train_lengths[0])): current_epoch_lengths = [x[i] for x in global_train_lengths] train_epoch_lengths.append(int(min(current_epoch_lengths))) print(f"Rank {local_rank}: Global lengths: {global_train_lengths} cut to {train_epoch_lengths}") # also sync the val subsets global_val_lengths = get_synced_values(local_val_steps) 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 total_train_steps = len(train_dataset) # log the number of training steps for each epoch train_subset_lengths = {f"epoch_{i}": len(subset) for i, subset in enumerate(train_subsets)} logger.info(f"Train_subsets: {train_subset_lengths}") logger.info(f"Rank {local_rank}: Total training steps: {total_train_steps}") # initialize loaders for non distributed if not torch.distributed.is_initialized(): # set the subsets for the datasets train_dataset.set_key_subset(train_subsets[0]) val_dataset.set_key_subset(val_subsets[0]) # create data loaders train_dataloader = DataLoader( train_dataset, batch_size=None, num_workers=num_dataloader_workers, ) val_dataloader = DataLoader( val_dataset, batch_size=None, 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 on {torch.cuda.device_count()} GPUs") logger.info( f"Rank {local_rank}: Current device: {torch.cuda.get_device_name(local_rank)} on local rank {local_rank}") 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, max_lr=learning_parameters["learning_rate"], # make sure to use length of full dataset here total_steps=total_train_steps) current_epoch = 1 loss_functions = training_configuration["loss_functions"] # loss_fn = nn.MSELoss() writer = SummaryWriter(log_dir=f'{log_dir}/{model_configuration["id"]}_{training_id}', ) best_val_loss = math.inf epochs_no_improve = 0 log_every_n_steps = max(len(train_dataset) // 500, 1) batch_loss_fn = model_configuration["batch_loss_fn"] for epoch in range(current_epoch, num_epochs + 1): logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs}") model.train() total_train_loss = 0 # on distributed training, reshuffle the data if torch.distributed.is_initialized(): # update the datasets with the new ids train_dataset.set_key_subset(train_subsets[epoch - 1]) val_dataset.set_key_subset(val_subsets[epoch - 1]) # recreate data loaders train_dataloader = DataLoader( train_dataset, batch_size=None, num_workers=num_dataloader_workers, # set multiprocessing start method to spawn # multiprocessing_context="forkserver", multiprocessing_context="spawn", ) train_length = train_epoch_lengths[epoch - 1] val_dataloader = DataLoader( val_dataset, batch_size=None, num_workers=num_dataloader_workers, # set multiprocessing start method to spawn # multiprocessing_context="forkserver", multiprocessing_context="spawn", ) val_length = val_epoch_lengths[epoch - 1] iterator = iter(train_dataloader) 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, continuing to next epoch.") break optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() total_train_loss += loss.item() if local_rank == 0: if step % log_every_n_steps == 0: writer.add_scalar("Loss/Train_Step", loss.item(), ((epoch - 1) * len(train_dataloader) + step) * training_configuration[ "batch_size"]) writer.add_scalar("LR", scheduler.get_last_lr()[0], ((epoch - 1) * len(train_dataloader) + step) * training_configuration[ "batch_size"]) writer.flush() avg_train_loss = total_train_loss / len(train_dataloader) 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(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, continuing to next epoch.") break total_val_loss += loss.item() if len(val_dataloader) == 0: logger.info(f"Rank {local_rank}: Validation set is empty, using 0 as validation loss.") 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() avg_train_loss_global = torch.tensor(avg_train_loss).to(device) torch.distributed.all_reduce(avg_train_loss_global) avg_train_loss_global /= torch.distributed.get_world_size() else: avg_val_loss_global = torch.tensor(avg_val_loss) avg_train_loss_global = torch.tensor(avg_train_loss) # only rank 0 checks for early stopping if local_rank == 0: 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() # Early stopping should_stop = False if avg_val_loss_global < best_val_loss: logger.info( f"Rank {local_rank}: Validation loss improved from {best_val_loss:.4f} to {avg_val_loss_global:.4f}.") best_val_loss = avg_val_loss_global epochs_no_improve = 0 # torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt")) save_fn = model_configuration["model_save_fn"] if torch.distributed.is_initialized(): save_fn(model.module, training_configuration) else: save_fn(model, training_configuration) else: epochs_no_improve += 1 logger.info( f"Rank {local_rank}: No improvement in validation loss, no-improve count: {epochs_no_improve}") if epochs_no_improve >= patience: logger.info("Early stopping triggered.") # broadcast stop signal to all gpus should_stop = True else: should_stop = None if torch.distributed.is_initialized(): if local_rank == 0: should_stop_tensor = torch.tensor([int(should_stop)], device=device) else: should_stop_tensor = torch.zeros(1, dtype=torch.uint8, device=device) # safe default torch.distributed.broadcast(should_stop_tensor, src=0) should_stop = bool(should_stop_tensor.item()) if should_stop: logger.info(f"Rank {local_rank}: Stopping training.") break # ensure sync between epochs if torch.distributed.is_initialized(): torch.distributed.barrier() # clean up del loss del train_dataloader del val_dataloader del model del optimizer del scheduler # free up memory torch.cuda.synchronize()