From 3c0a9fb7c1bd642f3370d406fca81ff50ae9bc82 Mon Sep 17 00:00:00 2001
From: shixian.shi <shixian.shi@alibaba-inc.com>
Date: 星期四, 27 四月 2023 12:10:41 +0800
Subject: [PATCH] fix name
---
funasr/train/trainer.py | 10 ++++++++--
1 files changed, 8 insertions(+), 2 deletions(-)
diff --git a/funasr/train/trainer.py b/funasr/train/trainer.py
index 9c4af41..7c187e9 100644
--- a/funasr/train/trainer.py
+++ b/funasr/train/trainer.py
@@ -582,10 +582,16 @@
if num_batch_updates % batch_interval == 0:
if options.use_pai and options.oss_bucket is not None:
buffer = BytesIO()
- torch.save(model.state_dict(), buffer)
+ if hasattr(model, "module"):
+ torch.save(model.module.state_dict(), buffer)
+ else:
+ torch.save(model.state_dict(), buffer)
options.oss_bucket.put_object(os.path.join(output_dir, f"{num_batch_updates}step.pb"), buffer.getvalue())
else:
- torch.save(model.state_dict(), os.path.join(output_dir, f"{num_batch_updates}step.pb"))
+ if hasattr(model, "module"):
+ torch.save(model.module.state_dict(), os.path.join(output_dir, f"{num_batch_updates}step.pb"))
+ else:
+ torch.save(model.state_dict(), os.path.join(output_dir, f"{num_batch_updates}step.pb"))
if distributed:
torch.distributed.all_reduce(iterator_stop, ReduceOp.SUM)
--
Gitblit v1.9.1