code update
This commit is contained in:
@@ -479,6 +479,10 @@ class LMDBIterableDataset(IterableDataset):
|
||||
self.key_subset = key_subset
|
||||
# reset length
|
||||
self.len = None
|
||||
# reset lmdb env
|
||||
if self.lmdb_env is not None:
|
||||
self.lmdb_env.close()
|
||||
self.lmdb_env = None
|
||||
|
||||
def get_length_of_data_subset(self, key_set: list[str]):
|
||||
"""
|
||||
@@ -522,12 +526,13 @@ class LMDBIterableDataset(IterableDataset):
|
||||
return num_steps
|
||||
|
||||
def __iter__(self):
|
||||
self.init_lmdb_env()
|
||||
|
||||
# if no key subset is set, use all keys
|
||||
if self.key_subset is None:
|
||||
self.key_subset = self.lmdb_keys
|
||||
|
||||
self.init_lmdb_env()
|
||||
|
||||
random.shuffle(self.key_subset)
|
||||
|
||||
batch = list()
|
||||
@@ -563,12 +568,12 @@ class LMDBIterableDataset(IterableDataset):
|
||||
meminit=False)
|
||||
|
||||
def __len__(self):
|
||||
self.init_lmdb_env()
|
||||
|
||||
# if no key subset is set, use all keys
|
||||
if self.key_subset is None:
|
||||
self.key_subset = self.lmdb_keys
|
||||
|
||||
self.init_lmdb_env()
|
||||
|
||||
if self.len is None:
|
||||
num_steps = self.get_length_of_data_subset(self.key_subset)
|
||||
self.len = num_steps
|
||||
|
||||
@@ -130,6 +130,8 @@ def evaluate_model(model_configuration: dict,
|
||||
if eval_fn is not None:
|
||||
eval_fn_name = eval_fn["name"]
|
||||
accumulation_fn = eval_fn["accumulation_fn"]
|
||||
if eval_fn_name not in errors:
|
||||
continue
|
||||
for key in errors[eval_fn_name]:
|
||||
if len(errors[eval_fn_name][key]) == 0:
|
||||
errors[eval_fn_name][key] = np.nan
|
||||
|
||||
+107
-22
@@ -104,7 +104,8 @@ def train_model(model: nn.Module,
|
||||
train_dataset: LMDBIterableDataset,
|
||||
val_dataset: LMDBIterableDataset,
|
||||
log_dir: str = "./logs",
|
||||
logger=None) -> torch.nn.Module:
|
||||
num_dataloader_workers: int = 4,
|
||||
logger=None) -> None:
|
||||
if logger is None:
|
||||
logger = get_logger(__name__, f"{log_dir}/{model_configuration['id']}_{training_configuration['id']}.log")
|
||||
|
||||
@@ -117,7 +118,6 @@ def train_model(model: nn.Module,
|
||||
# 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:
|
||||
@@ -132,10 +132,46 @@ 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
|
||||
total_train_steps = sum([train_dataset.get_length_of_data_subset(subset) for subset in train_subsets])
|
||||
local_steps = [train_dataset.get_length_of_data_subset(subset) for subset in train_subsets]
|
||||
total_train_steps = sum(local_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]
|
||||
|
||||
def get_synced_values(local_values):
|
||||
"""
|
||||
Sync values across all processes.
|
||||
Args:
|
||||
local_values: list of local values
|
||||
Returns:
|
||||
list of synced values
|
||||
"""
|
||||
if not isinstance(local_values, torch.Tensor):
|
||||
local_values = torch.tensor(local_values, device=device, dtype=torch.float32)
|
||||
else:
|
||||
local_values = local_values.to(device)
|
||||
gathered_values = [torch.zeros_like(local_values) for _ in range(world_size)]
|
||||
torch.distributed.all_gather(gathered_values, local_values)
|
||||
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]
|
||||
train_epoch_lengths = list()
|
||||
for i in range(len(global_lengths[0])):
|
||||
current_epoch_lengths = [x[i] for x in global_lengths]
|
||||
train_epoch_lengths.append(int(min(current_epoch_lengths)))
|
||||
print(f"Rank {local_rank}: Global lengths: {global_lengths} cut to {train_epoch_lengths}")
|
||||
|
||||
# also sync the val subsets
|
||||
global_val_lengths = get_synced_values(val_subset_lengths)
|
||||
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])):
|
||||
current_epoch_lengths = [x[i] for x in global_val_lengths]
|
||||
val_epoch_lengths.append(int(min(current_epoch_lengths)))
|
||||
print(f"Rank {local_rank}: Global val lengths: {global_val_lengths} cut to {val_epoch_lengths}")
|
||||
else:
|
||||
train_subsets = [train_dataset.lmdb_keys] * num_epochs
|
||||
val_subsets = [val_dataset.lmdb_keys] * num_epochs
|
||||
@@ -156,12 +192,12 @@ def train_model(model: nn.Module,
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=None,
|
||||
num_workers=4,
|
||||
num_workers=num_dataloader_workers,
|
||||
)
|
||||
val_dataloader = DataLoader(
|
||||
val_dataset,
|
||||
batch_size=None,
|
||||
num_workers=4,
|
||||
num_workers=num_dataloader_workers,
|
||||
)
|
||||
|
||||
logger.info(f"Rank {local_rank}: Training {training_id} with {num_epochs} epochs")
|
||||
@@ -171,6 +207,12 @@ def train_model(model: nn.Module,
|
||||
|
||||
model.to(device)
|
||||
|
||||
# wrap model in DDP if distributed training
|
||||
if torch.distributed.is_initialized():
|
||||
model = torch.nn.parallel.DistributedDataParallel(model,
|
||||
device_ids=[local_rank],
|
||||
output_device=local_rank)
|
||||
|
||||
# load training state from training configuration, if available
|
||||
optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"])
|
||||
scheduler = OneCycleLR(optimizer,
|
||||
@@ -204,21 +246,37 @@ def train_model(model: nn.Module,
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=None,
|
||||
num_workers=4,
|
||||
num_workers=num_dataloader_workers,
|
||||
# set multiprocessing start method to spawn
|
||||
# multiprocessing_context="forkserver",
|
||||
)
|
||||
train_length = train_epoch_lengths[epoch - 1]
|
||||
val_dataloader = DataLoader(
|
||||
val_dataset,
|
||||
batch_size=None,
|
||||
num_workers=4,
|
||||
num_workers=num_dataloader_workers,
|
||||
# set multiprocessing start method to spawn
|
||||
# multiprocessing_context="forkserver",
|
||||
)
|
||||
val_length = val_epoch_lengths[epoch - 1]
|
||||
|
||||
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)
|
||||
for step in tqdm(range(train_length)):
|
||||
try:
|
||||
loss = batch_loss_fn(model,
|
||||
iterator,
|
||||
loss_functions,
|
||||
device,
|
||||
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)
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
@@ -237,20 +295,35 @@ def train_model(model: nn.Module,
|
||||
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}")
|
||||
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()
|
||||
|
||||
# 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(len(val_dataloader))):
|
||||
loss = batch_loss_fn(model,
|
||||
val_iter,
|
||||
loss_functions,
|
||||
device,
|
||||
model_configuration)
|
||||
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, resetting it.")
|
||||
val_iter = iter(val_dataloader)
|
||||
loss = batch_loss_fn(model,
|
||||
val_iter,
|
||||
loss_functions,
|
||||
device,
|
||||
model_configuration)
|
||||
|
||||
total_val_loss += loss.item()
|
||||
|
||||
@@ -259,11 +332,20 @@ def train_model(model: nn.Module,
|
||||
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)
|
||||
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()
|
||||
|
||||
@@ -276,7 +358,7 @@ def train_model(model: nn.Module,
|
||||
|
||||
# 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}")
|
||||
logger.info(f"Rank {local_rank}: Overall val loss: {avg_val_loss_global:.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()
|
||||
@@ -290,7 +372,10 @@ 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"]
|
||||
save_fn(model, training_configuration)
|
||||
if torch.distributed.is_initialized():
|
||||
save_fn(model.module, training_configuration)
|
||||
else:
|
||||
save_fn(model, training_configuration)
|
||||
else:
|
||||
epochs_no_improve += 1
|
||||
logger.info(
|
||||
|
||||
Reference in New Issue
Block a user