Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
76 changes: 49 additions & 27 deletions funasr/train_utils/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,56 +208,78 @@ def save_checkpoint(
latest = Path(os.path.join(self.output_dir, f"model.pt"))
torch.save(state, latest)

if self.best_step_or_epoch == "":
if self.best_step_or_epoch == "" and ckpt_name in getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
):
self.best_step_or_epoch = ckpt_name

if self.avg_keep_nbest_models_type == "acc":
if (
self.val_acc_step_or_epoch[ckpt_name]
>= self.val_acc_step_or_epoch[self.best_step_or_epoch]
):
cur_acc = self.val_acc_step_or_epoch.get(ckpt_name)
best_acc = self.val_acc_step_or_epoch.get(self.best_step_or_epoch)
if cur_acc is not None and (best_acc is None or cur_acc >= best_acc):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
torch.save(state, best_ckpt)
logging.info(
f"Update best acc: {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best acc: {cur_acc:.4f}, {best_ckpt}"
)
elif cur_acc is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation acc yet; not considered for best."
)
else:
logging.info(
f"No improvement in acc: {self.val_acc_step_or_epoch[ckpt_name]:.4f} < {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in acc: {cur_acc:.4f} < {best_acc:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
elif self.avg_keep_nbest_models_type == "loss":
if (
self.val_loss_step_or_epoch[ckpt_name]
<= self.val_loss_step_or_epoch[self.best_step_or_epoch]
):
cur_loss = self.val_loss_step_or_epoch.get(ckpt_name)
best_loss = self.val_loss_step_or_epoch.get(self.best_step_or_epoch)
if cur_loss is not None and (best_loss is None or cur_loss <= best_loss):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
torch.save(state, best_ckpt)
logging.info(
f"Update best loss: {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best loss: {cur_loss:.4f}, {best_ckpt}"
)
elif cur_loss is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation loss yet; not considered for best."
)
else:
logging.info(
f"No improvement in loss: {self.val_loss_step_or_epoch[ckpt_name]:.4f} > {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in loss: {cur_loss:.4f} > {best_loss:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
else:
print("Undo")
self.saved_ckpts[ckpt_name] = getattr(
# Only checkpoints carrying the configured validation metric
# (acc or loss) are ranked. A checkpoint saved at an unvalidated
# step (e.g. save_checkpoint_interval is not a multiple of
# validate_interval) is kept on disk but excluded from saved_ckpts,
# so it never competes in best-model ranking or keep_nbest_models
# pruning with a fabricated score and cannot evict a validated
# best checkpoint.
metric_value = getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
)[ckpt_name]
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
os.remove(filename)
).get(ckpt_name)
if metric_value is None:
logging.info(
f"Checkpoint {ckpt_name} has no {self.avg_keep_nbest_models_type} metric; "
"kept on disk but excluded from keep_nbest_models ranking."
)
else:
self.saved_ckpts[ckpt_name] = metric_value
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
os.remove(filename)

if self.use_ddp or self.use_fsdp:
dist.barrier()
Expand Down
156 changes: 100 additions & 56 deletions funasr/train_utils/trainer_ds.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,14 +234,15 @@ def save_checkpoint(
# torch.save(state, latest)
with torch.no_grad():
model.save_checkpoint(save_dir=self.output_dir, tag=f"model.pt", client_state=state)
if self.best_step_or_epoch == "":
if self.best_step_or_epoch == "" and ckpt_name in getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
):
self.best_step_or_epoch = ckpt_name

if self.avg_keep_nbest_models_type == "acc":
if (
self.val_acc_step_or_epoch[ckpt_name]
>= self.val_acc_step_or_epoch[self.best_step_or_epoch]
):
cur_acc = self.val_acc_step_or_epoch.get(ckpt_name)
best_acc = self.val_acc_step_or_epoch.get(self.best_step_or_epoch)
if cur_acc is not None and (best_acc is None or cur_acc >= best_acc):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
# torch.save(state, best_ckpt)
Expand All @@ -250,17 +251,20 @@ def save_checkpoint(
save_dir=self.output_dir, tag=f"model.pt.best", client_state=state
)
logging.info(
f"Update best acc: {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best acc: {cur_acc:.4f}, {best_ckpt}"
)
elif cur_acc is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation acc yet; not considered for best."
)
else:
logging.info(
f"No improvement in acc: {self.val_acc_step_or_epoch[ckpt_name]:.4f} < {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in acc: {cur_acc:.4f} < {best_acc:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
elif self.avg_keep_nbest_models_type == "loss":
if (
self.val_loss_step_or_epoch[ckpt_name]
<= self.val_loss_step_or_epoch[self.best_step_or_epoch]
):
cur_loss = self.val_loss_step_or_epoch.get(ckpt_name)
best_loss = self.val_loss_step_or_epoch.get(self.best_step_or_epoch)
if cur_loss is not None and (best_loss is None or cur_loss <= best_loss):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
# torch.save(state, best_ckpt)
Expand All @@ -269,31 +273,49 @@ def save_checkpoint(
save_dir=self.output_dir, tag=f"model.pt.best", client_state=state
)
logging.info(
f"Update best loss: {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best loss: {cur_loss:.4f}, {best_ckpt}"
)
elif cur_loss is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation loss yet; not considered for best."
)
else:
logging.info(
f"No improvement in loss: {self.val_loss_step_or_epoch[ckpt_name]:.4f} > {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in loss: {cur_loss:.4f} > {best_loss:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
else:
print("Undo")
if self.rank == 0:
self.saved_ckpts[ckpt_name] = getattr(
# Only checkpoints carrying the configured validation metric
# (acc or loss) are ranked. A checkpoint saved at an unvalidated
# step (e.g. save_checkpoint_interval is not a multiple of
# validate_interval) is kept on disk but excluded from
# saved_ckpts, so it never competes in best-model ranking or
# keep_nbest_models pruning with a fabricated score and cannot
# evict a validated best checkpoint.
metric_value = getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
)[ckpt_name]
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
# os.remove(filename)
misc_utils.smart_remove(filename)
).get(ckpt_name)
if metric_value is None:
logging.info(
f"Checkpoint {ckpt_name} has no {self.avg_keep_nbest_models_type} metric; "
"kept on disk but excluded from keep_nbest_models ranking."
)
else:
self.saved_ckpts[ckpt_name] = metric_value
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
# os.remove(filename)
misc_utils.smart_remove(filename)

