aky15 2 лет назад
Родитель
Сommit
b3b4c1bc5b

+ 1 - 1
funasr/models/e2e_transducer.py → funasr/models/e2e_asr_transducer.py

@@ -13,7 +13,7 @@ from funasr.models.specaug.abs_specaug import AbsSpecAug
 from funasr.models.rnnt_predictor.abs_decoder import AbsDecoder
 from funasr.models.decoder.abs_decoder import AbsDecoder as AbsAttDecoder
 from funasr.models.encoder.conformer_encoder import ConformerChunkEncoder as Encoder
-from funasr.models.joint_network import JointNetwork
+from funasr.models.joint_net.joint_network import JointNetwork
 from funasr.modules.nets_utils import get_transducer_task_io
 from funasr.layers.abs_normalize import AbsNormalize
 from funasr.torch_utils.device_funcs import force_gatherable

+ 1 - 1
funasr/models/e2e_transducer_unified.py → funasr/models/e2e_asr_transducer_unified.py

@@ -12,7 +12,7 @@ from funasr.models.frontend.abs_frontend import AbsFrontend
 from funasr.models.specaug.abs_specaug import AbsSpecAug
 from funasr.models.rnnt_predictor.abs_decoder import AbsDecoder
 from funasr.models.encoder.conformer_encoder import ConformerChunkEncoder as Encoder
-from funasr.models.joint_network import JointNetwork
+from funasr.models.joint_net.joint_network import JointNetwork
 from funasr.modules.nets_utils import get_transducer_task_io
 from funasr.layers.abs_normalize import AbsNormalize
 from funasr.torch_utils.device_funcs import force_gatherable

+ 0 - 0
funasr/models/joint_network.py → funasr/models/joint_net/joint_network.py


+ 1 - 1
funasr/modules/beam_search/beam_search_transducer.py

@@ -7,7 +7,7 @@ import numpy as np
 import torch
 
 from funasr.models.rnnt_predictor.abs_decoder import AbsDecoder
-from funasr.models.joint_network import JointNetwork
+from funasr.models.joint_net.joint_network import JointNetwork
 
 
 @dataclass

+ 1 - 1
funasr/modules/e2e_asr_common.py

@@ -19,7 +19,7 @@ import torch
 
 from funasr.modules.beam_search.beam_search_transducer import BeamSearchTransducer
 from funasr.models.rnnt_predictor.abs_decoder import AbsDecoder
-from funasr.models.joint_network import JointNetwork
+from funasr.models.joint_net.joint_network import JointNetwork
 
 def end_detect(ended_hyps, i, M=3, D_end=np.log(1 * np.exp(-10))):
     """End detection.

+ 3 - 3
funasr/tasks/asr_transducer.py

@@ -25,9 +25,9 @@ from funasr.models.rnnt_predictor.abs_decoder import AbsDecoder
 from funasr.models.rnnt_predictor.rnn_decoder import RNNDecoder
 from funasr.models.rnnt_predictor.stateless_decoder import StatelessDecoder
 from funasr.models.encoder.conformer_encoder import ConformerChunkEncoder
-from funasr.models.e2e_transducer import TransducerModel
-from funasr.models.e2e_transducer_unified import UnifiedTransducerModel
-from funasr.models.joint_network import JointNetwork
+from funasr.models.e2e_asr_transducer import TransducerModel
+from funasr.models.e2e_asr_transducer_unified import UnifiedTransducerModel
+from funasr.models.joint_net.joint_network import JointNetwork
 from funasr.layers.abs_normalize import AbsNormalize
 from funasr.layers.global_mvn import GlobalMVN
 from funasr.layers.utterance_mvn import UtteranceMVN