From 02cb9c35b312af71644d61eff9b83650195d2af3 Mon Sep 17 00:00:00 2001 From: Sakura-RanChen <1908198662@qq.com> Date: Thu, 22 May 2025 17:14:58 +0800 Subject: [PATCH] =?UTF-8?q?update:=20=E8=AE=B0=E5=BF=86=E6=A8=A1=E5=9D=97?= =?UTF-8?q?=E4=BD=BF=E7=94=A8=E7=8B=AC=E7=AB=8BLLM=20openai=E5=A2=9E?= =?UTF-8?q?=E5=8A=A0=E8=B6=85=E5=8F=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/config.yaml | 4 +++ main/xiaozhi-server/core/connection.py | 28 +++++++++++++++++ .../xiaozhi-server/core/providers/llm/base.py | 6 ++-- .../core/providers/llm/openai/openai.py | 31 +++++++++++++------ .../core/providers/memory/base.py | 8 ++++- .../memory/mem_local_short/mem_local_short.py | 21 +++++++------ 6 files changed, 75 insertions(+), 23 deletions(-) diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index 44845df8..ebec6ccd 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -220,6 +220,10 @@ Memory: mem_local_short: # 本地记忆功能,通过selected_module的llm总结,数据保存在本地,不会上传到服务器 type: mem_local_short + # 配备记忆存储独立的思考模型 + # 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型 + # 如果你的不想使用selected_module.LLM记忆存储,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM + llm: ChatGLMLLM ASR: FunASR: diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 74b612fe..fd8900ed 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -452,6 +452,34 @@ class ConnectionHandler: save_to_file=not self.read_config_from_api, ) + # 获取记忆总结配置 + memory_config = self.config["Memory"] + memory_type = self.config["Memory"][self.config["selected_module"]["Memory"]][ + "type" + ] + # 如果使用 nomen,直接返回 + if memory_type == "nomem": + return + # 使用 mem_local_short 模式 + elif memory_type == "mem_local_short": + memory_llm_name = memory_config[self.config["selected_module"]["Memory"]]["llm"] + if memory_llm_name and memory_llm_name in self.config["LLM"]: + # 如果配置了专用LLM,则创建独立的LLM实例 + from core.utils import llm as llm_utils + memory_llm_config = self.config["LLM"][memory_llm_name] + memory_llm_type = memory_llm_config.get("type", memory_llm_name) + memory_llm = llm_utils.create_instance( + memory_llm_type, memory_llm_config + ) + self.logger.bind(tag=TAG).info( + f"为记忆总结创建了专用LLM: {memory_llm_name}, 类型: {memory_llm_type}" + ) + self.memory.set_llm(memory_llm) + else: + # 否则使用主LLM + self.memory.set_llm(self.llm) + self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型") + def _initialize_intent(self): self.intent_type = self.config["Intent"][ self.config["selected_module"]["Intent"] diff --git a/main/xiaozhi-server/core/providers/llm/base.py b/main/xiaozhi-server/core/providers/llm/base.py index 97d0d8e7..77786140 100644 --- a/main/xiaozhi-server/core/providers/llm/base.py +++ b/main/xiaozhi-server/core/providers/llm/base.py @@ -10,7 +10,7 @@ class LLMProviderBase(ABC): """LLM response generator""" pass - def response_no_stream(self, system_prompt, user_prompt): + def response_no_stream(self, system_prompt, user_prompt, **kwargs): try: # 构造对话格式 dialogue = [ @@ -18,7 +18,7 @@ class LLMProviderBase(ABC): {"role": "user", "content": user_prompt} ] result = "" - for part in self.response("", dialogue): + for part in self.response("", dialogue, **kwargs): result += part return result @@ -30,7 +30,7 @@ class LLMProviderBase(ABC): """ Default implementation for function calling (streaming) This should be overridden by providers that support function calls - + Returns: generator that yields either text tokens or a special function call token """ # For providers that don't support functions, just return regular response diff --git a/main/xiaozhi-server/core/providers/llm/openai/openai.py b/main/xiaozhi-server/core/providers/llm/openai/openai.py index a20067af..bc0e7f21 100644 --- a/main/xiaozhi-server/core/providers/llm/openai/openai.py +++ b/main/xiaozhi-server/core/providers/llm/openai/openai.py @@ -16,26 +16,37 @@ class LLMProvider(LLMProviderBase): self.base_url = config.get("base_url") else: self.base_url = config.get("url") - max_tokens = config.get("max_tokens") - if max_tokens is None or max_tokens == "": - max_tokens = 500 - try: - max_tokens = int(max_tokens) - except (ValueError, TypeError): - max_tokens = 500 - self.max_tokens = max_tokens + 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)) + } + + for param, (default, converter) in param_defaults.items(): + value = config.get(param) + try: + setattr(self, param, converter(value) if value not in (None, "") else default) + except (ValueError, TypeError): + setattr(self, param, default) + + logger.debug( + f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}") check_model_key("LLM", self.api_key) self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) - def response(self, session_id, dialogue): + def response(self, session_id, dialogue, **kwargs): try: responses = self.client.chat.completions.create( model=self.model_name, messages=dialogue, stream=True, - max_tokens=self.max_tokens, + 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), ) is_active = True diff --git a/main/xiaozhi-server/core/providers/memory/base.py b/main/xiaozhi-server/core/providers/memory/base.py index 19dcabb8..f404f15e 100644 --- a/main/xiaozhi-server/core/providers/memory/base.py +++ b/main/xiaozhi-server/core/providers/memory/base.py @@ -9,7 +9,13 @@ class MemoryProviderBase(ABC): def __init__(self, config): self.config = config self.role_id = None - self.llm = None + + def set_llm(self, llm): + self.llm = llm + # 获取模型名称和类型信息 + model_name = getattr(llm, "model_name", str(llm.__class__.__name__)) + # 记录更详细的日志 + logger.bind(tag=TAG).info(f"记忆总结设置LLM: {model_name}") @abstractmethod async def save_memory(self, msgs): diff --git a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py index a916e6c9..844035bd 100644 --- a/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py +++ b/main/xiaozhi-server/core/providers/memory/mem_local_short/mem_local_short.py @@ -107,7 +107,7 @@ TAG = __name__ class MemoryProvider(MemoryProviderBase): def __init__(self, config, summary_memory): super().__init__(config) - self.short_momery = "" + self.short_memory = "" self.save_to_file = True self.memory_path = get_project_dir() + "data/.memory.yaml" self.load_memory(summary_memory) @@ -122,7 +122,7 @@ class MemoryProvider(MemoryProviderBase): def load_memory(self, summary_memory): # api获取到总结记忆后直接返回 if summary_memory or not self.save_to_file: - self.short_momery = summary_memory + self.short_memory = summary_memory return all_memory = {} @@ -130,18 +130,21 @@ class MemoryProvider(MemoryProviderBase): with open(self.memory_path, "r", encoding="utf-8") as f: all_memory = yaml.safe_load(f) or {} if self.role_id in all_memory: - self.short_momery = all_memory[self.role_id] + self.short_memory = all_memory[self.role_id] def save_memory_to_file(self): all_memory = {} if os.path.exists(self.memory_path): with open(self.memory_path, "r", encoding="utf-8") as f: all_memory = yaml.safe_load(f) or {} - all_memory[self.role_id] = self.short_momery + all_memory[self.role_id] = self.short_memory with open(self.memory_path, "w", encoding="utf-8") as f: yaml.dump(all_memory, f, allow_unicode=True) async def save_memory(self, msgs): + # 打印使用的模型信息 + model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__)) + logger.bind(tag=TAG).debug(f"使用记忆保存模型: {model_info}") if self.llm is None: logger.bind(tag=TAG).error("LLM is not set for memory provider") return None @@ -155,9 +158,9 @@ class MemoryProvider(MemoryProviderBase): msgStr += f"User: {msg.content}\n" elif msg.role == "assistant": msgStr += f"Assistant: {msg.content}\n" - if self.short_momery and len(self.short_momery) > 0: + if self.short_memory and len(self.short_memory) > 0: msgStr += "历史记忆:\n" - msgStr += self.short_momery + msgStr += self.short_memory # 当前时间 time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) @@ -168,7 +171,7 @@ class MemoryProvider(MemoryProviderBase): json_str = extract_json_data(result) try: json.loads(json_str) # 检查json格式是否正确 - self.short_momery = json_str + self.short_memory = json_str self.save_memory_to_file() except Exception as e: print("Error:", e) @@ -179,7 +182,7 @@ class MemoryProvider(MemoryProviderBase): save_mem_local_short(self.role_id, result) logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}") - return self.short_momery + return self.short_memory async def query_memory(self, query: str) -> str: - return self.short_momery + return self.short_memory