Auto convert hf custom whisper to ct2(faster-whisper)
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user