code update

This commit is contained in:
Alex Blank
2025-05-20 23:39:23 +02:00
parent cb7896900b
commit d8b9ccfe99
4 changed files with 76 additions and 45 deletions
+15 -23
View File
@@ -132,12 +132,12 @@ def train_model(model: nn.Module,
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_steps = [train_dataset.get_length_of_data_subset(subset) for subset in train_subsets]
total_train_steps = sum(local_steps)
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)]
val_subset_lengths = [len(subset) for subset in val_subsets]
local_val_steps = [val_dataset.get_length_of_data_subset(subset) for subset in val_subsets]
def get_synced_values(local_values):
"""
@@ -156,16 +156,16 @@ def train_model(model: nn.Module,
return gathered_values
# sync lengths of subsets and adjust to minimum length for equal sized training lengths
global_lengths = get_synced_values(local_steps)
global_lengths = [x.cpu().numpy() for x in global_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_lengths[0])):
current_epoch_lengths = [x[i] for x in global_lengths]
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_lengths} cut to {train_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(val_subset_lengths)
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])):
@@ -249,6 +249,7 @@ def train_model(model: nn.Module,
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(
@@ -257,6 +258,7 @@ def train_model(model: nn.Module,
num_workers=num_dataloader_workers,
# set multiprocessing start method to spawn
# multiprocessing_context="forkserver",
multiprocessing_context="spawn",
)
val_length = val_epoch_lengths[epoch - 1]
@@ -270,13 +272,8 @@ def train_model(model: nn.Module,
model_configuration)
except StopIteration:
# if the iterator is exhausted, reset it
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
iterator = iter(train_dataloader)
loss = batch_loss_fn(model,
iterator,
loss_functions,
device,
model_configuration)
logger.info(f"Rank {local_rank}: Iterator exhausted, continuing to next epoch.")
break
optimizer.zero_grad()
loss.backward()
@@ -317,13 +314,8 @@ def train_model(model: nn.Module,
model_configuration)
except StopIteration:
# if the iterator is exhausted, reset it
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
val_iter = iter(val_dataloader)
loss = batch_loss_fn(model,
val_iter,
loss_functions,
device,
model_configuration)
logger.info(f"Rank {local_rank}: Iterator exhausted, continuing to next epoch.")
break
total_val_loss += loss.item()