added code

This commit is contained in:
2025-09-10 10:37:55 +02:00
parent 36901c736d
commit c78a68de80
199 changed files with 3561 additions and 22579 deletions
+150 -89
View File
@@ -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()