diff --git a/funasr/models/e2e_transducer.py b/funasr/models/e2e_asr_transducer.py similarity index 99% rename from funasr/models/e2e_transducer.py rename to funasr/models/e2e_asr_transducer.py index 460a6d796..6eb002320 100644 --- a/funasr/models/e2e_transducer.py +++ b/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 diff --git a/funasr/models/e2e_transducer_unified.py b/funasr/models/e2e_asr_transducer_unified.py similarity index 99% rename from funasr/models/e2e_transducer_unified.py rename to funasr/models/e2e_asr_transducer_unified.py index f79ba57c4..ad61d12c0 100644 --- a/funasr/models/e2e_transducer_unified.py +++ b/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 diff --git a/funasr/models/joint_network.py b/funasr/models/joint_net/joint_network.py similarity index 100% rename from funasr/models/joint_network.py rename to funasr/models/joint_net/joint_network.py diff --git a/funasr/modules/beam_search/beam_search_transducer.py b/funasr/modules/beam_search/beam_search_transducer.py index 49cce92a1..8b7e613fe 100644 --- a/funasr/modules/beam_search/beam_search_transducer.py +++ b/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 diff --git a/funasr/modules/e2e_asr_common.py b/funasr/modules/e2e_asr_common.py index 3746036ba..a01cd5ef1 100644 --- a/funasr/modules/e2e_asr_common.py +++ b/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. diff --git a/funasr/tasks/asr_transducer.py b/funasr/tasks/asr_transducer.py index 99b3d0c2e..d4136d068 100644 --- a/funasr/tasks/asr_transducer.py +++ b/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