Files
2025-09-10 10:37:55 +02:00

469 lines
20 KiB
Python

import math
import inspect
import torch
from torch import nn, Gradient
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 torch.amp import autocast, GradScaler
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,
# multiprocessing_context="forkserver",
)
val_dataloader = DataLoader(
val_dataset,
batch_size=None,
num_workers=num_dataloader_workers,
# multiprocessing_context="forkserver",
)
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)
# mixed precision grad scaler to prevent underflow
scaler = GradScaler()
current_epoch = 1
# get loss functions from feature config
used_targets = [feat for feat in model_configuration["feature_config"]["target_features"] if
feat not in model_configuration["feature_config"]["ignored_features"]]
loss_functions = [target_feat["loss_fn"] for target_feat in used_targets]
# instantiate class of loss functions if they are not already
for i, loss_fn in enumerate(loss_functions):
if inspect.isclass(loss_fn):
loss_functions[i] = loss_fn()
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
logger.info(f"Rank {local_rank}: Train subset length: {len(train_subsets[epoch - 1])}")
train_subset = train_subsets[epoch - 1]
train_dataset.set_key_subset(train_subset)
# 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]
logger.info(f"Rank {local_rank}: Current epoch train length: {train_length}")
logger.info(f"Rank {local_rank}: Val subset length: {len(val_subsets[epoch - 1])}")
else:
train_subset = train_dataset.lmdb_keys
train_length = len(train_dataset)
train_iter = iter(train_dataloader)
for step in tqdm(range(train_length)):
optimizer.zero_grad()
try:
with autocast("cuda"):
loss = batch_loss_fn(model,
train_iter,
loss_functions,
device,
model_configuration)
except StopIteration:
# if the iterator is exhausted, reset it
# recreate data loaders
train_dataset.set_key_subset(train_subset)
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_iter = iter(train_dataloader)
with autocast("cuda"):
loss = batch_loss_fn(model,
train_iter,
loss_functions,
device,
model_configuration)
# scale the loss and backpropagate
scaler.scale(loss).backward()
scaler.step(optimizer)
# loss.backward()
# optimizer.step()
scheduler.step()
scaler.update()
total_train_loss += loss.item()
if local_rank == 0:
if step % log_every_n_steps == 0:
normalized_step = ((epoch - 1) * train_length + step) * world_size * \
training_configuration["batch_size"]
writer.add_scalar("Loss/Train_Step", loss.item(),
normalized_step)
writer.add_scalar("LR", scheduler.get_last_lr()[0],
normalized_step)
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}")
# clear up memory and dataloaders
del train_iter
del train_dataloader
# Validation
try:
# create validation dataloader with proper subset
if torch.distributed.is_initialized():
val_subset = val_subsets[epoch - 1]
val_dataset.set_key_subset(val_subset)
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]
logger.info(f"Rank {local_rank}: Current epoch val length: {val_length}")
else:
val_subset = val_dataset.lmdb_keys
val_length = len(val_dataset)
# sync before validation
if torch.distributed.is_initialized():
torch.distributed.barrier()
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
val_dataset.set_key_subset(val_subset)
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_iter = iter(val_dataloader)
loss = batch_loss_fn(model,
val_iter,
loss_functions,
device,
model_configuration)
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)
# clean up
del val_iter
del val_dataloader
# 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)
except Exception as e:
logger.error(f"Rank {local_rank}: Validation failed: {e}")
raise e
# 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"]
logger.info(f"Rank {local_rank}: Saving model to {training_configuration['training_dir']}")
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} of max {patience}")
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()
# free up memory
torch.cuda.synchronize()