При adversarial review (gpt-5.5) обнаружено: пользователь с config
compute_type = "float32" и --device onnx падал на старте.
HARDCODED_DEFAULTS["compute_type"] = "float32", apply_device_defaults не
заменяет значение если оно есть в config — старая проверка
(compute_type in ("int8", "fp16", "float32")) пропускала "float32"
дальше как onnx_asr quantization, что заставляло искать несуществующий
файл с суффиксом _float32.
- _normalize_quantization() — explicit маппинг в onnx-asr quantization.
- float32/fp32 → None (unquantized loading в onnx-asr — это None, не строка).
- float16 → fp16 (CUDA-naming → onnx-asr-naming).
- int8/fp16 → pass-through.
- Неизвестные compute_type (например, int8_float32 от CTranslate2) → ValueError
вместо silent fallback на int8 — пользователь раньше получал не ту
модель без предупреждения.
4 новых теста: float32→None, fp32→None, float16→fp16, unknown→raises.
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
155 lines
5.6 KiB
Python
155 lines
5.6 KiB
Python
"""Бэкенд транскрипции на основе onnx-asr (GigaAM, Parakeet, FastConformer)."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Callable
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from local_transcriber.types import Segment, TranscribeResult
|
|
|
|
MODEL_ALIASES: dict[str, str] = {
|
|
"gigaam-v3": "gigaam-v3-ctc",
|
|
"parakeet-v3": "nemo-parakeet-tdt-0.6b-v3",
|
|
}
|
|
|
|
SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
|
|
|
|
# compute_type проекта → onnx-asr quantization (file suffix; None = unquantized).
|
|
_QUANTIZATION_MAP: dict[str, str | None] = {
|
|
"int8": "int8",
|
|
"fp16": "fp16",
|
|
"float16": "fp16",
|
|
"float32": None,
|
|
"fp32": None,
|
|
}
|
|
|
|
|
|
def _normalize_quantization(compute_type: str) -> str | None:
|
|
"""Маппит compute_type проекта в значение onnx-asr ``quantization``.
|
|
|
|
onnx-asr использует ``quantization`` как суффикс имени файла модели:
|
|
``int8``/``fp16`` подгружают квантизованные веса, ``None`` — unquantized
|
|
(float32). Передача ``"float32"`` строкой пытается найти несуществующий
|
|
файл с суффиксом ``_float32`` и приводит к ошибке загрузки.
|
|
"""
|
|
if compute_type not in _QUANTIZATION_MAP:
|
|
supported = ", ".join(sorted(_QUANTIZATION_MAP))
|
|
raise ValueError(
|
|
f"Неподдерживаемый compute_type '{compute_type}' для onnx-asr. "
|
|
f"Допустимо: {supported}."
|
|
)
|
|
return _QUANTIZATION_MAP[compute_type]
|
|
|
|
|
|
class OnnxAsrBackend:
|
|
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
|
|
|
def __init__(self):
|
|
self.actual_compute_type: str | None = None
|
|
self._resolved_model_id: str | None = None
|
|
self._vad: Any = None
|
|
|
|
def ensure_model_available(
|
|
self,
|
|
model_name: str,
|
|
compute_type: str,
|
|
on_status: Callable[[str], None] | None = None,
|
|
) -> str:
|
|
"""Resolves model alias and returns the onnx-asr model identifier.
|
|
|
|
onnx-asr downloads models automatically via load_model(),
|
|
so this just validates the alias and returns the identifier string.
|
|
"""
|
|
self.actual_compute_type = compute_type
|
|
self._resolved_model_id = self._resolve_model(model_name)
|
|
return self._resolved_model_id
|
|
|
|
def create_model(
|
|
self,
|
|
model_path: str,
|
|
device: str,
|
|
compute_type: str,
|
|
cpu_threads: int = 0,
|
|
) -> Any:
|
|
"""Creates onnx-asr model with VAD.
|
|
|
|
compute_type маппится в onnx-asr ``quantization`` — это суффикс файла
|
|
модели; для unquantized (float32/fp32) нужно None, не строку.
|
|
"""
|
|
import onnx_asr
|
|
|
|
quantization = _normalize_quantization(compute_type)
|
|
|
|
model = onnx_asr.load_model(
|
|
model=model_path,
|
|
quantization=quantization,
|
|
)
|
|
vad = onnx_asr.load_vad("silero")
|
|
self._vad = vad
|
|
return model.with_vad(vad)
|
|
|
|
def transcribe(
|
|
self,
|
|
model: Any,
|
|
file_path: Path,
|
|
language: str | None,
|
|
on_segment: Callable[[Segment], None] | None = None,
|
|
on_status: Callable[[str], None] | None = None,
|
|
) -> TranscribeResult:
|
|
"""Transcribes audio file using onnx-asr model with VAD.
|
|
|
|
model: result of create_model() — a SegmentResultsAsrAdapter.
|
|
file_path: path to audio/video file (any format supported by faster-whisper decode).
|
|
language: language code (e.g. "ru", "en") — only meaningful for multilingual models.
|
|
"""
|
|
from faster_whisper import decode_audio
|
|
|
|
_notify(on_status, "Загружаю аудио...")
|
|
audio_array = decode_audio(str(file_path), sampling_rate=16000)
|
|
duration = len(audio_array) / 16000.0
|
|
|
|
_notify(on_status, "Транскрибирую (onnx-asr)...")
|
|
segments: list[Segment] = []
|
|
detected_language = language or "unknown"
|
|
|
|
for vad_seg in model.recognize(audio_array, sample_rate=16000, language=language):
|
|
seg = Segment(
|
|
start=max(0.0, vad_seg.start),
|
|
end=max(0.0, vad_seg.end),
|
|
text=vad_seg.text,
|
|
)
|
|
if on_segment is not None:
|
|
on_segment(seg)
|
|
segments.append(seg)
|
|
_notify(
|
|
on_status,
|
|
f"Транскрибирую (onnx-asr)... [{len(segments)} сегм.]",
|
|
)
|
|
|
|
return TranscribeResult(
|
|
segments=segments,
|
|
language=detected_language,
|
|
language_probability=1.0 if language else 0.0,
|
|
duration=duration,
|
|
device_used="", # оркестратор проставит
|
|
)
|
|
|
|
def _resolve_model(self, model_name: str) -> str:
|
|
"""Resolve alias to onnx-asr model name. Raw names pass through."""
|
|
if model_name in MODEL_ALIASES:
|
|
return MODEL_ALIASES[model_name]
|
|
if "/" in model_name or model_name.count("-") >= 2:
|
|
# Looks like a raw onnx-asr name — allow passthrough
|
|
return model_name
|
|
raise ValueError(
|
|
f"Неподдерживаемая модель '{model_name}'. "
|
|
f"Доступные алиасы: {SUPPORTED_ALIASES}. "
|
|
f"Либо укажите полное имя модели onnx-asr."
|
|
)
|
|
|
|
|
|
def _notify(on_status: Callable[[str], None] | None, message: str) -> None:
|
|
if on_status is not None:
|
|
on_status(message)
|