From f25a546754828a11dc283553ebe35042038bbd1f Mon Sep 17 00:00:00 2001 From: Dmitry Dementev Date: Sat, 25 Apr 2026 21:20:35 +0300 Subject: [PATCH] =?UTF-8?q?feat(onnx-asr):=20=D1=80=D0=B5=D0=B0=D0=BB?= =?UTF-8?q?=D0=B8=D0=B7=D0=BE=D0=B2=D0=B0=D1=82=D1=8C=20ensure=5Fmodel=5Fa?= =?UTF-8?q?vailable?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Зачем: - Бэкенд должен резолвить алиасы моделей и сохранять compute_type для дальнейшего использования. - Что: - добавлен конструктор __init__ с полями actual_compute_type, _resolved_model_id, _vad. - метод ensure_model_available валидирует алиас и возвращает идентификатор модели onnx-asr. - написаны 3 теста на резолвинг и сохранение compute_type. - Проверка: - uv run pytest tests/test_onnx_asr.py -v --- src/local_transcriber/backends/onnx_asr.py | 23 ++++++++++++++++++++++ tests/test_onnx_asr.py | 18 +++++++++++++++++ 2 files changed, 41 insertions(+) diff --git a/src/local_transcriber/backends/onnx_asr.py b/src/local_transcriber/backends/onnx_asr.py index 190c6f2..770c412 100644 --- a/src/local_transcriber/backends/onnx_asr.py +++ b/src/local_transcriber/backends/onnx_asr.py @@ -2,6 +2,9 @@ from __future__ import annotations +from collections.abc import Callable +from typing import Any + MODEL_ALIASES: dict[str, str] = { "gigaam-v3": "gigaam-v3-ctc", "parakeet-v3": "nemo-parakeet-tdt-0.6b-v3", @@ -13,6 +16,26 @@ SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES) 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 _resolve_model(self, model_name: str) -> str: """Resolve alias to onnx-asr model name. Raw names pass through.""" if model_name in MODEL_ALIASES: diff --git a/tests/test_onnx_asr.py b/tests/test_onnx_asr.py index 4cceb3a..a1f842d 100644 --- a/tests/test_onnx_asr.py +++ b/tests/test_onnx_asr.py @@ -4,6 +4,24 @@ import pytest from local_transcriber.backends.onnx_asr import OnnxAsrBackend, MODEL_ALIASES +class TestEnsureModelAvailable: + def test_returns_model_id_for_gigaam(self): + backend = OnnxAsrBackend() + result = backend.ensure_model_available("gigaam-v3", "int8") + assert result == "gigaam-v3-ctc" + + def test_returns_model_id_for_parakeet(self): + backend = OnnxAsrBackend() + result = backend.ensure_model_available("parakeet-v3", "fp16") + assert result == "nemo-parakeet-tdt-0.6b-v3" + + def test_stores_compute_type(self): + backend = OnnxAsrBackend() + backend.ensure_model_available("gigaam-v3", "float32") + assert backend._resolved_model_id == "gigaam-v3-ctc" + assert backend.actual_compute_type == "float32" + + class TestModelAliases: def test_gigaam_v3_resolves(self): backend = OnnxAsrBackend()