added code
This commit is contained in:
+150
-89
@@ -1,11 +1,13 @@
|
||||
import math
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
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
|
||||
@@ -193,11 +195,13 @@ def train_model(model: nn.Module,
|
||||
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")
|
||||
@@ -219,10 +223,18 @@ def train_model(model: nn.Module,
|
||||
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
|
||||
|
||||
loss_functions = training_configuration["loss_functions"]
|
||||
# loss_fn = nn.MSELoss()
|
||||
# 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
|
||||
@@ -239,9 +251,9 @@ def train_model(model: nn.Module,
|
||||
# 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])
|
||||
|
||||
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,
|
||||
@@ -249,104 +261,160 @@ def train_model(model: nn.Module,
|
||||
num_workers=num_dataloader_workers,
|
||||
# set multiprocessing start method to spawn
|
||||
# multiprocessing_context="forkserver",
|
||||
multiprocessing_context="spawn",
|
||||
# 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]
|
||||
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)
|
||||
|
||||
iterator = iter(train_dataloader)
|
||||
train_iter = iter(train_dataloader)
|
||||
for step in tqdm(range(train_length)):
|
||||
optimizer.zero_grad()
|
||||
try:
|
||||
loss = batch_loss_fn(model,
|
||||
iterator,
|
||||
loss_functions,
|
||||
device,
|
||||
model_configuration)
|
||||
with autocast("cuda"):
|
||||
loss = batch_loss_fn(model,
|
||||
train_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
|
||||
# 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)
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
# 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(),
|
||||
((epoch - 1) * len(train_dataloader) + step) * training_configuration[
|
||||
"batch_size"])
|
||||
normalized_step)
|
||||
writer.add_scalar("LR", scheduler.get_last_lr()[0],
|
||||
((epoch - 1) * len(train_dataloader) + step) * training_configuration[
|
||||
"batch_size"])
|
||||
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}")
|
||||
|
||||
# sync before validation
|
||||
if torch.distributed.is_initialized():
|
||||
torch.distributed.barrier()
|
||||
# clear up memory and dataloaders
|
||||
del train_iter
|
||||
del train_dataloader
|
||||
|
||||
# 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)
|
||||
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:
|
||||
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()
|
||||
val_subset = val_dataset.lmdb_keys
|
||||
val_length = len(val_dataset)
|
||||
|
||||
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)
|
||||
# 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:
|
||||
@@ -364,6 +432,7 @@ def train_model(model: nn.Module,
|
||||
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:
|
||||
@@ -371,7 +440,7 @@ def train_model(model: nn.Module,
|
||||
else:
|
||||
epochs_no_improve += 1
|
||||
logger.info(
|
||||
f"Rank {local_rank}: No improvement in validation loss, no-improve count: {epochs_no_improve}")
|
||||
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
|
||||
@@ -395,13 +464,5 @@ def train_model(model: nn.Module,
|
||||
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