diff --git a/README.md b/README.md index c71f9297..4f7c96cd 100644 --- a/README.md +++ b/README.md @@ -15,7 +15,7 @@ ## 准备 -1. 提供 API 服务商的 API_KEY,并放置在 `src/.env` 文件中,参考 `src/.env.template`。默认使用的是智谱AI。需要配置 `ZHIPUAPI=`。 +1. 提供 API 服务商的 API_KEY,并放置在 `src/.env` 文件中,参考 `src/.env.template`。默认使用的是智谱AI。需要配置 `ZHIPUAI_API_KEY=`。 2. 配置 python 环境 `pip install -r requirements.txt`,python 版本应当小于 `3.12`。 3. 前端 UI 部分,需要安装 Node.js 环境,参考:[Download Node.js](https://nodejs.org/en/download/package-manager)。 diff --git a/src/.env.template b/src/.env.template index 31b78c9d..3dfedca2 100644 --- a/src/.env.template +++ b/src/.env.template @@ -1,7 +1,2 @@ -DIFYAPI=optional -DEEPSEEKAPI=optional -ZHIPUAPI=optional -QIANFAN_ACCESS_KEY=optional -QIANFAN_SECRET_KEY=optional -HF_ENDPOINT='https://hf-mirror.com' +ZHIPUAI_API_KEY= CUDA_VISIBLE_DEVICES=0 \ No newline at end of file diff --git a/src/models/README.md b/src/models/README.md index 1594a285..5027489c 100644 --- a/src/models/README.md +++ b/src/models/README.md @@ -7,8 +7,8 @@ |模型供应商(`config.model_provider`)|默认模型(`config.model_name`)|配置项目(`.env`)| |:-|:-|:-| |`qianfan`|`ernie_speed`|`QIANFAN_ACCESS_KEY`, `QIANFAN_SECRET_KEY`| -|`zhipu`(default)|`glm-4`|`ZHIPUAPI`| -|`deepseek`|`deepseek-chat`|`DEEPSEEKAPI`| +|`zhipu`(default)|`glm-4`|`ZHIPUAI_API_KEY`| +|`deepseek`|`deepseek-chat`|`DEEPSEEK_API_KEY`| |`vllm`|`vllm`|`VLLM_API_KEY`, `VLLM_API_BASE`| vllm 的具体配置项可以参考[这里](https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#named-arguments), 部署参考脚本: @@ -36,7 +36,7 @@ python -m vllm.entrypoints.openai.api_server \ |模型名称(`config.embed_model`)|默认路径/模型|需要配置项目(`config.model_local_paths`)| |:-|:-|:-| |`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`(*修改为本地路径)| -|`zhipu`|`embedding-2`|`ZHIPUAPI` (`.env`)| +|`zhipu`|`embedding-2`|`ZHIPUAI_API_KEY` (`.env`)| 例如(`saves/config/config.yaml`): diff --git a/src/models/__init__.py b/src/models/__init__.py index 8e1f6592..f941fbb5 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -32,6 +32,10 @@ def select_model(config): from src.models.chat_model import OpenModel return OpenModel(model_name) + elif model_provider == "siliconflow": + from src.models.chat_model import SiliconFlow + return SiliconFlow(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 2ed63c18..02e3a6d9 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -50,15 +50,15 @@ class OpenModel(OpenAIBase): class DeepSeek(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "deepseek-chat" - api_key = os.getenv("DEEPSEEKAPI") + api_key = os.getenv("DEEPSEEK_API_KEY") base_url = "https://api.deepseek.com" super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) class Zhipu(OpenAIBase): def __init__(self, model_name=None): - model_name = model_name or "glm-4" - api_key = os.getenv("ZHIPUAPI") + model_name = model_name or "glm-4-flash" + api_key = os.getenv("ZHIPUAI_API_KEY") base_url = "https://open.bigmodel.cn/api/paas/v4/" super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) @@ -69,7 +69,12 @@ class VLLM(OpenAIBase): base_url = os.getenv("VLLM_API_BASE") super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) - +class SiliconFlow(OpenAIBase): + def __init__(self, model_name=None): + model_name = model_name or "meta-llama/Meta-Llama-3.1-8B-Instruct" + api_key = os.getenv("SILICONFLOW_API_KEY") + base_url = "https://api.siliconflow.cn/v1" + super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) class GeneralResponse: @@ -119,11 +124,10 @@ class Qianfan: class DashScope: - def __init__(self, model_name="qwen-long") -> None: + def __init__(self, model_name="qwen-max-latest") -> None: self.model_name = model_name self.api_key= os.getenv("DASHSCOPE_API_KEY") - def predict(self, message, stream=False): if isinstance(message, str): messages=[{"role": "user", "content": message}] @@ -162,6 +166,6 @@ class DashScope: if __name__ == "__main__": - model = DashScope() + model = SiliconFlow() for a in model.predict("你好", stream=True): - print(a.content) \ No newline at end of file + print(a.content, end="") \ No newline at end of file diff --git a/src/models/embedding.py b/src/models/embedding.py index 1826ab27..55f55d32 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -52,7 +52,7 @@ class ZhipuEmbedding: def __init__(self, model_info, config) -> None: self.config = config self.model_info = model_info - self.client = ZhipuAI(api_key=os.getenv("ZHIPUAPI")) + self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) logger.info("Zhipu Embedding model loaded") self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:"