408 lines
17 KiB
Python
408 lines
17 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",
|
|
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()
|