embeddings.py 2.6 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071
  1. from typing import List
  2. import time
  3. import logging
  4. from langchain_core.embeddings import Embeddings
  5. from openai import OpenAI, RateLimitError
  6. from app.config import settings
  7. logger = logging.getLogger(__name__)
  8. class SiliconFlowEmbeddings(Embeddings):
  9. """OpenAI 兼容 embeddings 直连客户端(绕开 langchain tiktoken token 化)
  10. langchain 的 OpenAIEmbeddings 默认 tiktoken_enabled=True,会把文本编码成
  11. 整数 token 数组发送给 API(OpenAI 官方接受,但 siliconflow 等第三方网关
  12. 只接受字符串 input,导致 400 code:20015 parameter invalid)。
  13. 本类直接用 openai SDK 发送字符串数组,兼容 siliconflow 等网关。
  14. 内置 429 TPM 指数退避重试:SiliconFlow 全量同步易触 TPM 限额,429 是
  15. 临时性状态,退避重试能自动恢复完成,避免 sync 整段中断。
  16. """
  17. _BATCH_SIZE = 32
  18. _THROTTLE_SECONDS = 0.02
  19. _MAX_RETRY = 3
  20. _INITIAL_RETRY_SECONDS = 2.0
  21. def __init__(self, model: str, api_key: str, base_url: str):
  22. self.model = model
  23. self.client = OpenAI(api_key=api_key, base_url=base_url)
  24. def _call_with_retry(self, batch):
  25. wait = self._INITIAL_RETRY_SECONDS
  26. for attempt in range(self._MAX_RETRY + 1):
  27. try:
  28. return self.client.embeddings.create(model=self.model, input=batch)
  29. except RateLimitError as e:
  30. if attempt == self._MAX_RETRY:
  31. raise
  32. logger.warning("embeddings 429 TPM 限制,%.0fs 后重试 (%d/%d)", wait, attempt + 1, self._MAX_RETRY)
  33. time.sleep(wait)
  34. wait = min(wait * 2, 30.0)
  35. raise RuntimeError("unreachable")
  36. def embed_documents(self, texts: List[str]) -> List[List[float]]:
  37. result = []
  38. for i in range(0, len(texts), self._BATCH_SIZE):
  39. batch = texts[i:i + self._BATCH_SIZE]
  40. resp = self._call_with_retry(batch)
  41. result.extend(d.embedding for d in resp.data)
  42. if i + self._BATCH_SIZE < len(texts):
  43. time.sleep(self._THROTTLE_SECONDS)
  44. return result
  45. def embed_query(self, text: str) -> List[float]:
  46. resp = self._call_with_retry([text])
  47. return resp.data[0].embedding
  48. _embeddings = None
  49. def get_embeddings():
  50. global _embeddings
  51. if _embeddings is None:
  52. _embeddings = SiliconFlowEmbeddings(
  53. model=settings.embedding_model,
  54. api_key=settings.effective_embedding_api_key,
  55. base_url=settings.effective_embedding_base_url,
  56. )
  57. return _embeddings