update: 记忆模块使用独立LLM openai增加超参

This commit is contained in:
Sakura-RanChen
2025-05-22 17:14:58 +08:00
parent d934e0a69d
commit 02cb9c35b3
6 changed files with 75 additions and 23 deletions
+4
View File
@@ -220,6 +220,10 @@ Memory:
mem_local_short: mem_local_short:
# 本地记忆功能,通过selected_module的llm总结,数据保存在本地,不会上传到服务器 # 本地记忆功能,通过selected_module的llm总结,数据保存在本地,不会上传到服务器
type: mem_local_short type: mem_local_short
# 配备记忆存储独立的思考模型
# 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型
# 如果你的不想使用selected_module.LLM记忆存储,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM
llm: ChatGLMLLM
ASR: ASR:
FunASR: FunASR:
+28
View File
@@ -452,6 +452,34 @@ class ConnectionHandler:
save_to_file=not self.read_config_from_api, 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): def _initialize_intent(self):
self.intent_type = self.config["Intent"][ self.intent_type = self.config["Intent"][
self.config["selected_module"]["Intent"] self.config["selected_module"]["Intent"]
@@ -10,7 +10,7 @@ class LLMProviderBase(ABC):
"""LLM response generator""" """LLM response generator"""
pass pass
def response_no_stream(self, system_prompt, user_prompt): def response_no_stream(self, system_prompt, user_prompt, **kwargs):
try: try:
# 构造对话格式 # 构造对话格式
dialogue = [ dialogue = [
@@ -18,7 +18,7 @@ class LLMProviderBase(ABC):
{"role": "user", "content": user_prompt} {"role": "user", "content": user_prompt}
] ]
result = "" result = ""
for part in self.response("", dialogue): for part in self.response("", dialogue, **kwargs):
result += part result += part
return result return result
@@ -30,7 +30,7 @@ class LLMProviderBase(ABC):
""" """
Default implementation for function calling (streaming) Default implementation for function calling (streaming)
This should be overridden by providers that support function calls This should be overridden by providers that support function calls
Returns: generator that yields either text tokens or a special function call token Returns: generator that yields either text tokens or a special function call token
""" """
# For providers that don't support functions, just return regular response # For providers that don't support functions, just return regular response
@@ -16,26 +16,37 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url") self.base_url = config.get("base_url")
else: else:
self.base_url = config.get("url") self.base_url = config.get("url")
max_tokens = config.get("max_tokens")
if max_tokens is None or max_tokens == "":
max_tokens = 500
try: param_defaults = {
max_tokens = int(max_tokens) "max_tokens": (500, int),
except (ValueError, TypeError): "temperature": (0.7, lambda x: round(float(x), 1)),
max_tokens = 500 "top_p": (1.0, lambda x: round(float(x), 1)),
self.max_tokens = max_tokens "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) check_model_key("LLM", self.api_key)
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) 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: try:
responses = self.client.chat.completions.create( responses = self.client.chat.completions.create(
model=self.model_name, model=self.model_name,
messages=dialogue, messages=dialogue,
stream=True, 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 is_active = True
@@ -9,7 +9,13 @@ class MemoryProviderBase(ABC):
def __init__(self, config): def __init__(self, config):
self.config = config self.config = config
self.role_id = None 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 @abstractmethod
async def save_memory(self, msgs): async def save_memory(self, msgs):
@@ -107,7 +107,7 @@ TAG = __name__
class MemoryProvider(MemoryProviderBase): class MemoryProvider(MemoryProviderBase):
def __init__(self, config, summary_memory): def __init__(self, config, summary_memory):
super().__init__(config) super().__init__(config)
self.short_momery = "" self.short_memory = ""
self.save_to_file = True self.save_to_file = True
self.memory_path = get_project_dir() + "data/.memory.yaml" self.memory_path = get_project_dir() + "data/.memory.yaml"
self.load_memory(summary_memory) self.load_memory(summary_memory)
@@ -122,7 +122,7 @@ class MemoryProvider(MemoryProviderBase):
def load_memory(self, summary_memory): def load_memory(self, summary_memory):
# api获取到总结记忆后直接返回 # api获取到总结记忆后直接返回
if summary_memory or not self.save_to_file: if summary_memory or not self.save_to_file:
self.short_momery = summary_memory self.short_memory = summary_memory
return return
all_memory = {} all_memory = {}
@@ -130,18 +130,21 @@ class MemoryProvider(MemoryProviderBase):
with open(self.memory_path, "r", encoding="utf-8") as f: with open(self.memory_path, "r", encoding="utf-8") as f:
all_memory = yaml.safe_load(f) or {} all_memory = yaml.safe_load(f) or {}
if self.role_id in all_memory: 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): def save_memory_to_file(self):
all_memory = {} all_memory = {}
if os.path.exists(self.memory_path): if os.path.exists(self.memory_path):
with open(self.memory_path, "r", encoding="utf-8") as f: with open(self.memory_path, "r", encoding="utf-8") as f:
all_memory = yaml.safe_load(f) or {} 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: with open(self.memory_path, "w", encoding="utf-8") as f:
yaml.dump(all_memory, f, allow_unicode=True) yaml.dump(all_memory, f, allow_unicode=True)
async def save_memory(self, msgs): 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: if self.llm is None:
logger.bind(tag=TAG).error("LLM is not set for memory provider") logger.bind(tag=TAG).error("LLM is not set for memory provider")
return None return None
@@ -155,9 +158,9 @@ class MemoryProvider(MemoryProviderBase):
msgStr += f"User: {msg.content}\n" msgStr += f"User: {msg.content}\n"
elif msg.role == "assistant": elif msg.role == "assistant":
msgStr += f"Assistant: {msg.content}\n" 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 += "历史记忆:\n"
msgStr += self.short_momery msgStr += self.short_memory
# 当前时间 # 当前时间
time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) 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) json_str = extract_json_data(result)
try: try:
json.loads(json_str) # 检查json格式是否正确 json.loads(json_str) # 检查json格式是否正确
self.short_momery = json_str self.short_memory = json_str
self.save_memory_to_file() self.save_memory_to_file()
except Exception as e: except Exception as e:
print("Error:", e) print("Error:", e)
@@ -179,7 +182,7 @@ class MemoryProvider(MemoryProviderBase):
save_mem_local_short(self.role_id, result) save_mem_local_short(self.role_id, result)
logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}") 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: async def query_memory(self, query: str) -> str:
return self.short_momery return self.short_memory