From b7e6ab634aef318ce337715fb006198368a246a3 Mon Sep 17 00:00:00 2001 From: Dmitry Dementev Date: Wed, 18 Mar 2026 21:43:06 +0300 Subject: [PATCH] =?UTF-8?q?feat(config):=20=D0=B4=D0=BE=D0=B1=D0=B0=D0=B2?= =?UTF-8?q?=D0=BB=D0=B5=D0=BD=D1=8B=20device-aware=20=D0=B4=D0=B5=D1=84?= =?UTF-8?q?=D0=BE=D0=BB=D1=82=D1=8B=20=D0=B8=20=D1=80=D0=B5=D0=B7=D1=83?= =?UTF-8?q?=D0=BB=D1=8C=D1=82=D0=B0=D1=82=D1=8B=20=D1=82=D0=B5=D1=81=D1=82?= =?UTF-8?q?=D0=BE=D0=B2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Зачем: - тестирование на реальных записях показало, что int8 даёт галлюцинации на длинных файлах, auto-detect языка ошибается — нужны оптимальные дефолты по устройству. - Что: - дефолты: medium float16 (GPU), medium float32 (CPU), language=ru. - добавлены DEVICE_DEFAULTS и apply_device_defaults() в config.py. - убран preprocessor_config.json из обязательных файлов модели (отсутствует у medium). - README обновлён: таблицы скоростей, качества, результаты тестирования compute_type. - Проверка: - uv run pytest — 98 passed, 1 skipped. Co-Authored-By: Claude Opus 4.6 --- README.md | 91 +++++++++++++++++++--------- src/local_transcriber/cli.py | 18 +++--- src/local_transcriber/config.py | 31 +++++++++- src/local_transcriber/transcriber.py | 1 - tests/test_cli.py | 52 ++++++++-------- tests/test_config.py | 40 +++++++++++- 6 files changed, 163 insertions(+), 70 deletions(-) diff --git a/README.md b/README.md index 3187d5f..c7d4494 100644 --- a/README.md +++ b/README.md @@ -33,7 +33,7 @@ uv tool install . uv tool install --force . ``` -Модели скачиваются автоматически при первом запуске (~3 GB для large-v3), +Модели скачиваются автоматически при первом запуске (~1.5 GB для medium, ~3 GB для large-v3), нужен доступ в интернет (Hugging Face Hub). > **Если `transcribe: command not found`** — убедитесь, что директория @@ -43,14 +43,14 @@ uv tool install --force . ## Использование ```bash -# Простой запуск (large-v3, автодетект языка и устройства) +# Простой запуск (medium, русский, автодетект устройства) transcribe meeting.mp4 # Указать язык -transcribe lecture.mp3 --language ru +transcribe lecture.mp3 --language en -# Быстрая модель на CPU -transcribe podcast.wav --model small --device cpu +# Максимальное качество на GPU +transcribe podcast.wav --model large-v3 --compute-type float16 # Сохранить в конкретный файл transcribe interview.m4a --output result.md @@ -82,27 +82,34 @@ transcribe *.mp4 --force Дефолтные параметры можно задать в `.transcriber.toml`: ```toml -model = "small" language = "ru" -device = "cpu" -compute_type = "int8" +model = "large-v3" +compute_type = "float16" ``` Порядок поиска: 1. `.transcriber.toml` в текущей директории (проектный конфиг) 2. `~/.config/transcriber/config.toml` (глобальный конфиг пользователя) -Приоритет: **CLI-аргумент > конфиг > встроенный дефолт**. +Приоритет: **CLI-аргумент > конфиг > device-aware дефолт > встроенный дефолт**. + +Дефолты зависят от устройства (если не заданы явно): + +| Параметр | GPU (CUDA) | CPU | +|----------|-----------|-----| +| model | medium | medium | +| compute_type | float16 | float32 | +| language | ru | ru | ### Опции CLI | Опция | Сокращение | По умолчанию | Описание | |-------|-----------|-------------|----------| -| `--model` | `-m` | `large-v3` | Модель Whisper | -| `--language` | `-l` | `auto` | Язык (ru, en, auto) | +| `--model` | `-m` | `medium` | Модель Whisper | +| `--language` | `-l` | `ru` | Язык (ru, en, auto и др.) | | `--output` | `-o` | `<файл>-transcript.md` | Путь к выходному файлу | | `--device` | `-d` | `auto` | Устройство (auto, cpu, cuda) | -| `--compute-type` | — | `int8` | Тип вычислений | +| `--compute-type` | — | float16 (GPU) / float32 (CPU) | Тип вычислений | | `--force` | `-f` | — | Перезаписать существующие транскрипты | | `--verbose` | `-v` | — | Подробный вывод | @@ -118,14 +125,15 @@ compute_type = "int8" ### Типы квантизации (`--compute-type`) -| Тип | VRAM | Скорость | Качество | Когда использовать | -|-----|------|----------|----------|--------------------| -| `int8` | Низкое | Быстро | Почти без потерь | По умолчанию, GPU от 4 GB и CPU | -| `int8_float16` | Низкое | Быстро | Почти без потерь | GPU, чуть точнее int8 | -| `float16` | Среднее | Быстро | Без потерь | GPU от 6 GB | -| `float32` | Высокое | Медленно | Эталон | CPU (если int8 недоступен) | +| Тип | Устройство | VRAM/RAM | Качество | Когда использовать | +|-----|-----------|----------|----------|--------------------| +| `float16` | GPU | ~4.5-5 GB | Отлично | **По умолчанию для GPU** | +| `int8_float16` | GPU | ~4.7 GB | Отлично | GPU от 6 GB, альтернатива float16 | +| `int8` | GPU/CPU | Низкое | Хорошо, но бывают галлюцинации | GPU от 4 GB, CPU | +| `float32` | CPU | Среднее | Отлично | **По умолчанию для CPU** | -По умолчанию `int8` — универсален для GPU от 4 GB и CPU. +**Важно:** `int8` на длинных записях может давать галлюцинации (повтор фраз, потеря контента). +`float16` и `float32` значительно стабильнее на записях >20 минут с техническими терминами. ## Установка ffmpeg @@ -156,19 +164,44 @@ compute_type = "int8" ### Совместимость GPU -| GPU | VRAM | large-v3 int8 | Рекомендация | -|-----|------|--------------|--------------| -| RTX 3060 | 6 GB | ✅ | int8 | -| RTX 4050 | 6 GB | ✅ | int8 | -| Quadro M3000M | 4 GB | ✅ | int8 обязательно | +| GPU | VRAM | medium float16 | large-v3 float16 | Рекомендация | +|-----|------|---------------|-----------------|--------------| +| RTX 3060 | 6 GB | ✅ | ✅ | medium float16 (дефолт) | +| RTX 4050 | 6 GB | ✅ | ✅ | medium float16 | +| Quadro M3000M | 4 GB | ✅ | ⚠️ tight | medium float16 или int8 | ### Ожидаемая скорость -| Конфигурация | 1 час аудио ≈ | -|-------------|---------------| -| RTX 3060 + large-v3 | 4–6 мин | -| CPU + large-v3 | 60–120 мин | -| CPU + small | 12–20 мин | +Замеры на RTX 3060 Laptop (6 GB) и Intel CPU (WSL2): + +| Конфигурация | 16 мин файл | 42 мин файл | Отн. скорость | +|-------------|-------------|-------------|---------------| +| GPU + medium float16 | ~35с | ~133с | ~19x реалтайм | +| GPU + large-v3 float16 | ~90с | ~350с | ~7x реалтайм | +| CPU + medium float32 | 613с (10 мин) | ~26 мин* | ~1.5x реалтайм | +| CPU + large-v3 int8 | 839с (14 мин) | ~37 мин* | ~1:1 реалтайм | + +*Оценка на основе пропорции. + +### Результаты тестирования качества + +Тесты проведены на реальных записях рабочих созвонов (русский язык, технические термины: +SQL, PostgreSQL, Greenplum, Airflow, ClickHouse, Docker, CDR, GTP, MAP). + +| Конфигурация | Качество (длинная запись, 42 мин) | Проблемы | +|---|---|---| +| large-v3 int8 GPU | Плохо | Галлюцинации (фразы ×25), потеря контента | +| large-v3 float16 GPU | Отлично | — | +| medium float16 GPU | Хорошо | Редкие мелкие ляпы в терминах | +| medium float32 CPU | Хорошо | Сопоставимо с large-v3 int8, без галлюцинаций | +| large-v3 int8 CPU | Хорошо | Без галлюцинаций (на коротких файлах) | + +**Ключевые выводы:** + +1. **Указание языка (`--language ru`) критично** — auto-detect может ошибиться и выдать мусор +2. **float16/float32 стабильнее int8** — особенно на записях >20 минут +3. **medium + float16 на GPU — лучший баланс** скорости и качества для повседневного использования +4. **large-v3 + float16 на GPU** — для максимального качества важных записей ## Формат выходного файла diff --git a/src/local_transcriber/cli.py b/src/local_transcriber/cli.py index bf0f741..479d74e 100644 --- a/src/local_transcriber/cli.py +++ b/src/local_transcriber/cli.py @@ -6,7 +6,7 @@ import typer from rich.console import Console from rich.status import Status -from .config import load_config, resolve_defaults +from .config import apply_device_defaults, load_config, resolve_defaults from .formatter import format_transcript, write_transcript from .transcriber import ( Segment, @@ -34,27 +34,29 @@ console = Console(stderr=True) def main( files: list[Path] = typer.Argument(..., help="Пути к аудио/видеофайлам"), model: str | None = typer.Option( - None, "--model", "-m", show_default=False, help="Модель Whisper [по умолч.: large-v3]" + None, "--model", "-m", show_default=False, help="Модель Whisper [по умолч.: medium]" ), language: str | None = typer.Option( - None, "--language", "-l", show_default=False, help="Язык (ru|en|auto) [по умолч.: auto]" + None, "--language", "-l", show_default=False, help="Язык [по умолч.: ru]" ), output: Path | None = typer.Option(None, "--output", "-o", help="Путь к выходному файлу"), device: str | None = typer.Option( None, "--device", "-d", show_default=False, help="Устройство (auto|cpu|cuda) [по умолч.: auto]" ), compute_type: str | None = typer.Option( - None, "--compute-type", show_default=False, help="Тип вычислений [по умолч.: int8]" + None, "--compute-type", show_default=False, + help="Тип вычислений [по умолч.: float16 (GPU) / float32 (CPU)]" ), verbose: bool = typer.Option(False, "--verbose", "-v", help="Подробный вывод"), force: bool = typer.Option(False, "--force", "-f", help="Перезаписать существующие транскрипты"), ) -> None: try: config = load_config() - defaults = resolve_defaults( - {"model": model, "language": language, "device": device, "compute_type": compute_type}, - config, - ) + cli_values = {"model": model, "language": language, "device": device, "compute_type": compute_type} + defaults = resolve_defaults(cli_values, config) + + resolved_device = detect_device(defaults["device"]) + defaults = apply_device_defaults(defaults, resolved_device, cli_values, config) expanded = expand_globs(files) if not expanded: diff --git a/src/local_transcriber/config.py b/src/local_transcriber/config.py index 098381f..c4270ef 100644 --- a/src/local_transcriber/config.py +++ b/src/local_transcriber/config.py @@ -8,10 +8,15 @@ else: import tomli as tomllib HARDCODED_DEFAULTS: dict[str, str] = { - "model": "large-v3", - "language": "auto", + "model": "medium", + "language": "ru", "device": "auto", - "compute_type": "int8", + "compute_type": "float32", +} + +DEVICE_DEFAULTS: dict[str, dict[str, str]] = { + "cuda": {"model": "medium", "compute_type": "float16"}, + "cpu": {"model": "medium", "compute_type": "float32"}, } _VALID_KEYS = set(HARDCODED_DEFAULTS) @@ -83,3 +88,23 @@ def resolve_defaults( else: result[key] = HARDCODED_DEFAULTS[key] return result + + +def apply_device_defaults( + defaults: dict[str, str], + resolved_device: str, + cli_values: dict[str, str | None], + config: dict[str, str], +) -> dict[str, str]: + """Применяет device-aware дефолты для model и compute_type, + если они не были явно заданы через CLI или конфиг.""" + device_defs = DEVICE_DEFAULTS.get(resolved_device, {}) + if not device_defs: + return defaults + + result = dict(defaults) + for key in ("model", "compute_type"): + if cli_values.get(key) is None and key not in config: + if key in device_defs: + result[key] = device_defs[key] + return result diff --git a/src/local_transcriber/transcriber.py b/src/local_transcriber/transcriber.py index 3434d68..6050158 100644 --- a/src/local_transcriber/transcriber.py +++ b/src/local_transcriber/transcriber.py @@ -30,7 +30,6 @@ MODEL_ALLOW_PATTERNS = [ MODEL_REQUIRED_FILES = [ "config.json", - "preprocessor_config.json", "model.bin", "tokenizer.json", ] diff --git a/tests/test_cli.py b/tests/test_cli.py index 708853f..0f15324 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -43,7 +43,7 @@ def _single_patches(result=None, tmp_file=None, actual_device="cpu"): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", return_value=tmp_file), patch("local_transcriber.cli.detect_device", return_value=actual_device), - patch("local_transcriber.cli.ensure_model_available", return_value="/models/large-v3"), + patch("local_transcriber.cli.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, actual_device)), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -74,7 +74,7 @@ def test_cli_default_options_passed_to_transcribe(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli.write_transcript"), @@ -82,9 +82,9 @@ def test_cli_default_options_passed_to_transcribe(tmp_path): runner.invoke(app, [str(audio)]) call_kwargs = mock_transcribe_file.call_args[1] - assert call_kwargs["model_name"] == "/models/large-v3" - assert call_kwargs["compute_type"] == "int8" - assert call_kwargs["language"] is None # "auto" → None + assert call_kwargs["model_name"] == "/models/medium" + assert call_kwargs["compute_type"] == "float32" + assert call_kwargs["language"] == "ru" assert call_kwargs["on_segment"] is None # verbose=False @@ -134,7 +134,7 @@ def test_cli_verbose_passes_on_segment_callback(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli.write_transcript"), @@ -173,7 +173,7 @@ def test_cli_default_output_path(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript", mock_write), @@ -199,7 +199,7 @@ def test_cli_custom_output_path(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript", mock_write), @@ -236,7 +236,7 @@ def test_cli_passes_status_callback_to_transcribe(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli.write_transcript"), @@ -284,7 +284,7 @@ def test_cli_windows_cuda_diagnostic(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli.sys") as mock_sys, @@ -308,7 +308,7 @@ def test_cli_linux_cuda_error_no_windows_hint(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("CUDA error: no device")), patch("local_transcriber.cli.sys") as mock_sys, @@ -333,7 +333,7 @@ def test_cli_device_fallback_warning(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -357,7 +357,7 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cuda")), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli.write_transcript"), @@ -377,7 +377,7 @@ def test_cli_strict_device_passed_to_transcribe(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", mock_transcribe_file), patch("local_transcriber.cli.write_transcript"), @@ -398,7 +398,7 @@ def test_cli_keyboard_interrupt(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", side_effect=KeyboardInterrupt), patch("local_transcriber.cli.write_transcript"), @@ -432,7 +432,7 @@ def test_cli_unexpected_error_verbose_traceback(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli.write_transcript"), @@ -454,7 +454,7 @@ def test_cli_unexpected_error_no_verbose_hint(tmp_path): 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", side_effect=RuntimeError("unexpected boom")), patch("local_transcriber.cli.write_transcript"), @@ -484,7 +484,7 @@ def test_cli_batch_two_files(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -512,7 +512,7 @@ def test_cli_batch_skips_existing(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -563,7 +563,7 @@ def test_cli_batch_force_overwrites(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -598,7 +598,7 @@ def test_cli_batch_per_file_error(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", side_effect=transcribe_side_effect), patch("local_transcriber.cli.write_transcript"), @@ -630,7 +630,7 @@ def test_cli_batch_invalid_in_prescan(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=validate_side_effect), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -725,7 +725,7 @@ def test_cli_batch_fallback_warning(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), @@ -753,7 +753,7 @@ def test_cli_batch_empty_speech_warning(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model, "cpu")), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_empty, tfr_ok]), patch("local_transcriber.cli.write_transcript"), @@ -784,7 +784,7 @@ def test_cli_batch_midstream_fallback_warning(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", return_value=(model_gpu, "cuda")), patch("local_transcriber.cli._transcribe_file", side_effect=[tfr_fallback, tfr_ok]), patch("local_transcriber.cli.write_transcript"), @@ -811,7 +811,7 @@ def test_cli_batch_model_loaded_once(tmp_path): patch("local_transcriber.cli.check_ffmpeg"), patch("local_transcriber.cli.validate_input_file", side_effect=lambda p: p), 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.ensure_model_available", return_value="/models/medium"), patch("local_transcriber.cli.load_model", mock_load_model), patch("local_transcriber.cli._transcribe_file", return_value=tfr), patch("local_transcriber.cli.write_transcript"), diff --git a/tests/test_config.py b/tests/test_config.py index 20733f0..0fdd84e 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -4,6 +4,7 @@ from unittest.mock import patch import pytest from local_transcriber.config import ( + apply_device_defaults, find_config_file, load_config, resolve_defaults, @@ -89,8 +90,41 @@ def test_resolve_defaults_hardcoded_fallback(): {"model": None, "language": None, "device": None, "compute_type": None}, {} ) assert result == { - "model": "large-v3", - "language": "auto", + "model": "medium", + "language": "ru", "device": "auto", - "compute_type": "int8", + "compute_type": "float32", } + + +def test_apply_device_defaults_cuda(): + defaults = {"model": "medium", "language": "ru", "device": "auto", "compute_type": "float32"} + cli = {"model": None, "language": None, "device": None, "compute_type": None} + result = apply_device_defaults(defaults, "cuda", cli, {}) + assert result["model"] == "medium" + assert result["compute_type"] == "float16" + + +def test_apply_device_defaults_cpu(): + defaults = {"model": "medium", "language": "ru", "device": "auto", "compute_type": "float32"} + cli = {"model": None, "language": None, "device": None, "compute_type": None} + result = apply_device_defaults(defaults, "cpu", cli, {}) + assert result["model"] == "medium" + assert result["compute_type"] == "float32" + + +def test_apply_device_defaults_cli_overrides(): + defaults = {"model": "large-v3", "language": "ru", "device": "auto", "compute_type": "int8"} + cli = {"model": "large-v3", "language": None, "device": None, "compute_type": "int8"} + result = apply_device_defaults(defaults, "cuda", cli, {}) + assert result["model"] == "large-v3" + assert result["compute_type"] == "int8" + + +def test_apply_device_defaults_config_overrides(): + defaults = {"model": "small", "language": "ru", "device": "auto", "compute_type": "int8"} + cli = {"model": None, "language": None, "device": None, "compute_type": None} + config = {"model": "small", "compute_type": "int8"} + result = apply_device_defaults(defaults, "cuda", cli, config) + assert result["model"] == "small" + assert result["compute_type"] == "int8"