vllm supported
This commit is contained in:
parent
789ab07f51
commit
ddd4ce99c3
@ -4,9 +4,10 @@ name: base
|
|||||||
## model
|
## model
|
||||||
### model_provider, option in deepseek, zhipu
|
### model_provider, option in deepseek, zhipu
|
||||||
model_provider: zhipu
|
model_provider: zhipu
|
||||||
|
model_name: null # for default
|
||||||
|
|
||||||
## model dir 可以写相对路径和绝对路径
|
## model dir 可以写相对路径和绝对路径
|
||||||
### 相对路径是相对于环境变量中 MODEL_ROOT_DIR 的路径
|
### 相对路径是相对于环境变量 (.env) 中 MODEL_ROOT_DIR 的路径
|
||||||
model_local_paths:
|
model_local_paths:
|
||||||
bge-large-zh-v1.5: bge-large-zh-v1.5
|
bge-large-zh-v1.5: bge-large-zh-v1.5
|
||||||
oneke: oneke
|
oneke: oneke
|
||||||
42
src/models/README.md
Normal file
42
src/models/README.md
Normal 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
|
||||||
|
```
|
||||||
@ -20,6 +20,10 @@ def select_model(config):
|
|||||||
from models.chat_model import Qianfan
|
from models.chat_model import Qianfan
|
||||||
return Qianfan(model_name)
|
return Qianfan(model_name)
|
||||||
|
|
||||||
|
elif model_provider == "vllm":
|
||||||
|
from models.chat_model import VLLM
|
||||||
|
return VLLM(model_name)
|
||||||
|
|
||||||
elif model_provider is None:
|
elif model_provider is None:
|
||||||
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
|
raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`")
|
||||||
else:
|
else:
|
||||||
|
|||||||
@ -56,6 +56,13 @@ class Zhipu(OpenAIBase):
|
|||||||
base_url = "https://open.bigmodel.cn/api/paas/v4/"
|
base_url = "https://open.bigmodel.cn/api/paas/v4/"
|
||||||
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
|
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
|
import qianfan
|
||||||
|
|
||||||
@ -99,4 +106,6 @@ class Qianfan:
|
|||||||
messages=messages,
|
messages=messages,
|
||||||
stream=False,
|
stream=False,
|
||||||
)
|
)
|
||||||
return QianfanResponse(response["body"]["result"])
|
return QianfanResponse(response["body"]["result"])
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -44,13 +44,14 @@ def chat():
|
|||||||
def generate_response():
|
def generate_response():
|
||||||
content = ""
|
content = ""
|
||||||
for delta in model.predict(messages, stream=True):
|
for delta in model.predict(messages, stream=True):
|
||||||
content += delta.content
|
if delta.content:
|
||||||
response_chunk = json.dumps({
|
content += delta.content
|
||||||
"history": history_manager.update_ai(content),
|
response_chunk = json.dumps({
|
||||||
"response": content,
|
"history": history_manager.update_ai(content),
|
||||||
"refs": refs # TODO: 优化 refs,不需要每次都返回
|
"response": content,
|
||||||
}, ensure_ascii=False).encode('utf8') + b'\n'
|
"refs": refs # TODO: 优化 refs,不需要每次都返回
|
||||||
yield response_chunk
|
}, ensure_ascii=False).encode('utf8') + b'\n'
|
||||||
|
yield response_chunk
|
||||||
|
|
||||||
return Response(generate_response(), content_type='application/json', status=200)
|
return Response(generate_response(), content_type='application/json', status=200)
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user