From 637b528aadd401cee936ef87962303caaddf74fa Mon Sep 17 00:00:00 2001 From: Sakura-RanChen <1908198662@qq.com> Date: Thu, 13 Nov 2025 17:08:21 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=8E=BB=E9=99=A4=E9=BB=98=E8=AE=A4?= =?UTF-8?q?=E5=80=BC=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../core/providers/llm/openai/openai.py | 68 +++++++++++++------ 1 file changed, 46 insertions(+), 22 deletions(-) diff --git a/main/xiaozhi-server/core/providers/llm/openai/openai.py b/main/xiaozhi-server/core/providers/llm/openai/openai.py index 03bc57ed..2863c420 100644 --- a/main/xiaozhi-server/core/providers/llm/openai/openai.py +++ b/main/xiaozhi-server/core/providers/llm/openai/openai.py @@ -21,22 +21,22 @@ class LLMProvider(LLMProviderBase): self.timeout = int(timeout) if timeout else 300 param_defaults = { - "max_tokens": (500, int), - "temperature": (0.7, lambda x: round(float(x), 1)), - "top_p": (1.0, lambda x: round(float(x), 1)), - "frequency_penalty": (0, lambda x: round(float(x), 1)), + "max_tokens": int, + "temperature": lambda x: round(float(x), 1), + "top_p": lambda x: round(float(x), 1), + "frequency_penalty": lambda x: round(float(x), 1), } - for param, (default, converter) in param_defaults.items(): + for param, converter in param_defaults.items(): value = config.get(param) try: setattr( self, param, - converter(value) if value not in (None, "") else default, + converter(value) if value not in (None, "") else None, ) except (ValueError, TypeError): - setattr(self, param, default) + setattr(self, param, None) logger.debug( f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}" @@ -59,17 +59,25 @@ class LLMProvider(LLMProviderBase): try: dialogue = self.normalize_dialogue(dialogue) - responses = self.client.chat.completions.create( - model=self.model_name, - messages=dialogue, - stream=True, - max_tokens=kwargs.get("max_tokens", self.max_tokens), - temperature=kwargs.get("temperature", self.temperature), - top_p=kwargs.get("top_p", self.top_p), - frequency_penalty=kwargs.get( - "frequency_penalty", self.frequency_penalty - ), - ) + request_params = { + "model": self.model_name, + "messages": dialogue, + "stream": True, + } + + # 添加可选参数,只有当参数不为None时才添加 + optional_params = { + "max_tokens": kwargs.get("max_tokens", self.max_tokens), + "temperature": kwargs.get("temperature", self.temperature), + "top_p": kwargs.get("top_p", self.top_p), + "frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty), + } + + for key, value in optional_params.items(): + if value is not None: + request_params[key] = value + + responses = self.client.chat.completions.create(**request_params) is_active = True for chunk in responses: @@ -91,13 +99,29 @@ class LLMProvider(LLMProviderBase): except Exception as e: logger.bind(tag=TAG).error(f"Error in response generation: {e}") - def response_with_functions(self, session_id, dialogue, functions=None): + def response_with_functions(self, session_id, dialogue, functions=None, **kwargs): try: dialogue = self.normalize_dialogue(dialogue) - stream = self.client.chat.completions.create( - model=self.model_name, messages=dialogue, stream=True, tools=functions - ) + request_params = { + "model": self.model_name, + "messages": dialogue, + "stream": True, + "tools": functions, + } + + optional_params = { + "max_tokens": kwargs.get("max_tokens", self.max_tokens), + "temperature": kwargs.get("temperature", self.temperature), + "top_p": kwargs.get("top_p", self.top_p), + "frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty), + } + + for key, value in optional_params.items(): + if value is not None: + request_params[key] = value + + stream = self.client.chat.completions.create(**request_params) for chunk in stream: if getattr(chunk, "choices", None):