From 47c14aff79bb902482b0a953a55d31bf130c7b04 Mon Sep 17 00:00:00 2001
From: 游雁 <zhifu.gzf@alibaba-inc.com>
Date: 星期一, 20 五月 2024 16:20:45 +0800
Subject: [PATCH] add
---
funasr/train_utils/trainer_ds.py | 11 +++++++----
1 files changed, 7 insertions(+), 4 deletions(-)
diff --git a/funasr/train_utils/trainer_ds.py b/funasr/train_utils/trainer_ds.py
index bb9fca6..db92bc8 100644
--- a/funasr/train_utils/trainer_ds.py
+++ b/funasr/train_utils/trainer_ds.py
@@ -15,6 +15,7 @@
from funasr.train_utils.recursive_op import recursive_average
from funasr.train_utils.average_nbest_models import average_checkpoints
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
+import funasr.utils.misc as misc_utils
try:
import wandb
@@ -268,7 +269,8 @@
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
- os.remove(filename)
+ # os.remove(filename)
+ misc_utils.smart_remove(filename)
elif self.use_fsdp:
pass
@@ -360,7 +362,8 @@
filename = os.path.join(self.output_dir, key)
logging.info(f"Delete: {filename}")
if os.path.exists(filename):
- os.remove(filename)
+ # os.remove(filename)
+ misc_utils.smart_remove(filename)
if self.use_ddp or self.use_fsdp:
dist.barrier()
@@ -709,8 +712,8 @@
"data_split_i": kwargs.get("data_split_i", 0),
"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": batch_idx,
+ "batch_total": batch_idx + 1,
+ "step_in_epoch": batch_idx + 1,
"lr": 0.0,
}
--
Gitblit v1.9.1