refactor(embedding): 重构嵌入模型代码,添加异步支持并改进错误处理
- 将BaseEmbeddingModel改为抽象基类(ABC) - 添加异步预测方法apredict - 重构批量编码方法,添加任务状态跟踪 - 为Ollama和Other嵌入实现添加异步支持 - 改进错误处理和日志记录 - 添加类型注解提高代码可读性 docs: 更新changelog记录已知问题
This commit is contained in:
parent
f214f41a75
commit
92fe65fecf
@ -7,6 +7,8 @@
|
|||||||
|
|
||||||
🐛**BUGs**
|
🐛**BUGs**
|
||||||
- [x] LlightRAG 知识库中,点击边,没有显示,但是在全屏的时候却又能够显示出来。
|
- [x] LlightRAG 知识库中,点击边,没有显示,但是在全屏的时候却又能够显示出来。
|
||||||
|
- [ ] 部分 doc 格式的文件支持有问题
|
||||||
|
- [ ] 当出现不支持的文件类型的时候,前端没有限制
|
||||||
|
|
||||||
💯 **More**:
|
💯 **More**:
|
||||||
|
|
||||||
|
|||||||
@ -100,6 +100,9 @@ def chunk(text_or_path, params=None):
|
|||||||
|
|
||||||
def pdfreader(file_path, params=None):
|
def pdfreader(file_path, params=None):
|
||||||
"""读取PDF文件并返回text文本"""
|
"""读取PDF文件并返回text文本"""
|
||||||
|
if isinstance(file_path, str):
|
||||||
|
file_path = Path(file_path)
|
||||||
|
|
||||||
assert file_path.exists(), "File not found"
|
assert file_path.exists(), "File not found"
|
||||||
assert file_path.suffix.lower() == ".pdf", "File format not supported"
|
assert file_path.suffix.lower() == ".pdf", "File format not supported"
|
||||||
|
|
||||||
|
|||||||
@ -1,16 +1,15 @@
|
|||||||
import os
|
import os
|
||||||
import json
|
import json
|
||||||
|
import httpx
|
||||||
import requests
|
import requests
|
||||||
import asyncio
|
import asyncio
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod, ABC
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
|
|
||||||
from src import config
|
from src import config
|
||||||
from src.utils import hashstr, logger, get_docker_safe_url
|
from src.utils import hashstr, logger, get_docker_safe_url
|
||||||
|
|
||||||
|
|
||||||
class BaseEmbeddingModel:
|
class BaseEmbeddingModel(ABC):
|
||||||
embed_state = {}
|
|
||||||
|
|
||||||
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None):
|
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None):
|
||||||
"""
|
"""
|
||||||
@ -27,30 +26,38 @@ class BaseEmbeddingModel:
|
|||||||
self.dimension = dimension
|
self.dimension = dimension
|
||||||
self.base_url = get_docker_safe_url(base_url)
|
self.base_url = get_docker_safe_url(base_url)
|
||||||
self.api_key = os.getenv(api_key, api_key)
|
self.api_key = os.getenv(api_key, api_key)
|
||||||
|
self.embed_state = {}
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def predict(self, message):
|
def predict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
|
"""同步编码"""
|
||||||
raise NotImplementedError("Subclasses must implement this method")
|
raise NotImplementedError("Subclasses must implement this method")
|
||||||
|
|
||||||
def encode(self, message):
|
@abstractmethod
|
||||||
|
async def apredict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
|
"""异步编码"""
|
||||||
|
raise NotImplementedError("Subclasses must implement this method")
|
||||||
|
|
||||||
|
def encode(self, message: list[str] | str) -> list[list[float]]:
|
||||||
|
"""等同于predict"""
|
||||||
return self.predict(message)
|
return self.predict(message)
|
||||||
|
|
||||||
def encode_queries(self, queries):
|
def encode_queries(self, queries: list[str] | str) -> list[list[float]]:
|
||||||
|
"""等同于predict"""
|
||||||
return self.predict(queries)
|
return self.predict(queries)
|
||||||
|
|
||||||
async def aencode(self, message):
|
async def aencode(self, message: list[str] | str) -> list[list[float]]:
|
||||||
return await asyncio.to_thread(self.encode, message)
|
"""等同于apredict"""
|
||||||
|
return await self.apredict(message)
|
||||||
|
|
||||||
async def aencode_queries(self, queries):
|
async def aencode_queries(self, queries: list[str] | str) -> list[list[float]]:
|
||||||
return await asyncio.to_thread(self.encode_queries, queries)
|
"""等同于apredict"""
|
||||||
|
return await self.apredict(queries)
|
||||||
|
|
||||||
async def abatch_encode(self, messages, batch_size=40):
|
def batch_encode(self, messages: list[str], batch_size: int = 40) -> list[list[float]]:
|
||||||
return await asyncio.to_thread(self.batch_encode, messages, batch_size)
|
|
||||||
|
|
||||||
def batch_encode(self, messages, batch_size=40):
|
|
||||||
# logger.info(f"Batch encoding {len(messages)} messages")
|
# logger.info(f"Batch encoding {len(messages)} messages")
|
||||||
data = []
|
data = []
|
||||||
|
task_id = None
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
task_id = hashstr(messages)
|
task_id = hashstr(messages)
|
||||||
self.embed_state[task_id] = {
|
self.embed_state[task_id] = {
|
||||||
@ -63,10 +70,36 @@ class BaseEmbeddingModel:
|
|||||||
group_msg = messages[i:i+batch_size]
|
group_msg = messages[i:i+batch_size]
|
||||||
logger.info(f"Encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
|
logger.info(f"Encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
|
||||||
response = self.encode(group_msg)
|
response = self.encode(group_msg)
|
||||||
# logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}")
|
|
||||||
data.extend(response)
|
data.extend(response)
|
||||||
|
if task_id:
|
||||||
|
self.embed_state[task_id]['progress'] = i + len(group_msg)
|
||||||
|
|
||||||
|
if task_id:
|
||||||
|
self.embed_state[task_id]['status'] = 'completed'
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
async def abatch_encode(self, messages: list[str], batch_size: int = 40) -> list[list[float]]:
|
||||||
|
data = []
|
||||||
|
task_id = None
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
|
task_id = hashstr(messages)
|
||||||
|
self.embed_state[task_id] = {
|
||||||
|
'status': 'in-progress',
|
||||||
|
'total': len(messages),
|
||||||
|
'progress': 0
|
||||||
|
}
|
||||||
|
|
||||||
|
tasks = []
|
||||||
|
for i in range(0, len(messages), batch_size):
|
||||||
|
group_msg = messages[i:i+batch_size]
|
||||||
|
tasks.append(self.aencode(group_msg))
|
||||||
|
|
||||||
|
results = await asyncio.gather(*tasks)
|
||||||
|
for res in results:
|
||||||
|
data.extend(res)
|
||||||
|
|
||||||
|
if task_id:
|
||||||
self.embed_state[task_id]['progress'] = len(messages)
|
self.embed_state[task_id]['progress'] = len(messages)
|
||||||
self.embed_state[task_id]['status'] = 'completed'
|
self.embed_state[task_id]['status'] = 'completed'
|
||||||
|
|
||||||
@ -81,18 +114,38 @@ class OllamaEmbedding(BaseEmbeddingModel):
|
|||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self.base_url = self.base_url or get_docker_safe_url("http://localhost:11434/api/embed")
|
self.base_url = self.base_url or get_docker_safe_url("http://localhost:11434/api/embed")
|
||||||
|
|
||||||
def predict(self, message: list[str] | str):
|
def predict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
if isinstance(message, str):
|
if isinstance(message, str):
|
||||||
message = [message]
|
message = [message]
|
||||||
|
|
||||||
payload = {
|
payload = {"model": self.model, "input": message}
|
||||||
"model": self.model,
|
try:
|
||||||
"input": message,
|
response = requests.post(self.base_url, json=payload, timeout=60)
|
||||||
}
|
response.raise_for_status()
|
||||||
response = requests.request("POST", self.base_url, json=payload)
|
result = response.json()
|
||||||
response = json.loads(response.text)
|
if "embeddings" not in result:
|
||||||
assert response.get("embeddings"), f"Ollama Embedding failed: {response}"
|
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
|
||||||
return response["embeddings"]
|
return result["embeddings"]
|
||||||
|
except (requests.RequestException, json.JSONDecodeError) as e:
|
||||||
|
logger.error(f"Ollama Embedding request failed: {e}")
|
||||||
|
raise ValueError(f"Ollama Embedding request failed: {e}")
|
||||||
|
|
||||||
|
async def apredict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
|
if isinstance(message, str):
|
||||||
|
message = [message]
|
||||||
|
|
||||||
|
payload = {"model": self.model, "input": message}
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
try:
|
||||||
|
response = await client.post(self.base_url, json=payload, timeout=60)
|
||||||
|
response.raise_for_status()
|
||||||
|
result = response.json()
|
||||||
|
if "embeddings" not in result:
|
||||||
|
raise ValueError(f"Ollama Embedding failed: Invalid response format {result}")
|
||||||
|
return result["embeddings"]
|
||||||
|
except (httpx.RequestError, json.JSONDecodeError) as e:
|
||||||
|
logger.error(f"Ollama Embedding async request failed: {e}")
|
||||||
|
raise ValueError(f"Ollama Embedding async request failed: {e}")
|
||||||
|
|
||||||
|
|
||||||
class OtherEmbedding(BaseEmbeddingModel):
|
class OtherEmbedding(BaseEmbeddingModel):
|
||||||
@ -104,16 +157,32 @@ class OtherEmbedding(BaseEmbeddingModel):
|
|||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json"
|
||||||
}
|
}
|
||||||
|
|
||||||
def predict(self, message):
|
def build_payload(self, message: list[str] | str) -> dict:
|
||||||
payload = self.build_payload(message)
|
return {"model": self.model, "input": message}
|
||||||
response = requests.request("POST", self.base_url, json=payload, headers=self.headers)
|
|
||||||
response = json.loads(response.text)
|
|
||||||
assert response["data"], f"Other Embedding failed: {response}"
|
|
||||||
data = [a["embedding"] for a in response["data"]]
|
|
||||||
return data
|
|
||||||
|
|
||||||
def build_payload(self, message):
|
def predict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
return {
|
payload = self.build_payload(message)
|
||||||
"model": self.model,
|
try:
|
||||||
"input": message,
|
response = requests.post(self.base_url, json=payload, headers=self.headers, timeout=60)
|
||||||
}
|
response.raise_for_status()
|
||||||
|
result = response.json()
|
||||||
|
if not isinstance(result, dict) or "data" not in result:
|
||||||
|
raise ValueError(f"Other Embedding failed: Invalid response format {result}")
|
||||||
|
return [item["embedding"] for item in result["data"]]
|
||||||
|
except (requests.RequestException, json.JSONDecodeError) as e:
|
||||||
|
logger.error(f"Other Embedding request failed: {e}")
|
||||||
|
raise ValueError(f"Other Embedding request failed: {e}")
|
||||||
|
|
||||||
|
async def apredict(self, message: list[str] | str) -> list[list[float]]:
|
||||||
|
payload = self.build_payload(message)
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
try:
|
||||||
|
response = await client.post(self.base_url, json=payload, headers=self.headers, timeout=60)
|
||||||
|
response.raise_for_status()
|
||||||
|
result = response.json()
|
||||||
|
if not isinstance(result, dict) or "data" not in result:
|
||||||
|
raise ValueError(f"Other Embedding failed: Invalid response format {result}")
|
||||||
|
return [item["embedding"] for item in result["data"]]
|
||||||
|
except (httpx.RequestError, json.JSONDecodeError) as e:
|
||||||
|
logger.error(f"Other Embedding async request failed: {e}")
|
||||||
|
raise ValueError(f"Other Embedding async request failed: {e}")
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user