fixes
This commit is contained in:
@@ -0,0 +1,330 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user