From fdf016254d760b3340faf191e90efb0d9c6f4389 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 7 Jul 2024 17:21:07 +0800 Subject: [PATCH] add wenxin support and fix double-predict bugs --- .gitignore | 3 +++ src/cli.py | 6 ++--- src/config/base.yaml | 2 +- src/models/__init__.py | 8 ++++++ src/models/chat_model.py | 56 ++++++++++++++++++++++++++++++++++------ 5 files changed, 63 insertions(+), 12 deletions(-) diff --git a/.gitignore b/.gitignore index 3b4a39e6..88bf7442 100644 --- a/.gitignore +++ b/.gitignore @@ -19,3 +19,6 @@ log logs *.log.* *.db + +### IDE +.vscode diff --git a/src/cli.py b/src/cli.py index abe9f33a..0edf9352 100644 --- a/src/cli.py +++ b/src/cli.py @@ -2,16 +2,16 @@ import os from dotenv import load_dotenv from core.history import HistoryManager from config import Config -from models.chat_model import DeepSeek, Zhipu +from models import select_model load_dotenv() if __name__ == "__main__": config = Config("config/base.yaml") - model = Zhipu() + model = select_model(config) - print("[CLI] Type 'exit' to quit") + print(f"[{config.model_provider}:{config.get('model_name', 'default')}] Type 'exit' to quit") history_manager = HistoryManager() while True: diff --git a/src/config/base.yaml b/src/config/base.yaml index 4a50c9cb..ba90d2ce 100644 --- a/src/config/base.yaml +++ b/src/config/base.yaml @@ -2,7 +2,7 @@ name: base ## model ### model_provider, option in deepseek, zhipu -model_provider: zhipu +model_provider: wenxin ## startup stream: True \ No newline at end of file diff --git a/src/models/__init__.py b/src/models/__init__.py index 83e7ea0d..d2bfca61 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -8,9 +8,17 @@ def select_model(config): model_name = config.model_name if model_provider == "deepseek": + from models.chat_model import DeepSeek return DeepSeek(model_name) + elif model_provider == "zhipu": + from models.chat_model import Zhipu return Zhipu(model_name) + + elif model_provider == "wenxin": + from models.chat_model import Wenxin + return Wenxin(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 1f0399d6..56d8f8bf 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -18,13 +18,6 @@ class OpenAIBase(): else: messages = message - response = self.client.chat.completions.create( - model=self.model_name, - messages=messages, - stream=stream, - ) - logger.debug(response) - if stream: return self._stream_response(messages) else: @@ -61,4 +54,51 @@ class Zhipu(OpenAIBase): model_name = model_name or "glm-4" api_key = os.getenv("ZHIPUAPI") base_url = "https://open.bigmodel.cn/api/paas/v4/" - super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) \ No newline at end of file + super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) + + +import qianfan + + +class WenxinResponse: + def __init__(self, content): + self.content = content + + +class Wenxin: + + def __init__(self, model_name="ernie-lite-8k") -> None: + self.model_name = model_name + access_key = os.getenv("QIANFAN_ACCESS_KEY") + secret_key = os.getenv("QIANFAN_SECRET_KEY") + self.client = qianfan.ChatCompletion(ak=access_key, sk=secret_key) + + def predict(self, message, stream=False): + + logger.debug(message) + if isinstance(message, str): + messages=[{"role": "user", "content": message}] + else: + messages = message + + if stream: + return self._stream_response(messages) + else: + return self._get_response(messages) + + def _stream_response(self, messages): + response = self.client.do( + model=self.model_name, + messages=messages, + stream=True, + ) + for chunk in response: + yield WenxinResponse(chunk["body"]["result"]) + + def _get_response(self, messages): + response = self.client.do( + model=self.model_name, + messages=messages, + stream=False, + ) + return WenxinResponse(response["body"]["result"]) \ No newline at end of file