Initial commit

This commit is contained in:
2026-07-18 15:51:01 +08:00
commit f2c80c5a9c
799 changed files with 133475 additions and 0 deletions
+52
View File
@@ -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",
]
+154
View File
@@ -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()
+60
View File
@@ -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]
+109
View File
@@ -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)
+49
View File
@@ -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
+183
View File
@@ -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()