elif self.use_fsdp:
raise NotImplementedError(
Expand Down Expand Up @@ -356,57 +378,79 @@ def save_checkpoint(
logging.info(f"\nCheckpoint saved to {filename}\n")
latest = Path(os.path.join(self.output_dir, f"model.pt"))
torch.save(state, latest)
if self.best_step_or_epoch == "":
if self.best_step_or_epoch == "" and ckpt_name in getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
):
self.best_step_or_epoch = ckpt_name

if self.avg_keep_nbest_models_type == "acc":
if (
self.val_acc_step_or_epoch[ckpt_name]
>= self.val_acc_step_or_epoch[self.best_step_or_epoch]
):
cur_acc = self.val_acc_step_or_epoch.get(ckpt_name)
best_acc = self.val_acc_step_or_epoch.get(self.best_step_or_epoch)
if cur_acc is not None and (best_acc is None or cur_acc >= best_acc):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
torch.save(state, best_ckpt)
logging.info(
f"Update best acc: {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best acc: {cur_acc:.4f}, {best_ckpt}"
)
elif cur_acc is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation acc yet; not considered for best."
)
else:
logging.info(
f"No improvement in acc: {self.val_acc_step_or_epoch[ckpt_name]:.4f} < {self.val_acc_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in acc: {cur_acc:.4f} < {best_acc:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
elif self.avg_keep_nbest_models_type == "loss":
if (
self.val_loss_step_or_epoch[ckpt_name]
<= self.val_loss_step_or_epoch[self.best_step_or_epoch]
):
cur_loss = self.val_loss_step_or_epoch.get(ckpt_name)
best_loss = self.val_loss_step_or_epoch.get(self.best_step_or_epoch)
if cur_loss is not None and (best_loss is None or cur_loss <= best_loss):
self.best_step_or_epoch = ckpt_name
best_ckpt = Path(os.path.join(self.output_dir, f"model.pt.best"))
torch.save(state, best_ckpt)
logging.info(
f"Update best loss: {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {best_ckpt}"
f"Update best loss: {cur_loss:.4f}, {best_ckpt}"
)
elif cur_loss is None:
logging.info(
f"Checkpoint {ckpt_name} saved at a step with no validation loss yet; not considered for best."
)
else:
logging.info(
f"No improvement in loss: {self.val_loss_step_or_epoch[ckpt_name]:.4f} > {self.val_loss_step_or_epoch[self.best_step_or_epoch]:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
f"No improvement in loss: {cur_loss:.4f} > {best_loss:.4f}, {os.path.join(self.output_dir, self.best_step_or_epoch)}"
)
else:
print("Undo")
self.saved_ckpts[ckpt_name] = getattr(
# Only checkpoints carrying the configured validation metric
# (acc or loss) are ranked. A checkpoint saved at an unvalidated
# step (e.g. save_checkpoint_interval is not a multiple of
# validate_interval) is kept on disk but excluded from saved_ckpts,
# so it never competes in best-model ranking or keep_nbest_models
# pruning with a fabricated score and cannot evict a validated
# best checkpoint.
metric_value = getattr(
self, f"val_{self.avg_keep_nbest_models_type}_step_or_epoch"
)[ckpt_name]
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
# os.remove(filename)
misc_utils.smart_remove(filename)
).get(ckpt_name)
if metric_value is None:
logging.info(
f"Checkpoint {ckpt_name} has no {self.avg_keep_nbest_models_type} metric; "
"kept on disk but excluded from keep_nbest_models ranking."
)
else:
self.saved_ckpts[ckpt_name] = metric_value
if self.keep_nbest_models > 0:
if len(self.saved_ckpts) > self.keep_nbest_models:
if self.avg_keep_nbest_models_type == "acc":
key = min(self.saved_ckpts, key=self.saved_ckpts.get)
else:
key = max(self.saved_ckpts, key=self.saved_ckpts.get)
if key in self.saved_ckpts:
del self.saved_ckpts[key]
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
# os.remove(filename)
misc_utils.smart_remove(filename)

if self.use_ddp or self.use_fsdp:
dist.barrier()
Expand Down
Loading