refactor: strategy-pattern embedder (Local/API), remove global env vars, add batch embedding
This commit is contained in:
+77
-52
@@ -1,79 +1,104 @@
|
|||||||
"""嵌入模型抽象层."""
|
"""嵌入模型抽象层 — 策略模式.
|
||||||
|
|
||||||
|
使用方式:
|
||||||
|
from src.core.config import EmbedConfig
|
||||||
|
from src.core.embedder import create_embedder, batch_embed
|
||||||
|
|
||||||
|
embedder = create_embedder(EmbedConfig(mode="local"))
|
||||||
|
vectors = embedder.embed(["文本1", "文本2"])
|
||||||
|
# 或分批嵌入避免 OOM:
|
||||||
|
vectors = batch_embed(embedder, long_text_list)
|
||||||
|
"""
|
||||||
import os
|
import os
|
||||||
|
import logging
|
||||||
# 解决 Windows SSL 证书问题: 强制使用本地缓存模型, 避免联网验证
|
from typing import Protocol
|
||||||
# 模型首次下载需要先临时取消这些环境变量 (或手动运行一次下载脚本)
|
|
||||||
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
|
||||||
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
|
|
||||||
|
|
||||||
from sentence_transformers import SentenceTransformer
|
|
||||||
|
|
||||||
from src.core.config import EmbedConfig
|
from src.core.config import EmbedConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
class Embedder:
|
# 模型下载镜像 (仅首次下载时使用)
|
||||||
"""文本嵌入器, 支持本地模型和 API 两种模式."""
|
_HF_MIRROR = "https://hf-mirror.com"
|
||||||
|
|
||||||
|
|
||||||
|
class Embedder(Protocol):
|
||||||
|
"""嵌入器接口."""
|
||||||
|
@property
|
||||||
|
def dimension(self) -> int: ...
|
||||||
|
def embed(self, texts: list[str]) -> list[list[float]]: ...
|
||||||
|
|
||||||
|
|
||||||
|
class LocalEmbedder:
|
||||||
|
"""本地 sentence-transformers 模型嵌入器."""
|
||||||
|
|
||||||
def __init__(self, config: EmbedConfig):
|
def __init__(self, config: EmbedConfig):
|
||||||
|
from sentence_transformers import SentenceTransformer
|
||||||
self._config = config
|
self._config = config
|
||||||
if config.mode == "local":
|
try:
|
||||||
# 优先从本地缓存加载, 若模型未下载则通过镜像下载
|
self._model = SentenceTransformer(
|
||||||
|
config.local_model, local_files_only=True
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.info("模型未缓存, 通过镜像下载 %s", config.local_model)
|
||||||
|
old_endpoint = os.environ.get("HF_ENDPOINT")
|
||||||
|
os.environ["HF_ENDPOINT"] = _HF_MIRROR
|
||||||
try:
|
try:
|
||||||
self._model = SentenceTransformer(
|
|
||||||
config.local_model, local_files_only=True
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
# 模型未缓存: 临时开启网络, 使用国内镜像
|
|
||||||
os.environ.pop("HF_HUB_OFFLINE", None)
|
|
||||||
os.environ.pop("TRANSFORMERS_OFFLINE", None)
|
|
||||||
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
|
|
||||||
self._model = SentenceTransformer(config.local_model)
|
self._model = SentenceTransformer(config.local_model)
|
||||||
# 恢复离线设置, 后续加载走缓存
|
finally:
|
||||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
if old_endpoint is not None:
|
||||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
os.environ["HF_ENDPOINT"] = old_endpoint
|
||||||
elif config.mode == "api":
|
else:
|
||||||
self._model = None # 延迟初始化, 需要 openai 包
|
os.environ.pop("HF_ENDPOINT", None)
|
||||||
self._api_base = config.api_base
|
|
||||||
self._api_key = config.api_key
|
|
||||||
else:
|
|
||||||
raise ValueError(f"不支持的嵌入模式: {config.mode}")
|
|
||||||
|
|
||||||
@property
|
|
||||||
def mode(self) -> str:
|
|
||||||
"""当前嵌入模式."""
|
|
||||||
return self._config.mode
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def dimension(self) -> int:
|
def dimension(self) -> int:
|
||||||
"""嵌入向量维度."""
|
return self._model.get_embedding_dimension()
|
||||||
if self.mode == "local":
|
|
||||||
return self._model.get_embedding_dimension()
|
|
||||||
else:
|
|
||||||
# 默认 OpenAI text-embedding-ada-002 / text-embedding-3-small 维度
|
|
||||||
return 1536
|
|
||||||
|
|
||||||
def embed(self, texts: list[str]) -> list[list[float]]:
|
def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
"""对文本列表进行嵌入, 返回向量列表."""
|
|
||||||
if not texts:
|
if not texts:
|
||||||
raise ValueError("文本列表不能为空")
|
raise ValueError("文本列表不能为空")
|
||||||
if self.mode == "local":
|
embeddings = self._model.encode(texts, normalize_embeddings=True)
|
||||||
embeddings = self._model.encode(texts, normalize_embeddings=True)
|
return embeddings.tolist()
|
||||||
return embeddings.tolist()
|
|
||||||
else:
|
|
||||||
return self._embed_via_api(texts)
|
|
||||||
|
|
||||||
def _embed_via_api(self, texts: list[str]) -> list[list[float]]:
|
|
||||||
"""通过 OpenAI 兼容 API 嵌入(延迟导入 openai)."""
|
class APIEmbedder:
|
||||||
|
"""OpenAI 兼容 API 嵌入器."""
|
||||||
|
|
||||||
|
def __init__(self, config: EmbedConfig):
|
||||||
|
self._api_base = config.api_base
|
||||||
|
self._api_key = config.api_key
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dimension(self) -> int:
|
||||||
|
return 1536
|
||||||
|
|
||||||
|
def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
if not texts:
|
||||||
|
raise ValueError("文本列表不能为空")
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
||||||
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
|
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
|
||||||
response = client.embeddings.create(
|
response = client.embeddings.create(
|
||||||
model="text-embedding-3-small",
|
model="text-embedding-3-small", input=texts,
|
||||||
input=texts,
|
|
||||||
)
|
)
|
||||||
return [d.embedding for d in response.data]
|
return [d.embedding for d in response.data]
|
||||||
|
|
||||||
|
|
||||||
def create_embedder(config: EmbedConfig) -> Embedder:
|
def create_embedder(config: EmbedConfig) -> Embedder:
|
||||||
"""工厂函数: 根据配置创建嵌入器."""
|
"""工厂函数: 根据配置创建嵌入器."""
|
||||||
return Embedder(config)
|
if config.mode == "local":
|
||||||
|
return LocalEmbedder(config)
|
||||||
|
elif config.mode == "api":
|
||||||
|
return APIEmbedder(config)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"不支持的嵌入模式: {config.mode}")
|
||||||
|
|
||||||
|
|
||||||
|
def batch_embed(
|
||||||
|
embedder: Embedder, texts: list[str], batch_size: int = 32
|
||||||
|
) -> list[list[float]]:
|
||||||
|
"""分批嵌入, 避免一次性传入过多文本导致 OOM."""
|
||||||
|
all_embeddings = []
|
||||||
|
for i in range(0, len(texts), batch_size):
|
||||||
|
batch = texts[i:i + batch_size]
|
||||||
|
all_embeddings.extend(embedder.embed(batch))
|
||||||
|
return all_embeddings
|
||||||
|
|||||||
+3
-3
@@ -4,7 +4,7 @@ from dataclasses import dataclass
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from src.core.db import VectorDB
|
from src.core.db import VectorDB
|
||||||
from src.core.embedder import Embedder
|
from src.core.embedder import Embedder, batch_embed
|
||||||
|
|
||||||
|
|
||||||
class MarkdownSplitter:
|
class MarkdownSplitter:
|
||||||
@@ -190,9 +190,9 @@ class DocumentIngestor:
|
|||||||
if not chunks:
|
if not chunks:
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
# 嵌入
|
# 分批嵌入 (避免大文档 OOM)
|
||||||
texts = [c["content"] for c in chunks]
|
texts = [c["content"] for c in chunks]
|
||||||
embeddings = self.embedder.embed(texts)
|
embeddings = batch_embed(self.embedder, texts)
|
||||||
|
|
||||||
# 入库
|
# 入库
|
||||||
ids = [f"{file_name}_{i}" for i in range(len(chunks))]
|
ids = [f"{file_name}_{i}" for i in range(len(chunks))]
|
||||||
|
|||||||
+32
-11
@@ -2,15 +2,11 @@
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.core.config import EmbedConfig
|
from src.core.config import EmbedConfig
|
||||||
from src.core.embedder import Embedder, create_embedder
|
from src.core.embedder import LocalEmbedder, APIEmbedder, create_embedder, batch_embed
|
||||||
|
|
||||||
|
|
||||||
class TestEmbedder:
|
class TestLocalEmbedder:
|
||||||
"""Embedder 单元测试.
|
"""本地嵌入器测试."""
|
||||||
|
|
||||||
注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB),
|
|
||||||
首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。
|
|
||||||
"""
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def local_config(self):
|
def local_config(self):
|
||||||
@@ -21,6 +17,7 @@ class TestEmbedder:
|
|||||||
embedder = create_embedder(local_config)
|
embedder = create_embedder(local_config)
|
||||||
assert embedder.dimension > 0
|
assert embedder.dimension > 0
|
||||||
assert isinstance(embedder.dimension, int)
|
assert isinstance(embedder.dimension, int)
|
||||||
|
assert isinstance(embedder, LocalEmbedder)
|
||||||
|
|
||||||
def test_embed_single_text(self, local_config):
|
def test_embed_single_text(self, local_config):
|
||||||
"""嵌入单条文本返回正确维度向量."""
|
"""嵌入单条文本返回正确维度向量."""
|
||||||
@@ -45,7 +42,31 @@ class TestEmbedder:
|
|||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
embedder.embed([])
|
embedder.embed([])
|
||||||
|
|
||||||
def test_create_embedder_from_factory_function(self, local_config):
|
|
||||||
"""工厂函数正确创建 Embedder 实例."""
|
class TestAPIEmbedder:
|
||||||
embedder = create_embedder(local_config)
|
"""API 嵌入器测试."""
|
||||||
assert isinstance(embedder, Embedder)
|
|
||||||
|
def test_api_embedder_init(self):
|
||||||
|
"""API 嵌入器初始化."""
|
||||||
|
cfg = EmbedConfig(mode="api", api_base="https://api.test.com", api_key="sk-test")
|
||||||
|
emb = APIEmbedder(cfg)
|
||||||
|
assert emb.dimension == 1536
|
||||||
|
|
||||||
|
def test_api_embedder_empty_raises(self):
|
||||||
|
"""API 嵌入器空列表抛异常."""
|
||||||
|
cfg = EmbedConfig(mode="api")
|
||||||
|
emb = APIEmbedder(cfg)
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
emb.embed([])
|
||||||
|
|
||||||
|
|
||||||
|
class TestBatchEmbed:
|
||||||
|
"""分批嵌入测试."""
|
||||||
|
|
||||||
|
def test_batch_embed(self):
|
||||||
|
"""分批嵌入返回正确数量."""
|
||||||
|
embedder = create_embedder(EmbedConfig(mode="local"))
|
||||||
|
texts = ["测试文本"] * 70 # 超过 batch_size 32
|
||||||
|
results = batch_embed(embedder, texts, batch_size=32)
|
||||||
|
assert len(results) == 70
|
||||||
|
assert len(results[0]) == embedder.dimension
|
||||||
|
|||||||
Reference in New Issue
Block a user