feat(cuda): добавлен прозрачный GPU runtime для Linux/WSL2
- Зачем:
- ctranslate2 требует libcublas.so.12 для CUDA, но не бандлит её в wheel —
без системного CUDA toolkit GPU не работает из коробки.
- Что:
- добавлена зависимость nvidia-cublas-cu12 (Linux x86_64).
- создан _cuda_bootstrap.py: preload libcublas через ctypes.CDLL(RTLD_GLOBAL)
до импорта ctranslate2 (LD_LIBRARY_PATH не работает — glibc кеширует пути).
- добавлен strict_device в transcriber: --device cuda/cpu не делает silent fallback.
- CLI: диагностика requested vs resolved device, Windows CUDA-подсказка.
- Проверка:
- uv run pytest -v (52 passed, 1 skipped).
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from local_transcriber.cli import app
|
||||
@@ -234,3 +235,101 @@ def test_cli_resolves_model_before_transcribe(tmp_path):
|
||||
mock_ensure_model.assert_called_once()
|
||||
call_kwargs = mock_transcribe.call_args[1]
|
||||
assert call_kwargs["model_name"] == "/models/large-v3"
|
||||
|
||||
|
||||
def test_cli_windows_cuda_diagnostic(tmp_path):
|
||||
"""CUDA error on Windows prints choco/winget install hint."""
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.check_ffmpeg"),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"),
|
||||
patch("local_transcriber.cli.transcribe", side_effect=RuntimeError("CUDA error: no device")),
|
||||
patch("local_transcriber.cli.sys") as mock_sys,
|
||||
):
|
||||
mock_sys.platform = "win32"
|
||||
out = runner.invoke(app, [str(audio), "--device", "cuda"])
|
||||
|
||||
assert out.exit_code == 1
|
||||
assert "choco install cuda" in out.output
|
||||
assert "winget install" in out.output
|
||||
|
||||
|
||||
def test_cli_linux_cuda_error_no_windows_hint(tmp_path):
|
||||
"""CUDA error on Linux does NOT print Windows-specific hint."""
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.check_ffmpeg"),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"),
|
||||
patch("local_transcriber.cli.transcribe", side_effect=RuntimeError("CUDA error: no device")),
|
||||
patch("local_transcriber.cli.sys") as mock_sys,
|
||||
):
|
||||
mock_sys.platform = "linux"
|
||||
out = runner.invoke(app, [str(audio), "--device", "cuda"])
|
||||
|
||||
assert out.exit_code == 1
|
||||
assert "choco install cuda" not in out.output
|
||||
|
||||
|
||||
def test_cli_device_fallback_warning(tmp_path):
|
||||
"""When auto-detected device differs from actual, show fallback warning."""
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result(device_used="cpu")
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.check_ffmpeg"),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"),
|
||||
patch("local_transcriber.cli.transcribe", return_value=result),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
# --device auto (default) -> detect_device returns "cuda" but result is "cpu"
|
||||
out = runner.invoke(app, [str(audio)])
|
||||
|
||||
assert "fallback" in out.output
|
||||
|
||||
|
||||
def test_cli_strict_device_passed_to_transcribe(tmp_path):
|
||||
"""--device cuda passes strict_device=True; default auto passes False."""
|
||||
audio = tmp_path / "test.mp3"
|
||||
audio.write_bytes(b"fake")
|
||||
result = _make_result(device_used="cuda")
|
||||
mock_transcribe = MagicMock(return_value=result)
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.check_ffmpeg"),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cuda"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"),
|
||||
patch("local_transcriber.cli.transcribe", mock_transcribe),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
patch("local_transcriber.cli.get_gpu_name", return_value="RTX 3060"),
|
||||
):
|
||||
runner.invoke(app, [str(audio), "--device", "cuda"])
|
||||
|
||||
assert mock_transcribe.call_args[1]["strict_device"] is True
|
||||
|
||||
mock_transcribe.reset_mock()
|
||||
result_cpu = _make_result(device_used="cpu")
|
||||
mock_transcribe.return_value = result_cpu
|
||||
|
||||
with (
|
||||
patch("local_transcriber.cli.check_ffmpeg"),
|
||||
patch("local_transcriber.cli.validate_input_file", return_value=audio),
|
||||
patch("local_transcriber.cli.detect_device", return_value="cpu"),
|
||||
patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"),
|
||||
patch("local_transcriber.cli.transcribe", mock_transcribe),
|
||||
patch("local_transcriber.cli.write_transcript"),
|
||||
):
|
||||
runner.invoke(app, [str(audio)])
|
||||
|
||||
assert mock_transcribe.call_args[1]["strict_device"] is False
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
import ctypes
|
||||
import glob
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from local_transcriber._cuda_bootstrap import ensure_cublas_loadable, is_cublas_available
|
||||
|
||||
|
||||
def test_ensure_cublas_no_nvidia_package(monkeypatch):
|
||||
"""Без nvidia-cublas-cu12 -- ничего не падает."""
|
||||
monkeypatch.setattr(sys, "platform", "linux")
|
||||
monkeypatch.setitem(sys.modules, "nvidia.cublas", None)
|
||||
|
||||
ensure_cublas_loadable() # не должно бросать исключений
|
||||
|
||||
|
||||
def test_ensure_cublas_loads_library(monkeypatch, tmp_path):
|
||||
"""С nvidia.cublas -- вызывает ctypes.CDLL с полным путём и RTLD_GLOBAL."""
|
||||
monkeypatch.setattr(sys, "platform", "linux")
|
||||
|
||||
# Создаём фейковый nvidia.cublas с lib/libcublas.so.12
|
||||
lib_dir = tmp_path / "lib"
|
||||
lib_dir.mkdir()
|
||||
fake_so = lib_dir / "libcublas.so.12"
|
||||
fake_so.touch()
|
||||
|
||||
# Мокаем родительский пакет nvidia (иначе import nvidia.cublas упадёт)
|
||||
fake_nvidia = types.ModuleType("nvidia")
|
||||
fake_nvidia.__path__ = [str(tmp_path)]
|
||||
|
||||
fake_cublas = types.ModuleType("nvidia.cublas")
|
||||
fake_cublas.__path__ = [str(tmp_path)]
|
||||
fake_nvidia.cublas = fake_cublas
|
||||
|
||||
monkeypatch.setitem(sys.modules, "nvidia", fake_nvidia)
|
||||
monkeypatch.setitem(sys.modules, "nvidia.cublas", fake_cublas)
|
||||
|
||||
calls = []
|
||||
monkeypatch.setattr(ctypes, "CDLL", lambda path, mode=0: calls.append((path, mode)))
|
||||
|
||||
ensure_cublas_loadable()
|
||||
|
||||
assert len(calls) == 1
|
||||
assert calls[0][0] == str(fake_so)
|
||||
assert calls[0][1] == ctypes.RTLD_GLOBAL
|
||||
|
||||
|
||||
def test_ensure_cublas_skips_non_linux(monkeypatch):
|
||||
"""На не-Linux платформах -- no-op."""
|
||||
monkeypatch.setattr(sys, "platform", "win32")
|
||||
ensure_cublas_loadable() # не должно бросать исключений
|
||||
|
||||
|
||||
def _nvidia_cublas_installed() -> bool:
|
||||
"""Проверяет, что pip-пакет nvidia-cublas-cu12 установлен."""
|
||||
try:
|
||||
import nvidia.cublas # type: ignore[import-untyped]
|
||||
|
||||
cublas_paths = getattr(nvidia.cublas, "__path__", None)
|
||||
if not cublas_paths:
|
||||
return False
|
||||
lib_dir = os.path.join(cublas_paths[0], "lib")
|
||||
return any(glob.glob(os.path.join(lib_dir, "libcublas.so.12*")))
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
|
||||
def _system_cublas_available() -> bool:
|
||||
"""Проверяет, что libcublas.so.12 доступна через системный линкер (без bootstrap)."""
|
||||
try:
|
||||
ctypes.CDLL("libcublas.so.12")
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.platform != "linux",
|
||||
reason="CUDA bootstrap только для Linux",
|
||||
)
|
||||
@pytest.mark.skipif(
|
||||
not _nvidia_cublas_installed(),
|
||||
reason="nvidia-cublas-cu12 не установлен",
|
||||
)
|
||||
def test_bootstrap_makes_cublas_resolvable():
|
||||
"""Bootstrap из pip-пакета делает libcublas.so.12 резолвимой.
|
||||
|
||||
Тест проходит ТОЛЬКО если:
|
||||
1. nvidia-cublas-cu12 установлен (иначе skip)
|
||||
2. libcublas НЕ доступна через системный линкер до bootstrap
|
||||
(иначе skip -- тест не может доказать, что сработал именно bootstrap)
|
||||
3. После ensure_cublas_loadable() -- libcublas доступна
|
||||
"""
|
||||
if _system_cublas_available():
|
||||
pytest.skip(
|
||||
"libcublas.so.12 уже доступна через системный линкер -- "
|
||||
"невозможно проверить, что сработал именно bootstrap"
|
||||
)
|
||||
|
||||
ensure_cublas_loadable()
|
||||
assert is_cublas_available(), (
|
||||
"nvidia-cublas-cu12 установлен, но после bootstrap "
|
||||
"libcublas.so.12 всё ещё не резолвится через dlopen"
|
||||
)
|
||||
@@ -355,3 +355,61 @@ def test_ensure_model_available_rejects_incomplete_local_directory(tmp_path):
|
||||
|
||||
with pytest.raises(ValueError, match="Неполная локальная модель"):
|
||||
ensure_model_available(str(model_dir))
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_strict_cuda_error(mock_model_cls):
|
||||
"""strict_device=True + CUDA error -> raise, без fallback."""
|
||||
mock_model_cls.side_effect = RuntimeError("CUDA out of memory")
|
||||
|
||||
with pytest.raises(RuntimeError, match="CUDA out of memory"):
|
||||
transcribe(
|
||||
file_path=Path("test.mp3"),
|
||||
model_name="tiny",
|
||||
device="cuda",
|
||||
strict_device=True,
|
||||
)
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_non_strict_cuda_fallback(mock_model_cls):
|
||||
"""strict_device=False + CUDA error -> fallback на CPU."""
|
||||
raw_segments = _make_raw_segments(2)
|
||||
info = _make_info()
|
||||
|
||||
cpu_instance = MagicMock()
|
||||
cpu_instance.transcribe.return_value = (iter(raw_segments), info)
|
||||
|
||||
def model_side_effect(model_name, device, compute_type):
|
||||
if device == "cuda":
|
||||
raise RuntimeError("CUDA out of memory")
|
||||
return cpu_instance
|
||||
|
||||
mock_model_cls.side_effect = model_side_effect
|
||||
|
||||
with pytest.warns(UserWarning, match="Переключение на CPU"):
|
||||
result = transcribe(
|
||||
file_path=Path("test.mp3"),
|
||||
model_name="tiny",
|
||||
device="cuda",
|
||||
strict_device=False,
|
||||
)
|
||||
|
||||
assert result.device_used == "cpu"
|
||||
assert len(result.segments) == 2
|
||||
|
||||
|
||||
@patch("local_transcriber.transcriber.WhisperModel")
|
||||
def test_transcribe_strict_cuda_error_during_transcription(mock_model_cls):
|
||||
"""strict_device=True + CUDA error during transcription -> raise."""
|
||||
cuda_instance = MagicMock()
|
||||
cuda_instance.transcribe.side_effect = RuntimeError("CUDA error during transcription")
|
||||
mock_model_cls.return_value = cuda_instance
|
||||
|
||||
with pytest.raises(RuntimeError, match="CUDA error during transcription"):
|
||||
transcribe(
|
||||
file_path=Path("test.mp3"),
|
||||
model_name="tiny",
|
||||
device="cuda",
|
||||
strict_device=True,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user