2024-07-21 18:15:28 +08:00
|
|
|
|
import os
|
2025-02-23 16:38:56 +08:00
|
|
|
|
import json
|
|
|
|
|
|
import requests
|
2025-04-23 21:18:39 +08:00
|
|
|
|
import asyncio
|
2025-05-09 23:45:16 +08:00
|
|
|
|
from abc import abstractmethod
|
|
|
|
|
|
from langchain_huggingface import HuggingFaceEmbeddings
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
from src import config
|
2025-03-04 13:49:00 +08:00
|
|
|
|
from src.utils import hashstr, logger, get_docker_safe_url
|
2024-07-09 05:04:20 +08:00
|
|
|
|
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
class BaseEmbeddingModel:
|
2025-03-03 21:57:27 +08:00
|
|
|
|
embed_state = {}
|
2025-03-20 19:51:46 +08:00
|
|
|
|
|
2025-07-26 03:36:54 +08:00
|
|
|
|
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Args:
|
|
|
|
|
|
model: 模型名称,冗余设计,同name
|
|
|
|
|
|
name: 模型名称,冗余设计,同model
|
|
|
|
|
|
dimension: 维度
|
|
|
|
|
|
url: 请求URL,冗余设计,同base_url
|
|
|
|
|
|
base_url: 基础URL,请求URL,冗余设计,同url
|
|
|
|
|
|
api_key: 请求API密钥
|
|
|
|
|
|
"""
|
|
|
|
|
|
base_url = base_url or url
|
|
|
|
|
|
self.model = model or name
|
|
|
|
|
|
self.dimension = dimension
|
|
|
|
|
|
self.base_url = get_docker_safe_url(base_url)
|
|
|
|
|
|
self.api_key = os.getenv(api_key, api_key)
|
2025-07-02 02:38:36 +08:00
|
|
|
|
|
2025-05-09 23:45:16 +08:00
|
|
|
|
@abstractmethod
|
|
|
|
|
|
def predict(self, message):
|
|
|
|
|
|
raise NotImplementedError("Subclasses must implement this method")
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
def encode(self, message):
|
|
|
|
|
|
return self.predict(message)
|
|
|
|
|
|
|
|
|
|
|
|
def encode_queries(self, queries):
|
|
|
|
|
|
return self.predict(queries)
|
2025-03-03 21:57:27 +08:00
|
|
|
|
|
2025-04-23 21:18:39 +08:00
|
|
|
|
async def aencode(self, message):
|
|
|
|
|
|
return await asyncio.to_thread(self.encode, message)
|
|
|
|
|
|
|
|
|
|
|
|
async def aencode_queries(self, queries):
|
|
|
|
|
|
return await asyncio.to_thread(self.encode_queries, queries)
|
|
|
|
|
|
|
2025-07-26 03:36:54 +08:00
|
|
|
|
async def abatch_encode(self, messages, batch_size=40):
|
2025-04-23 21:18:39 +08:00
|
|
|
|
return await asyncio.to_thread(self.batch_encode, messages, batch_size)
|
|
|
|
|
|
|
2025-07-26 03:36:54 +08:00
|
|
|
|
def batch_encode(self, messages, batch_size=40):
|
2025-07-29 12:58:13 +08:00
|
|
|
|
# logger.info(f"Batch encoding {len(messages)} messages")
|
2025-03-03 21:57:27 +08:00
|
|
|
|
data = []
|
|
|
|
|
|
|
|
|
|
|
|
if len(messages) > batch_size:
|
|
|
|
|
|
task_id = hashstr(messages)
|
|
|
|
|
|
self.embed_state[task_id] = {
|
|
|
|
|
|
'status': 'in-progress',
|
|
|
|
|
|
'total': len(messages),
|
|
|
|
|
|
'progress': 0
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
for i in range(0, len(messages), batch_size):
|
|
|
|
|
|
group_msg = messages[i:i+batch_size]
|
2025-07-29 12:58:13 +08:00
|
|
|
|
logger.info(f"Encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
|
2025-03-11 16:26:55 +08:00
|
|
|
|
response = self.encode(group_msg)
|
2025-05-09 23:45:16 +08:00
|
|
|
|
# logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}")
|
2025-03-03 21:57:27 +08:00
|
|
|
|
data.extend(response)
|
|
|
|
|
|
|
|
|
|
|
|
if len(messages) > batch_size:
|
|
|
|
|
|
self.embed_state[task_id]['progress'] = len(messages)
|
|
|
|
|
|
self.embed_state[task_id]['status'] = 'completed'
|
|
|
|
|
|
|
|
|
|
|
|
return data
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
class OllamaEmbedding(BaseEmbeddingModel):
|
2025-07-02 02:38:36 +08:00
|
|
|
|
"""
|
|
|
|
|
|
Ollama Embedding Model
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
2025-07-26 03:36:54 +08:00
|
|
|
|
def __init__(self, **kwargs) -> None:
|
|
|
|
|
|
super().__init__(**kwargs)
|
|
|
|
|
|
self.base_url = self.base_url or get_docker_safe_url("http://localhost:11434/api/embed")
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
def predict(self, message: list[str] | str):
|
|
|
|
|
|
if isinstance(message, str):
|
|
|
|
|
|
message = [message]
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
payload = {
|
|
|
|
|
|
"model": self.model,
|
|
|
|
|
|
"input": message,
|
|
|
|
|
|
}
|
2025-07-26 03:36:54 +08:00
|
|
|
|
response = requests.request("POST", self.base_url, json=payload)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
response = json.loads(response.text)
|
|
|
|
|
|
assert response.get("embeddings"), f"Ollama Embedding failed: {response}"
|
|
|
|
|
|
return response["embeddings"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class OtherEmbedding(BaseEmbeddingModel):
|
2025-02-23 16:38:56 +08:00
|
|
|
|
|
2025-07-26 03:36:54 +08:00
|
|
|
|
def __init__(self, **kwargs) -> None:
|
|
|
|
|
|
super().__init__(**kwargs)
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.headers = {
|
2025-03-04 13:49:00 +08:00
|
|
|
|
"Authorization": f"Bearer {self.api_key}",
|
2025-02-23 16:38:56 +08:00
|
|
|
|
"Content-Type": "application/json"
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-03-04 13:49:00 +08:00
|
|
|
|
def predict(self, message):
|
2025-02-23 16:38:56 +08:00
|
|
|
|
payload = self.build_payload(message)
|
2025-07-26 03:36:54 +08:00
|
|
|
|
response = requests.request("POST", self.base_url, json=payload, headers=self.headers)
|
2025-02-23 16:38:56 +08:00
|
|
|
|
response = json.loads(response.text)
|
2025-03-04 13:49:00 +08:00
|
|
|
|
assert response["data"], f"Other Embedding failed: {response}"
|
2025-02-23 16:38:56 +08:00
|
|
|
|
data = [a["embedding"] for a in response["data"]]
|
|
|
|
|
|
return data
|
|
|
|
|
|
|
|
|
|
|
|
def build_payload(self, message):
|
|
|
|
|
|
return {
|
|
|
|
|
|
"model": self.model,
|
|
|
|
|
|
"input": message,
|
|
|
|
|
|
}
|