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):