mirror of
https://github.com/linyqh/NarratoAI.git
synced 2026-08-01 10:55:54 +00:00
功能: 新增字幕任务专用高效模型支持
新增高推理模型和高效模型的中英文国际化配置字符串 新增高效文本模型配置项,支持空值时自动回退到高推理模型 添加resolve_text_model_name工具函数用于根据偏好解析对应文本模型 更新字幕校对和翻译服务,默认使用高效模型处理批量字幕任务 优化OpenAI兼容提供商的模型参数处理逻辑,支持通过参数覆盖模型名称 更新WebUI基础设置页面,新增高效模型配置项并支持测试两个模型的连接 同步更新示例配置文件、默认配置项与相关测试用例
This commit is contained in:
parent
022b8bbea3
commit
0a5dcf5f21
@ -10,6 +10,7 @@ DEFAULT_VISION_OPENAI_MODEL_NAME = "Qwen/Qwen3.5-122B-A10B"
|
|||||||
|
|
||||||
DEFAULT_TEXT_LLM_PROVIDER = DEFAULT_OPENAI_COMPATIBLE_PROVIDER
|
DEFAULT_TEXT_LLM_PROVIDER = DEFAULT_OPENAI_COMPATIBLE_PROVIDER
|
||||||
DEFAULT_TEXT_OPENAI_MODEL_NAME = "Pro/zai-org/GLM-5"
|
DEFAULT_TEXT_OPENAI_MODEL_NAME = "Pro/zai-org/GLM-5"
|
||||||
|
DEFAULT_TEXT_OPENAI_FAST_MODEL_NAME = ""
|
||||||
|
|
||||||
DEFAULT_LLM_GENERATION_CONFIG = {
|
DEFAULT_LLM_GENERATION_CONFIG = {
|
||||||
"temperature": 1.0,
|
"temperature": 1.0,
|
||||||
@ -33,6 +34,7 @@ DEFAULT_LLM_APP_CONFIG = {
|
|||||||
"vision_openai_base_url": DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
"vision_openai_base_url": DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
||||||
"text_llm_provider": DEFAULT_TEXT_LLM_PROVIDER,
|
"text_llm_provider": DEFAULT_TEXT_LLM_PROVIDER,
|
||||||
"text_openai_model_name": DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
"text_openai_model_name": DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
||||||
|
"text_openai_fast_model_name": DEFAULT_TEXT_OPENAI_FAST_MODEL_NAME,
|
||||||
"text_openai_api_key": "",
|
"text_openai_api_key": "",
|
||||||
"text_openai_base_url": DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
"text_openai_base_url": DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
||||||
"tavily_api_key": "",
|
"tavily_api_key": "",
|
||||||
@ -69,6 +71,31 @@ def normalize_openai_compatible_model_name(
|
|||||||
return normalized
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_text_model_name(
|
||||||
|
app_config: dict,
|
||||||
|
provider: str = DEFAULT_OPENAI_COMPATIBLE_PROVIDER,
|
||||||
|
*,
|
||||||
|
prefer_fast: bool = False,
|
||||||
|
) -> str:
|
||||||
|
"""Resolve the configured reasoning or fast text model with legacy fallback."""
|
||||||
|
provider = (provider or DEFAULT_OPENAI_COMPATIBLE_PROVIDER).strip().lower()
|
||||||
|
reasoning_model = normalize_openai_compatible_model_name(
|
||||||
|
str(app_config.get(f"text_{provider}_model_name") or ""),
|
||||||
|
provider=provider,
|
||||||
|
)
|
||||||
|
if not reasoning_model and provider == DEFAULT_OPENAI_COMPATIBLE_PROVIDER:
|
||||||
|
reasoning_model = DEFAULT_TEXT_OPENAI_MODEL_NAME
|
||||||
|
|
||||||
|
if not prefer_fast:
|
||||||
|
return reasoning_model
|
||||||
|
|
||||||
|
fast_model = normalize_openai_compatible_model_name(
|
||||||
|
str(app_config.get(f"text_{provider}_fast_model_name") or ""),
|
||||||
|
provider=provider,
|
||||||
|
)
|
||||||
|
return fast_model or reasoning_model
|
||||||
|
|
||||||
|
|
||||||
def get_openai_compatible_ui_values(
|
def get_openai_compatible_ui_values(
|
||||||
full_model_name: str,
|
full_model_name: str,
|
||||||
default_model: str,
|
default_model: str,
|
||||||
|
|||||||
@ -12,6 +12,7 @@ from app.config import config as cfg
|
|||||||
from app.config.defaults import (
|
from app.config.defaults import (
|
||||||
get_openai_compatible_ui_values,
|
get_openai_compatible_ui_values,
|
||||||
normalize_openai_compatible_model_name,
|
normalize_openai_compatible_model_name,
|
||||||
|
resolve_text_model_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@ -82,11 +83,13 @@ hide_config = true
|
|||||||
self.assertEqual(0.95, config_data["app"]["vision_openai_top_p"])
|
self.assertEqual(0.95, config_data["app"]["vision_openai_top_p"])
|
||||||
self.assertEqual("openai", config_data["app"]["text_llm_provider"])
|
self.assertEqual("openai", config_data["app"]["text_llm_provider"])
|
||||||
self.assertEqual("Pro/zai-org/GLM-5", config_data["app"]["text_openai_model_name"])
|
self.assertEqual("Pro/zai-org/GLM-5", config_data["app"]["text_openai_model_name"])
|
||||||
|
self.assertEqual("", config_data["app"]["text_openai_fast_model_name"])
|
||||||
self.assertEqual("https://api.siliconflow.cn/v1", config_data["app"]["text_openai_base_url"])
|
self.assertEqual("https://api.siliconflow.cn/v1", config_data["app"]["text_openai_base_url"])
|
||||||
self.assertEqual(1.0, config_data["app"]["text_openai_temperature"])
|
self.assertEqual(1.0, config_data["app"]["text_openai_temperature"])
|
||||||
self.assertEqual(0.95, config_data["app"]["text_openai_top_p"])
|
self.assertEqual(0.95, config_data["app"]["text_openai_top_p"])
|
||||||
self.assertEqual("Qwen/Qwen3.5-122B-A10B", saved_config["app"]["vision_openai_model_name"])
|
self.assertEqual("Qwen/Qwen3.5-122B-A10B", saved_config["app"]["vision_openai_model_name"])
|
||||||
self.assertEqual("Pro/zai-org/GLM-5", saved_config["app"]["text_openai_model_name"])
|
self.assertEqual("Pro/zai-org/GLM-5", saved_config["app"]["text_openai_model_name"])
|
||||||
|
self.assertEqual("", saved_config["app"]["text_openai_fast_model_name"])
|
||||||
self.assertTrue(saved_config["app"]["hide_config"])
|
self.assertTrue(saved_config["app"]["hide_config"])
|
||||||
|
|
||||||
def test_legacy_indextts2_config_is_migrated_to_indextts_15(self):
|
def test_legacy_indextts2_config_is_migrated_to_indextts_15(self):
|
||||||
@ -127,6 +130,23 @@ hide_config = true
|
|||||||
|
|
||||||
|
|
||||||
class OpenAICompatibleModelDefaultsTests(unittest.TestCase):
|
class OpenAICompatibleModelDefaultsTests(unittest.TestCase):
|
||||||
|
def test_fast_text_model_falls_back_to_reasoning_model(self):
|
||||||
|
app_config = {
|
||||||
|
"text_openai_model_name": "reasoning-model",
|
||||||
|
"text_openai_fast_model_name": "",
|
||||||
|
}
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
"reasoning-model",
|
||||||
|
resolve_text_model_name(app_config, "openai", prefer_fast=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
app_config["text_openai_fast_model_name"] = "fast-model"
|
||||||
|
self.assertEqual(
|
||||||
|
"fast-model",
|
||||||
|
resolve_text_model_name(app_config, "openai", prefer_fast=True),
|
||||||
|
)
|
||||||
|
|
||||||
def test_ui_keeps_full_model_name_and_openai_provider(self):
|
def test_ui_keeps_full_model_name_and_openai_provider(self):
|
||||||
provider, model_name = get_openai_compatible_ui_values(
|
provider, model_name = get_openai_compatible_ui_values(
|
||||||
"Qwen/Qwen3.5-122B-A10B",
|
"Qwen/Qwen3.5-122B-A10B",
|
||||||
|
|||||||
@ -164,11 +164,13 @@ class LLMConfigValidator:
|
|||||||
config_prefix = f"text_{provider_name}"
|
config_prefix = f"text_{provider_name}"
|
||||||
api_key = config.app.get(f'{config_prefix}_api_key')
|
api_key = config.app.get(f'{config_prefix}_api_key')
|
||||||
model_name = config.app.get(f'{config_prefix}_model_name')
|
model_name = config.app.get(f'{config_prefix}_model_name')
|
||||||
|
fast_model_name = config.app.get(f'{config_prefix}_fast_model_name')
|
||||||
base_url = config.app.get(f'{config_prefix}_base_url')
|
base_url = config.app.get(f'{config_prefix}_base_url')
|
||||||
|
|
||||||
result["config"] = {
|
result["config"] = {
|
||||||
"api_key": "***" if api_key else None,
|
"api_key": "***" if api_key else None,
|
||||||
"model_name": model_name,
|
"model_name": model_name,
|
||||||
|
"fast_model_name": fast_model_name,
|
||||||
"base_url": base_url
|
"base_url": base_url
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -241,7 +243,8 @@ class LLMConfigValidator:
|
|||||||
f"text_{provider}_model_name"
|
f"text_{provider}_model_name"
|
||||||
],
|
],
|
||||||
"optional_configs": [
|
"optional_configs": [
|
||||||
f"text_{provider}_base_url"
|
f"text_{provider}_base_url",
|
||||||
|
f"text_{provider}_fast_model_name",
|
||||||
],
|
],
|
||||||
"example_models": LLMConfigValidator._get_example_models(provider, "text")
|
"example_models": LLMConfigValidator._get_example_models(provider, "text")
|
||||||
}
|
}
|
||||||
|
|||||||
@ -258,8 +258,9 @@ class OpenAICompatibleTextProvider(_OpenAICompatibleBase, TextModelProvider):
|
|||||||
response_format: Optional[str],
|
response_format: Optional[str],
|
||||||
kwargs: Dict[str, Any],
|
kwargs: Dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
model_name = _normalize_model_name(self.model_name)
|
|
||||||
generation_kwargs = dict(kwargs)
|
generation_kwargs = dict(kwargs)
|
||||||
|
model_override = generation_kwargs.pop("model", None) or generation_kwargs.pop("model_name", None)
|
||||||
|
model_name = _normalize_model_name(model_override or self.model_name)
|
||||||
temperature_override = generation_kwargs.pop("temperature", None)
|
temperature_override = generation_kwargs.pop("temperature", None)
|
||||||
if temperature_override is None and temperature != 1.0:
|
if temperature_override is None and temperature != 1.0:
|
||||||
temperature_override = temperature
|
temperature_override = temperature
|
||||||
|
|||||||
@ -149,6 +149,20 @@ class OpenAICompatGenerationOptionTests(unittest.TestCase):
|
|||||||
self.assertEqual(65536, options["max_tokens"])
|
self.assertEqual(65536, options["max_tokens"])
|
||||||
self.assertNotIn("extra_body", options)
|
self.assertNotIn("extra_body", options)
|
||||||
|
|
||||||
|
def test_text_request_can_override_model_for_fast_tasks(self):
|
||||||
|
provider = OpenAICompatibleTextProvider(api_key="k", model_name="reasoning-model")
|
||||||
|
|
||||||
|
options = provider._build_text_completion_kwargs(
|
||||||
|
messages=[{"role": "user", "content": "hello"}],
|
||||||
|
temperature=0.2,
|
||||||
|
max_tokens=None,
|
||||||
|
response_format=None,
|
||||||
|
kwargs={"model": "fast-model", "thinking_level": "off"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual("fast-model", options["model"])
|
||||||
|
self.assertNotIn("extra_body", options)
|
||||||
|
|
||||||
def test_build_options_uses_per_model_generation_config(self):
|
def test_build_options_uses_per_model_generation_config(self):
|
||||||
provider = OpenAICompatibleTextProvider(api_key="k", model_name="m")
|
provider = OpenAICompatibleTextProvider(api_key="k", model_name="m")
|
||||||
config.app.update(
|
config.app.update(
|
||||||
|
|||||||
@ -10,6 +10,8 @@ from typing import Any
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
|
from app.config import config
|
||||||
|
from app.config.defaults import resolve_text_model_name
|
||||||
from app.services.llm.manager import LLMServiceManager
|
from app.services.llm.manager import LLMServiceManager
|
||||||
from app.services.llm.migration_adapter import _run_async_safely
|
from app.services.llm.migration_adapter import _run_async_safely
|
||||||
from app.services.llm.unified_service import UnifiedLLMService
|
from app.services.llm.unified_service import UnifiedLLMService
|
||||||
@ -174,12 +176,16 @@ def correct_srt_content(
|
|||||||
provider: str = "",
|
provider: str = "",
|
||||||
api_key: str = "",
|
api_key: str = "",
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
|
model_name: str = "",
|
||||||
temperature: float = 0.1,
|
temperature: float = 0.1,
|
||||||
) -> str:
|
) -> str:
|
||||||
blocks = parse_srt_blocks(srt_content)
|
blocks = parse_srt_blocks(srt_content)
|
||||||
_ensure_llm_providers_registered()
|
_ensure_llm_providers_registered()
|
||||||
|
|
||||||
logger.info(f"开始校准字幕,共 {len(blocks)} 条")
|
resolved_model_name = str(
|
||||||
|
model_name or resolve_text_model_name(config.app, provider, prefer_fast=True)
|
||||||
|
).strip()
|
||||||
|
logger.info(f"开始使用高效率模型 {resolved_model_name} 校准字幕,共 {len(blocks)} 条")
|
||||||
prompt = _build_correction_prompt(blocks)
|
prompt = _build_correction_prompt(blocks)
|
||||||
raw_output = _run_async_safely(
|
raw_output = _run_async_safely(
|
||||||
UnifiedLLMService.generate_text,
|
UnifiedLLMService.generate_text,
|
||||||
@ -190,6 +196,8 @@ def correct_srt_content(
|
|||||||
response_format="json",
|
response_format="json",
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
api_base=base_url,
|
api_base=base_url,
|
||||||
|
model=resolved_model_name,
|
||||||
|
thinking_level="off",
|
||||||
)
|
)
|
||||||
corrections = _parse_corrections(raw_output, {block.order for block in blocks})
|
corrections = _parse_corrections(raw_output, {block.order for block in blocks})
|
||||||
corrected_srt = _render_srt(blocks, corrections)
|
corrected_srt = _render_srt(blocks, corrections)
|
||||||
@ -215,6 +223,7 @@ def correct_subtitle_file(
|
|||||||
provider: str = "",
|
provider: str = "",
|
||||||
api_key: str = "",
|
api_key: str = "",
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
|
model_name: str = "",
|
||||||
temperature: float = 0.1,
|
temperature: float = 0.1,
|
||||||
) -> str:
|
) -> str:
|
||||||
if not subtitle_file or not os.path.isfile(subtitle_file):
|
if not subtitle_file or not os.path.isfile(subtitle_file):
|
||||||
@ -226,6 +235,7 @@ def correct_subtitle_file(
|
|||||||
provider=provider,
|
provider=provider,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
|
model_name=model_name,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
)
|
)
|
||||||
return write_srt_file(corrected_srt, output_file)
|
return write_srt_file(corrected_srt, output_file)
|
||||||
|
|||||||
@ -11,6 +11,7 @@ from typing import Any, Callable
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from app.config import config
|
from app.config import config
|
||||||
|
from app.config.defaults import resolve_text_model_name
|
||||||
from app.services.llm.migration_adapter import _run_async_safely
|
from app.services.llm.migration_adapter import _run_async_safely
|
||||||
from app.services.llm.unified_service import UnifiedLLMService
|
from app.services.llm.unified_service import UnifiedLLMService
|
||||||
from app.services.subtitle_corrector import (
|
from app.services.subtitle_corrector import (
|
||||||
@ -151,6 +152,7 @@ def _translate_chunk(
|
|||||||
provider: str,
|
provider: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
|
model_name: str,
|
||||||
temperature: float,
|
temperature: float,
|
||||||
max_repair_attempts: int,
|
max_repair_attempts: int,
|
||||||
) -> dict[int, str]:
|
) -> dict[int, str]:
|
||||||
@ -189,6 +191,8 @@ def _translate_chunk(
|
|||||||
response_format="json",
|
response_format="json",
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
api_base=base_url,
|
api_base=base_url,
|
||||||
|
model=model_name,
|
||||||
|
thinking_level="off",
|
||||||
)
|
)
|
||||||
last_output = str(raw_output or "")
|
last_output = str(raw_output or "")
|
||||||
try:
|
try:
|
||||||
@ -243,6 +247,7 @@ def translate_srt_content(
|
|||||||
provider: str = "",
|
provider: str = "",
|
||||||
api_key: str = "",
|
api_key: str = "",
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
|
model_name: str = "",
|
||||||
temperature: float = 0.2,
|
temperature: float = 0.2,
|
||||||
batch_size: int | None = None,
|
batch_size: int | None = None,
|
||||||
max_workers: int | None = None,
|
max_workers: int | None = None,
|
||||||
@ -251,6 +256,9 @@ def translate_srt_content(
|
|||||||
target_language = str(target_language or "").strip() or "中文"
|
target_language = str(target_language or "").strip() or "中文"
|
||||||
blocks = parse_srt_blocks(srt_content)
|
blocks = parse_srt_blocks(srt_content)
|
||||||
_ensure_llm_providers_registered()
|
_ensure_llm_providers_registered()
|
||||||
|
resolved_model_name = str(
|
||||||
|
model_name or resolve_text_model_name(config.app, provider, prefer_fast=True)
|
||||||
|
).strip()
|
||||||
|
|
||||||
resolved_batch_size = _resolve_batch_size(batch_size)
|
resolved_batch_size = _resolve_batch_size(batch_size)
|
||||||
chunks = _split_blocks(blocks, resolved_batch_size)
|
chunks = _split_blocks(blocks, resolved_batch_size)
|
||||||
@ -260,7 +268,8 @@ def translate_srt_content(
|
|||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"开始批量翻译字幕: 共 {total_blocks} 条, {total_chunks} 批, "
|
f"开始批量翻译字幕: 共 {total_blocks} 条, {total_chunks} 批, "
|
||||||
f"每批最多 {resolved_batch_size} 条, 并发 {resolved_max_workers}, 目标语言: {target_language}"
|
f"每批最多 {resolved_batch_size} 条, 并发 {resolved_max_workers}, "
|
||||||
|
f"目标语言: {target_language}, 高效率模型: {resolved_model_name}"
|
||||||
)
|
)
|
||||||
|
|
||||||
translations: dict[int, str] = {}
|
translations: dict[int, str] = {}
|
||||||
@ -282,6 +291,7 @@ def translate_srt_content(
|
|||||||
provider=provider,
|
provider=provider,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
|
model_name=resolved_model_name,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
max_repair_attempts=DEFAULT_MAX_REPAIR_ATTEMPTS,
|
max_repair_attempts=DEFAULT_MAX_REPAIR_ATTEMPTS,
|
||||||
)
|
)
|
||||||
@ -301,6 +311,7 @@ def translate_srt_content(
|
|||||||
provider=provider,
|
provider=provider,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
|
model_name=resolved_model_name,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
max_repair_attempts=DEFAULT_MAX_REPAIR_ATTEMPTS,
|
max_repair_attempts=DEFAULT_MAX_REPAIR_ATTEMPTS,
|
||||||
)
|
)
|
||||||
@ -347,6 +358,7 @@ def translate_subtitle_file(
|
|||||||
provider: str = "",
|
provider: str = "",
|
||||||
api_key: str = "",
|
api_key: str = "",
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
|
model_name: str = "",
|
||||||
temperature: float = 0.2,
|
temperature: float = 0.2,
|
||||||
batch_size: int | None = None,
|
batch_size: int | None = None,
|
||||||
max_workers: int | None = None,
|
max_workers: int | None = None,
|
||||||
@ -362,6 +374,7 @@ def translate_subtitle_file(
|
|||||||
provider=provider,
|
provider=provider,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
|
model_name=model_name,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
batch_size=batch_size,
|
batch_size=batch_size,
|
||||||
max_workers=max_workers,
|
max_workers=max_workers,
|
||||||
|
|||||||
@ -27,6 +27,14 @@ class SubtitleCorrectorTests(unittest.TestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
with (
|
with (
|
||||||
|
mock.patch.dict(
|
||||||
|
corrector.config.app,
|
||||||
|
{
|
||||||
|
"text_openai_model_name": "reasoning-model",
|
||||||
|
"text_openai_fast_model_name": "fast-subtitle-model",
|
||||||
|
},
|
||||||
|
clear=False,
|
||||||
|
),
|
||||||
mock.patch("app.services.subtitle_corrector._ensure_llm_providers_registered"),
|
mock.patch("app.services.subtitle_corrector._ensure_llm_providers_registered"),
|
||||||
mock.patch(
|
mock.patch(
|
||||||
"app.services.subtitle_corrector._run_async_safely",
|
"app.services.subtitle_corrector._run_async_safely",
|
||||||
@ -49,6 +57,8 @@ class SubtitleCorrectorTests(unittest.TestCase):
|
|||||||
self.assertEqual("openai", call_kwargs["provider"])
|
self.assertEqual("openai", call_kwargs["provider"])
|
||||||
self.assertEqual("sk-test", call_kwargs["api_key"])
|
self.assertEqual("sk-test", call_kwargs["api_key"])
|
||||||
self.assertEqual("https://llm.example/v1", call_kwargs["api_base"])
|
self.assertEqual("https://llm.example/v1", call_kwargs["api_base"])
|
||||||
|
self.assertEqual("fast-subtitle-model", call_kwargs["model"])
|
||||||
|
self.assertEqual("off", call_kwargs["thinking_level"])
|
||||||
self.assertEqual("json", call_kwargs["response_format"])
|
self.assertEqual("json", call_kwargs["response_format"])
|
||||||
self.assertIn("多语言字幕校对员", call_kwargs["system_prompt"])
|
self.assertIn("多语言字幕校对员", call_kwargs["system_prompt"])
|
||||||
self.assertIn("保持原语言", call_kwargs["prompt"])
|
self.assertIn("保持原语言", call_kwargs["prompt"])
|
||||||
|
|||||||
@ -48,6 +48,14 @@ class SubtitleTranslatorTests(unittest.TestCase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
with (
|
with (
|
||||||
|
mock.patch.dict(
|
||||||
|
translator.config.app,
|
||||||
|
{
|
||||||
|
"text_openai_model_name": "reasoning-model",
|
||||||
|
"text_openai_fast_model_name": "fast-subtitle-model",
|
||||||
|
},
|
||||||
|
clear=False,
|
||||||
|
),
|
||||||
mock.patch("app.services.subtitle_translator._ensure_llm_providers_registered"),
|
mock.patch("app.services.subtitle_translator._ensure_llm_providers_registered"),
|
||||||
mock.patch(
|
mock.patch(
|
||||||
"app.services.subtitle_translator._run_async_safely",
|
"app.services.subtitle_translator._run_async_safely",
|
||||||
@ -71,6 +79,8 @@ class SubtitleTranslatorTests(unittest.TestCase):
|
|||||||
self.assertEqual("openai", call_kwargs["provider"])
|
self.assertEqual("openai", call_kwargs["provider"])
|
||||||
self.assertEqual("sk-test", call_kwargs["api_key"])
|
self.assertEqual("sk-test", call_kwargs["api_key"])
|
||||||
self.assertEqual("https://llm.example/v1", call_kwargs["api_base"])
|
self.assertEqual("https://llm.example/v1", call_kwargs["api_base"])
|
||||||
|
self.assertEqual("fast-subtitle-model", call_kwargs["model"])
|
||||||
|
self.assertEqual("off", call_kwargs["thinking_level"])
|
||||||
self.assertEqual("json", call_kwargs["response_format"])
|
self.assertEqual("json", call_kwargs["response_format"])
|
||||||
self.assertIn("专业字幕翻译员", call_kwargs["system_prompt"])
|
self.assertIn("专业字幕翻译员", call_kwargs["system_prompt"])
|
||||||
self.assertIn("翻译为中文", call_kwargs["prompt"])
|
self.assertIn("翻译为中文", call_kwargs["prompt"])
|
||||||
|
|||||||
@ -53,7 +53,8 @@
|
|||||||
# - Qwen: qwen/qwen-plus, qwen/qwen-turbo
|
# - Qwen: qwen/qwen-plus, qwen/qwen-turbo
|
||||||
# - SiliconFlow: siliconflow/deepseek-ai/DeepSeek-R1
|
# - SiliconFlow: siliconflow/deepseek-ai/DeepSeek-R1
|
||||||
# - Moonshot: moonshot/moonshot-v1-8k
|
# - Moonshot: moonshot/moonshot-v1-8k
|
||||||
text_openai_model_name = "Pro/zai-org/GLM-5"
|
text_openai_model_name = "Pro/zai-org/GLM-5" # 高推理模型:剧情分析、文案生成、脚本匹配
|
||||||
|
text_openai_fast_model_name = "" # 高效率模型:字幕翻译、字幕校准;留空时回退到高推理模型
|
||||||
text_openai_api_key = "" # 填入对应 provider 的 API key
|
text_openai_api_key = "" # 填入对应 provider 的 API key
|
||||||
text_openai_base_url = "https://api.siliconflow.cn/v1" # 可选:自定义 API base URL;界面会提示 API key 将发送到对应端点
|
text_openai_base_url = "https://api.siliconflow.cn/v1" # 可选:自定义 API base URL;界面会提示 API key 将发送到对应端点
|
||||||
text_openai_temperature = 1.0
|
text_openai_temperature = 1.0
|
||||||
|
|||||||
@ -9,6 +9,7 @@ from app.config.defaults import (
|
|||||||
DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
DEFAULT_OPENAI_COMPATIBLE_BASE_URL,
|
||||||
DEFAULT_OPENAI_COMPATIBLE_PROVIDER,
|
DEFAULT_OPENAI_COMPATIBLE_PROVIDER,
|
||||||
DEFAULT_TEXT_LLM_PROVIDER,
|
DEFAULT_TEXT_LLM_PROVIDER,
|
||||||
|
DEFAULT_TEXT_OPENAI_FAST_MODEL_NAME,
|
||||||
DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
||||||
DEFAULT_VISION_LLM_PROVIDER,
|
DEFAULT_VISION_LLM_PROVIDER,
|
||||||
DEFAULT_VISION_OPENAI_MODEL_NAME,
|
DEFAULT_VISION_OPENAI_MODEL_NAME,
|
||||||
@ -876,6 +877,10 @@ def render_text_llm_settings(tr):
|
|||||||
|
|
||||||
# 获取已保存的配置
|
# 获取已保存的配置
|
||||||
full_text_model_name = config.app.get("text_openai_model_name") or DEFAULT_TEXT_OPENAI_MODEL_NAME
|
full_text_model_name = config.app.get("text_openai_model_name") or DEFAULT_TEXT_OPENAI_MODEL_NAME
|
||||||
|
full_fast_model_name = (
|
||||||
|
config.app.get("text_openai_fast_model_name")
|
||||||
|
or DEFAULT_TEXT_OPENAI_FAST_MODEL_NAME
|
||||||
|
)
|
||||||
text_api_key = config.app.get("text_openai_api_key", "")
|
text_api_key = config.app.get("text_openai_api_key", "")
|
||||||
text_base_url = config.app.get("text_openai_base_url", DEFAULT_OPENAI_COMPATIBLE_BASE_URL)
|
text_base_url = config.app.get("text_openai_base_url", DEFAULT_OPENAI_COMPATIBLE_BASE_URL)
|
||||||
|
|
||||||
@ -885,10 +890,14 @@ def render_text_llm_settings(tr):
|
|||||||
DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
DEFAULT_TEXT_OPENAI_MODEL_NAME,
|
||||||
provider=DEFAULT_TEXT_LLM_PROVIDER,
|
provider=DEFAULT_TEXT_LLM_PROVIDER,
|
||||||
)
|
)
|
||||||
|
current_fast_model = normalize_openai_compatible_model_id(
|
||||||
|
full_fast_model_name,
|
||||||
|
provider=DEFAULT_TEXT_LLM_PROVIDER,
|
||||||
|
)
|
||||||
selected_provider = DEFAULT_TEXT_LLM_PROVIDER
|
selected_provider = DEFAULT_TEXT_LLM_PROVIDER
|
||||||
|
|
||||||
# 渲染配置输入框
|
# 渲染配置输入框
|
||||||
col1, col2 = st.columns([1, 2])
|
col1, col2, col3 = st.columns([1, 2, 2])
|
||||||
with col1:
|
with col1:
|
||||||
render_openai_compatible_protocol_field(
|
render_openai_compatible_protocol_field(
|
||||||
tr,
|
tr,
|
||||||
@ -897,11 +906,13 @@ def render_text_llm_settings(tr):
|
|||||||
)
|
)
|
||||||
|
|
||||||
with col2:
|
with col2:
|
||||||
model_name_input = st.text_input(
|
reasoning_model_name_input = st.text_input(
|
||||||
tr("Text Model Name"),
|
tr("High Reasoning Model Name"),
|
||||||
value=current_model,
|
value=current_model,
|
||||||
help=(
|
help=(
|
||||||
tr("Model Name Input Help")
|
tr("High Reasoning Model Help")
|
||||||
|
+ "\n\n"
|
||||||
|
+ tr("Model Name Input Help")
|
||||||
+ "\n\n"
|
+ "\n\n"
|
||||||
+ "• Pro/zai-org/GLM-5\n"
|
+ "• Pro/zai-org/GLM-5\n"
|
||||||
+ "• deepseek/deepseek-chat\n"
|
+ "• deepseek/deepseek-chat\n"
|
||||||
@ -912,8 +923,24 @@ def render_text_llm_settings(tr):
|
|||||||
key="text_model_input"
|
key="text_model_input"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
with col3:
|
||||||
|
fast_model_name_input = st.text_input(
|
||||||
|
tr("High Efficiency Model Name"),
|
||||||
|
value=current_fast_model,
|
||||||
|
help=(
|
||||||
|
tr("High Efficiency Model Help")
|
||||||
|
+ "\n\n"
|
||||||
|
+ "• Qwen/Qwen3.5-32B\n"
|
||||||
|
+ "• gpt-4o-mini\n"
|
||||||
|
+ "• gemini-2.5-flash\n"
|
||||||
|
+ "• deepseek/deepseek-chat"
|
||||||
|
),
|
||||||
|
key="text_fast_model_input",
|
||||||
|
)
|
||||||
|
|
||||||
# 组合完整的模型名称
|
# 组合完整的模型名称
|
||||||
st_text_model_name = normalize_openai_compatible_model_name(model_name_input)
|
st_text_model_name = normalize_openai_compatible_model_name(reasoning_model_name_input)
|
||||||
|
st_text_fast_model_name = normalize_openai_compatible_model_name(fast_model_name_input)
|
||||||
|
|
||||||
st_text_api_key = st.text_input(
|
st_text_api_key = st.text_input(
|
||||||
tr("Text API Key"),
|
tr("Text API Key"),
|
||||||
@ -952,7 +979,7 @@ def render_text_llm_settings(tr):
|
|||||||
test_errors = []
|
test_errors = []
|
||||||
if not st_text_api_key:
|
if not st_text_api_key:
|
||||||
test_errors.append(tr("Please enter API key"))
|
test_errors.append(tr("Please enter API key"))
|
||||||
if not model_name_input:
|
if not reasoning_model_name_input:
|
||||||
test_errors.append(tr("Please enter model name"))
|
test_errors.append(tr("Please enter model name"))
|
||||||
|
|
||||||
if test_errors:
|
if test_errors:
|
||||||
@ -961,17 +988,25 @@ def render_text_llm_settings(tr):
|
|||||||
else:
|
else:
|
||||||
with st.spinner(tr("Testing connection...")):
|
with st.spinner(tr("Testing connection...")):
|
||||||
try:
|
try:
|
||||||
success, message = test_openai_compatible_text_model(
|
test_targets = [
|
||||||
api_key=st_text_api_key,
|
(tr("High Reasoning Model Name"), st_text_model_name),
|
||||||
base_url=st_text_base_url,
|
]
|
||||||
model_name=st_text_model_name,
|
if st_text_fast_model_name:
|
||||||
tr=tr
|
test_targets.append((
|
||||||
)
|
tr("High Efficiency Model Name"),
|
||||||
|
st_text_fast_model_name,
|
||||||
if success:
|
))
|
||||||
st.success(message)
|
for label, target_model in test_targets:
|
||||||
else:
|
success, message = test_openai_compatible_text_model(
|
||||||
st.error(message)
|
api_key=st_text_api_key,
|
||||||
|
base_url=st_text_base_url,
|
||||||
|
model_name=target_model,
|
||||||
|
tr=tr,
|
||||||
|
)
|
||||||
|
if success:
|
||||||
|
st.success(f"{label}: {message}")
|
||||||
|
else:
|
||||||
|
st.error(f"{label}: {message}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"{tr('Connection test error')}: {str(e)}")
|
st.error(f"{tr('Connection test error')}: {str(e)}")
|
||||||
logger.error(f"OpenAI 兼容 文案生成模型连接测试失败: {str(e)}")
|
logger.error(f"OpenAI 兼容 文案生成模型连接测试失败: {str(e)}")
|
||||||
@ -992,6 +1027,25 @@ def render_text_llm_settings(tr):
|
|||||||
else:
|
else:
|
||||||
text_validation_errors.append(error_msg)
|
text_validation_errors.append(error_msg)
|
||||||
|
|
||||||
|
if st_text_fast_model_name:
|
||||||
|
is_valid, error_msg = validate_openai_compatible_model_name(
|
||||||
|
st_text_fast_model_name,
|
||||||
|
"高效率文案生成",
|
||||||
|
)
|
||||||
|
if is_valid:
|
||||||
|
text_config_changed |= update_app_config_if_changed(
|
||||||
|
"text_openai_fast_model_name",
|
||||||
|
st_text_fast_model_name,
|
||||||
|
)
|
||||||
|
st.session_state["text_openai_fast_model_name"] = st_text_fast_model_name
|
||||||
|
else:
|
||||||
|
text_validation_errors.append(error_msg)
|
||||||
|
else:
|
||||||
|
text_config_changed |= update_app_config_if_changed(
|
||||||
|
"text_openai_fast_model_name",
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
|
||||||
# 验证 API 密钥
|
# 验证 API 密钥
|
||||||
if st_text_api_key:
|
if st_text_api_key:
|
||||||
is_valid, error_msg = validate_api_key(st_text_api_key, "文案生成")
|
is_valid, error_msg = validate_api_key(st_text_api_key, "文案生成")
|
||||||
@ -1027,7 +1081,7 @@ def render_text_llm_settings(tr):
|
|||||||
config.save_config()
|
config.save_config()
|
||||||
# 清除缓存,确保下次使用新配置
|
# 清除缓存,确保下次使用新配置
|
||||||
UnifiedLLMService.clear_cache()
|
UnifiedLLMService.clear_cache()
|
||||||
if st_text_api_key or st_text_base_url or st_text_model_name:
|
if st_text_api_key or st_text_base_url or st_text_model_name or st_text_fast_model_name:
|
||||||
st.success(tr("Text model config saved"))
|
st.success(tr("Text model config saved"))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
st.error(f"{tr('Failed to save config')}: {str(e)}")
|
st.error(f"{tr('Failed to save config')}: {str(e)}")
|
||||||
|
|||||||
@ -205,6 +205,10 @@
|
|||||||
"Text API Key": "Text API Key",
|
"Text API Key": "Text API Key",
|
||||||
"Text Base URL": "Text Base URL",
|
"Text Base URL": "Text Base URL",
|
||||||
"Text Model Name": "Text Model Name",
|
"Text Model Name": "Text Model Name",
|
||||||
|
"High Reasoning Model Name": "High-Reasoning Model Name",
|
||||||
|
"High Efficiency Model Name": "High-Efficiency Model Name",
|
||||||
|
"High Reasoning Model Help": "Used for plot analysis, copy generation, and script generation and matching.",
|
||||||
|
"High Efficiency Model Help": "Used for subtitle translation and calibration; falls back to the high-reasoning model when empty.",
|
||||||
"Top P": "Top P",
|
"Top P": "Top P",
|
||||||
"Top K": "Top K",
|
"Top K": "Top K",
|
||||||
"Max Output Tokens": "Max Output Tokens",
|
"Max Output Tokens": "Max Output Tokens",
|
||||||
|
|||||||
@ -194,6 +194,10 @@
|
|||||||
"Text API Key": "文案生成 API 密钥",
|
"Text API Key": "文案生成 API 密钥",
|
||||||
"Text Base URL": "文案生成接口地址",
|
"Text Base URL": "文案生成接口地址",
|
||||||
"Text Model Name": "文案生成模型名称",
|
"Text Model Name": "文案生成模型名称",
|
||||||
|
"High Reasoning Model Name": "高推理模型名称",
|
||||||
|
"High Efficiency Model Name": "高效率模型名称",
|
||||||
|
"High Reasoning Model Help": "用于剧情分析、文案生成、脚本生成与匹配等复杂任务。",
|
||||||
|
"High Efficiency Model Help": "用于字幕翻译、字幕校准等批量任务;留空时自动使用高推理模型。",
|
||||||
"Top P": "Top P",
|
"Top P": "Top P",
|
||||||
"Top K": "Top K",
|
"Top K": "Top K",
|
||||||
"Max Output Tokens": "最大输出 Token",
|
"Max Output Tokens": "最大输出 Token",
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user