code update
This commit is contained in:
+15
-23
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user