From d9c9aefdb385cb7db8730e64be3c8e3b95d560d7 Mon Sep 17 00:00:00 2001 From: Dmitry Dementev Date: Sat, 25 Apr 2026 21:22:02 +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=20create=5Fmodel=20?= =?UTF-8?q?=D1=81=20Silero=20VAD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Зачем: - onnx-asr модель должна создаваться с квантизацией и VAD для разбивки аудио на сегменты. - Что: - метод create_model вызывает onnx_asr.load_model с квантизацией и cpu_preprocessing=True. - подгружается Silero VAD, прикрепляется к модели через with_vad. - написаны 3 теста: проверка аргументов load_model, загрузка VAD, передача fp16. - Проверка: - uv run pytest tests/test_onnx_asr.py -v --- src/local_transcriber/backends/onnx_asr.py | 26 ++++++++ tests/test_onnx_asr.py | 73 ++++++++++++++++++++++ 2 files changed, 99 insertions(+) diff --git a/src/local_transcriber/backends/onnx_asr.py b/src/local_transcriber/backends/onnx_asr.py index 770c412..7411767 100644 --- a/src/local_transcriber/backends/onnx_asr.py +++ b/src/local_transcriber/backends/onnx_asr.py @@ -36,6 +36,32 @@ class OnnxAsrBackend: 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. + + model_path: onnx-asr model identifier (e.g. "gigaam-v3-ctc"). + compute_type: "int8", "fp16", or "float32" — passed as quantization. + cpu_threads: not used by onnx-asr (onnxruntime manages threads internally). + """ + import onnx_asr + + ct = compute_type if compute_type in ("int8", "fp16", "float32") else "int8" + + model = onnx_asr.load_model( + model=model_path, + quantization=ct, + cpu_preprocessing=True, + ) + vad = onnx_asr.load_vad("silero") + self._vad = vad + return model.with_vad(vad) + 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 a1f842d..7a6a81b 100644 --- a/tests/test_onnx_asr.py +++ b/tests/test_onnx_asr.py @@ -22,6 +22,79 @@ class TestEnsureModelAvailable: assert backend.actual_compute_type == "float32" +class TestCreateModel: + def test_calls_load_model_with_correct_args(self, monkeypatch): + """Verify create_model passes correct args to onnx_asr.load_model.""" + calls = [] + + def fake_load_model(model=None, path=None, quantization=None, + cpu_preprocessing=None, **kwargs): + calls.append({ + "model": model, "path": path, "quantization": quantization, + "cpu_preprocessing": cpu_preprocessing, + }) + return FakeAsrAdapter() + + class FakeAsrAdapter: + def with_vad(self, vad): + return self + + monkeypatch.setattr("onnx_asr.load_model", fake_load_model) + + backend = OnnxAsrBackend() + backend.actual_compute_type = "int8" + model = backend.create_model("gigaam-v3-ctc", "onnx", "int8") + + assert len(calls) == 1 + assert calls[0]["quantization"] == "int8" + assert calls[0]["cpu_preprocessing"] is True + assert model is not None + + def test_loads_silero_vad(self, monkeypatch): + """Verify Silero VAD is loaded and attached to model.""" + vad_calls = [] + + def fake_load_vad(model, **kwargs): + vad_calls.append(model) + return "fake_vad" + + def fake_load_model(**kwargs): + return FakeAsrAdapter() + + class FakeAsrAdapter: + def with_vad(self, vad): + self._vad = vad + return self + + monkeypatch.setattr("onnx_asr.load_model", fake_load_model) + monkeypatch.setattr("onnx_asr.load_vad", fake_load_vad) + + backend = OnnxAsrBackend() + model = backend.create_model("gigaam-v3-ctc", "onnx", "int8") + + assert vad_calls == ["silero"] + + def test_fp16_compute_type(self, monkeypatch): + """Verify fp16 compute_type is passed through.""" + calls = [] + + def fake_load_model(model=None, quantization=None, **kwargs): + calls.append(quantization) + return FakeAsrAdapter() + + class FakeAsrAdapter: + def with_vad(self, vad): + return self + + monkeypatch.setattr("onnx_asr.load_model", fake_load_model) + monkeypatch.setattr("onnx_asr.load_vad", lambda model, **kw: None) + + backend = OnnxAsrBackend() + backend.create_model("parakeet-v3", "onnx", "fp16") + + assert calls == ["fp16"] + + class TestModelAliases: def test_gigaam_v3_resolves(self): backend = OnnxAsrBackend()