From 9cfa437b3338be27a4fd0260c658175b23286e8c Mon Sep 17 00:00:00 2001 From: Dmitry Dementev Date: Sun, 26 Apr 2026 00:06:22 +0300 Subject: [PATCH] =?UTF-8?q?fix(onnx-asr):=20=D0=BA=D0=BE=D1=80=D1=80=D0=B5?= =?UTF-8?q?=D0=BA=D1=82=D0=BD=D1=8B=D0=B9=20=D0=BC=D0=B0=D0=BF=D0=BF=D0=B8?= =?UTF-8?q?=D0=BD=D0=B3=20compute=5Ftype=20=D0=B2=20onnx-asr=20quantizatio?= =?UTF-8?q?n?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit При 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) --- src/local_transcriber/backends/onnx_asr.py | 35 +++++++++-- tests/test_onnx_asr.py | 73 ++++++++++++++++++++++ 2 files changed, 103 insertions(+), 5 deletions(-) diff --git a/src/local_transcriber/backends/onnx_asr.py b/src/local_transcriber/backends/onnx_asr.py index 664cba0..f3d7b87 100644 --- a/src/local_transcriber/backends/onnx_asr.py +++ b/src/local_transcriber/backends/onnx_asr.py @@ -15,6 +15,32 @@ MODEL_ALIASES: dict[str, str] = { 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).""" @@ -48,17 +74,16 @@ class OnnxAsrBackend: ) -> 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). + compute_type маппится в onnx-asr ``quantization`` — это суффикс файла + модели; для unquantized (float32/fp32) нужно None, не строку. """ import onnx_asr - ct = compute_type if compute_type in ("int8", "fp16", "float32") else "int8" + quantization = _normalize_quantization(compute_type) model = onnx_asr.load_model( model=model_path, - quantization=ct, + quantization=quantization, ) vad = onnx_asr.load_vad("silero") self._vad = vad diff --git a/tests/test_onnx_asr.py b/tests/test_onnx_asr.py index 207e60f..ff192eb 100644 --- a/tests/test_onnx_asr.py +++ b/tests/test_onnx_asr.py @@ -104,6 +104,79 @@ class TestCreateModel: assert calls == ["fp16"] + def test_float32_maps_to_none(self, monkeypatch): + """compute_type='float32' маппится в quantization=None (unquantized). + + onnx-asr использует quantization как суффикс файла; для float32 нужен None, + строка "float32" приведёт к попытке загрузить несуществующий файл. + """ + calls = [] + + def fake_load_model(model=None, quantization="MISSING", **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("gigaam-v3-ctc", "onnx", "float32") + + assert calls == [None] + + def test_fp32_maps_to_none(self, monkeypatch): + """compute_type='fp32' тоже маппится в quantization=None.""" + calls = [] + + def fake_load_model(model=None, quantization="MISSING", **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("gigaam-v3-ctc", "onnx", "fp32") + + assert calls == [None] + + def test_float16_alias_maps_to_fp16(self, monkeypatch): + """compute_type='float16' (CUDA-naming) маппится в onnx-asr 'fp16'.""" + 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("gigaam-v3-ctc", "onnx", "float16") + + assert calls == ["fp16"] + + def test_unknown_compute_type_raises(self, monkeypatch): + """Неподдерживаемый compute_type → ValueError, не silent fallback.""" + monkeypatch.setattr("onnx_asr.load_model", lambda **kw: None) + monkeypatch.setattr("onnx_asr.load_vad", lambda model, **kw: None) + + backend = OnnxAsrBackend() + with pytest.raises(ValueError, match="Неподдерживаемый compute_type"): + backend.create_model("gigaam-v3-ctc", "onnx", "int8_float32") + class TestTranscribe: def test_transcribe_collects_segments(self, monkeypatch, tmp_path):