Files
intl_news/embedding/client.py
T
simon 9aa44e610d feat: LLM/Embedding 重试次数纳入 system.yaml 统一配置
- system.yaml: llm.max_retries(死配置) 改为 llm.max_attempts;embedding 段新增 max_attempts
- llm/client.py: LLMConfig 新增 max_attempts(默认 3 兜底),load_llm_config 读取配置
- llm/extractor.py: translate_and_extract(_async) max_attempts=None 时取 config.max_attempts
- embedding/client.py: load_embedding_config 读取 embedding.max_attempts
- scheduler/reporter.py: _call_llm_simple max_retries=None 时取 max_attempts-1
- 新增 2 个配置读取测试;已同步 pi5 验证
2026-08-05 08:48:43 +08:00

201 lines
6.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""DashScope Embedding 客户端。
通过 OpenAI 兼容接口调用阿里百炼 text-embedding-v3
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
model: text-embedding-v31024 维)
配置来源:
- .env → DASHSCOPE_API_KEY / QWEN_BASE_URLOpenAI 兼容端点)
- configs/system.yaml → embedding 段(model / dimension / batch_size / timeout
"""
import logging
import os
import time
from dataclasses import dataclass
from pathlib import Path
import yaml
from openai import OpenAI
from embedding.models import EmbeddingError
logger = logging.getLogger(__name__)
# 默认值
DASHSCOPE_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
DASHSCOPE_DEFAULT_MODEL = "text-embedding-v3"
DASHSCOPE_DEFAULT_DIM = 1024
DASHSCOPE_BATCH_LIMIT = 10 # 百炼实测单批上限
# 重试
DEFAULT_MAX_ATTEMPTS = 3
RETRY_BASE_WAIT_SEC = 1.0
RETRY_MAX_WAIT_SEC = 8.0
def _load_embedding_config() -> dict:
"""从 system.yaml 加载 embedding 段配置。"""
config_path = Path("configs/system.yaml")
if config_path.exists():
try:
with open(config_path, encoding="utf-8") as f:
raw = yaml.safe_load(f)
return raw.get("embedding", {})
except Exception:
logger.warning("加载 embedding 配置失败")
return {}
@dataclass
class EmbeddingConfig:
"""Embedding 调用配置。"""
provider: str = "dashscope"
model: str = DASHSCOPE_DEFAULT_MODEL
api_key: str = ""
base_url: str = DASHSCOPE_DEFAULT_BASE
dimension: int = DASHSCOPE_DEFAULT_DIM
batch_size: int = DASHSCOPE_BATCH_LIMIT
timeout_sec: float = 30.0
max_attempts: int = DEFAULT_MAX_ATTEMPTS
def __post_init__(self) -> None:
if not self.api_key:
raise EmbeddingError("DASHSCOPE_API_KEY 未配置,请检查 .env")
def load_embedding_config(
*,
model: str | None = None,
) -> EmbeddingConfig:
"""根据配置构造 EmbeddingConfig。
优先级: system.yaml > .env 默认值 > 硬编码默认值
"""
sys_cfg = _load_embedding_config()
# API key: 从环境变量读取
api_key = os.environ.get("DASHSCOPE_API_KEY", "")
if not api_key:
api_key = os.environ.get("QWEN_API_KEY", "")
# Base URL: 优先用 QWEN_BASE_URLOpenAI 兼容),
# DASHSCOPE_BASE_URL 通常是旧版非兼容端点,不作为默认
base_url = (
os.environ.get("QWEN_BASE_URL")
or os.environ.get("DASHSCOPE_EMBEDDING_BASE_URL")
or DASHSCOPE_DEFAULT_BASE
)
m = model or sys_cfg.get("dashscope_model", DASHSCOPE_DEFAULT_MODEL)
dimension = int(sys_cfg.get("dimension", DASHSCOPE_DEFAULT_DIM))
batch_size = min(int(sys_cfg.get("batch_size", DASHSCOPE_BATCH_LIMIT)), DASHSCOPE_BATCH_LIMIT)
timeout = float(sys_cfg.get("timeout_sec", 30.0))
max_attempts = int(sys_cfg.get("max_attempts", DEFAULT_MAX_ATTEMPTS))
if not api_key:
raise EmbeddingError("DASHSCOPE_API_KEY 未配置,请检查 .env")
return EmbeddingConfig(
provider="dashscope",
model=m,
api_key=api_key,
base_url=base_url,
dimension=dimension,
batch_size=batch_size,
timeout_sec=timeout,
max_attempts=max_attempts,
)
def make_embedding_client(config: EmbeddingConfig) -> OpenAI:
"""构造同步 OpenAI 客户端(指向 DashScope 兼容端点)。"""
logger.info(
"初始化 Embedding 客户端: provider=%s model=%s base_url=%s dim=%d",
config.provider, config.model, config.base_url, config.dimension,
)
return OpenAI(
api_key=config.api_key,
base_url=config.base_url,
timeout=config.timeout_sec,
)
def _chunked(items: list[str], size: int) -> list[list[str]]:
"""把列表按 size 分块。"""
return [items[i : i + size] for i in range(0, len(items), size)]
def embed_batch(
client: OpenAI,
config: EmbeddingConfig,
texts: list[str],
) -> list[list[float]]:
"""批量嵌入,自动分块+重试。
Args:
client: OpenAI 客户端
config: Embedding 配置
texts: 待嵌入文本列表
Returns:
与 texts 等长的向量列表,每个为 1024 维 float 列表
"""
if not texts:
return []
all_results: list[list[float]] = []
chunks = _chunked(texts, config.batch_size)
for chunk_idx, chunk in enumerate(chunks):
result = _call_with_retry(client, config, chunk, chunk_idx, len(chunks))
all_results.extend(result)
return all_results
def _call_with_retry(
client: OpenAI,
config: EmbeddingConfig,
batch: list[str],
chunk_idx: int,
total_chunks: int,
) -> list[list[float]]:
"""单批嵌入调用,带指数退避重试。"""
last_err: Exception | None = None
for attempt in range(1, config.max_attempts + 1):
try:
resp = client.embeddings.create(model=config.model, input=batch)
vectors = [d.embedding for d in resp.data]
# 维度校验
if vectors and len(vectors[0]) != config.dimension:
logger.warning(
"实际维度 %d 与预期 %d 不一致",
len(vectors[0]), config.dimension,
)
logger.debug(
"Embedding chunk %d/%d 完成(%d 条,attempt %d",
chunk_idx + 1, total_chunks, len(batch), attempt,
)
return vectors
except Exception as e:
last_err = e
logger.warning(
"DashScope embed 失败 chunk %d/%d 尝试 %d/%d: %s: %s",
chunk_idx + 1, total_chunks, attempt, config.max_attempts,
type(e).__name__, e,
)
if attempt < config.max_attempts:
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
time.sleep(wait)
raise EmbeddingError(
f"DashScope embed 放弃({config.max_attempts} 次): {last_err}",
attempts=config.max_attempts,
)