From e5849ca1af1aa7c9fe4a482c815f9565c7cce7b1 Mon Sep 17 00:00:00 2001 From: pumpkinperson996 Date: Mon, 20 Jul 2026 20:19:44 -0500 Subject: [PATCH] Honor explicit Whisper model selection --- .../app/transcriber/transcriber_provider.py | 8 ++-- .../tests/test_whisper_config_forwarding.py | 38 +++++++++++++++++++ 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/backend/app/transcriber/transcriber_provider.py b/backend/app/transcriber/transcriber_provider.py index f4a7b50..af29c17 100644 --- a/backend/app/transcriber/transcriber_provider.py +++ b/backend/app/transcriber/transcriber_provider.py @@ -83,13 +83,13 @@ def get_mlx_whisper_transcriber(model_size="base"): return _init_transcriber(TranscriberType.MLX_WHISPER, MLXWhisperTranscriber, model_size=model_size) # 通用入口 -def get_transcriber(transcriber_type="fast-whisper", model_size="base", device="cuda"): +def get_transcriber(transcriber_type="fast-whisper", model_size=None, device="cuda"): """ 获取指定类型的转录器实例 参数: transcriber_type: 支持 "fast-whisper", "mlx-whisper", "bcut", "kuaishou", "groq" - model_size: 模型大小,适用于 whisper 类 + model_size: 模型大小,适用于 whisper 类;未提供时才读取环境变量默认值 device: 设备类型(如 cuda / cpu),仅 whisper 使用 返回: @@ -103,7 +103,9 @@ def get_transcriber(transcriber_type="fast-whisper", model_size="base", device=" logger.warning(f'未知转录器类型 "{transcriber_type}",默认使用 fast-whisper') transcriber_enum = TranscriberType.FAST_WHISPER - whisper_model_size = os.environ.get("WHISPER_MODEL_SIZE", model_size) + # The explicit value normally comes from the persisted frontend setting and + # must take precedence over Docker's startup default. + whisper_model_size = model_size or os.environ.get("WHISPER_MODEL_SIZE", "base") if transcriber_enum == TranscriberType.FAST_WHISPER: return get_whisper_transcriber(whisper_model_size, device=device) diff --git a/backend/tests/test_whisper_config_forwarding.py b/backend/tests/test_whisper_config_forwarding.py index 7d2ec81..2887541 100644 --- a/backend/tests/test_whisper_config_forwarding.py +++ b/backend/tests/test_whisper_config_forwarding.py @@ -80,3 +80,41 @@ def test_whisper_cache_is_rebuilt_when_model_size_changes(): assert turbo_again is turbo """ ) + + +def test_explicit_model_size_wins_over_environment_default(): + _run_isolated( + """ + import os + + from app.transcriber import transcriber_provider as provider + + class FakeWhisperTranscriber: + def __init__(self, model_size, device): + self.model_size = model_size + self.device = device + + os.environ["WHISPER_MODEL_SIZE"] = "tiny" + provider._transcribers = {key: None for key in provider._transcribers} + provider._transcriber_configs = { + key: None for key in provider._transcribers + } + provider.WhisperTranscriber = FakeWhisperTranscriber + + transcriber = provider.get_transcriber( + transcriber_type="fast-whisper", + model_size="large-v3-turbo", + device="cpu", + ) + + assert transcriber.model_size == "large-v3-turbo" + + fallback = provider.get_transcriber( + transcriber_type="fast-whisper", + model_size=None, + device="cpu", + ) + + assert fallback.model_size == "tiny" + """ + )