Merge pull request #239 from msghik/feature/fine-tuned-model-support
feat: Allow loading fine-tuned models in simulstreaming
This commit is contained in:
commit
40bff38933
1 changed files with 7 additions and 3 deletions
|
|
@ -210,11 +210,15 @@ class SimulStreamingASR():
|
||||||
else:
|
else:
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
|
|
||||||
self.model_name = os.path.basename(self.cfg.model_path).replace(".pt", "")
|
if model_dir:
|
||||||
self.model_path = os.path.dirname(os.path.abspath(self.cfg.model_path))
|
self.model_name = model_dir
|
||||||
|
self.model_path = None
|
||||||
|
else:
|
||||||
|
self.model_name = os.path.basename(self.cfg.model_path).replace(".pt", "")
|
||||||
|
self.model_path = os.path.dirname(os.path.abspath(self.cfg.model_path))
|
||||||
|
|
||||||
self.mlx_encoder, self.fw_encoder = None, None
|
self.mlx_encoder, self.fw_encoder = None, None
|
||||||
if not self.disable_fast_encoder:
|
if not self.disable_fast_encoder and not model_dir:
|
||||||
if HAS_MLX_WHISPER:
|
if HAS_MLX_WHISPER:
|
||||||
print('Simulstreaming will use MLX whisper for a faster encoder.')
|
print('Simulstreaming will use MLX whisper for a faster encoder.')
|
||||||
mlx_model_name = mlx_model_mapping[self.model_name]
|
mlx_model_name = mlx_model_mapping[self.model_name]
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue