From ad9e5ce449ae6dea0e7f0accc73b50043ede1fb7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E8=88=AA=E5=AE=87?= <3364451258@qq.com> Date: Sun, 5 Jul 2026 01:35:54 +0800 Subject: [PATCH] refactor: strategy-pattern embedder (Local/API), remove global env vars, add batch embedding --- src/core/embedder.py | 129 ++++++++++++++++++++++++----------------- src/core/ingest.py | 6 +- tests/test_embedder.py | 43 ++++++++++---- 3 files changed, 112 insertions(+), 66 deletions(-) diff --git a/src/core/embedder.py b/src/core/embedder.py index cfa5ba4..437a1f8 100644 --- a/src/core/embedder.py +++ b/src/core/embedder.py @@ -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 - -# 解决 Windows SSL 证书问题: 强制使用本地缓存模型, 避免联网验证 -# 模型首次下载需要先临时取消这些环境变量 (或手动运行一次下载脚本) -os.environ.setdefault("HF_HUB_OFFLINE", "1") -os.environ.setdefault("TRANSFORMERS_OFFLINE", "1") - -from sentence_transformers import SentenceTransformer +import logging +from typing import Protocol 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): + from sentence_transformers import SentenceTransformer 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: - 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) - # 恢复离线设置, 后续加载走缓存 - os.environ["HF_HUB_OFFLINE"] = "1" - os.environ["TRANSFORMERS_OFFLINE"] = "1" - 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 + finally: + if old_endpoint is not None: + os.environ["HF_ENDPOINT"] = old_endpoint + else: + os.environ.pop("HF_ENDPOINT", None) @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 + return self._model.get_embedding_dimension() 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) + embeddings = self._model.encode(texts, normalize_embeddings=True) + return embeddings.tolist() - 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 - client = OpenAI(base_url=self._api_base, api_key=self._api_key) response = client.embeddings.create( - model="text-embedding-3-small", - input=texts, + model="text-embedding-3-small", input=texts, ) return [d.embedding for d in response.data] 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 diff --git a/src/core/ingest.py b/src/core/ingest.py index fb56393..b7a14b7 100644 --- a/src/core/ingest.py +++ b/src/core/ingest.py @@ -4,7 +4,7 @@ from dataclasses import dataclass from pathlib import Path from src.core.db import VectorDB -from src.core.embedder import Embedder +from src.core.embedder import Embedder, batch_embed class MarkdownSplitter: @@ -190,9 +190,9 @@ class DocumentIngestor: if not chunks: return 0 - # 嵌入 + # 分批嵌入 (避免大文档 OOM) 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))] diff --git a/tests/test_embedder.py b/tests/test_embedder.py index 1fc6152..da9ecbf 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -2,15 +2,11 @@ import pytest 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: - """Embedder 单元测试. - - 注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB), - 首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。 - """ +class TestLocalEmbedder: + """本地嵌入器测试.""" @pytest.fixture def local_config(self): @@ -21,6 +17,7 @@ class TestEmbedder: embedder = create_embedder(local_config) assert embedder.dimension > 0 assert isinstance(embedder.dimension, int) + assert isinstance(embedder, LocalEmbedder) def test_embed_single_text(self, local_config): """嵌入单条文本返回正确维度向量.""" @@ -45,7 +42,31 @@ class TestEmbedder: 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) + +class TestAPIEmbedder: + """API 嵌入器测试.""" + + 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