From e71546b06d0bddd30adf7edde9db7bee52b42b9f Mon Sep 17 00:00:00 2001
From: shixian <shixian@U-09RYG5WD-2244.local>
Date: 星期四, 05 十二月 2024 15:14:47 +0800
Subject: [PATCH] debug

---
 funasr/bin/train_ds.py                                            |    2 +-
 examples/industrial_data_pretraining/seaco_paraformer/finetune.sh |    1 +
 2 files changed, 2 insertions(+), 1 deletions(-)

diff --git a/examples/industrial_data_pretraining/seaco_paraformer/finetune.sh b/examples/industrial_data_pretraining/seaco_paraformer/finetune.sh
index b07d513..bac0ac1 100644
--- a/examples/industrial_data_pretraining/seaco_paraformer/finetune.sh
+++ b/examples/industrial_data_pretraining/seaco_paraformer/finetune.sh
@@ -78,5 +78,6 @@
 ++train_conf.avg_nbest_model=10 \
 ++train_conf.use_deepspeed=false \
 ++train_conf.deepspeed_config=${deepspeed_config} \
+++train_conf.find_unused_parameters=true \
 ++optim_conf.lr=0.0002 \
 ++output_dir="${output_dir}" &> ${log_file}
\ No newline at end of file
diff --git a/funasr/bin/train_ds.py b/funasr/bin/train_ds.py
index 5b1eeaa..dc7fb42 100644
--- a/funasr/bin/train_ds.py
+++ b/funasr/bin/train_ds.py
@@ -134,7 +134,7 @@
         **kwargs.get("train_conf"),
     )
 
-    model = trainer.warp_model(model)
+    model = trainer.warp_model(model, **kwargs)
 
     kwargs["device"] = int(os.environ.get("LOCAL_RANK", 0))
     trainer.device = int(os.environ.get("LOCAL_RANK", 0))

--
Gitblit v1.9.1