embeddings.py 1.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940
  1. from typing import List
  2. from langchain_core.embeddings import Embeddings
  3. from openai import OpenAI
  4. from app.config import settings
  5. class SiliconFlowEmbeddings(Embeddings):
  6. """OpenAI 兼容 embeddings 直连客户端(绕开 langchain tiktoken token 化)
  7. langchain 的 OpenAIEmbeddings 默认 tiktoken_enabled=True,会把文本编码成
  8. 整数 token 数组发送给 API(OpenAI 官方接受,但 siliconflow 等第三方网关
  9. 只接受字符串 input,导致 400 code:20015 parameter invalid)。
  10. 本类直接用 openai SDK 发送字符串数组,兼容 siliconflow / siliconflow 等网关。
  11. """
  12. def __init__(self, model: str, api_key: str, base_url: str):
  13. self.model = model
  14. self.client = OpenAI(api_key=api_key, base_url=base_url)
  15. def embed_documents(self, texts: List[str]) -> List[List[float]]:
  16. resp = self.client.embeddings.create(model=self.model, input=texts)
  17. return [d.embedding for d in resp.data]
  18. def embed_query(self, text: str) -> List[float]:
  19. resp = self.client.embeddings.create(model=self.model, input=[text])
  20. return resp.data[0].embedding
  21. _embeddings = None
  22. def get_embeddings():
  23. global _embeddings
  24. if _embeddings is None:
  25. _embeddings = SiliconFlowEmbeddings(
  26. model=settings.embedding_model,
  27. api_key=settings.effective_embedding_api_key,
  28. base_url=settings.effective_embedding_base_url,
  29. )
  30. return _embeddings