feat: Update transcriber to support large-v3 model with 128 mel filters

This commit is contained in:
华晨
2024-01-01 20:22:37 +08:00
parent 5918b5ed42
commit 02793a93f8
+20 -1
View File
@@ -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],