feat: Update transcriber to support large-v3 model with 128 mel filters
This commit is contained in:
@@ -4,6 +4,8 @@ import itertools
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import zlib
|
import zlib
|
||||||
|
import json
|
||||||
|
from inspect import signature
|
||||||
|
|
||||||
from typing import BinaryIO, Iterable, List, NamedTuple, Optional, Tuple, Union
|
from typing import BinaryIO, Iterable, List, NamedTuple, Optional, Tuple, Union
|
||||||
|
|
||||||
@@ -144,7 +146,8 @@ class WhisperModel:
|
|||||||
"openai/whisper-tiny" + ("" if self.model.is_multilingual else ".en")
|
"openai/whisper-tiny" + ("" if self.model.is_multilingual else ".en")
|
||||||
)
|
)
|
||||||
|
|
||||||
self.feature_extractor = FeatureExtractor()
|
self.feat_kwargs = self._get_feature_kwargs(model_path)
|
||||||
|
self.feature_extractor = FeatureExtractor(**self.feat_kwargs)
|
||||||
self.num_samples_per_token = self.feature_extractor.hop_length * 2
|
self.num_samples_per_token = self.feature_extractor.hop_length * 2
|
||||||
self.frames_per_second = (
|
self.frames_per_second = (
|
||||||
self.feature_extractor.sampling_rate // self.feature_extractor.hop_length
|
self.feature_extractor.sampling_rate // self.feature_extractor.hop_length
|
||||||
@@ -161,6 +164,22 @@ class WhisperModel:
|
|||||||
"""The languages supported by the model."""
|
"""The languages supported by the model."""
|
||||||
return list(_LANGUAGE_CODES) if self.model.is_multilingual else ["en"]
|
return list(_LANGUAGE_CODES) if self.model.is_multilingual else ["en"]
|
||||||
|
|
||||||
|
def _get_feature_kwargs(self, model_path) -> dict:
|
||||||
|
preprocessor_config_file = os.path.join(model_path, "preprocessor_config.json")
|
||||||
|
config = {}
|
||||||
|
if os.path.isfile(preprocessor_config_file):
|
||||||
|
try:
|
||||||
|
with open(preprocessor_config_file, "r", encoding="utf-8") as json_file:
|
||||||
|
config = json.load(json_file)
|
||||||
|
valid_keys = signature(FeatureExtractor.__init__).parameters.keys()
|
||||||
|
config = {k: v for k, v in config.items() if k in valid_keys}
|
||||||
|
except json.JSONDecodeError as e:
|
||||||
|
self.logger.warning(
|
||||||
|
"Could not load preprocessor_config.json: %s", str(e)
|
||||||
|
)
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
def transcribe(
|
def transcribe(
|
||||||
self,
|
self,
|
||||||
audio: Union[str, BinaryIO, np.ndarray],
|
audio: Union[str, BinaryIO, np.ndarray],
|
||||||
|
|||||||
Reference in New Issue
Block a user