feat(backends): add GPU detection, runtime fallback and the OpenVINO backend
Detection is two-phase by construction: the registry only calls probe_toolchain once probe_hardware confirms the vendor, which is what structurally keeps torch from being imported on an Intel-only box. A test asserts exactly that. Hardware probing reads each render node's bound driver rather than loaded kernel modules: /sys/module/xe exists here with zero bound devices while i915 owns the card, so a module-presence check false-positives. Fallback is runtime, not detection-time -- construction and the first generate() sit in the same try, because the render-node permission failure and the OpenCL JIT failure both surface there rather than at device enumeration. A failure demotes the backend process-wide so a 50-job batch does not retry it 50 times, and an explicit --device never falls back silently. NVIDIA and AMD are interface-only: detection is real and the error names the module to implement and the model format required. The cached model is OpenVINO IR and cannot load on CUDA or ROCm, and CTranslate2 has no ROCm support, so those are two separate paths rather than one parameterized one. Model resolution is offline-first, and CACHE_DIR is anchored under XDG rather than the working directory, which the benchmark scripts depend on. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -62,8 +62,15 @@ venv = ".venv"
|
||||
strict = true
|
||||
files = ["src", "tests"]
|
||||
|
||||
[[tool.mypy.overrides]]
|
||||
# These ship no py.typed marker.
|
||||
module = ["openvino.*", "openvino_genai.*", "huggingface_hub.*", "yt_dlp.*"]
|
||||
ignore_missing_imports = true
|
||||
|
||||
[tool.pylint.design]
|
||||
max-attributes = 15
|
||||
# Keyword-only arguments do not carry the readability cost this check guards against.
|
||||
max-args = 9
|
||||
|
||||
[tool.pylint.main]
|
||||
# W0621: pytest fixtures shadow their names by design.
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
"""AMD backend -- interface only.
|
||||
|
||||
Detection and routing are real: if an AMD GPU is present the registry reports
|
||||
it and explains why it is stepping down. The inference path is not implemented,
|
||||
because the model cached on this machine is OpenVINO IR and cannot be loaded by
|
||||
ROCm, and nothing here could be tested against real hardware.
|
||||
|
||||
Implementing it means transformers + a ROCm build of torch. CTranslate2 has no
|
||||
ROCm support, so this cannot share the NVIDIA path.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, ClassVar
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends import probe
|
||||
from ccn_transcribe.backends.base import ToolchainStatus
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from ccn_transcribe.backends.base import TranscribeRequest
|
||||
from ccn_transcribe.transcript import TranscriptResult
|
||||
|
||||
_NOT_IMPLEMENTED = "an AMD GPU was detected, but the AMD backend is not implemented in this build"
|
||||
_HINT = (
|
||||
"Implement src/ccn_transcribe/backends/amd_backend.py using transformers with a "
|
||||
"ROCm build of torch (pip --index-url https://download.pytorch.org/whl/rocm6.2). "
|
||||
"The cached OpenVINO IR model cannot be loaded on ROCm."
|
||||
)
|
||||
|
||||
|
||||
class AmdBackend:
|
||||
name: ClassVar[str] = "amd"
|
||||
device: ClassVar[str] = "hip"
|
||||
default_model: ClassVar[str] = "openai/whisper-large-v3-turbo"
|
||||
|
||||
def __init__(self, model_id: str | None = None, cache_dir: Path | None = None) -> None:
|
||||
self.model_id = model_id or self.default_model
|
||||
self.cache_dir = cache_dir
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool:
|
||||
return probe.has_amd_gpu()
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus:
|
||||
return ToolchainStatus(ok=False, reason=_NOT_IMPLEMENTED)
|
||||
|
||||
def load(self) -> None:
|
||||
raise errors.BackendNotImplementedError(_NOT_IMPLEMENTED, hint=_HINT)
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
pcm: npt.NDArray[np.float32], # noqa: ARG002 - fixed by the Backend protocol
|
||||
request: TranscribeRequest, # noqa: ARG002 - fixed by the Backend protocol
|
||||
) -> TranscriptResult:
|
||||
raise errors.BackendNotImplementedError(_NOT_IMPLEMENTED, hint=_HINT)
|
||||
|
||||
def close(self) -> None:
|
||||
return
|
||||
@@ -0,0 +1,57 @@
|
||||
"""The backend interface.
|
||||
|
||||
Two-phase detection is the contract: ``probe_hardware`` must import nothing
|
||||
heavy, and the registry only calls ``probe_toolchain`` once hardware is confirmed.
|
||||
That ordering is what keeps torch off an Intel-only machine.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, ClassVar, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from ccn_transcribe.transcript import TranscriptResult
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolchainStatus:
|
||||
ok: bool
|
||||
reason: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TranscribeRequest:
|
||||
language: str | None = None
|
||||
task: str = "transcribe"
|
||||
num_beams: int = 1
|
||||
initial_prompt: str | None = None
|
||||
hotwords: str | None = None
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Backend(Protocol):
|
||||
name: ClassVar[str]
|
||||
device: ClassVar[str]
|
||||
default_model: ClassVar[str]
|
||||
|
||||
def __init__(self, model_id: str | None = None, cache_dir: Path | None = None) -> None: ...
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool: ...
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus: ...
|
||||
|
||||
def load(self) -> None: ...
|
||||
|
||||
def transcribe(
|
||||
self, pcm: npt.NDArray[np.float32], request: TranscribeRequest
|
||||
) -> TranscriptResult: ...
|
||||
|
||||
def close(self) -> None: ...
|
||||
@@ -0,0 +1,67 @@
|
||||
"""NVIDIA backend -- interface only.
|
||||
|
||||
Detection and routing are real: if an NVIDIA GPU is present the registry reports
|
||||
it and explains why it is stepping down. The inference path is not implemented,
|
||||
because the model cached on this machine is OpenVINO IR and cannot be loaded by
|
||||
CUDA, and nothing here could be tested against real hardware.
|
||||
|
||||
Implementing it means CTranslate2 (faster-whisper) with a CTranslate2 model; note
|
||||
that CTranslate2 has no ROCm support, so AMD is a genuinely separate path.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, ClassVar
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends import probe
|
||||
from ccn_transcribe.backends.base import ToolchainStatus
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from ccn_transcribe.backends.base import TranscribeRequest
|
||||
from ccn_transcribe.transcript import TranscriptResult
|
||||
|
||||
_NOT_IMPLEMENTED = (
|
||||
"an NVIDIA GPU was detected, but the NVIDIA backend is not implemented in this build"
|
||||
)
|
||||
_HINT = (
|
||||
"Implement src/ccn_transcribe/backends/nvidia_backend.py using faster-whisper "
|
||||
"(CTranslate2) with a CTranslate2 model. The cached OpenVINO IR model cannot "
|
||||
"be loaded on CUDA."
|
||||
)
|
||||
|
||||
|
||||
class NvidiaBackend:
|
||||
name: ClassVar[str] = "nvidia"
|
||||
device: ClassVar[str] = "cuda"
|
||||
default_model: ClassVar[str] = "large-v3-turbo"
|
||||
|
||||
def __init__(self, model_id: str | None = None, cache_dir: Path | None = None) -> None:
|
||||
self.model_id = model_id or self.default_model
|
||||
self.cache_dir = cache_dir
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool:
|
||||
return probe.has_nvidia_gpu()
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus:
|
||||
return ToolchainStatus(ok=False, reason=_NOT_IMPLEMENTED)
|
||||
|
||||
def load(self) -> None:
|
||||
raise errors.BackendNotImplementedError(_NOT_IMPLEMENTED, hint=_HINT)
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
pcm: npt.NDArray[np.float32], # noqa: ARG002 - fixed by the Backend protocol
|
||||
request: TranscribeRequest, # noqa: ARG002 - fixed by the Backend protocol
|
||||
) -> TranscriptResult:
|
||||
raise errors.BackendNotImplementedError(_NOT_IMPLEMENTED, hint=_HINT)
|
||||
|
||||
def close(self) -> None:
|
||||
return
|
||||
@@ -0,0 +1,229 @@
|
||||
"""OpenVINO GenAI backend -- the path this machine actually runs.
|
||||
|
||||
Model resolution is offline-first so a warm cache never touches the network, and
|
||||
CACHE_DIR is anchored under XDG rather than the working directory.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends import probe
|
||||
from ccn_transcribe.backends.base import ToolchainStatus
|
||||
from ccn_transcribe.transcript import TranscriptResult, normalize_segments
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from ccn_transcribe.backends.base import TranscribeRequest
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MODEL = "OpenVINO/whisper-large-v3-turbo-int8-ov"
|
||||
DEFAULT_REVISION = "b568445dd5dc8c695bde596f8acbb4694fd6ba64"
|
||||
SAMPLE_RATE = 16_000
|
||||
|
||||
|
||||
def cache_root() -> Path:
|
||||
base = os.environ.get("XDG_CACHE_HOME") or str(Path.home() / ".cache")
|
||||
return Path(base) / "ccn-transcribe" / "ov_cache"
|
||||
|
||||
|
||||
def language_token(code: str | None, allowed: set[str] | None = None) -> str | None:
|
||||
"""Whisper wants ``<|en|>``, not ``en``."""
|
||||
if code is None:
|
||||
return None
|
||||
token = code if code.startswith("<|") else f"<|{code}|>"
|
||||
if allowed is not None and token not in allowed:
|
||||
valid = ", ".join(sorted(t.strip("<|>") for t in allowed)[:12])
|
||||
raise errors.ConfigError(
|
||||
f"unknown language {code!r}",
|
||||
hint=f"Valid codes include: {valid}, ...",
|
||||
)
|
||||
return token
|
||||
|
||||
|
||||
def supported_languages(model_dir: Path) -> set[str] | None:
|
||||
config = model_dir / "generation_config.json"
|
||||
if not config.exists():
|
||||
return None
|
||||
try:
|
||||
doc = json.loads(config.read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError):
|
||||
return None
|
||||
mapping = doc.get("lang_to_id")
|
||||
return set(mapping) if mapping else None
|
||||
|
||||
|
||||
def _snapshot(model_id: str, revision: str | None, *, local_only: bool) -> str:
|
||||
"""One typed boundary around huggingface_hub, which ships no type information."""
|
||||
# Imported lazily so --help and doctor do not pay for it.
|
||||
from huggingface_hub import ( # pylint: disable=import-outside-toplevel
|
||||
snapshot_download, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
|
||||
path: object = snapshot_download(model_id, revision=revision, local_files_only=local_only)
|
||||
return str(path)
|
||||
|
||||
|
||||
def resolve_model(model_id: str, revision: str | None) -> Path:
|
||||
"""Prefer the local cache; only reach the network when we must.
|
||||
|
||||
A warm cache therefore never makes a network call, which is what lets the
|
||||
pipeline run fully offline.
|
||||
"""
|
||||
try:
|
||||
cached = _snapshot(model_id, revision, local_only=True)
|
||||
except Exception: # pylint: disable=broad-exception-caught
|
||||
# Any failure here just means "not in the cache"; fall through to the network.
|
||||
log.info("%s is not cached; downloading (~790 MB for the default)", model_id)
|
||||
else:
|
||||
return Path(cached)
|
||||
|
||||
try:
|
||||
fetched = _snapshot(model_id, revision, local_only=False)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
raise errors.ModelUnavailableError(
|
||||
f"could not obtain model {model_id}",
|
||||
hint=f"Fetch it manually with: hf download {model_id}",
|
||||
) from exc
|
||||
return Path(fetched)
|
||||
|
||||
|
||||
class _OpenvinoBackend:
|
||||
name: ClassVar[str] = "openvino"
|
||||
device: ClassVar[str] = "CPU"
|
||||
default_model: ClassVar[str] = DEFAULT_MODEL
|
||||
|
||||
def __init__(self, model_id: str | None = None, cache_dir: Path | None = None) -> None:
|
||||
self.model_id = model_id or self.default_model
|
||||
self.revision = DEFAULT_REVISION if self.model_id == DEFAULT_MODEL else None
|
||||
self.cache_dir = cache_dir or (cache_root() / self.device.lower())
|
||||
self.pipe: Any | None = None
|
||||
self.languages: set[str] | None = None
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus:
|
||||
try:
|
||||
import openvino as ov # pylint: disable=import-outside-toplevel
|
||||
except ImportError as exc:
|
||||
return ToolchainStatus(ok=False, reason=f"openvino is not importable: {exc}")
|
||||
devices = ov.Core().available_devices
|
||||
if not cls.device_present(devices):
|
||||
return ToolchainStatus(
|
||||
ok=False,
|
||||
reason=(
|
||||
f"OpenVINO reports {devices}, without {cls.device}. The driver is "
|
||||
"fine or clinfo would have failed; this is a plugin problem."
|
||||
),
|
||||
)
|
||||
return ToolchainStatus(ok=True)
|
||||
|
||||
@classmethod
|
||||
def device_present(cls, devices: list[str]) -> bool:
|
||||
return cls.device in devices
|
||||
|
||||
def load(self) -> None:
|
||||
# Lazy: constructing a pipeline is the only thing that needs this.
|
||||
import openvino_genai as ov_genai # pylint: disable=import-outside-toplevel
|
||||
|
||||
model_dir = resolve_model(self.model_id, self.revision)
|
||||
self.languages = supported_languages(model_dir)
|
||||
self.cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
self.pipe = ov_genai.WhisperPipeline(
|
||||
str(model_dir), self.device, CACHE_DIR=str(self.cache_dir)
|
||||
)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
raise errors.BackendRuntimeError(
|
||||
f"OpenVINO could not build a pipeline on {self.device}: {exc}",
|
||||
hint=self.load_hint(),
|
||||
) from exc
|
||||
|
||||
def load_hint(self) -> str:
|
||||
if self.device != "GPU":
|
||||
return "Try --device cpu, or run `ccn-transcribe doctor`."
|
||||
nodes = [n for n in probe.render_nodes() if n.vendor == "intel"]
|
||||
if nodes and not nodes[0].writable:
|
||||
return (
|
||||
f"{nodes[0].path} is not writable by this user. "
|
||||
"Run: sudo usermod -aG render $USER, then log out and back in."
|
||||
)
|
||||
return "Run `ccn-transcribe doctor` to separate driver problems from plugin ones."
|
||||
|
||||
def transcribe(
|
||||
self, pcm: npt.NDArray[np.float32], request: TranscribeRequest
|
||||
) -> TranscriptResult:
|
||||
if self.pipe is None:
|
||||
self.load()
|
||||
assert self.pipe is not None # noqa: S101 - load() raises otherwise
|
||||
kwargs: dict[str, Any] = {
|
||||
"task": request.task,
|
||||
"return_timestamps": True,
|
||||
"num_beams": request.num_beams,
|
||||
}
|
||||
token = language_token(request.language, self.languages)
|
||||
if token is not None:
|
||||
kwargs["language"] = token
|
||||
if request.initial_prompt:
|
||||
kwargs["initial_prompt"] = request.initial_prompt
|
||||
if request.hotwords:
|
||||
kwargs["hotwords"] = request.hotwords
|
||||
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
raw = self.pipe.generate(pcm, **kwargs)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
raise errors.BackendRuntimeError(
|
||||
f"{self.name}/{self.device} failed during generation: {exc}",
|
||||
hint=self.load_hint(),
|
||||
) from exc
|
||||
wall = time.perf_counter() - started
|
||||
|
||||
audio_s = len(pcm) / SAMPLE_RATE
|
||||
raw_chunks: list[Any] = list(getattr(raw, "chunks", None) or [])
|
||||
chunks: list[tuple[float | None, float | None, str]] = [
|
||||
(c.start_ts, c.end_ts, c.text) for c in raw_chunks
|
||||
]
|
||||
segments = normalize_segments(chunks, audio_s) if chunks else ()
|
||||
return TranscriptResult(
|
||||
segments=segments,
|
||||
language=request.language,
|
||||
backend=self.name,
|
||||
device=self.device,
|
||||
model_id=self.model_id,
|
||||
model_revision=self.revision,
|
||||
audio_s=audio_s,
|
||||
wall_s=wall,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self.pipe = None
|
||||
|
||||
|
||||
class OpenvinoGpuBackend(_OpenvinoBackend):
|
||||
device: ClassVar[str] = "GPU"
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool:
|
||||
return probe.has_intel_gpu()
|
||||
|
||||
@classmethod
|
||||
def device_present(cls, devices: list[str]) -> bool:
|
||||
# Multi-GPU systems enumerate GPU.0, GPU.1, ...
|
||||
return any(d == "GPU" or d.startswith("GPU.") for d in devices)
|
||||
|
||||
|
||||
class OpenvinoCpuBackend(_OpenvinoBackend):
|
||||
device: ClassVar[str] = "CPU"
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Hardware probing: cheap, and it imports nothing heavier than pathlib.
|
||||
|
||||
The registry calls these before any toolchain import, which is what structurally
|
||||
guarantees torch is never imported on a machine with no NVIDIA or AMD GPU.
|
||||
|
||||
Detection reads each render node's *bound driver*, not loaded kernel modules:
|
||||
/sys/module/xe can exist with zero bound devices while i915 actually owns the
|
||||
card, so a module-presence check false-positives.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import stat
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
VENDORS = {"0x8086": "intel", "0x10de": "nvidia", "0x1002": "amd"}
|
||||
INTEL_DRIVERS = ("i915", "xe")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SysfsRoots:
|
||||
drm: Path
|
||||
dev_dri: Path
|
||||
kfd: Path
|
||||
|
||||
|
||||
DEFAULT_ROOTS = SysfsRoots(
|
||||
drm=Path("/sys/class/drm"),
|
||||
dev_dri=Path("/dev/dri"),
|
||||
kfd=Path("/dev/kfd"),
|
||||
)
|
||||
|
||||
DEFAULT_PROC_NVIDIA = Path("/proc/driver/nvidia/version")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RenderNode:
|
||||
path: Path
|
||||
vendor: str
|
||||
driver: str
|
||||
mode: int
|
||||
writable: bool
|
||||
|
||||
|
||||
def render_nodes(roots: SysfsRoots = DEFAULT_ROOTS) -> list[RenderNode]:
|
||||
if not roots.drm.is_dir():
|
||||
return []
|
||||
nodes: list[RenderNode] = []
|
||||
for link in sorted(roots.drm.glob("renderD*")):
|
||||
try:
|
||||
vendor_id = (link / "device" / "vendor").read_text().strip()
|
||||
driver = (link / "device" / "driver").resolve().name
|
||||
except OSError:
|
||||
continue
|
||||
device = roots.dev_dri / link.name
|
||||
mode = 0
|
||||
writable = False
|
||||
if device.exists():
|
||||
mode = stat.S_IMODE(device.stat().st_mode)
|
||||
writable = os.access(device, os.R_OK | os.W_OK)
|
||||
nodes.append(
|
||||
RenderNode(
|
||||
path=device,
|
||||
vendor=VENDORS.get(vendor_id, vendor_id),
|
||||
driver=driver,
|
||||
mode=mode,
|
||||
writable=writable,
|
||||
)
|
||||
)
|
||||
return nodes
|
||||
|
||||
|
||||
def has_intel_gpu(roots: SysfsRoots = DEFAULT_ROOTS) -> bool:
|
||||
return any(n.vendor == "intel" and n.driver in INTEL_DRIVERS for n in render_nodes(roots))
|
||||
|
||||
|
||||
def has_amd_gpu(roots: SysfsRoots = DEFAULT_ROOTS) -> bool:
|
||||
if roots.kfd.exists():
|
||||
return True
|
||||
return any(n.vendor == "amd" for n in render_nodes(roots))
|
||||
|
||||
|
||||
def has_nvidia_gpu(
|
||||
roots: SysfsRoots = DEFAULT_ROOTS, proc_nvidia: Path = DEFAULT_PROC_NVIDIA
|
||||
) -> bool:
|
||||
if proc_nvidia.exists():
|
||||
return True
|
||||
return any(n.vendor == "nvidia" for n in render_nodes(roots))
|
||||
@@ -0,0 +1,151 @@
|
||||
"""Backend selection: detection chain, runtime fallback, and demotion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends.amd_backend import AmdBackend
|
||||
from ccn_transcribe.backends.nvidia_backend import NvidiaBackend
|
||||
from ccn_transcribe.backends.openvino_backend import OpenvinoCpuBackend, OpenvinoGpuBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from ccn_transcribe.backends.base import Backend, TranscribeRequest
|
||||
from ccn_transcribe.transcript import TranscriptResult
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_CHAIN: tuple[type[Backend], ...] = (
|
||||
NvidiaBackend,
|
||||
AmdBackend,
|
||||
OpenvinoGpuBackend,
|
||||
OpenvinoCpuBackend,
|
||||
)
|
||||
|
||||
# Process-wide: once a backend fails, a 50-job batch must not retry it 50 times.
|
||||
_DEMOTED: set[str] = set()
|
||||
|
||||
|
||||
def key(backend: type[Backend]) -> str:
|
||||
return f"{backend.name}/{backend.device}"
|
||||
|
||||
|
||||
def demote(backend: type[Backend]) -> None:
|
||||
_DEMOTED.add(key(backend))
|
||||
|
||||
|
||||
def reset_demotions() -> None:
|
||||
_DEMOTED.clear()
|
||||
|
||||
|
||||
def _select_requested(
|
||||
requested: str, chain: tuple[type[Backend], ...]
|
||||
) -> tuple[type[Backend], ...]:
|
||||
matches = tuple(b for b in chain if requested in (b.name, b.device.lower(), key(b).lower()))
|
||||
if not matches:
|
||||
names = ", ".join(sorted({b.name for b in chain} | {"auto"}))
|
||||
raise errors.ConfigError(f"unknown backend {requested!r}", hint=f"Choose one of: {names}")
|
||||
return matches
|
||||
|
||||
|
||||
def candidates(
|
||||
requested: str | None = None, chain: tuple[type[Backend], ...] = DEFAULT_CHAIN
|
||||
) -> Iterator[type[Backend]]:
|
||||
"""Backends worth trying, in order."""
|
||||
explicit = requested is not None and requested != "auto"
|
||||
pool = _select_requested(requested, chain) if explicit and requested else chain
|
||||
|
||||
for backend in pool:
|
||||
label = key(backend)
|
||||
if label in _DEMOTED:
|
||||
continue
|
||||
if not backend.probe_hardware():
|
||||
log.debug("skipping %s: no hardware", label)
|
||||
continue
|
||||
# Only now may a backend import its toolchain.
|
||||
status = backend.probe_toolchain()
|
||||
if not status.ok:
|
||||
log.warning("skipping %s: %s", label, status.reason)
|
||||
continue
|
||||
yield backend
|
||||
|
||||
|
||||
class BackendHolder:
|
||||
"""Keeps a loaded backend alive across jobs (~7s cold, ~0.7s warm)."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._loaded: dict[str, Backend] = {}
|
||||
|
||||
def get(
|
||||
self,
|
||||
backend: type[Backend],
|
||||
model_id: str | None = None,
|
||||
cache_dir: Path | None = None,
|
||||
) -> Backend:
|
||||
label = key(backend)
|
||||
instance = self._loaded.get(label)
|
||||
if instance is None:
|
||||
instance = backend(model_id, cache_dir)
|
||||
instance.load()
|
||||
self._loaded[label] = instance
|
||||
return instance
|
||||
|
||||
def drop(self, backend: type[Backend]) -> None:
|
||||
instance = self._loaded.pop(key(backend), None)
|
||||
if instance is not None:
|
||||
instance.close()
|
||||
|
||||
def close_all(self) -> None:
|
||||
for instance in self._loaded.values():
|
||||
instance.close()
|
||||
self._loaded.clear()
|
||||
|
||||
|
||||
def transcribe_with_fallback(
|
||||
pcm: npt.NDArray[np.float32],
|
||||
request: TranscribeRequest,
|
||||
*,
|
||||
holder: BackendHolder | None = None,
|
||||
requested: str | None = None,
|
||||
model_id: str | None = None,
|
||||
cache_dir: Path | None = None,
|
||||
chain: tuple[type[Backend], ...] = DEFAULT_CHAIN,
|
||||
) -> TranscriptResult:
|
||||
"""Transcribe, stepping down the chain on any runtime failure.
|
||||
|
||||
Construction and the first generate() are both inside the try: a render-node
|
||||
permission problem and an OpenCL JIT failure surface at one of those points,
|
||||
never at device enumeration.
|
||||
"""
|
||||
owned = holder or BackendHolder()
|
||||
explicit = requested is not None and requested != "auto"
|
||||
last: Exception | None = None
|
||||
|
||||
for backend in candidates(requested, chain):
|
||||
label = key(backend)
|
||||
try:
|
||||
instance = owned.get(backend, model_id, cache_dir)
|
||||
return instance.transcribe(pcm, request)
|
||||
except Exception as exc: # pylint: disable=broad-exception-caught
|
||||
# Any failure -- construction or first generate -- means step down.
|
||||
owned.drop(backend)
|
||||
if explicit:
|
||||
raise errors.TranscribeError(
|
||||
f"{label} failed and --device was explicit, so no fallback was tried",
|
||||
hint=str(exc),
|
||||
) from exc
|
||||
demote(backend)
|
||||
log.warning("backend %s failed (%s); demoting it and falling back", label, exc)
|
||||
last = exc
|
||||
|
||||
raise errors.TranscribeError(
|
||||
"no usable backend remains",
|
||||
hint=str(last) if last else "Run `ccn-transcribe doctor` to see what was detected.",
|
||||
) from last
|
||||
@@ -0,0 +1,298 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import builtins
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import openvino_genai
|
||||
import pytest
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends import openvino_backend as ovb
|
||||
from ccn_transcribe.backends import probe as probe_mod
|
||||
from ccn_transcribe.backends.base import TranscribeRequest
|
||||
|
||||
LIVE = pytest.mark.skipif(
|
||||
not os.environ.get("CCN_LIVE"), reason="set CCN_LIVE=1 to run against the real GPU"
|
||||
)
|
||||
|
||||
|
||||
class TestLanguageToken:
|
||||
def test_bare_code_becomes_a_token(self) -> None:
|
||||
assert ovb.language_token("en") == "<|en|>"
|
||||
|
||||
def test_an_existing_token_is_left_alone(self) -> None:
|
||||
assert ovb.language_token("<|de|>") == "<|de|>"
|
||||
|
||||
def test_none_stays_none_for_autodetect(self) -> None:
|
||||
assert ovb.language_token(None) is None
|
||||
|
||||
def test_an_unknown_code_is_rejected_with_suggestions(self) -> None:
|
||||
with pytest.raises(errors.ConfigError, match="unknown language"):
|
||||
ovb.language_token("zz", allowed={"<|en|>", "<|de|>"})
|
||||
|
||||
def test_a_known_code_passes_validation(self) -> None:
|
||||
assert ovb.language_token("en", allowed={"<|en|>"}) == "<|en|>"
|
||||
|
||||
|
||||
class TestSupportedLanguages:
|
||||
def test_none_when_the_config_is_absent(self, tmp_path: Path) -> None:
|
||||
assert ovb.supported_languages(tmp_path) is None
|
||||
|
||||
def test_none_when_the_config_is_unreadable(self, tmp_path: Path) -> None:
|
||||
(tmp_path / "generation_config.json").write_text("{broken")
|
||||
assert ovb.supported_languages(tmp_path) is None
|
||||
|
||||
def test_none_when_there_is_no_language_map(self, tmp_path: Path) -> None:
|
||||
(tmp_path / "generation_config.json").write_text('{"max_length": 448}')
|
||||
assert ovb.supported_languages(tmp_path) is None
|
||||
|
||||
def test_reads_the_language_map(self, tmp_path: Path) -> None:
|
||||
(tmp_path / "generation_config.json").write_text('{"lang_to_id": {"<|en|>": 1}}')
|
||||
assert ovb.supported_languages(tmp_path) == {"<|en|>"}
|
||||
|
||||
def test_reads_the_real_cached_model(self) -> None:
|
||||
languages = ovb.supported_languages(ovb.resolve_model(ovb.DEFAULT_MODEL, None))
|
||||
assert languages is not None
|
||||
assert "<|en|>" in languages
|
||||
assert len(languages) > 90
|
||||
|
||||
|
||||
class TestCacheRoot:
|
||||
def test_respects_xdg_cache_home(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
monkeypatch.setenv("XDG_CACHE_HOME", str(tmp_path))
|
||||
assert ovb.cache_root() == tmp_path / "ccn-transcribe/ov_cache"
|
||||
|
||||
def test_falls_back_to_home_cache(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("XDG_CACHE_HOME", raising=False)
|
||||
assert ovb.cache_root() == Path.home() / ".cache/ccn-transcribe/ov_cache"
|
||||
|
||||
def test_is_not_relative_to_the_working_directory(self) -> None:
|
||||
# The benchmark scripts use a relative .ov_cache_* and only work from the
|
||||
# repo root; the pipeline must not inherit that.
|
||||
assert ovb.cache_root().is_absolute()
|
||||
|
||||
|
||||
class TestResolveModel:
|
||||
def test_finds_the_cached_model_without_network(self) -> None:
|
||||
path = ovb.resolve_model(ovb.DEFAULT_MODEL, ovb.DEFAULT_REVISION)
|
||||
assert (path / "openvino_encoder_model.xml").exists()
|
||||
|
||||
def test_a_missing_model_raises_with_a_manual_command(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def boom(*_a: object, **_k: object) -> None:
|
||||
raise OSError("no such model")
|
||||
|
||||
monkeypatch.setattr(ovb, "_snapshot", boom)
|
||||
with pytest.raises(errors.ModelUnavailableError) as caught:
|
||||
ovb.resolve_model("nope/nope", None)
|
||||
assert "hf download" in (caught.value.hint or "")
|
||||
|
||||
|
||||
class TestDeviceSelection:
|
||||
def test_gpu_backend_probes_intel_hardware(self) -> None:
|
||||
assert ovb.OpenvinoGpuBackend.probe_hardware() is True
|
||||
|
||||
def test_cpu_backend_always_has_hardware(self) -> None:
|
||||
assert ovb.OpenvinoCpuBackend.probe_hardware() is True
|
||||
|
||||
def test_gpu_toolchain_is_available_here(self) -> None:
|
||||
assert ovb.OpenvinoGpuBackend.probe_toolchain().ok is True
|
||||
|
||||
def test_cpu_toolchain_is_available_here(self) -> None:
|
||||
assert ovb.OpenvinoCpuBackend.probe_toolchain().ok is True
|
||||
|
||||
def test_multi_gpu_enumeration_is_accepted(self) -> None:
|
||||
assert ovb.OpenvinoGpuBackend.device_present(["CPU", "GPU.0", "GPU.1"]) is True
|
||||
|
||||
def test_a_cpu_only_enumeration_is_rejected_for_gpu(self) -> None:
|
||||
assert ovb.OpenvinoGpuBackend.device_present(["CPU"]) is False
|
||||
|
||||
def test_a_missing_gpu_plugin_is_reported_as_a_plugin_problem(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def absent(_cls: type[ovb.OpenvinoGpuBackend], _devices: list[str]) -> bool:
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(ovb.OpenvinoGpuBackend, "device_present", classmethod(absent))
|
||||
status = ovb.OpenvinoGpuBackend.probe_toolchain()
|
||||
assert status.ok is False
|
||||
assert "plugin problem" in status.reason
|
||||
|
||||
def test_an_unimportable_openvino_is_reported(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
real = builtins.__import__
|
||||
|
||||
def no_openvino(name: str, *a: object, **k: object) -> Any:
|
||||
if name == "openvino":
|
||||
raise ImportError("gone")
|
||||
return real(name, *a, **k) # type: ignore[arg-type]
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", no_openvino)
|
||||
assert ovb.OpenvinoCpuBackend.probe_toolchain().ok is False
|
||||
|
||||
|
||||
class TestTranscribeWiring:
|
||||
class _Pipe:
|
||||
def __init__(self, chunks: list[Any] | None = None, boom: bool = False) -> None:
|
||||
self.chunks = chunks
|
||||
self.boom = boom
|
||||
self.seen: dict[str, Any] = {}
|
||||
|
||||
def generate(self, _pcm: object, **kwargs: object) -> Any:
|
||||
if self.boom:
|
||||
raise RuntimeError("kernel compile failed")
|
||||
self.seen = dict(kwargs)
|
||||
return type("R", (), {"chunks": self.chunks})()
|
||||
|
||||
def _chunk(self, start: float, end: float, text: str) -> Any:
|
||||
return type("C", (), {"start_ts": start, "end_ts": end, "text": text})()
|
||||
|
||||
def _backend(self, pipe: object) -> ovb.OpenvinoCpuBackend:
|
||||
backend = ovb.OpenvinoCpuBackend()
|
||||
backend.pipe = pipe
|
||||
backend.languages = {"<|en|>"}
|
||||
return backend
|
||||
|
||||
def test_passes_timestamps_and_task(self) -> None:
|
||||
pipe = self._Pipe(chunks=[])
|
||||
self._backend(pipe).transcribe(np.zeros(16000, np.float32), TranscribeRequest())
|
||||
assert pipe.seen["return_timestamps"] is True
|
||||
assert pipe.seen["task"] == "transcribe"
|
||||
|
||||
def test_converts_the_language_to_token_form(self) -> None:
|
||||
pipe = self._Pipe(chunks=[])
|
||||
self._backend(pipe).transcribe(
|
||||
np.zeros(16000, np.float32), TranscribeRequest(language="en")
|
||||
)
|
||||
assert pipe.seen["language"] == "<|en|>"
|
||||
|
||||
def test_omits_language_when_autodetecting(self) -> None:
|
||||
pipe = self._Pipe(chunks=[])
|
||||
self._backend(pipe).transcribe(np.zeros(16000, np.float32), TranscribeRequest())
|
||||
assert "language" not in pipe.seen
|
||||
|
||||
def test_forwards_prompt_and_hotwords_only_when_set(self) -> None:
|
||||
pipe = self._Pipe(chunks=[])
|
||||
self._backend(pipe).transcribe(
|
||||
np.zeros(16000, np.float32),
|
||||
TranscribeRequest(initial_prompt="P", hotwords="H"),
|
||||
)
|
||||
assert pipe.seen["initial_prompt"] == "P"
|
||||
assert pipe.seen["hotwords"] == "H"
|
||||
|
||||
def test_normalizes_chunks_into_segments(self) -> None:
|
||||
pipe = self._Pipe(chunks=[self._chunk(0.0, -1.0, "hi")])
|
||||
got = self._backend(pipe).transcribe(np.zeros(32000, np.float32), TranscribeRequest())
|
||||
assert got.segments[0].end == pytest.approx(2.0)
|
||||
|
||||
def test_reports_audio_duration_and_device(self) -> None:
|
||||
pipe = self._Pipe(chunks=[self._chunk(0.0, 1.0, "hi")])
|
||||
got = self._backend(pipe).transcribe(np.zeros(32000, np.float32), TranscribeRequest())
|
||||
assert got.audio_s == pytest.approx(2.0)
|
||||
assert got.device == "CPU"
|
||||
assert got.backend == "openvino"
|
||||
|
||||
def test_empty_chunks_give_no_segments(self) -> None:
|
||||
pipe = self._Pipe(chunks=None)
|
||||
assert (
|
||||
self._backend(pipe)
|
||||
.transcribe(np.zeros(16000, np.float32), TranscribeRequest())
|
||||
.segments
|
||||
== ()
|
||||
)
|
||||
|
||||
def test_a_generate_failure_becomes_a_backend_runtime_error(self) -> None:
|
||||
pipe = self._Pipe(boom=True)
|
||||
with pytest.raises(errors.BackendRuntimeError, match="during generation"):
|
||||
self._backend(pipe).transcribe(np.zeros(16000, np.float32), TranscribeRequest())
|
||||
|
||||
def test_transcribe_loads_lazily(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
backend = ovb.OpenvinoCpuBackend()
|
||||
calls: list[int] = []
|
||||
|
||||
def fake_load(self: ovb.OpenvinoCpuBackend) -> None:
|
||||
calls.append(1)
|
||||
self.pipe = TestTranscribeWiring._Pipe(chunks=[])
|
||||
|
||||
monkeypatch.setattr(ovb.OpenvinoCpuBackend, "load", fake_load)
|
||||
backend.transcribe(np.zeros(16000, np.float32), TranscribeRequest())
|
||||
assert calls == [1]
|
||||
|
||||
def test_close_releases_the_pipeline(self) -> None:
|
||||
backend = self._backend(self._Pipe())
|
||||
backend.close()
|
||||
assert backend.pipe is None
|
||||
|
||||
|
||||
class TestLoadHints:
|
||||
def test_cpu_hint_points_at_doctor(self) -> None:
|
||||
assert "doctor" in ovb.OpenvinoCpuBackend().load_hint()
|
||||
|
||||
def test_gpu_hint_names_the_render_group_when_the_node_is_unwritable(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
node = probe_mod.RenderNode(
|
||||
path=Path("/dev/dri/renderD128"),
|
||||
vendor="intel",
|
||||
driver="i915",
|
||||
mode=0o660,
|
||||
writable=False,
|
||||
)
|
||||
|
||||
def one_node(*_a: object, **_k: object) -> list[probe_mod.RenderNode]:
|
||||
return [node]
|
||||
|
||||
monkeypatch.setattr(probe_mod, "render_nodes", one_node)
|
||||
assert "usermod -aG render" in ovb.OpenvinoGpuBackend().load_hint()
|
||||
|
||||
def test_gpu_hint_falls_back_to_doctor_when_permissions_are_fine(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def no_nodes(*_a: object, **_k: object) -> list[probe_mod.RenderNode]:
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(probe_mod, "render_nodes", no_nodes)
|
||||
assert "doctor" in ovb.OpenvinoGpuBackend().load_hint()
|
||||
|
||||
def test_a_pipeline_construction_failure_is_wrapped(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def boom(*_a: object, **_k: object) -> None:
|
||||
raise RuntimeError("no device")
|
||||
|
||||
monkeypatch.setattr(openvino_genai, "WhisperPipeline", boom)
|
||||
with pytest.raises(errors.BackendRuntimeError, match="could not build a pipeline"):
|
||||
ovb.OpenvinoCpuBackend().load()
|
||||
|
||||
|
||||
@LIVE
|
||||
class TestLive:
|
||||
def test_real_transcription_on_the_default_device(self, sine_wav: Path) -> None:
|
||||
from ccn_transcribe.media import ffmpeg
|
||||
|
||||
backend = ovb.OpenvinoGpuBackend()
|
||||
backend.load()
|
||||
result = backend.transcribe(ffmpeg.decode_16k_mono(sine_wav), TranscribeRequest())
|
||||
assert result.device == "GPU"
|
||||
|
||||
|
||||
def test_a_cache_miss_falls_through_to_the_network(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
calls: list[bool] = []
|
||||
|
||||
def snapshot(_model: str, _revision: str | None, *, local_only: bool) -> str:
|
||||
calls.append(local_only)
|
||||
if local_only:
|
||||
raise OSError("not cached")
|
||||
return str(tmp_path)
|
||||
|
||||
monkeypatch.setattr(ovb, "_snapshot", snapshot)
|
||||
with caplog.at_level(logging.INFO):
|
||||
assert ovb.resolve_model("some/model", None) == tmp_path
|
||||
assert calls == [True, False], "the cache must be consulted before the network"
|
||||
assert "not cached" in caplog.text
|
||||
@@ -0,0 +1,130 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pytest
|
||||
|
||||
from ccn_transcribe.backends import probe
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
INTEL = "0x8086"
|
||||
NVIDIA = "0x10de"
|
||||
AMD = "0x1002"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sysfs(tmp_path: Path) -> probe.SysfsRoots:
|
||||
(tmp_path / "sys").mkdir()
|
||||
(tmp_path / "dev").mkdir()
|
||||
return probe.SysfsRoots(drm=tmp_path / "sys", dev_dri=tmp_path / "dev", kfd=tmp_path / "kfd")
|
||||
|
||||
|
||||
def add_node(roots: probe.SysfsRoots, name: str, vendor: str, driver: str) -> None:
|
||||
node = roots.drm / name
|
||||
device = node / "device"
|
||||
device.mkdir(parents=True)
|
||||
(device / "vendor").write_text(vendor + "\n")
|
||||
# The bound driver is a symlink; its basename is the driver name.
|
||||
driver_dir = roots.drm.parent / "drivers" / driver
|
||||
driver_dir.mkdir(parents=True, exist_ok=True)
|
||||
(device / "driver").symlink_to(driver_dir)
|
||||
(roots.dev_dri / name).write_text("")
|
||||
|
||||
|
||||
class TestRenderNodes:
|
||||
def test_empty_when_there_is_no_drm_tree(self, tmp_path: Path) -> None:
|
||||
roots = probe.SysfsRoots(drm=tmp_path / "nope", dev_dri=tmp_path / "d", kfd=tmp_path / "k")
|
||||
assert probe.render_nodes(roots) == []
|
||||
|
||||
def test_reads_vendor_and_bound_driver(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
[node] = probe.render_nodes(sysfs)
|
||||
assert node.vendor == "intel"
|
||||
assert node.driver == "i915"
|
||||
|
||||
def test_unknown_vendor_id_is_passed_through(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", "0xbeef", "weird")
|
||||
assert probe.render_nodes(sysfs)[0].vendor == "0xbeef"
|
||||
|
||||
def test_nodes_without_a_vendor_file_are_skipped(self, sysfs: probe.SysfsRoots) -> None:
|
||||
(sysfs.drm / "renderD200" / "device").mkdir(parents=True)
|
||||
assert probe.render_nodes(sysfs) == []
|
||||
|
||||
def test_only_render_nodes_are_listed(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
add_node(sysfs, "card1", INTEL, "i915")
|
||||
assert [n.path.name for n in probe.render_nodes(sysfs)] == ["renderD128"]
|
||||
|
||||
def test_reports_permissions(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
(sysfs.dev_dri / "renderD128").chmod(0o666)
|
||||
node = probe.render_nodes(sysfs)[0]
|
||||
assert node.mode == 0o666
|
||||
assert node.writable is True
|
||||
|
||||
def test_a_node_with_no_device_file_is_not_writable(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
(sysfs.dev_dri / "renderD128").unlink()
|
||||
assert probe.render_nodes(sysfs)[0].writable is False
|
||||
|
||||
|
||||
class TestVendorProbes:
|
||||
def test_intel_detected_via_i915(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
assert probe.has_intel_gpu(sysfs) is True
|
||||
|
||||
def test_intel_detected_via_xe(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "xe")
|
||||
assert probe.has_intel_gpu(sysfs) is True
|
||||
|
||||
def test_intel_not_detected_when_the_driver_is_unbound(self, sysfs: probe.SysfsRoots) -> None:
|
||||
# /sys/module/xe can exist with zero bound devices, which is why the
|
||||
# bound-driver symlink is the probe rather than module presence.
|
||||
add_node(sysfs, "renderD128", INTEL, "vfio-pci")
|
||||
assert probe.has_intel_gpu(sysfs) is False
|
||||
|
||||
def test_no_intel_when_the_box_has_none(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", AMD, "amdgpu")
|
||||
assert probe.has_intel_gpu(sysfs) is False
|
||||
|
||||
def test_amd_detected_via_render_node(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", AMD, "amdgpu")
|
||||
assert probe.has_amd_gpu(sysfs) is True
|
||||
|
||||
def test_amd_detected_via_kfd(self, sysfs: probe.SysfsRoots) -> None:
|
||||
sysfs.kfd.write_text("")
|
||||
assert probe.has_amd_gpu(sysfs) is True
|
||||
|
||||
def test_no_amd_on_an_intel_box(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
assert probe.has_amd_gpu(sysfs) is False
|
||||
|
||||
def test_nvidia_detected_via_proc(self, sysfs: probe.SysfsRoots, tmp_path: Path) -> None:
|
||||
proc = tmp_path / "nvidia_version"
|
||||
proc.write_text("NVRM version ...")
|
||||
assert probe.has_nvidia_gpu(sysfs, proc_nvidia=proc) is True
|
||||
|
||||
def test_nvidia_detected_via_render_node(self, sysfs: probe.SysfsRoots) -> None:
|
||||
# The proprietary stack does not guarantee a DRM render node, so sysfs is
|
||||
# the fallback rather than the primary signal.
|
||||
add_node(sysfs, "renderD128", NVIDIA, "nvidia-drm")
|
||||
assert probe.has_nvidia_gpu(sysfs, proc_nvidia=sysfs.drm / "absent") is True
|
||||
|
||||
def test_no_nvidia_on_an_intel_box(self, sysfs: probe.SysfsRoots) -> None:
|
||||
add_node(sysfs, "renderD128", INTEL, "i915")
|
||||
assert probe.has_nvidia_gpu(sysfs, proc_nvidia=sysfs.drm / "absent") is False
|
||||
|
||||
|
||||
class TestAgainstTheRealMachine:
|
||||
def test_defaults_point_at_the_real_system(self) -> None:
|
||||
assert probe.DEFAULT_ROOTS.drm.name == "drm"
|
||||
|
||||
def test_intel_igpu_is_detected_here(self) -> None:
|
||||
# This box is an Iris Xe on i915; if this fails the probe is broken.
|
||||
assert probe.has_intel_gpu() is True
|
||||
|
||||
def test_no_nvidia_or_amd_here(self) -> None:
|
||||
assert probe.has_nvidia_gpu() is False
|
||||
assert probe.has_amd_gpu() is False
|
||||
@@ -0,0 +1,227 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from typing import TYPE_CHECKING, ClassVar
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends import registry
|
||||
from ccn_transcribe.backends.base import ToolchainStatus, TranscribeRequest
|
||||
from ccn_transcribe.transcript import Segment, TranscriptResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
import numpy.typing as npt
|
||||
|
||||
PCM = np.zeros(16_000, dtype=np.float32)
|
||||
REQ = TranscribeRequest()
|
||||
|
||||
|
||||
class FakeBackend:
|
||||
name: ClassVar[str] = "fake"
|
||||
device: ClassVar[str] = "FAKE"
|
||||
default_model: ClassVar[str] = "fake-model"
|
||||
|
||||
hardware: ClassVar[bool] = True
|
||||
toolchain: ClassVar[ToolchainStatus] = ToolchainStatus(ok=True)
|
||||
fail_on_load: ClassVar[bool] = False
|
||||
fail_on_transcribe: ClassVar[bool] = False
|
||||
loads: ClassVar[int] = 0
|
||||
|
||||
def __init__(self, model_id: str | None = None, cache_dir: Path | None = None) -> None:
|
||||
self.model_id = model_id or self.default_model
|
||||
self.cache_dir = cache_dir
|
||||
self.closed = False
|
||||
|
||||
@classmethod
|
||||
def probe_hardware(cls) -> bool:
|
||||
return cls.hardware
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus:
|
||||
return cls.toolchain
|
||||
|
||||
def load(self) -> None:
|
||||
type(self).loads += 1
|
||||
if self.fail_on_load:
|
||||
raise errors.BackendRuntimeError(f"{self.name} failed to load")
|
||||
|
||||
def transcribe(
|
||||
self, pcm: npt.NDArray[np.float32], request: TranscribeRequest
|
||||
) -> TranscriptResult:
|
||||
if self.fail_on_transcribe:
|
||||
raise errors.BackendRuntimeError(f"{self.name} blew up mid-generate")
|
||||
return TranscriptResult(
|
||||
segments=(Segment(0.0, 1.0, self.name),),
|
||||
language=request.language,
|
||||
backend=self.name,
|
||||
device=self.device,
|
||||
model_id=self.model_id,
|
||||
model_revision=None,
|
||||
audio_s=len(pcm) / 16_000,
|
||||
wall_s=0.1,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def make_backend(label: str, **attrs: object) -> type[FakeBackend]:
|
||||
return type(
|
||||
f"{label}Backend", (FakeBackend,), {"name": label, "device": label.upper(), **attrs}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean() -> None:
|
||||
registry.reset_demotions()
|
||||
|
||||
|
||||
class TestCandidates:
|
||||
def test_skips_backends_with_no_hardware(self) -> None:
|
||||
a = make_backend("a", hardware=False)
|
||||
b = make_backend("b")
|
||||
assert list(registry.candidates(chain=(a, b))) == [b]
|
||||
|
||||
def test_skips_backends_whose_toolchain_is_absent(self) -> None:
|
||||
a = make_backend("a", toolchain=ToolchainStatus(ok=False, reason="not installed"))
|
||||
b = make_backend("b")
|
||||
assert list(registry.candidates(chain=(a, b))) == [b]
|
||||
|
||||
def test_toolchain_is_never_probed_without_hardware(self) -> None:
|
||||
# This ordering is what keeps torch from being imported on an Intel box.
|
||||
probed: list[str] = []
|
||||
|
||||
class Guard(FakeBackend):
|
||||
name: ClassVar[str] = "guard"
|
||||
hardware: ClassVar[bool] = False
|
||||
|
||||
@classmethod
|
||||
def probe_toolchain(cls) -> ToolchainStatus:
|
||||
probed.append(cls.name)
|
||||
return ToolchainStatus(ok=True)
|
||||
|
||||
list(registry.candidates(chain=(Guard,)))
|
||||
assert probed == []
|
||||
|
||||
def test_chain_order_is_preserved(self) -> None:
|
||||
a, b, c = make_backend("a"), make_backend("b"), make_backend("c")
|
||||
assert list(registry.candidates(chain=(a, b, c))) == [a, b, c]
|
||||
|
||||
def test_an_explicit_request_selects_only_that_backend(self) -> None:
|
||||
a, b = make_backend("a"), make_backend("b")
|
||||
assert list(registry.candidates(requested="b", chain=(a, b))) == [b]
|
||||
|
||||
def test_an_unknown_request_is_a_config_error(self) -> None:
|
||||
with pytest.raises(errors.ConfigError, match="unknown"):
|
||||
list(registry.candidates(requested="nope", chain=(make_backend("a"),)))
|
||||
|
||||
def test_auto_is_not_treated_as_a_name(self) -> None:
|
||||
a = make_backend("a")
|
||||
assert list(registry.candidates(requested="auto", chain=(a,))) == [a]
|
||||
|
||||
def test_demoted_backends_are_skipped(self) -> None:
|
||||
a, b = make_backend("a"), make_backend("b")
|
||||
registry.demote(a)
|
||||
assert list(registry.candidates(chain=(a, b))) == [b]
|
||||
|
||||
|
||||
class TestFallback:
|
||||
def test_uses_the_first_working_backend(self) -> None:
|
||||
a, b = make_backend("a"), make_backend("b")
|
||||
got = registry.transcribe_with_fallback(PCM, REQ, chain=(a, b))
|
||||
assert got.backend == "a"
|
||||
|
||||
def test_falls_back_when_construction_fails(self) -> None:
|
||||
a = make_backend("a", fail_on_load=True)
|
||||
b = make_backend("b")
|
||||
assert registry.transcribe_with_fallback(PCM, REQ, chain=(a, b)).backend == "b"
|
||||
|
||||
def test_falls_back_when_the_first_generate_fails(self) -> None:
|
||||
# The permission and JIT failures both surface here, not at detection.
|
||||
a = make_backend("a", fail_on_transcribe=True)
|
||||
b = make_backend("b")
|
||||
assert registry.transcribe_with_fallback(PCM, REQ, chain=(a, b)).backend == "b"
|
||||
|
||||
def test_a_failure_demotes_for_the_rest_of_the_run(self) -> None:
|
||||
a = make_backend("a", fail_on_transcribe=True)
|
||||
b = make_backend("b")
|
||||
registry.transcribe_with_fallback(PCM, REQ, chain=(a, b))
|
||||
a.loads = 0
|
||||
registry.transcribe_with_fallback(PCM, REQ, chain=(a, b))
|
||||
assert a.loads == 0, "a 50-job batch must not retry a dead backend 50 times"
|
||||
|
||||
def test_exhausting_the_chain_raises(self) -> None:
|
||||
a = make_backend("a", fail_on_load=True)
|
||||
with pytest.raises(errors.TranscribeError, match="no usable backend"):
|
||||
registry.transcribe_with_fallback(PCM, REQ, chain=(a,))
|
||||
|
||||
def test_no_candidates_at_all_raises(self) -> None:
|
||||
a = make_backend("a", hardware=False)
|
||||
with pytest.raises(errors.TranscribeError, match="no usable backend"):
|
||||
registry.transcribe_with_fallback(PCM, REQ, chain=(a,))
|
||||
|
||||
def test_an_explicit_device_never_falls_back(self) -> None:
|
||||
# Asking for the GPU and silently getting the CPU makes a benchmark a lie.
|
||||
a = make_backend("a", fail_on_transcribe=True)
|
||||
b = make_backend("b")
|
||||
with pytest.raises(errors.TranscribeError):
|
||||
registry.transcribe_with_fallback(PCM, REQ, requested="a", chain=(a, b))
|
||||
|
||||
def test_warnings_name_the_failing_backend(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
a = make_backend("a", fail_on_transcribe=True)
|
||||
b = make_backend("b")
|
||||
registry.transcribe_with_fallback(PCM, REQ, chain=(a, b))
|
||||
assert "a" in caplog.text.lower()
|
||||
|
||||
|
||||
class TestHolder:
|
||||
def test_reuses_a_loaded_backend_across_jobs(self) -> None:
|
||||
a = make_backend("a")
|
||||
holder = registry.BackendHolder()
|
||||
holder.get(a)
|
||||
holder.get(a)
|
||||
assert a.loads == 1, "model load is ~7s cold; reloading per job is pure waste"
|
||||
|
||||
def test_close_all_closes_instances(self) -> None:
|
||||
a = make_backend("a")
|
||||
holder = registry.BackendHolder()
|
||||
instance = holder.get(a)
|
||||
holder.close_all()
|
||||
assert isinstance(instance, FakeBackend)
|
||||
assert instance.closed is True
|
||||
|
||||
def test_dropping_forces_a_reload(self) -> None:
|
||||
a = make_backend("a")
|
||||
holder = registry.BackendHolder()
|
||||
holder.get(a)
|
||||
holder.drop(a)
|
||||
holder.get(a)
|
||||
assert a.loads == 2
|
||||
|
||||
|
||||
class TestRealChain:
|
||||
def test_default_chain_order(self) -> None:
|
||||
assert [b.name for b in registry.DEFAULT_CHAIN] == [
|
||||
"nvidia",
|
||||
"amd",
|
||||
"openvino",
|
||||
"openvino",
|
||||
]
|
||||
|
||||
def test_cpu_is_last_and_terminal(self) -> None:
|
||||
assert registry.DEFAULT_CHAIN[-1].device == "CPU"
|
||||
|
||||
def test_this_machine_resolves_to_the_intel_gpu(self) -> None:
|
||||
chosen = next(iter(registry.candidates()))
|
||||
assert (chosen.name, chosen.device) == ("openvino", "GPU")
|
||||
|
||||
def test_torch_is_never_imported_on_this_intel_only_box(self) -> None:
|
||||
# The structural guarantee, made executable.
|
||||
for module in ("torch", "faster_whisper", "ctranslate2", "transformers"):
|
||||
sys.modules.pop(module, None)
|
||||
list(registry.candidates())
|
||||
assert not {"torch", "faster_whisper", "ctranslate2"} & set(sys.modules)
|
||||
@@ -0,0 +1,50 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from ccn_transcribe import errors
|
||||
from ccn_transcribe.backends.amd_backend import AmdBackend
|
||||
from ccn_transcribe.backends.base import TranscribeRequest
|
||||
from ccn_transcribe.backends.nvidia_backend import NvidiaBackend
|
||||
|
||||
STUBS = [NvidiaBackend, AmdBackend]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("backend", STUBS)
|
||||
class TestStubs:
|
||||
def test_hardware_probe_is_real(self, backend: type[NvidiaBackend]) -> None:
|
||||
# No NVIDIA or AMD on this machine, so detection must say so honestly.
|
||||
assert backend.probe_hardware() is False
|
||||
|
||||
def test_toolchain_reports_not_implemented(self, backend: type[NvidiaBackend]) -> None:
|
||||
status = backend.probe_toolchain()
|
||||
assert status.ok is False
|
||||
assert "not implemented" in status.reason
|
||||
|
||||
def test_load_raises_with_an_actionable_hint(self, backend: type[NvidiaBackend]) -> None:
|
||||
with pytest.raises(errors.BackendNotImplementedError) as caught:
|
||||
backend().load()
|
||||
hint = caught.value.hint or ""
|
||||
assert "backends/" in hint
|
||||
|
||||
def test_transcribe_raises_rather_than_returning_nonsense(
|
||||
self, backend: type[NvidiaBackend]
|
||||
) -> None:
|
||||
with pytest.raises(errors.BackendNotImplementedError):
|
||||
backend().transcribe(np.zeros(16, np.float32), TranscribeRequest())
|
||||
|
||||
def test_close_is_safe(self, backend: type[NvidiaBackend]) -> None:
|
||||
backend().close()
|
||||
|
||||
def test_default_model_is_not_the_openvino_ir(self, backend: type[NvidiaBackend]) -> None:
|
||||
# OpenVINO IR cannot load on CUDA or ROCm, so each backend owns its default.
|
||||
assert "-ov" not in backend.default_model
|
||||
|
||||
def test_a_custom_model_id_is_kept(self, backend: type[NvidiaBackend]) -> None:
|
||||
assert backend("custom/model").model_id == "custom/model"
|
||||
|
||||
|
||||
def test_the_two_stubs_use_different_runtimes() -> None:
|
||||
# CTranslate2 has no ROCm support, so these cannot share one implementation.
|
||||
assert NvidiaBackend.device != AmdBackend.device
|
||||
Reference in New Issue
Block a user