From c312a15e3e67339ec67875c3c6debee2033c1e3c Mon Sep 17 00:00:00 2001 From: martsforever Date: Sun, 7 Sep 2025 00:00:36 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20LLamaIndexEmbeddings=EF=BC=8C=E5=90=8C?= =?UTF-8?q?=E6=97=B6=E6=94=AF=E6=8C=81=E5=90=8C=E6=AD=A5=E5=BC=82=E6=AD=A5?= =?UTF-8?q?=E7=9A=84embedding=E8=B0=83=E7=94=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/utils/CustomEmbeddings.py | 72 ------------------------------- app/utils/LLamaIndexEmbeddings.py | 49 +++++++++++++++++++++ 2 files changed, 49 insertions(+), 72 deletions(-) delete mode 100644 app/utils/CustomEmbeddings.py create mode 100644 app/utils/LLamaIndexEmbeddings.py diff --git a/app/utils/CustomEmbeddings.py b/app/utils/CustomEmbeddings.py deleted file mode 100644 index 98db2f2..0000000 --- a/app/utils/CustomEmbeddings.py +++ /dev/null @@ -1,72 +0,0 @@ -import json - -import requests -from langchain_core.embeddings import Embeddings - - -class CustomEmbeddings(Embeddings): - """ - 自定义文本嵌入类,用于将文本转换为向量表示 - 继承自LangChain的Embeddings基类 - """ - - def __init__(self, base_url, api_key, model): - """ - 初始化自定义嵌入类 - - 参数: - base_url: API的基础URL - api_key: 访问API所需的密钥 - model: 要使用的嵌入模型名称 - """ - self.base_url = base_url # API基础URL - self.api_key = api_key # API访问密钥 - self.model = model # 嵌入模型名称 - - def embed_documents(self, texts): - """ - 将多个文档转换为嵌入向量 - - 参数: - texts: 包含多个文本的列表 - - 返回: - 包含每个文本对应嵌入向量的列表 - """ - # 设置请求头,包括内容类型和认证信息 - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {self.api_key}" - } - - # 构建请求负载 - payload = {"input": texts, "model": self.model, "encoding_format": "float"} - - # 发送POST请求到嵌入API - response = requests.post( - f"{self.base_url}/embeddings", - headers=headers, - data=json.dumps(payload) - ) - - # 检查请求是否成功,如果失败则抛出异常 - response.raise_for_status() - - # 解析响应JSON数据 - json_data = response.json() - - # 从响应数据中提取嵌入向量并返回 - return [item["embedding"] for item in json_data["data"]] - - def embed_query(self, text): - """ - 将单个查询文本转换为嵌入向量 - - 参数: - text: 查询文本 - - 返回: - 对应的嵌入向量 - """ - # 调用embed_documents处理单个文本,并返回第一个结果 - return self.embed_documents([text])[0] diff --git a/app/utils/LLamaIndexEmbeddings.py b/app/utils/LLamaIndexEmbeddings.py new file mode 100644 index 0000000..0b0977f --- /dev/null +++ b/app/utils/LLamaIndexEmbeddings.py @@ -0,0 +1,49 @@ +import json + +import httpx +import requests +from llama_index.core.base.embeddings.base import BaseEmbedding, Embedding + + +class LLamaIndexEmbeddings(BaseEmbedding): + base_url: str = "" + api_key: str = "" + model: str = "" + def __init__(self, base_url, api_key, model): + super().__init__() + self.base_url = base_url # API基础URL + self.api_key = api_key # API访问密钥 + self.model = model # 嵌入模型名称 + + def embed_documents(self, texts): + headers = {"Content-Type": "application/json","Authorization": f"Bearer {self.api_key}"} + payload = {"input": texts, "model": self.model, "encoding_format": "float"} + response = requests.post(f"{self.base_url}/embeddings",headers=headers,data=json.dumps(payload)) + response.raise_for_status() + json_data = response.json() + return [item["embedding"] for item in json_data["data"]] + + def embed_query(self, text): + return self.embed_documents([text])[0] + + async def aembed_documents(self, texts): + headers = {"Content-Type": "application/json","Authorization": f"Bearer {self.api_key}"} + payload = {"input": texts, "model": self.model, "encoding_format": "float"} + async with httpx.AsyncClient() as client: + response = await client.post(url=f"{self.base_url}/embeddings",headers=headers,json=payload) + json_data = response.json() + return [item["embedding"] for item in json_data["data"]] + + async def aembed_query(self, text): + return (await self.aembed_documents([text]))[0] + + + def _get_query_embedding(self, query: str) -> Embedding: + raise Exception("_get_query_embedding: LLamaIndexEmbeddings only support async event loop.") + + async def _aget_query_embedding(self, query: str) -> Embedding: + return await self.aembed_query(query) + + def _get_text_embedding(self, text: str) -> Embedding: + return self.embed_query(text) +