feat(onnx-asr): реализовать create_model с Silero VAD
- Зачем: - 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
This commit is contained in:
@@ -36,6 +36,32 @@ class OnnxAsrBackend:
|
|||||||
self._resolved_model_id = self._resolve_model(model_name)
|
self._resolved_model_id = self._resolve_model(model_name)
|
||||||
return self._resolved_model_id
|
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:
|
def _resolve_model(self, model_name: str) -> str:
|
||||||
"""Resolve alias to onnx-asr model name. Raw names pass through."""
|
"""Resolve alias to onnx-asr model name. Raw names pass through."""
|
||||||
if model_name in MODEL_ALIASES:
|
if model_name in MODEL_ALIASES:
|
||||||
|
|||||||
@@ -22,6 +22,79 @@ class TestEnsureModelAvailable:
|
|||||||
assert backend.actual_compute_type == "float32"
|
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:
|
class TestModelAliases:
|
||||||
def test_gigaam_v3_resolves(self):
|
def test_gigaam_v3_resolves(self):
|
||||||
backend = OnnxAsrBackend()
|
backend = OnnxAsrBackend()
|
||||||
|
|||||||
Reference in New Issue
Block a user