embeddings.py 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990
  1. from __future__ import annotations
  2. import hashlib
  3. import math
  4. import re
  5. from typing import Protocol
  6. from app.config import Settings
  7. class EmbeddingProvider(Protocol):
  8. @property
  9. def dimension(self) -> int: ...
  10. def embed_query(self, text: str) -> list[float]: ...
  11. def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
  12. class HashEmbeddingProvider:
  13. """无外部模型依赖的确定性向量,仅用于课程演示和自动化测试。"""
  14. def __init__(self, dimension: int = 256) -> None:
  15. self._dimension = dimension
  16. @property
  17. def dimension(self) -> int:
  18. return self._dimension
  19. @staticmethod
  20. def _features(text: str) -> list[str]:
  21. normalized = re.sub(r"\s+", "", text.lower())
  22. chars = list(normalized)
  23. bigrams = [normalized[index : index + 2] for index in range(len(normalized) - 1)]
  24. words = re.findall(r"[a-z]+\d+|\d+|[a-z]+", normalized)
  25. return chars + bigrams + words
  26. def _embed(self, text: str) -> list[float]:
  27. vector = [0.0] * self._dimension
  28. for feature in self._features(text):
  29. # 稳定哈希保证相同文本在不同进程中生成相同演示向量。
  30. digest = hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest()
  31. raw = int.from_bytes(digest, "big")
  32. index = raw % self._dimension
  33. sign = 1.0 if raw & 1 else -1.0
  34. vector[index] += sign
  35. norm = math.sqrt(sum(value * value for value in vector))
  36. if norm == 0:
  37. return vector
  38. # 单位化后可直接使用 COSINE 距离进行课程检索演示。
  39. return [value / norm for value in vector]
  40. def embed_query(self, text: str) -> list[float]:
  41. return self._embed(text)
  42. def embed_documents(self, texts: list[str]) -> list[list[float]]:
  43. return [self._embed(text) for text in texts]
  44. class SentenceTransformerEmbeddingProvider:
  45. def __init__(self, model_name: str) -> None:
  46. try:
  47. from sentence_transformers import SentenceTransformer
  48. except ImportError as exc:
  49. raise RuntimeError(
  50. "缺少 sentence-transformers,请执行 uv sync --extra embeddings"
  51. ) from exc
  52. self._model = SentenceTransformer(model_name)
  53. self._dimension = self._model.get_sentence_embedding_dimension()
  54. @property
  55. def dimension(self) -> int:
  56. return int(self._dimension)
  57. def embed_query(self, text: str) -> list[float]:
  58. return self._model.encode(text, normalize_embeddings=True).tolist()
  59. def embed_documents(self, texts: list[str]) -> list[list[float]]:
  60. return self._model.encode(texts, normalize_embeddings=True).tolist()
  61. def create_embedding_provider(settings: Settings) -> EmbeddingProvider:
  62. # Provider 工厂隔离模型选择,Milvus Store 只依赖统一向量接口。
  63. if settings.embedding_provider == "hash":
  64. return HashEmbeddingProvider(settings.hash_embedding_dim)
  65. if settings.embedding_provider == "sentence_transformers":
  66. return SentenceTransformerEmbeddingProvider(settings.embedding_model_name)
  67. raise ValueError(f"不支持的 EMBEDDING_PROVIDER: {settings.embedding_provider}")