From 438c4663d2094ed7ce1762fa4f16cf89401b8bec Mon Sep 17 00:00:00 2001
From: shixian.shi <shixian.shi@alibaba-inc.com>
Date: 星期三, 28 六月 2023 16:02:49 +0800
Subject: [PATCH] update build model for tp

---
 funasr/build_utils/build_asr_model.py |    5 +++++
 tests/test_tp_pipeline.py             |    1 +
 2 files changed, 6 insertions(+), 0 deletions(-)

diff --git a/funasr/build_utils/build_asr_model.py b/funasr/build_utils/build_asr_model.py
index 200395d..a76b204 100644
--- a/funasr/build_utils/build_asr_model.py
+++ b/funasr/build_utils/build_asr_model.py
@@ -408,10 +408,15 @@
             **args.model_conf,
         )
     elif args.model == "timestamp_prediction":
+        # predictor
+        predictor_class = predictor_choices.get_class(args.predictor)
+        predictor = predictor_class(**args.predictor_conf)
+        
         model_class = model_choices.get_class(args.model)
         model = model_class(
             frontend=frontend,
             encoder=encoder,
+            predictor=predictor,
             token_list=token_list,
             **args.model_conf,
         )
diff --git a/tests/test_tp_pipeline.py b/tests/test_tp_pipeline.py
index 21284a6..07084f2 100644
--- a/tests/test_tp_pipeline.py
+++ b/tests/test_tp_pipeline.py
@@ -23,6 +23,7 @@
             text_in='涓� 涓� 涓� 澶� 骞� 娲� 鍥� 瀹� 涓� 浠� 涔� 璺� 鍒� 瑗� 澶� 骞� 娲� 鏉� 浜� 鍛�',)
         print(rec_result)
         logger.info("punctuation inference result: {0}".format(rec_result))
+        assert rec_result=={'text': '<sil> 0.000 0.380;涓� 0.380 0.560;涓� 0.560 0.800;涓� 0.800 0.980;澶� 0.980 1.140;骞� 1.140 1.260;娲� 1.260 1.440;鍥� 1.440 1.680;瀹� 1.680 1.920;<sil> 1.920 2.040;涓� 2.040 2.200;浠� 2.200 2.320;涔� 2.320 2.500;璺� 2.500 2.680;鍒� 2.680 2.860;瑗� 2.860 3.040;澶� 3.040 3.200;骞� 3.200 3.380;娲� 3.380 3.500;鏉� 3.500 3.640;浜� 3.640 3.800;鍛� 3.800 4.150;<sil> 4.150 4.440;', 'timestamp': [[380, 560], [560, 800], [800, 980], [980, 1140], [1140, 1260], [1260, 1440], [1440, 1680], [1680, 1920], [2040, 2200], [2200, 2320], [2320, 2500], [2500, 2680], [2680, 2860], [2860, 3040], [3040, 3200], [3200, 3380], [3380, 3500], [3500, 3640], [3640, 3800], [3800, 4150]]}
 
 
 if __name__ == '__main__':

--
Gitblit v1.9.1