331 lines
13 KiB
Python
331 lines
13 KiB
Python
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",
|
|
logger=None) -> torch.nn.Module:
|
|
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()
|
|
torch.cuda.set_device(local_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
|
|
total_train_steps = sum([train_dataset.get_length_of_data_subset(subset) for subset in train_subsets])
|
|
|
|
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)]
|
|
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=4,
|
|
)
|
|
val_dataloader = DataLoader(
|
|
val_dataset,
|
|
batch_size=None,
|
|
num_workers=4,
|
|
)
|
|
|
|
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)
|
|
|
|
# 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=4,
|
|
)
|
|
val_dataloader = DataLoader(
|
|
val_dataset,
|
|
batch_size=None,
|
|
num_workers=4,
|
|
)
|
|
|
|
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)
|
|
|
|
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. Train loss: {avg_train_loss:.4f}")
|
|
|
|
# Validation
|
|
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)
|
|
|
|
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)
|
|
|
|
# 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)
|
|
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}: Epoch {epoch}/{num_epochs} done. Val loss: {avg_val_loss:.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"]
|
|
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()
|