优化本地 embedding 的GPU 处理问题,添加 GPU 部署映射
This commit is contained in:
parent
e422fd3dcf
commit
95b8b64170
@ -57,6 +57,8 @@ ZHIPUAI_API_KEY=<API_KEY> # 如果配置 智谱清言 添加此行,并替换
|
|||||||
|
|
||||||
### 启动服务
|
### 启动服务
|
||||||
|
|
||||||
|
> 确保已经安装了 [docker](https://docs.docker.com/engine/install/ubuntu/) 以及 [nvidia-container-toolkit](https://docs.nvidia.com/datacenter/cloud-native/container-toolkit/latest/install-guide.html)
|
||||||
|
|
||||||
**开发环境启动**(源代码修改会自动更新):
|
**开发环境启动**(源代码修改会自动更新):
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|||||||
@ -11,6 +11,13 @@ services:
|
|||||||
- ./src:/app/src
|
- ./src:/app/src
|
||||||
- ./saves:/app/saves
|
- ./saves:/app/saves
|
||||||
- ${MODEL_DIR:-./models}:/models # 使用默认值处理未定义的环境变量
|
- ${MODEL_DIR:-./models}:/models # 使用默认值处理未定义的环境变量
|
||||||
|
deploy:
|
||||||
|
resources:
|
||||||
|
reservations:
|
||||||
|
devices:
|
||||||
|
- driver: nvidia
|
||||||
|
count: 1
|
||||||
|
capabilities: [gpu]
|
||||||
ports:
|
ports:
|
||||||
- "5050:5050"
|
- "5050:5050"
|
||||||
networks:
|
networks:
|
||||||
|
|||||||
@ -10,6 +10,7 @@ dependencies = [
|
|||||||
"fastapi>=0.115.12",
|
"fastapi>=0.115.12",
|
||||||
"flagembedding>=1.3.4",
|
"flagembedding>=1.3.4",
|
||||||
"langchain-community>=0.3.22",
|
"langchain-community>=0.3.22",
|
||||||
|
"langchain-huggingface>=0.2.0",
|
||||||
"langchain-openai>=0.3.14",
|
"langchain-openai>=0.3.14",
|
||||||
"langgraph>=0.3.34",
|
"langgraph>=0.3.34",
|
||||||
"langsmith>=0.3.37",
|
"langsmith>=0.3.37",
|
||||||
|
|||||||
@ -2,8 +2,10 @@ import os
|
|||||||
import json
|
import json
|
||||||
import requests
|
import requests
|
||||||
import asyncio
|
import asyncio
|
||||||
from FlagEmbedding import FlagModel
|
import json
|
||||||
|
from abc import abstractmethod
|
||||||
from zhipuai import ZhipuAI
|
from zhipuai import ZhipuAI
|
||||||
|
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
|
||||||
@ -12,6 +14,10 @@ from src.utils import hashstr, logger, get_docker_safe_url
|
|||||||
class BaseEmbeddingModel:
|
class BaseEmbeddingModel:
|
||||||
embed_state = {}
|
embed_state = {}
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def predict(self, message):
|
||||||
|
raise NotImplementedError("Subclasses must implement this method")
|
||||||
|
|
||||||
def get_dimension(self):
|
def get_dimension(self):
|
||||||
if hasattr(self, "dimension"):
|
if hasattr(self, "dimension"):
|
||||||
return self.dimension
|
return self.dimension
|
||||||
@ -52,7 +58,7 @@ class BaseEmbeddingModel:
|
|||||||
group_msg = messages[i:i+batch_size]
|
group_msg = messages[i:i+batch_size]
|
||||||
logger.info(f"Encoding {i} to {i+batch_size} with {len(messages)} messages")
|
logger.info(f"Encoding {i} to {i+batch_size} with {len(messages)} messages")
|
||||||
response = self.encode(group_msg)
|
response = self.encode(group_msg)
|
||||||
logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}")
|
# logger.debug(f"Response: {len(response)=}, {len(group_msg)=}, {len(response[0])=}")
|
||||||
data.extend(response)
|
data.extend(response)
|
||||||
|
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
@ -61,21 +67,8 @@ class BaseEmbeddingModel:
|
|||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel):
|
class LocalEmbeddingModel(BaseEmbeddingModel):
|
||||||
def __init__(self, config, **kwargs):
|
def __init__(self, config, **kwargs):
|
||||||
"""
|
|
||||||
对于本地模型,也可以在 src/static/models.private.yaml 中配置对应的 local_path 路径
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
EMBED_MODEL_INFO:
|
|
||||||
local/BAAI/bge-m3:
|
|
||||||
dimension: 1024
|
|
||||||
name: BAAI/bge-m3
|
|
||||||
local_path: /path/to/bge-m3
|
|
||||||
```
|
|
||||||
|
|
||||||
但是也要确保在 docker-compose 中映射了 MODEL_DIR 到 /models 目录
|
|
||||||
"""
|
|
||||||
info = config.embed_model_names[config.embed_model]
|
info = config.embed_model_names[config.embed_model]
|
||||||
|
|
||||||
self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
|
self.model = config.model_local_paths.get(info["name"], info.get("local_path"))
|
||||||
@ -91,17 +84,32 @@ class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel):
|
|||||||
|
|
||||||
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`,"
|
logger.info(f"Loading local model `{info['name']}` from `{self.model}` with device `{config.device}`,"
|
||||||
f"如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
|
f"如果没配置任何路径的话,正常情况下会自动从 Huggingface 下载模型,如果遇到下载失败,可以尝试使用 HF_MIRROR 环境变量;"
|
||||||
f"如果还是不行,建议手动下载到某个文件夹比如 /path/to/models/BAAI/bge-m3 目录下;"
|
f"如果还是不行,建议手动下载到某个文件夹,比如 {os.getenv('MODEL_DIR', '/models')}/BAAI/bge-m3 目录下;")
|
||||||
f"然后配置 src/.env 文件中的 MODEL_DIR 环境变量到 /path/to/models 目录;"
|
|
||||||
f"如果是在 docker 中运行,请确保 docker-compose 文件(line 12 左右)中映射了 MODEL_DIR 到 /models 目录")
|
|
||||||
|
|
||||||
super().__init__(self.model,
|
self.model = HuggingFaceEmbeddings(
|
||||||
query_instruction_for_retrieval=info.get("query_instruction", None),
|
model_name=self.model,
|
||||||
use_fp16=False,
|
model_kwargs={'device': config.device},
|
||||||
device=config.device,
|
encode_kwargs={
|
||||||
**kwargs)
|
'normalize_embeddings': True,
|
||||||
|
'prompt_name': info.get("query_instruction", None),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
logger.info(f"Embedding model {info['name']} loaded")
|
logger.info(f"Embedding model {info['name']} loaded, {self.model=}")
|
||||||
|
|
||||||
|
def predict(self, message):
|
||||||
|
return self.model.embed_documents(message)
|
||||||
|
|
||||||
|
async def aencode(self, message):
|
||||||
|
return await self.model.aembed_documents(message)
|
||||||
|
|
||||||
|
def encode_queries(self, queries):
|
||||||
|
logger.warning(f"Huggingface Model 不支持批量 encode queries,因此使用训练实现")
|
||||||
|
data = []
|
||||||
|
for q in queries:
|
||||||
|
data.append(self.predict(q))
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class ZhipuEmbedding(BaseEmbeddingModel):
|
class ZhipuEmbedding(BaseEmbeddingModel):
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user