diff --git a/src/config/base.yaml b/src/config/base.yaml index f7e34f54..9ff69532 100644 --- a/src/config/base.yaml +++ b/src/config/base.yaml @@ -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 \ No newline at end of file diff --git a/src/models/README.md b/src/models/README.md new file mode 100644 index 00000000..732f26bd --- /dev/null +++ b/src/models/README.md @@ -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 +``` diff --git a/src/models/__init__.py b/src/models/__init__.py index bae3a78e..35cd0c5d 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -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: diff --git a/src/models/chat_model.py b/src/models/chat_model.py index f6b8a8fc..94e4f7f0 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -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"]) \ No newline at end of file + return QianfanResponse(response["body"]["result"]) + + diff --git a/src/views/common_view.py b/src/views/common_view.py index 8e0b70dd..0436ead6 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -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)