| 12345678910111213141516171819202122232425262728293031323334353637383940 |
- from typing import List
- from langchain_core.embeddings import Embeddings
- from openai import OpenAI
- from app.config import settings
- class SiliconFlowEmbeddings(Embeddings):
- """OpenAI 兼容 embeddings 直连客户端(绕开 langchain tiktoken token 化)
- langchain 的 OpenAIEmbeddings 默认 tiktoken_enabled=True,会把文本编码成
- 整数 token 数组发送给 API(OpenAI 官方接受,但 siliconflow 等第三方网关
- 只接受字符串 input,导致 400 code:20015 parameter invalid)。
- 本类直接用 openai SDK 发送字符串数组,兼容 siliconflow / siliconflow 等网关。
- """
- def __init__(self, model: str, api_key: str, base_url: str):
- self.model = model
- self.client = OpenAI(api_key=api_key, base_url=base_url)
- def embed_documents(self, texts: List[str]) -> List[List[float]]:
- resp = self.client.embeddings.create(model=self.model, input=texts)
- return [d.embedding for d in resp.data]
- def embed_query(self, text: str) -> List[float]:
- resp = self.client.embeddings.create(model=self.model, input=[text])
- return resp.data[0].embedding
- _embeddings = None
- def get_embeddings():
- global _embeddings
- if _embeddings is None:
- _embeddings = SiliconFlowEmbeddings(
- model=settings.embedding_model,
- api_key=settings.effective_embedding_api_key,
- base_url=settings.effective_embedding_base_url,
- )
- return _embeddings
|