Huge refactoring
This commit is contained in:
83
asr.py
Normal file
83
asr.py
Normal file
@@ -0,0 +1,83 @@
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
from pydantic import BaseModel
|
||||
|
||||
import whisper
|
||||
|
||||
class AsrRawSegment(BaseModel):
|
||||
"""Segment produced by audio recognition engine"""
|
||||
|
||||
start: float
|
||||
"""Start of the segment"""
|
||||
|
||||
end: float
|
||||
"""End of the segment"""
|
||||
|
||||
text: str
|
||||
"""Text of the segment"""
|
||||
|
||||
engine: dict[str, Any]
|
||||
"""Engine-related data"""
|
||||
|
||||
class AsrRawResult(BaseModel):
|
||||
"""Result of transcribing"""
|
||||
|
||||
engine: str
|
||||
"""Name of the engine that was used for transcribing"""
|
||||
|
||||
segments: list[AsrRawSegment]
|
||||
"""Segments produced by the engine"""
|
||||
|
||||
class Asr:
|
||||
"""This class performs transcription of the audio file."""
|
||||
@staticmethod
|
||||
def get_models_list() -> list[str]:
|
||||
"""Returns list of allowed `model` values for constructor."""
|
||||
return whisper.available_models()
|
||||
|
||||
def __init__(self, model: str, **kwargs) -> None:
|
||||
"""Create the transcriber instance.
|
||||
|
||||
Args:
|
||||
- model - model to use (call `get_models_list` to get the list of
|
||||
available models)
|
||||
- **kwargs are passed to `whisper.load_model(...)`
|
||||
"""
|
||||
allowed_models = self.get_models_list()
|
||||
if model not in allowed_models:
|
||||
raise ValueError(
|
||||
f"Model `{model}` is not available. "
|
||||
f"Use one of {', '.join(f'`{s}`' for s in allowed_models)}"
|
||||
)
|
||||
self._model = whisper.load_model(model, **kwargs)
|
||||
|
||||
def recognize(self, path: str, **kwargs) -> AsrRawResult:
|
||||
"""Transcribe audiofile. The operation will take a lot of time for large
|
||||
files.
|
||||
|
||||
Args:
|
||||
- path - path to the file to transcribe.
|
||||
- **kwargs - passed to `transcribe()`
|
||||
|
||||
Returns:
|
||||
- result of transcribing
|
||||
"""
|
||||
raw_segments: list[dict]
|
||||
raw_segments = self._model.transcribe(path, **kwargs)["segments"] # type: ignore
|
||||
result = AsrRawResult(
|
||||
engine="whisper",
|
||||
segments=[]
|
||||
)
|
||||
for raw_segment in raw_segments:
|
||||
e = AsrRawSegment(
|
||||
start=float(raw_segment["start"]),
|
||||
end=float(raw_segment["end"]),
|
||||
text=str(raw_segment["text"]),
|
||||
engine={
|
||||
"whisper_temperature": float(raw_segment["temperature"]),
|
||||
"whisper_avg_logprob": float(raw_segment["avg_logprob"]),
|
||||
"whisper_no_speech_prob": float(raw_segment["no_speech_prob"])
|
||||
}
|
||||
)
|
||||
result.segments.append(e)
|
||||
return result
|
||||
Reference in New Issue
Block a user