Initial commit
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
"""Embedding 向量化模块 (M5)。
|
||||
|
||||
公共 API:
|
||||
- EmbeddingResult / EmbeddingError / EmbeddingProviderType
|
||||
- EmbeddingProvider / AsyncEmbeddingProvider (ABC)
|
||||
- compose_text (统一文本组装策略)
|
||||
- DashScopeEmbeddingProvider / DashScopeAsyncEmbeddingProvider
|
||||
- LocalBGEEmbeddingProvider / LocalBGEAsyncEmbeddingProvider (可选)
|
||||
- resolve_provider_type / make_sync_provider / make_async_provider
|
||||
"""
|
||||
|
||||
from .base import (
|
||||
MAX_TEXT_CHARS,
|
||||
AsyncEmbeddingProvider,
|
||||
EmbeddingProvider,
|
||||
compose_text,
|
||||
)
|
||||
from .factory import (
|
||||
make_async_provider,
|
||||
make_sync_provider,
|
||||
resolve_provider_type,
|
||||
)
|
||||
from .models import (
|
||||
EmbeddingError,
|
||||
EmbeddingProviderType,
|
||||
EmbeddingResult,
|
||||
)
|
||||
from .remote import (
|
||||
DASHSCOPE_BATCH_LIMIT,
|
||||
DASHSCOPE_DEFAULT_DIM,
|
||||
DASHSCOPE_DEFAULT_MODEL,
|
||||
DashScopeAsyncEmbeddingProvider,
|
||||
DashScopeEmbeddingProvider,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DASHSCOPE_BATCH_LIMIT",
|
||||
"DASHSCOPE_DEFAULT_DIM",
|
||||
"DASHSCOPE_DEFAULT_MODEL",
|
||||
"MAX_TEXT_CHARS",
|
||||
"AsyncEmbeddingProvider",
|
||||
"DashScopeAsyncEmbeddingProvider",
|
||||
"DashScopeEmbeddingProvider",
|
||||
"EmbeddingError",
|
||||
"EmbeddingProvider",
|
||||
"EmbeddingProviderType",
|
||||
"EmbeddingResult",
|
||||
"compose_text",
|
||||
"make_async_provider",
|
||||
"make_sync_provider",
|
||||
"resolve_provider_type",
|
||||
]
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Embedding provider 抽象接口与文本组装工具。
|
||||
|
||||
文本组装策略 (compose_text):
|
||||
优先组装 ExtractedEvent 时,把"语义浓缩"信息前置:
|
||||
title | sentiment+importance+event_type | summary | content[截断]
|
||||
退化为 Article 时:
|
||||
title | content[截断]
|
||||
超长截断保护:默认 4000 字符(BGE-M3 max_seq=8192,远程也保守取值)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from extractor import Article
|
||||
|
||||
# 拼接后送入 embedder 的字符数上限
|
||||
MAX_TEXT_CHARS = 4000
|
||||
|
||||
|
||||
def _from_event_dict(event_obj: dict[str, Any]) -> tuple[Article, str, str | None]:
|
||||
"""从 ExtractedEvent JSON dict 中,取出 Article 元数据 + 加权 head 段。
|
||||
|
||||
返回 (article, head, summary):
|
||||
head 为"事件标签摘要",会作为前缀拼到嵌入文本前;
|
||||
summary 为 event.summary。
|
||||
"""
|
||||
title = event_obj.get("title") or ""
|
||||
url = event_obj.get("url") or ""
|
||||
url_hash = event_obj.get("url_hash") or ""
|
||||
source_id = event_obj.get("source_id") or ""
|
||||
publish_time_raw = event_obj.get("publish_time")
|
||||
publish_time = (
|
||||
datetime.fromisoformat(publish_time_raw)
|
||||
if isinstance(publish_time_raw, str) and publish_time_raw
|
||||
else None
|
||||
)
|
||||
|
||||
ev = event_obj.get("event") or {}
|
||||
sentiment = ev.get("sentiment") or "neutral"
|
||||
importance = ev.get("importance") or 0
|
||||
event_type = ev.get("event_type") or "其他"
|
||||
stock_codes = ev.get("stock_codes") or []
|
||||
company_names = ev.get("company_names") or []
|
||||
industries = ev.get("industries") or []
|
||||
summary = ev.get("summary") or ""
|
||||
|
||||
head_parts = [
|
||||
f"sentiment={sentiment}",
|
||||
f"importance={importance}",
|
||||
f"event_type={event_type}",
|
||||
]
|
||||
if company_names:
|
||||
head_parts.append("公司=" + ",".join(company_names[:5]))
|
||||
if industries:
|
||||
head_parts.append("行业=" + ",".join(industries[:5]))
|
||||
if stock_codes:
|
||||
head_parts.append("代码=" + ",".join(stock_codes[:5]))
|
||||
head = "[" + " ".join(head_parts) + "]"
|
||||
|
||||
# 用 ExtractedEvent 中存在的字段构造一个最小 Article 让下游兼容
|
||||
article = Article(
|
||||
source_id=source_id,
|
||||
url=url,
|
||||
url_hash=url_hash,
|
||||
title=title,
|
||||
content=ev.get("summary") or title, # 占位,真正的正文从原 article 文件读
|
||||
publish_time=publish_time,
|
||||
word_count=0,
|
||||
)
|
||||
return article, head, summary
|
||||
|
||||
|
||||
def compose_text(
|
||||
article: Article,
|
||||
*,
|
||||
head: str | None = None,
|
||||
summary: str | None = None,
|
||||
max_chars: int = MAX_TEXT_CHARS,
|
||||
) -> str:
|
||||
"""把 Article 组装成单段嵌入文本。
|
||||
|
||||
参数:
|
||||
article: 输入文章(用于 title + content)
|
||||
head: 可选事件标签摘要(由 ExtractedEvent 提取),前置可提高检索信号
|
||||
summary: 可选 LLM 生成的一句话摘要,前置 head 之后
|
||||
max_chars: 整段最大字符数,超出截断 content
|
||||
"""
|
||||
parts: list[str] = [f"标题:{article.title}"]
|
||||
if head:
|
||||
parts.append(head)
|
||||
if summary:
|
||||
parts.append(f"摘要:{summary}")
|
||||
body = article.content or ""
|
||||
parts.append("正文:" + body)
|
||||
text = "\n".join(parts)
|
||||
if len(text) > max_chars:
|
||||
text = text[:max_chars]
|
||||
return text
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Provider 抽象接口
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
||||
class EmbeddingProvider(ABC):
|
||||
"""同步嵌入 provider 抽象。"""
|
||||
|
||||
name: str
|
||||
model: str
|
||||
dim: int
|
||||
|
||||
@abstractmethod
|
||||
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
"""批量嵌入,返回与输入等长的向量列表。"""
|
||||
|
||||
def embed_one(self, text: str) -> list[float]:
|
||||
"""单条嵌入,默认走 batch=1。"""
|
||||
return self.embed_batch([text])[0]
|
||||
|
||||
def close(self) -> None: # noqa: B027 - 默认空实现,子类按需覆盖
|
||||
"""释放资源(如 HTTP client / 模型)。"""
|
||||
|
||||
def __enter__(self) -> EmbeddingProvider:
|
||||
return self
|
||||
|
||||
def __exit__(self, *_: object) -> None:
|
||||
self.close()
|
||||
|
||||
|
||||
class AsyncEmbeddingProvider(ABC):
|
||||
"""异步嵌入 provider(用于批处理高并发)。"""
|
||||
|
||||
name: str
|
||||
model: str
|
||||
dim: int
|
||||
|
||||
@abstractmethod
|
||||
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
...
|
||||
|
||||
async def embed_one(self, text: str) -> list[float]:
|
||||
return (await self.embed_batch([text]))[0]
|
||||
|
||||
async def close(self) -> None: # noqa: B027 - 默认空实现
|
||||
...
|
||||
|
||||
async def __aenter__(self) -> AsyncEmbeddingProvider:
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_: object) -> None:
|
||||
await self.close()
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Embedding provider 工厂:根据环境变量构造合适后端。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from .base import AsyncEmbeddingProvider, EmbeddingProvider
|
||||
from .models import EmbeddingError, EmbeddingProviderType
|
||||
from .remote import (
|
||||
DashScopeAsyncEmbeddingProvider,
|
||||
DashScopeEmbeddingProvider,
|
||||
)
|
||||
|
||||
|
||||
def _read_env(key: str, default: str | None = None) -> str | None:
|
||||
val = os.environ.get(key)
|
||||
if val is None or val.strip() == "":
|
||||
return default
|
||||
return val.strip()
|
||||
|
||||
|
||||
def resolve_provider_type(provider: str | None = None) -> EmbeddingProviderType:
|
||||
"""根据 provider 参数 / env 解析出 EmbeddingProviderType。
|
||||
|
||||
映射:
|
||||
dashscope / qwen / remote -> DASHSCOPE
|
||||
local / local-bge / bge / bge-m3 -> LOCAL_BGE
|
||||
默认 dashscope。
|
||||
"""
|
||||
p = (provider or _read_env("EMBEDDING_PROVIDER", "dashscope") or "dashscope").lower()
|
||||
if p in ("dashscope", "qwen", "remote"):
|
||||
return EmbeddingProviderType.DASHSCOPE
|
||||
if p in ("local", "local-bge", "bge", "bge-m3"):
|
||||
return EmbeddingProviderType.LOCAL_BGE
|
||||
raise EmbeddingError(f"未知 embedding provider: {provider!r}")
|
||||
|
||||
|
||||
def make_sync_provider(
|
||||
provider: str | None = None,
|
||||
**kwargs: object,
|
||||
) -> EmbeddingProvider:
|
||||
"""构造同步 provider。"""
|
||||
pt = resolve_provider_type(provider)
|
||||
if pt == EmbeddingProviderType.DASHSCOPE:
|
||||
return DashScopeEmbeddingProvider(**kwargs) # type: ignore[arg-type]
|
||||
# 本地后端
|
||||
from .local import LocalBGEEmbeddingProvider
|
||||
return LocalBGEEmbeddingProvider(**kwargs) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def make_async_provider(
|
||||
provider: str | None = None,
|
||||
**kwargs: object,
|
||||
) -> AsyncEmbeddingProvider:
|
||||
"""构造异步 provider。"""
|
||||
pt = resolve_provider_type(provider)
|
||||
if pt == EmbeddingProviderType.DASHSCOPE:
|
||||
return DashScopeAsyncEmbeddingProvider(**kwargs) # type: ignore[arg-type]
|
||||
from .local import LocalBGEAsyncEmbeddingProvider
|
||||
return LocalBGEAsyncEmbeddingProvider(**kwargs) # type: ignore[arg-type]
|
||||
@@ -0,0 +1,109 @@
|
||||
"""本地 BGE-M3 嵌入实现(可选)。
|
||||
|
||||
需要额外安装本地依赖:
|
||||
uv sync --extra local-embedding
|
||||
|
||||
模型 BAAI/bge-m3 首次加载约 2.3 GB(从 HuggingFace 自动下载)。
|
||||
默认 1024 维,与 DashScope 兼容。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from .base import AsyncEmbeddingProvider, EmbeddingProvider
|
||||
from .models import EmbeddingError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sentence_transformers import SentenceTransformer
|
||||
|
||||
LOCAL_DEFAULT_MODEL = "BAAI/bge-m3"
|
||||
LOCAL_DEFAULT_DIM = 1024
|
||||
|
||||
|
||||
def _read_env(key: str, default: str | None = None) -> str | None:
|
||||
val = os.environ.get(key)
|
||||
if val is None or val.strip() == "":
|
||||
return default
|
||||
return val.strip()
|
||||
|
||||
|
||||
def _try_import_st() -> type[SentenceTransformer]:
|
||||
try:
|
||||
from sentence_transformers import SentenceTransformer # noqa: F811
|
||||
except ImportError as e:
|
||||
raise EmbeddingError(
|
||||
"本地 BGE-M3 后端需要 sentence-transformers,请执行: "
|
||||
"uv sync --extra local-embedding"
|
||||
) from e
|
||||
return SentenceTransformer
|
||||
|
||||
|
||||
class LocalBGEEmbeddingProvider(EmbeddingProvider):
|
||||
"""本地 BGE-M3 同步实现。"""
|
||||
|
||||
name = "local-bge"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
device: str | None = None,
|
||||
normalize: bool = True,
|
||||
) -> None:
|
||||
st_cls = _try_import_st()
|
||||
# 模型优先级:CLI/参数 > LOCAL_EMBEDDING_MODEL > 默认值
|
||||
# 不读全局 EMBEDDING_MODEL,避免与 DashScope 冲突
|
||||
self.model = (
|
||||
model
|
||||
or _read_env("LOCAL_EMBEDDING_MODEL", LOCAL_DEFAULT_MODEL)
|
||||
or LOCAL_DEFAULT_MODEL
|
||||
)
|
||||
self.dim = LOCAL_DEFAULT_DIM
|
||||
self._device = device # None -> 让 ST 自选 cpu/cuda
|
||||
self._normalize = normalize
|
||||
logger.info("加载本地嵌入模型 {} (device={}),首次会下载...", self.model, device or "auto")
|
||||
self._st = st_cls(self.model, device=device)
|
||||
# 实际维度自检
|
||||
actual_dim = self._st.get_sentence_embedding_dimension()
|
||||
if actual_dim and actual_dim != self.dim:
|
||||
logger.warning(
|
||||
"BGE 模型实际维度 {} 与默认 {} 不一致,以实际为准", actual_dim, self.dim
|
||||
)
|
||||
self.dim = actual_dim
|
||||
|
||||
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
try:
|
||||
arr = self._st.encode(
|
||||
texts,
|
||||
normalize_embeddings=self._normalize,
|
||||
convert_to_numpy=True,
|
||||
show_progress_bar=False,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise EmbeddingError(f"BGE-M3 推理失败: {type(e).__name__}: {e}") from e
|
||||
return [v.tolist() for v in arr]
|
||||
|
||||
|
||||
class LocalBGEAsyncEmbeddingProvider(AsyncEmbeddingProvider):
|
||||
"""本地 BGE-M3 异步包装。
|
||||
|
||||
sentence-transformers 本身是同步的,这里用 asyncio.to_thread 包一层,
|
||||
主要为了让批处理脚本能用统一的 async 接口。
|
||||
"""
|
||||
|
||||
name = "local-bge"
|
||||
|
||||
def __init__(self, **kwargs: object) -> None:
|
||||
self._sync = LocalBGEEmbeddingProvider(**kwargs) # type: ignore[arg-type]
|
||||
self.model = self._sync.model
|
||||
self.dim = self._sync.dim
|
||||
|
||||
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
return await asyncio.to_thread(self._sync.embed_batch, texts)
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Embedding 模块的数据模型 (M5)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from enum import StrEnum
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class EmbeddingProviderType(StrEnum):
|
||||
"""支持的 embedding provider 标识。"""
|
||||
|
||||
DASHSCOPE = "dashscope" # 远程 Qwen / 百炼
|
||||
LOCAL_BGE = "local-bge" # 本地 BGE-M3
|
||||
|
||||
|
||||
class EmbeddingResult(BaseModel):
|
||||
"""单篇文章的嵌入结果(落盘格式)。"""
|
||||
|
||||
url_hash: str = Field(..., description="主键,与 Article.url_hash 一致")
|
||||
source_id: str = Field(..., description="来源源 id")
|
||||
title: str = Field(..., description="原文标题(便于人工检索)")
|
||||
text: str = Field(
|
||||
..., description="实际送入 embedder 的文本(已截断/拼接)"
|
||||
)
|
||||
vector: list[float] = Field(..., description="嵌入向量")
|
||||
dim: int = Field(..., gt=0, description="向量维度")
|
||||
|
||||
provider: str = Field(..., description="dashscope / local-bge")
|
||||
model: str = Field(..., description="嵌入模型名")
|
||||
embedded_at: datetime = Field(default_factory=datetime.now)
|
||||
char_count: int = Field(default=0, ge=0, description="text 字符数,便于排查")
|
||||
publish_time: datetime | None = None
|
||||
|
||||
def short_summary(self) -> str:
|
||||
return (
|
||||
f"[{self.source_id}] {self.title[:30]} "
|
||||
f"dim={self.dim} provider={self.provider}"
|
||||
)
|
||||
|
||||
|
||||
class EmbeddingError(Exception):
|
||||
"""嵌入调用失败。"""
|
||||
|
||||
def __init__(self, reason: str, *, attempts: int = 0) -> None:
|
||||
super().__init__(reason)
|
||||
self.reason = reason
|
||||
self.attempts = attempts
|
||||
@@ -0,0 +1,183 @@
|
||||
"""DashScope / Qwen 远程嵌入实现。
|
||||
|
||||
通过 OpenAI 兼容接口调用阿里百炼的 text-embedding-v3:
|
||||
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||
model: text-embedding-v3 (1024 维)
|
||||
限制: 单次请求 input ≤ 25 条
|
||||
|
||||
环境变量:
|
||||
DASHSCOPE_API_KEY
|
||||
QWEN_BASE_URL (默认百炼兼容路径)
|
||||
EMBEDDING_MODEL (默认 text-embedding-v3)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from loguru import logger
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
from .base import AsyncEmbeddingProvider, EmbeddingProvider
|
||||
from .models import EmbeddingError
|
||||
|
||||
DASHSCOPE_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
DASHSCOPE_DEFAULT_MODEL = "text-embedding-v3"
|
||||
DASHSCOPE_DEFAULT_DIM = 1024
|
||||
DASHSCOPE_BATCH_LIMIT = 10 # 百炼实测单批上限(2026-06,文档曾标 25 但 API 报 400)
|
||||
|
||||
# 重试策略
|
||||
DEFAULT_MAX_ATTEMPTS = 3
|
||||
RETRY_BASE_WAIT_SEC = 1.0
|
||||
RETRY_MAX_WAIT_SEC = 8.0
|
||||
|
||||
|
||||
def _read_env(key: str, default: str | None = None) -> str | None:
|
||||
val = os.environ.get(key)
|
||||
if val is None or val.strip() == "":
|
||||
return default
|
||||
return val.strip()
|
||||
|
||||
|
||||
def _resolve_config() -> tuple[str, str, str]:
|
||||
"""读取 API key / base_url / model,返回 (api_key, base_url, model)。
|
||||
|
||||
模型名优先级:DASHSCOPE_EMBEDDING_MODEL > 默认值。
|
||||
不再读全局 EMBEDDING_MODEL,避免与 LOCAL provider 冲突。
|
||||
"""
|
||||
api_key = _read_env("DASHSCOPE_EMBEDDING_API_KEY") or _read_env("DASHSCOPE_API_KEY") or ""
|
||||
if not api_key:
|
||||
raise EmbeddingError("DASHSCOPE_EMBEDDING_API_KEY 或 DASHSCOPE_API_KEY 未配置")
|
||||
# DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认
|
||||
base_url = (
|
||||
_read_env("DASHSCOPE_EMBEDDING_BASE_URL")
|
||||
or _read_env("QWEN_BASE_URL")
|
||||
or DASHSCOPE_DEFAULT_BASE
|
||||
)
|
||||
model = (
|
||||
_read_env("DASHSCOPE_EMBEDDING_MODEL", DASHSCOPE_DEFAULT_MODEL)
|
||||
or DASHSCOPE_DEFAULT_MODEL
|
||||
)
|
||||
return api_key, base_url, model
|
||||
|
||||
|
||||
def _chunked(items: list[str], size: int) -> list[list[str]]:
|
||||
"""把列表按 size 分块。"""
|
||||
return [items[i : i + size] for i in range(0, len(items), size)]
|
||||
|
||||
|
||||
class DashScopeEmbeddingProvider(EmbeddingProvider):
|
||||
"""同步实现,主要用于测试/单条调用。"""
|
||||
|
||||
name = "dashscope"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
timeout_sec: float = 60.0,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
) -> None:
|
||||
env_key, env_base, env_model = _resolve_config()
|
||||
self.model = model or env_model
|
||||
self.dim = DASHSCOPE_DEFAULT_DIM
|
||||
self.max_attempts = max_attempts
|
||||
self._client = OpenAI(
|
||||
api_key=api_key or env_key,
|
||||
base_url=base_url or env_base,
|
||||
timeout=timeout_sec,
|
||||
)
|
||||
|
||||
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
results: list[list[float]] = []
|
||||
for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT):
|
||||
results.extend(self._call_with_retry(chunk))
|
||||
return results
|
||||
|
||||
def _call_with_retry(self, batch: list[str]) -> list[list[float]]:
|
||||
import time
|
||||
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(1, self.max_attempts + 1):
|
||||
try:
|
||||
resp = self._client.embeddings.create(model=self.model, input=batch)
|
||||
return [d.embedding for d in resp.data]
|
||||
except Exception as e: # noqa: BLE001
|
||||
last_err = e
|
||||
logger.warning(
|
||||
"DashScope embed 失败 尝试 {}/{}: {}: {}",
|
||||
attempt, self.max_attempts, type(e).__name__, e,
|
||||
)
|
||||
if attempt < self.max_attempts:
|
||||
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
|
||||
time.sleep(wait)
|
||||
raise EmbeddingError(
|
||||
f"DashScope embed 放弃 {self.max_attempts} 次: {last_err}",
|
||||
attempts=self.max_attempts,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self._client.close()
|
||||
|
||||
|
||||
class DashScopeAsyncEmbeddingProvider(AsyncEmbeddingProvider):
|
||||
"""异步实现,用于批处理。"""
|
||||
|
||||
name = "dashscope"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
timeout_sec: float = 60.0,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
) -> None:
|
||||
env_key, env_base, env_model = _resolve_config()
|
||||
self.model = model or env_model
|
||||
self.dim = DASHSCOPE_DEFAULT_DIM
|
||||
self.max_attempts = max_attempts
|
||||
self._client = AsyncOpenAI(
|
||||
api_key=api_key or env_key,
|
||||
base_url=base_url or env_base,
|
||||
timeout=timeout_sec,
|
||||
)
|
||||
|
||||
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
results: list[list[float]] = []
|
||||
for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT):
|
||||
results.extend(await self._call_with_retry(chunk))
|
||||
return results
|
||||
|
||||
async def _call_with_retry(self, batch: list[str]) -> list[list[float]]:
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(1, self.max_attempts + 1):
|
||||
try:
|
||||
resp = await self._client.embeddings.create(
|
||||
model=self.model, input=batch
|
||||
)
|
||||
return [d.embedding for d in resp.data]
|
||||
except Exception as e: # noqa: BLE001
|
||||
last_err = e
|
||||
logger.warning(
|
||||
"DashScope embed 失败 尝试 {}/{}: {}: {}",
|
||||
attempt, self.max_attempts, type(e).__name__, e,
|
||||
)
|
||||
if attempt < self.max_attempts:
|
||||
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
|
||||
await asyncio.sleep(wait)
|
||||
raise EmbeddingError(
|
||||
f"DashScope embed 放弃 {self.max_attempts} 次: {last_err}",
|
||||
attempts=self.max_attempts,
|
||||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
await self._client.close()
|
||||
Reference in New Issue
Block a user