vllm supported

This commit is contained in:
Wenjie Zhang 2024-07-20 15:27:33 +08:00
parent 789ab07f51
commit ddd4ce99c3
5 changed files with 66 additions and 9 deletions

View File

@ -4,9 +4,10 @@ name: base
## model
### model_provider, option in deepseek, zhipu
model_provider: zhipu
model_name: null # for default
## model dir 可以写相对路径和绝对路径
### 相对路径是相对于环境变量中 MODEL_ROOT_DIR 的路径
### 相对路径是相对于环境变量 (.env) 中 MODEL_ROOT_DIR 的路径
model_local_paths:
bge-large-zh-v1.5: bge-large-zh-v1.5
oneke: oneke

42
src/models/README.md Normal file
View File

@ -0,0 +1,42 @@
## 模型说明
### 1. 对话模型支持
模型仅支持通过API调用的模型如果是需要运行本地模型则建议使用 vllm 转成 API 之后使用。
|模型供应商(`config.model_provider`)|默认模型(`model.model_name`)|配置项目(`.env`)|
|:-|:-|:-|
|`qianfan`|`ernie_speed`|`QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY`|
|`zhipu`|`glm-4`|`ZHIPUAPI`|
|`deepseek`|`deepseek-chat`|`DEEPSEEKAPI`|
|`vllm`|`vllm`|`VLLM_API_KEY`, `VLLM_API_BASE`|
vllm 部署参考脚本:
```bash
python -m vllm.entrypoints.openai.api_server --model ~/models/Meta-Llama-3-8B-Instruct --served-model-name vllm --trust-remote-code
```
*openai 没条件测,不知道
### 2. 向量模型支持
|模型名称(`config.embed_model`)|默认路径|可配置项目(`config.model_local_paths`|
|:-|:-|:-|
|`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`|
### 3. 重排序模型支持
例如:
```yaml
model_provider: qianfan
model_name: null # for default
## model dir 可以写相对路径和绝对路径
### 相对路径是相对于环境变量 (.env) 中 MODEL_ROOT_DIR 的路径
model_local_paths:
bge-large-zh-v1.5: bge-large-zh-v1.5
```

View File

@ -20,6 +20,10 @@ def select_model(config):
from models.chat_model import Qianfan
return Qianfan(model_name)
elif model_provider == "vllm":
from models.chat_model import VLLM
return VLLM(model_name)
elif model_provider is None:
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
else:

View File

@ -56,6 +56,13 @@ class Zhipu(OpenAIBase):
base_url = "https://open.bigmodel.cn/api/paas/v4/"
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class VLLM(OpenAIBase):
def __init__(self, model_name=None):
model_name = model_name or "vllm"
api_key = os.getenv("VLLM_API_KEY")
base_url = os.getenv("VLLM_API_BASE")
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
import qianfan
@ -99,4 +106,6 @@ class Qianfan:
messages=messages,
stream=False,
)
return QianfanResponse(response["body"]["result"])
return QianfanResponse(response["body"]["result"])

View File

@ -44,13 +44,14 @@ def chat():
def generate_response():
content = ""
for delta in model.predict(messages, stream=True):
content += delta.content
response_chunk = json.dumps({
"history": history_manager.update_ai(content),
"response": content,
"refs": refs # TODO: 优化 refs不需要每次都返回
}, ensure_ascii=False).encode('utf8') + b'\n'
yield response_chunk
if delta.content:
content += delta.content
response_chunk = json.dumps({
"history": history_manager.update_ai(content),
"response": content,
"refs": refs # TODO: 优化 refs不需要每次都返回
}, ensure_ascii=False).encode('utf8') + b'\n'
yield response_chunk
return Response(generate_response(), content_type='application/json', status=200)