fix(onnx-asr): корректный маппинг compute_type в onnx-asr quantization
При 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>
This commit is contained in:
@@ -15,6 +15,32 @@ MODEL_ALIASES: dict[str, str] = {
|
|||||||
|
|
||||||
SUPPORTED_ALIASES = ", ".join(MODEL_ALIASES)
|
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:
|
class OnnxAsrBackend:
|
||||||
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
"""Бэкенд транскрипции через onnx-asr (ONNX Runtime)."""
|
||||||
@@ -48,17 +74,16 @@ class OnnxAsrBackend:
|
|||||||
) -> Any:
|
) -> Any:
|
||||||
"""Creates onnx-asr model with VAD.
|
"""Creates onnx-asr model with VAD.
|
||||||
|
|
||||||
model_path: onnx-asr model identifier (e.g. "gigaam-v3-ctc").
|
compute_type маппится в onnx-asr ``quantization`` — это суффикс файла
|
||||||
compute_type: "int8", "fp16", or "float32" — passed as quantization.
|
модели; для unquantized (float32/fp32) нужно None, не строку.
|
||||||
cpu_threads: not used by onnx-asr (onnxruntime manages threads internally).
|
|
||||||
"""
|
"""
|
||||||
import onnx_asr
|
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 = onnx_asr.load_model(
|
||||||
model=model_path,
|
model=model_path,
|
||||||
quantization=ct,
|
quantization=quantization,
|
||||||
)
|
)
|
||||||
vad = onnx_asr.load_vad("silero")
|
vad = onnx_asr.load_vad("silero")
|
||||||
self._vad = vad
|
self._vad = vad
|
||||||
|
|||||||
@@ -104,6 +104,79 @@ class TestCreateModel:
|
|||||||
|
|
||||||
assert calls == ["fp16"]
|
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:
|
class TestTranscribe:
|
||||||
def test_transcribe_collects_segments(self, monkeypatch, tmp_path):
|
def test_transcribe_collects_segments(self, monkeypatch, tmp_path):
|
||||||
|
|||||||
Reference in New Issue
Block a user