feat: add embedder with local/API dual mode
This commit is contained in:
@@ -0,0 +1,59 @@
|
|||||||
|
"""嵌入模型抽象层."""
|
||||||
|
from sentence_transformers import SentenceTransformer
|
||||||
|
|
||||||
|
from src.core.config import EmbedConfig
|
||||||
|
|
||||||
|
|
||||||
|
class Embedder:
|
||||||
|
"""文本嵌入器, 支持本地模型和 API 两种模式."""
|
||||||
|
|
||||||
|
def __init__(self, config: EmbedConfig):
|
||||||
|
self._config = config
|
||||||
|
if config.mode == "local":
|
||||||
|
self._model = SentenceTransformer(config.local_model)
|
||||||
|
elif config.mode == "api":
|
||||||
|
self._model = None # 延迟初始化, 需要 openai 包
|
||||||
|
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
|
||||||
|
def dimension(self) -> int:
|
||||||
|
"""嵌入向量维度."""
|
||||||
|
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]]:
|
||||||
|
"""对文本列表进行嵌入, 返回向量列表."""
|
||||||
|
if not texts:
|
||||||
|
raise ValueError("文本列表不能为空")
|
||||||
|
if self.mode == "local":
|
||||||
|
embeddings = self._model.encode(texts, normalize_embeddings=True)
|
||||||
|
return embeddings.tolist()
|
||||||
|
else:
|
||||||
|
return self._embed_via_api(texts)
|
||||||
|
|
||||||
|
def _embed_via_api(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
"""通过 OpenAI 兼容 API 嵌入(延迟导入 openai)."""
|
||||||
|
from openai import OpenAI
|
||||||
|
|
||||||
|
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
|
||||||
|
response = client.embeddings.create(
|
||||||
|
model="text-embedding-3-small",
|
||||||
|
input=texts,
|
||||||
|
)
|
||||||
|
return [d.embedding for d in response.data]
|
||||||
|
|
||||||
|
|
||||||
|
def create_embedder(config: EmbedConfig) -> Embedder:
|
||||||
|
"""工厂函数: 根据配置创建嵌入器."""
|
||||||
|
return Embedder(config)
|
||||||
@@ -0,0 +1,51 @@
|
|||||||
|
"""嵌入模型测试."""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.core.config import EmbedConfig
|
||||||
|
from src.core.embedder import Embedder, create_embedder
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbedder:
|
||||||
|
"""Embedder 单元测试.
|
||||||
|
|
||||||
|
注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB),
|
||||||
|
首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。
|
||||||
|
"""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def local_config(self):
|
||||||
|
return EmbedConfig(mode="local")
|
||||||
|
|
||||||
|
def test_create_local_embedder(self, local_config):
|
||||||
|
"""创建本地嵌入器,验证维度正确."""
|
||||||
|
embedder = create_embedder(local_config)
|
||||||
|
assert embedder.dimension > 0
|
||||||
|
assert isinstance(embedder.dimension, int)
|
||||||
|
|
||||||
|
def test_embed_single_text(self, local_config):
|
||||||
|
"""嵌入单条文本返回正确维度向量."""
|
||||||
|
embedder = create_embedder(local_config)
|
||||||
|
result = embedder.embed(["你好世界"])
|
||||||
|
assert len(result) == 1
|
||||||
|
assert len(result[0]) == embedder.dimension
|
||||||
|
assert all(isinstance(v, float) for v in result[0])
|
||||||
|
|
||||||
|
def test_embed_multiple_texts(self, local_config):
|
||||||
|
"""嵌入多条文本返回对应数量的向量."""
|
||||||
|
embedder = create_embedder(local_config)
|
||||||
|
texts = ["第一段文本", "第二段文本", "第三段文本"]
|
||||||
|
result = embedder.embed(texts)
|
||||||
|
assert len(result) == 3
|
||||||
|
for vec in result:
|
||||||
|
assert len(vec) == embedder.dimension
|
||||||
|
|
||||||
|
def test_embed_empty_list_raises(self, local_config):
|
||||||
|
"""空列表应抛出异常."""
|
||||||
|
embedder = create_embedder(local_config)
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
embedder.embed([])
|
||||||
|
|
||||||
|
def test_create_embedder_from_factory_function(self, local_config):
|
||||||
|
"""工厂函数正确创建 Embedder 实例."""
|
||||||
|
embedder = create_embedder(local_config)
|
||||||
|
assert isinstance(embedder, Embedder)
|
||||||
Reference in New Issue
Block a user