From 73f3f2f91b8549371d8a62ca41355a301d6fcc50 Mon Sep 17 00:00:00 2001
From: 游雁 <zhifu.gzf@alibaba-inc.com>
Date: 星期一, 20 五月 2024 15:32:40 +0800
Subject: [PATCH] Merge branch 'dev_gzf_deepspeed' of github.com:alibaba-damo-academy/FunASR into dev_gzf_deepspeed merge
---
funasr/train_utils/trainer_ds.py | 12 +++++-------
1 files changed, 5 insertions(+), 7 deletions(-)
diff --git a/funasr/train_utils/trainer_ds.py b/funasr/train_utils/trainer_ds.py
index 8a52746..bb9fca6 100644
--- a/funasr/train_utils/trainer_ds.py
+++ b/funasr/train_utils/trainer_ds.py
@@ -168,8 +168,7 @@
"""
step_in_epoch = None if step is None else step_in_epoch
if self.use_deepspeed:
- with torch.no_grad():
- model.save_checkpoint(save_dir=model_dir, tag=tag, client_state=info_dict)
+
logging.info(f"Save checkpoint: {epoch}, rank: {self.local_rank}\n")
# self.step_or_epoch += 1
state = {
@@ -273,8 +272,7 @@
elif self.use_fsdp:
pass
- step_in_epoch = None if step is None else step_in_epoch
- if self.rank == 0:
+ elif self.rank == 0:
logging.info(f"Save checkpoint: {epoch}, rank: {self.local_rank}\n")
# self.step_or_epoch += 1
state = {
@@ -385,8 +383,8 @@
if self.use_deepspeed:
ckpt = os.path.join(self.output_dir, "model.pt")
- if os.path.isfile(ckpt):
- _, checkpoint = model_engine.load_checkpoint(self.output_dir, "model.pt")
+ if os.path.exists(ckpt):
+ _, checkpoint = model.load_checkpoint(self.output_dir, "model.pt")
self.saved_ckpts = checkpoint["saved_ckpts"]
self.val_acc_step_or_eoch = (
@@ -712,7 +710,7 @@
"data_split_num": kwargs.get("data_split_num", 1),
"log_step": batch_idx + kwargs.get("start_step", 0),
"batch_total": batch_idx,
- "step_in_epoch": step_in_epoch,
+ "step_in_epoch": batch_idx,
"lr": 0.0,
}
--
Gitblit v1.9.1