- 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 验证
201 lines
6.0 KiB
Python
201 lines
6.0 KiB
Python
"""DashScope Embedding 客户端。
|
||
|
||
通过 OpenAI 兼容接口调用阿里百炼 text-embedding-v3:
|
||
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
||
model: text-embedding-v3(1024 维)
|
||
|
||
配置来源:
|
||
- .env → DASHSCOPE_API_KEY / QWEN_BASE_URL(OpenAI 兼容端点)
|
||
- 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_URL(OpenAI 兼容),
|
||
# 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,
|
||
)
|