| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990 |
- from __future__ import annotations
- import hashlib
- import math
- import re
- from typing import Protocol
- from app.config import Settings
- class EmbeddingProvider(Protocol):
- @property
- def dimension(self) -> int: ...
- def embed_query(self, text: str) -> list[float]: ...
- def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
- class HashEmbeddingProvider:
- """无外部模型依赖的确定性向量,仅用于课程演示和自动化测试。"""
- def __init__(self, dimension: int = 256) -> None:
- self._dimension = dimension
- @property
- def dimension(self) -> int:
- return self._dimension
- @staticmethod
- def _features(text: str) -> list[str]:
- normalized = re.sub(r"\s+", "", text.lower())
- chars = list(normalized)
- bigrams = [normalized[index : index + 2] for index in range(len(normalized) - 1)]
- words = re.findall(r"[a-z]+\d+|\d+|[a-z]+", normalized)
- return chars + bigrams + words
- def _embed(self, text: str) -> list[float]:
- vector = [0.0] * self._dimension
- for feature in self._features(text):
- # 稳定哈希保证相同文本在不同进程中生成相同演示向量。
- digest = hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest()
- raw = int.from_bytes(digest, "big")
- index = raw % self._dimension
- sign = 1.0 if raw & 1 else -1.0
- vector[index] += sign
- norm = math.sqrt(sum(value * value for value in vector))
- if norm == 0:
- return vector
- # 单位化后可直接使用 COSINE 距离进行课程检索演示。
- return [value / norm for value in vector]
- def embed_query(self, text: str) -> list[float]:
- return self._embed(text)
- def embed_documents(self, texts: list[str]) -> list[list[float]]:
- return [self._embed(text) for text in texts]
- class SentenceTransformerEmbeddingProvider:
- def __init__(self, model_name: str) -> None:
- try:
- from sentence_transformers import SentenceTransformer
- except ImportError as exc:
- raise RuntimeError(
- "缺少 sentence-transformers,请执行 uv sync --extra embeddings"
- ) from exc
- self._model = SentenceTransformer(model_name)
- self._dimension = self._model.get_sentence_embedding_dimension()
- @property
- def dimension(self) -> int:
- return int(self._dimension)
- def embed_query(self, text: str) -> list[float]:
- return self._model.encode(text, normalize_embeddings=True).tolist()
- def embed_documents(self, texts: list[str]) -> list[list[float]]:
- return self._model.encode(texts, normalize_embeddings=True).tolist()
- def create_embedding_provider(settings: Settings) -> EmbeddingProvider:
- # Provider 工厂隔离模型选择,Milvus Store 只依赖统一向量接口。
- if settings.embedding_provider == "hash":
- return HashEmbeddingProvider(settings.hash_embedding_dim)
- if settings.embedding_provider == "sentence_transformers":
- return SentenceTransformerEmbeddingProvider(settings.embedding_model_name)
- raise ValueError(f"不支持的 EMBEDDING_PROVIDER: {settings.embedding_provider}")
|