Auto convert hf custom whisper to ct2(faster-whisper)

Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
This commit is contained in:
makaveli10
2025-05-29 09:23:21 +00:00
parent 4b46371dac
commit bff88ed3e7
2 changed files with 42 additions and 25 deletions
+1
View File
@@ -11,6 +11,7 @@ evaluate
numpy<2 numpy<2
openai-whisper==20240930 openai-whisper==20240930
tokenizers==0.20.3 tokenizers==0.20.3
transformers[torch]
# openvino # openvino
librosa librosa
+41 -25
View File
@@ -1,8 +1,11 @@
import os
import json import json
import logging import logging
import threading import threading
import time import time
import torch import torch
import ctranslate2
from huggingface_hub import snapshot_download
from whisper_live.transcriber.transcriber_faster_whisper import WhisperModel from whisper_live.transcriber.transcriber_faster_whisper import WhisperModel
from whisper_live.backend.base import ServeClientBase from whisper_live.backend.base import ServeClientBase
@@ -118,38 +121,51 @@ class ServeClientFasterWhisper(ServeClientBase):
def create_model(self, device): def create_model(self, device):
""" """
Instantiates a new model, sets it as the transcriber. Instantiates a new model, sets it as the transcriber. If model is a huggingface model_id
then it is automatically converted to ctranslate2(faster_whisper) format.
""" """
model_ref = self.model_size_or_path
if model_ref in self.model_sizes:
model_to_load = model_ref
else:
logging.info(f"Model not in model_sizes")
if os.path.isdir(model_ref) and ctranslate2.contains_model(model_ref):
model_to_load = model_ref
else:
local_snapshot = snapshot_download(
repo_id = model_ref,
repo_type = "model",
)
if ctranslate2.contains_model(local_snapshot):
model_to_load = local_snapshot
else:
cache_root = os.path.expanduser("~/.cache/whisper-live/whisper-ct2-models/")
os.makedirs(cache_root, exist_ok=True)
safe_name = model_ref.replace("/", "--")
ct2_dir = os.path.join(cache_root, safe_name)
if not ctranslate2.contains_model(ct2_dir):
logging.info(f"Converting '{model_ref}' to CTranslate2 @ {ct2_dir}")
ct2_converter = ctranslate2.converters.TransformersConverter(
local_snapshot,
copy_files=["tokenizer.json", "preprocessor_config.json"]
)
ct2_converter.convert(
output_dir=ct2_dir,
quantization=self.compute_type,
force=False, # skip if already up-to-date
)
model_to_load = ct2_dir
logging.info(f"Loading model: {model_to_load}")
self.transcriber = WhisperModel( self.transcriber = WhisperModel(
self.model_size_or_path, model_to_load,
device=device, device=device,
compute_type=self.compute_type, compute_type=self.compute_type,
local_files_only=False, local_files_only=False,
) )
def check_valid_model(self, model_size):
"""
Check if it's a valid whisper model size.
Args:
model_size (str): The name of the model size to check.
Returns:
str: The model size if valid, None otherwise.
"""
if model_size not in self.model_sizes:
self.websocket.send(
json.dumps(
{
"uid": self.client_uid,
"status": "ERROR",
"message": f"Invalid model size {model_size}. Available choices: {self.model_sizes}"
}
)
)
return None
return model_size
def set_language(self, info): def set_language(self, info):
""" """
Updates the language attribute based on the detected language information. Updates the language attribute based on the detected language information.