feat: Allow loading fine-tuned models in simulstreaming
This change modifies the `simulstreaming` backend to support loading fine-tuned Whisper models via the `--model_dir` argument. The `SimulStreamingASR` class has been updated to: - Use the `model_dir` path directly to load the model, which is the correct procedure for fine-tuned `.pt` files. - Automatically disable the `faster-whisper` and `mlx-whisper` fast encoders when `model_dir` is used, as they are not compatible with standard fine-tuned models. The call site in `core.py` already passed the `model_dir` argument, so no changes were needed there. This change makes the `simulstreaming` backend more flexible and allows users to leverage their own custom models.
This commit is contained in:
parent
d55490cd27
commit
70e854b346
1 changed files with 7 additions and 3 deletions
|
|
@ -210,11 +210,15 @@ class SimulStreamingASR():
|
||||||
else:
|
else:
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
|
|
||||||
|
if model_dir:
|
||||||
|
self.model_name = model_dir
|
||||||
|
self.model_path = None
|
||||||
|
else:
|
||||||
self.model_name = os.path.basename(self.cfg.model_path).replace(".pt", "")
|
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.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