From d590f98abd2dfb40a26b14bcdeb724b8b6f7e26c Mon Sep 17 00:00:00 2001 From: hs Date: Tue, 2 Jul 2024 12:30:38 +0800 Subject: [PATCH] summay_long --- memory_scope/constants/common_constants.py | 14 -- memory_scope/constants/language_constants.py | 53 ++-- .../worker/summary/get_insight_prompt.yaml | 11 + .../worker/summary/get_insight_worker.py | 160 ++++++++++++ .../worker/summary/get_reflection_prompt.yaml | 9 + .../worker/summary/get_reflection_worker.py | 95 +++++++ .../summary/long_contra_repeat_prompt.yaml | 12 + .../summary/long_contra_repeat_worker.py | 133 ++++++++++ .../worker/summary/summary_collect_worker.py | 46 ++++ .../worker/summary/update_insight_prompt.yaml | 67 +++++ .../worker/summary/update_insight_worker.py | 173 +++++++++++++ .../worker/summary/update_profile_prompt.yaml | 123 +++++++++ .../worker/summary/update_profile_worker.py | 236 ++++++++++++++++++ .../worker/write/get_observation_worker.py | 2 +- .../worker/write/store_memory_worker.py | 2 +- memory_scope/scheme/memory_node.py | 6 +- 16 files changed, 1101 insertions(+), 41 deletions(-) create mode 100644 memory_scope/memory/worker/summary/get_insight_prompt.yaml create mode 100644 memory_scope/memory/worker/summary/get_insight_worker.py create mode 100644 memory_scope/memory/worker/summary/get_reflection_prompt.yaml create mode 100644 memory_scope/memory/worker/summary/get_reflection_worker.py create mode 100644 memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml create mode 100644 memory_scope/memory/worker/summary/long_contra_repeat_worker.py create mode 100644 memory_scope/memory/worker/summary/summary_collect_worker.py create mode 100644 memory_scope/memory/worker/summary/update_insight_prompt.yaml create mode 100644 memory_scope/memory/worker/summary/update_insight_worker.py create mode 100644 memory_scope/memory/worker/summary/update_profile_prompt.yaml create mode 100644 memory_scope/memory/worker/summary/update_profile_worker.py diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index a0f4e9e4..b14d0520 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -12,11 +12,6 @@ RETRIEVE_MEMORY_NODES = "retrieve_memory_nodes" RANKED_MEMORY_NODES = "ranked_memory_nodes" - - - - - PIPELINE = "pipeline" WORKER = "worker" @@ -58,8 +53,6 @@ INSIGHT_KEY = "insight_key" INSIGHT_VALUE = "insight_value" -DT = "dt" - MSG_TIME = "msg_time" NEW = "new" @@ -68,8 +61,6 @@ TIME_INFER = "time_infer" KEY_WORD = "key_word" -REFLECTED = "reflected" - NEW_USER_PROFILE = "new_user_profile" RECALL_TYPE = "recall_type" @@ -82,9 +73,6 @@ TIME_MATCHED = "time_matched" QUERY_KEYWORDS = "query_keywords" - - - TIME_FORMAT_V1 = "{year}年{month}月{day}日{weekday}{hour}点" DATATIME_KEY_MAP = { @@ -94,5 +82,3 @@ DATATIME_KEY_MAP = { "周": "week", "星期几": "weekday", } - -CONTENT_MODIFIED = "content_modified" \ No newline at end of file diff --git a/memory_scope/constants/language_constants.py b/memory_scope/constants/language_constants.py index c91514af..caa5aa9a 100644 --- a/memory_scope/constants/language_constants.py +++ b/memory_scope/constants/language_constants.py @@ -1,30 +1,29 @@ from memory_scope.enumeration.language_enum import LanguageEnum DATATIME_WORD_LIST = { - LanguageEnum.CN: - [ - "天", - "周", - "月", - "年", - "星期", - "点", - "分钟", - "小时", - "秒", - "上午", - "下午", - "早上", - "早晨", - "晚上", - "中午", - "日", - "夜", - "清晨", - "傍晚", - "凌晨", - "岁", - ], + LanguageEnum.CN: [ + "天", + "周", + "月", + "年", + "星期", + "点", + "分钟", + "小时", + "秒", + "上午", + "下午", + "早上", + "早晨", + "晚上", + "中午", + "日", + "夜", + "清晨", + "傍晚", + "凌晨", + "岁", + ], LanguageEnum.EN: [ ] @@ -69,3 +68,9 @@ COLON_WORD = { LanguageEnum.CN: ":", LanguageEnum.EN: ":" } + + +COMMA_WORD = { + LanguageEnum.CN: ",", + LanguageEnum.EN: "," +} diff --git a/memory_scope/memory/worker/summary/get_insight_prompt.yaml b/memory_scope/memory/worker/summary/get_insight_prompt.yaml new file mode 100644 index 00000000..419057c9 --- /dev/null +++ b/memory_scope/memory/worker/summary/get_insight_prompt.yaml @@ -0,0 +1,11 @@ +system_prompt: + cn: + +few_shot_prompt: + cn: + +user_query_prompt: + cn: + +content_format: + cn: "用户的{insight_key}:{insight_value}" diff --git a/memory_scope/memory/worker/summary/get_insight_worker.py b/memory_scope/memory/worker/summary/get_insight_worker.py new file mode 100644 index 00000000..ada7ad81 --- /dev/null +++ b/memory_scope/memory/worker/summary/get_insight_worker.py @@ -0,0 +1,160 @@ +from datetime import datetime +from typing import List + +from memory_scope.utils.tool_functions import time_to_formatted_str, get_datetime_info_dict, prompt_to_msg +from memory_scope.constants.common_constants import ( + NEW_INSIGHT_NODES, + DT, + NOT_REFLECTED_MERGE_NODES, + NEW_INSIGHT_KEYS, + INSIGHT_KEY, + INSIGHT_VALUE, +) +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.constants.language_constants import NONE_WORD + + +class GetInsightWorker(MemoryBaseWorker): + def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode: + created_dt = datetime.now() + obs_dt = time_to_formatted_str(time=created_dt) + + # 组合meta_data + meta_data = { + INSIGHT_KEY: insight_key, + INSIGHT_VALUE: insight_value, + } + meta_data.update( + {k: str(v) for k, v in get_datetime_info_dict(created_dt).items()} + ) + + content = self.prompt_handler.content_format.format(insight_key=insight_key, insight_value=insight_value) + return MemoryNode( + content=content, + user_name=self.user_name, + target_name=self.target_name, + memory_type=MemoryTypeEnum.INSIGHT.value, + meta_data=meta_data, + status=MemoryNodeStatus.ACTIVE.value, + obs_dt=obs_dt, + obs_updated=True, + ) + + def reflect_new_insight_key( + self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode] + ) -> MemoryNode | None: + + # 检索历史memory + related_nodes = self.vector_store.similar_search( + text=insight_key, + size=self.es_insight_similar_top_k, + exact_filters={ + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [ + MemoryTypeEnum.OBSERVATION.value, + MemoryTypeEnum.OBS_CUSTOMIZED.value, + ], + }, + ) + + # 合并新增nodes + related_nodes.extend(not_reflected_merge_nodes) + + # content去重 + related_node_dict = {n.memory_node.content: n for n in related_nodes} + related_nodes = sorted( + list(related_node_dict.values()), key=lambda x: x.memory_node.memory_id # memory_id or use_id? + ) + documents = [n.memory_node.content for n in related_nodes] + + # 重排所有记忆 + result = self.rank_model.call(query=insight_key, documents=documents) + if not result: + self.add_run_info( + f"reflect insight_key={insight_key} call rerank client failed!" + ) + return + + # 根据打分过滤 + for index, score in result.rank_scores.items(): + related_nodes[index].score_rank = score + + related_nodes_sorted = sorted( + related_nodes, key=lambda x: x.score_rank, reverse=True + )[: self.insight_obs_max_cnt] + + # prepare prompt + user_query_list = [x.memory_node.content for x in related_nodes_sorted] + get_insight_message = prompt_to_msg( + system_prompt=self.prompt_handler.system_prompt, + few_shot=self.prompt_handler.few_shot_prompt, + user_query=self.prompt_handler.user_query_prompt.format( + insight_key=insight_key, user_query="\n".join(user_query_list) + ) + ) + + self.logger.info(f"get_insight_message={get_insight_message}") + + # call LLM, 提取insight + response = self.generation_model.call( + messages=get_insight_message, + model_name=self.get_insight_model, + max_token=self.get_insight_max_token, + temperature=self.get_insight_temperature, + top_k=self.get_insight_top_k, + ) + # return if empty + if not response: + self.add_run_info("reflect_upon_user_attr call llm failed!") + return + response_text = response.message.content.strip() + if response_text in [self.get_language_value(NONE_WORD)]: + return + return self.new_insight_node( + insight_key=insight_key, insight_value=response_text + ) + + def _run(self): + new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS) + if not new_insight_keys: + self.add_run_info("new_insight_keys is empty! stop insight.") + return + + not_reflected_merge_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_MERGE_NODES + ) + if not not_reflected_merge_nodes: + self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.") + return + + # submit insight task + for insight_key in new_insight_keys: + self.submit_thread( + self.reflect_new_insight_key, + sleep_time=1, + insight_key=insight_key, + not_reflected_merge_nodes=not_reflected_merge_nodes, + ) + + # save output + new_insight_nodes: List[MemoryNode] = [] + for result in self.join_threads(): + if result: + new_insight_nodes.append(result) + assert isinstance(result, MemoryNode) + insight_key = result.meta_data.get(INSIGHT_KEY, "") + insight_value = result.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"after_get_insight insight_key={insight_key} insight_value={insight_value}" + ) + + self.set_context(NEW_INSIGHT_NODES, new_insight_nodes) + + # set REFLECTED + for node in not_reflected_merge_nodes: + node.obs_reflected = True diff --git a/memory_scope/memory/worker/summary/get_reflection_prompt.yaml b/memory_scope/memory/worker/summary/get_reflection_prompt.yaml new file mode 100644 index 00000000..a5763f19 --- /dev/null +++ b/memory_scope/memory/worker/summary/get_reflection_prompt.yaml @@ -0,0 +1,9 @@ +system_prompt: + cn: + +few_shot_prompt: + cn: + +user_query_prompt: + cn: + \ No newline at end of file diff --git a/memory_scope/memory/worker/summary/get_reflection_worker.py b/memory_scope/memory/worker/summary/get_reflection_worker.py new file mode 100644 index 00000000..5b14eac0 --- /dev/null +++ b/memory_scope/memory/worker/summary/get_reflection_worker.py @@ -0,0 +1,95 @@ +from typing import List + +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.constants.common_constants import ( + NEW_OBS_NODES, + NOT_REFLECTED_OBS_NODES, + INSIGHT_NODES, + INSIGHT_KEY, + NEW_INSIGHT_KEYS, + NOT_REFLECTED_MERGE_NODES, +) +from memory_scope.constants.language_constants import COLON_WORD, COMMA_WORD +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.utils.tool_functions import prompt_to_msg + + +class GetReflectionWorker(MemoryBaseWorker): + def _run(self): + # 过滤得到 not_reflected_merge_nodes + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_OBS_NODES + ) + not_reflected_merge_nodes: List[MemoryNode] = [] + if new_obs_nodes: + not_reflected_merge_nodes.extend(new_obs_nodes) + if not_reflected_nodes: + not_reflected_merge_nodes.extend(not_reflected_nodes) + not_reflected_merge_nodes = [ + node + for node in not_reflected_merge_nodes + if not node.obs_reflected + ] + + # count + not_reflected_count = len(not_reflected_merge_nodes) + if not_reflected_count <= self.reflect_obs_cnt_threshold: + self.logger.info( + f"not_reflected_count={not_reflected_count} is not enough, stop reflect." + ) + return + + # save context + self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes) + + # get profile_keys + exist_keys: List[str] = [] + profile_keys: List[str] = list(self.user_profile_dict.keys()) + exist_keys.extend(profile_keys) + self.logger.info(f"profile_keys={profile_keys}") + + # get insight_keys + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + if insight_nodes: + insight_keys = [ + n.meta_data.get(INSIGHT_KEY) for n in insight_nodes + ] + insight_keys = [x.strip() for x in insight_keys if x] + exist_keys.extend(insight_keys) + self.logger.info(f"insight_keys={insight_keys}") + + # gen reflect prompt + user_query_list = [n.content for n in not_reflected_merge_nodes] + reflect_message = prompt_to_msg( + system_prompt=self.prompt_handler.system_prompt.format( + num_questions=self.reflect_num_questions + ), + few_shot=self.prompt_handler.few_shot_prompt, + user_query=self.prompt_handler.user_query_prompt.format( + exist_keys=self.get_language_value(COMMA_WORD).join(exist_keys), user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"reflect_message={reflect_message}") + + # # call LLM + response = self.generation_model.call( + messages=reflect_message, + model_name=self.reflect_obs_model, + max_token=self.reflect_obs_max_token, + temperature=self.reflect_obs_temperature, + top_k=self.reflect_obs_top_k, + ) + + # return if empty + if not response: + self.add_run_info("reflect_obs_questions call llm failed!") + return + + # parse text & save + new_insight_keys = ResponseTextParser(response.message.content).parse_v2( + "get_insight_keys" + ) + if new_insight_keys: + self.set_context(NEW_INSIGHT_KEYS, new_insight_keys) diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml b/memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml new file mode 100644 index 00000000..82a85de3 --- /dev/null +++ b/memory_scope/memory/worker/summary/long_contra_repeat_prompt.yaml @@ -0,0 +1,12 @@ +systemp_prompt: + cn: + en: + +few_shot_prompt: + cn: + en: + +user_query_prompt: + cn: + en: + diff --git a/memory_scope/memory/worker/summary/long_contra_repeat_worker.py b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py new file mode 100644 index 00000000..c99217ce --- /dev/null +++ b/memory_scope/memory/worker/summary/long_contra_repeat_worker.py @@ -0,0 +1,133 @@ +from typing import List + +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.constants.common_constants import ( + NEW_OBS_NODES, + MSG_TIME, + MODIFIED_MEMORIES, +) +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.utils.tool_functions import prompt_to_msg +from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD + + +class LongContraRepeatWorker(MemoryBaseWorker): + + def _run(self): + # 合并当前的obs和今日的obs + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + all_obs_nodes: List[MemoryNode] = [] + for new_obs_node in new_obs_nodes: + text = new_obs_node.content + related_nodes = self.vector_store.similar_search( + text=text, + size=self.es_contra_repeat_similar_top_k, + exact_filters={ + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [ + MemoryTypeEnum.OBSERVATION.value, + MemoryTypeEnum.OBS_CUSTOMIZED.value, + ], + }, + ) + + has_match = False + for related_node in related_nodes: + if related_node.score_similar < self.long_contra_repeat_threshold: + continue + else: + has_match = True + all_obs_nodes.append(related_node) + if has_match: + all_obs_nodes.append(new_obs_node) + + if not all_obs_nodes: + self.add_run_info("all_obs_nodes is empty!") + return + + # gene prompt + user_query_list = [] + all_obs_nodes = sorted( + all_obs_nodes, + key=lambda x: x.meta_data.get(MSG_TIME, ""), + reverse=True, + ) + for i, n in enumerate(all_obs_nodes): + user_query_list.append(f"{i + 1} {n.content}") + merge_obs_message = prompt_to_msg( + system_prompt=self.prompt_handler.system_prompt.format( + num_obs=len(user_query_list) + ), + few_shot=self.prompt_handler.few_shot_prompt, + user_query=self.prompt_handler.user_query_prompt.format( + user_query="\n".join(user_query_list) + ), + ) + self.logger.info(f"merge_obs_message={merge_obs_message}") + + # call LLM + response = self.generation_model.call( + messages=merge_obs_message, + model_name=self.merge_obs_model, + max_token=self.merge_obs_max_token, + temperature=self.merge_obs_temperature, + top_k=self.merge_obs_top_k, + ) + + # return if empty + if not response: + self.add_run_info("contra repeat call llm failed!") + return + + # parse text + idx_merge_obs_list = ResponseTextParser(response.message.content).parse_v1("contra_repeat") + if len(idx_merge_obs_list) <= 0: + self.add_run_info("idx_merge_obs_list is empty!") + return + + # add merged obs + merge_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_merge_obs_list: + if not obs_content_list: + continue + + # [6, 逃课] + if len(obs_content_list) != 2: + self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") + continue + + idx, keep_flag = obs_content_list + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + # 序号需要修正-1 + idx = int(idx) - 1 + if idx >= len(all_obs_nodes): + self.logger.warning(f"idx={idx} is invalid!") + continue + + long_contra_repeat_keep_flag = [ + self.get_language_value(CONTRADICTORY_WORD), + self.get_language_value(INCLUDED_WORD), + self.get_language_value(NONE_WORD), + ] + + if keep_flag not in long_contra_repeat_keep_flag.values(): + self.logger.warning(f"keep_flag={keep_flag} is invalid!") + continue + + node: MemoryNode = all_obs_nodes[idx] + if keep_flag != self.get_language_value(NONE_WORD): + node.status = MemoryNodeStatus.EXPIRED.value + merge_obs_nodes.append(node) + self.logger.info(f"after contra repeat: {node.content} {node.status}") + + # save context + self.set_context(MODIFIED_MEMORIES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/summary/summary_collect_worker.py b/memory_scope/memory/worker/summary/summary_collect_worker.py new file mode 100644 index 00000000..a77644ea --- /dev/null +++ b/memory_scope/memory/worker/summary/summary_collect_worker.py @@ -0,0 +1,46 @@ +from typing import List, Dict + +from memory_scope.constants.common_constants import ( + NEW_INSIGHT_NODES, + MODIFIED_MEMORIES, + INSIGHT_NODES, + NEW_OBS_NODES, + NOT_REFLECTED_OBS_NODES, + NEW, + NOT_REFLECTED_MERGE_NODES, +) +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker + + +class SummaryCollectWorker(MemoryBaseWorker): + + def _run(self): + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES) + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + not_reflected_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_OBS_NODES + ) + not_reflected_merge_nodes: List[MemoryNode] = self.get_context( + NOT_REFLECTED_MERGE_NODES + ) + + # 合并逻辑,复杂,务必check + all_node_dict: Dict[str, MemoryNode] = {} + if insight_nodes: + all_node_dict.update( + {n.id: n for n in insight_nodes if n.obs_updated} + ) + if new_insight_nodes: + all_node_dict.update({n.content: n for n in new_insight_nodes}) + if new_obs_nodes: + # 设置为非新 + for n in new_obs_nodes: + n.obs_updated = "0" + all_node_dict.update({n.content: n for n in new_obs_nodes}) + if not_reflected_merge_nodes and not_reflected_nodes: + # 进入reflect阶段 + all_node_dict.update({n.id: n for n in not_reflected_nodes}) + + self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values())) diff --git a/memory_scope/memory/worker/summary/update_insight_prompt.yaml b/memory_scope/memory/worker/summary/update_insight_prompt.yaml new file mode 100644 index 00000000..f5d6e362 --- /dev/null +++ b/memory_scope/memory/worker/summary/update_insight_prompt.yaml @@ -0,0 +1,67 @@ +system_prompt: + cn: "从下面的句子中提取出给定类别的用户资料信息,并判断与已有信息是否矛盾,若矛盾以新信息为准。整合已有信息和新信息并输出。若已有信息为空则直接输出提取的 +信息。若无法提取出给定类别的用户资料信息则输出“无”。 +请一步步思考,并按如下格式输出: +思考: 思考的依据和过程,150字以内。 +用户资料: <信息或无>, 一定加<>" + en: + +few_shot_promp: + cn: " +示例1: +句子:因为昨天成都下大雨,用户全身都被淋湿了。 +句子:用户关心明天成都的天气预报。 +类别:用户所在地区 +已有信息:用户所在地区: +思考:从第一句句子可以得出用户在成都。第二句句子没有直接透露用户所在地信息,但与第一句句子用户在成都的信息吻合。已有信息为空,直接输出得出的信息。 +用户资料:<成都> + +示例2: +句子:用户最近养好了肠胃。 +句子:用户关注中医养生。 +类别:用户健康状况 +已有信息:用户健康状况: 肠胃不好,高血压 +思考:从第一句句子可以得出用户最近养好了肠胃,与已有信息矛盾,以新信息为准。第二句句子与用户健康状况无关。整合已有信息和新信息得到用户健康状况是肠胃健康,高血压。 +用户资料:<肠胃健康,高血压> + +示例3: +句子:用户刚刚毕业,第一份工作是银行前台。 +句子:用户的理想工作是职业游戏选手。 +类别:用户职业 +已有信息:用户职业:在招商银行工作 +思考:整合已有信息和第一句句子的信息可以得出用户的现在的职业是招商银行前台。第二句句子说明了用户的理想工作但并不是现在的职业。 +用户资料:<招商银行前台> + +示例4: +句子:用户大学期间接触过优化算法的研究。 +类别:用户学习专业 +已有信息:用户学习专业:与人工智能相关 +思考:从句子可以得出用户大学学习的专业与优化算法相关,这与已有信息(用户学习专业与人工智能相关)不矛盾,整合可以得出用户大学学习的专业与人工智能和优化算法相关。 +用户资料:<与人工智能和优化算法相关> + +示例5: +句子:用户单身。 +句子:用户受到一名18岁男生的追求,但不想接受又不想伤害他。 +句子:用户喜欢成熟且情绪稳定的男生。 +类别:用户情感状况 +已有信息:用户情感状况:有男朋友 +思考:从第一句句子可以得出用户现在单身,与已有信息矛盾,以新信息为准。从第二句句子得出用户受到一名18岁男生的追求但并不喜欢他。第三句话表达了用户理想的伴侣类型但与用户 +情感状况无关。整合得出用户情感状况为单身,受到一名18岁男生的追求但并不喜欢他。 +用户资料:<单身,受到一名18岁男生的追求但并不喜欢他。> +" + en: + +user_query_prompt: + cn: " +{user_query} +类别:{insight_key} +已有信息:{insight_key_value} +" + en: + +user_query: + cn: "句子:{content}" + +insight_value: + cn: ["无", "重复"] + en: \ No newline at end of file diff --git a/memory_scope/memory/worker/summary/update_insight_worker.py b/memory_scope/memory/worker/summary/update_insight_worker.py new file mode 100644 index 00000000..756d75c9 --- /dev/null +++ b/memory_scope/memory/worker/summary/update_insight_worker.py @@ -0,0 +1,173 @@ +from typing import List + +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.constants.common_constants import ( + INSIGHT_NODES, + NEW_OBS_NODES, + INSIGHT_KEY, + INSIGHT_VALUE, +) +from memory_scope.utils.tool_functions import prompt_to_msg +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.constants.language_constants import COMMA_WORD, COLON_WORD + + +class UpdateInsightWorker(MemoryBaseWorker): + + def filter_obs_nodes( + self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode] + ) -> (MemoryNode, List[MemoryNode], float): + max_score: float = 0 + filtered_nodes: List[MemoryNode] = [] + + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + if not insight_key or not insight_value: + self.logger.warning( + f"insight_key={insight_key} insight_value={insight_value} is empty!" + ) + return insight_node, filtered_nodes, max_score + + result = self.rank_model.call( + query=insight_key, documents=[x.content for x in new_obs_nodes] + ) + + if not result: + self.add_run_info(f"update_insight={insight_key} call rerank failed!") + return insight_node, filtered_nodes, max_score + + # 找到大于阈值的obs node + + for index, score in result.rank_scores.items(): + node = new_obs_nodes[index] + keep_flag = "filtered" + if score >= self.update_insight_threshold: + filtered_nodes.append(node) + keep_flag = "keep" + max_score = max(max_score, score) + self.logger.info( + f"insight_key={insight_key} insight_value={insight_value} " + f"score={score} keep_flag={keep_flag}" + ) + + if not filtered_nodes: + self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!") + + return insight_node, filtered_nodes, max_score + + def update_insight( + self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode] + ) -> MemoryNode: + + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"update_insight insight_key={insight_key} insight_value={insight_value} " + f"doc.size={len(filtered_nodes)}" + ) + + # gen prompt + user_query_list = [] + for node in filtered_nodes: + user_query_list.append(self.prompt_handler.user_query.format(content=node.content)) + update_insight_message = prompt_to_msg( + system_prompt=self.prompt_handler.system_prompt, + few_shot=self.prompt_handler.few_shot_prompt, + user_query=self.prompt_handler.user_query_prompt.format( + user_query="\n".join(user_query_list), + insight_key=insight_key, + insight_key_value=insight_key + self.get_language_value(COLON_WORD) + insight_value, + ), + ) + self.logger.info(f"update_insight_message={update_insight_message}") + + # call LLM + response: str = self.generation_model.call( + messages=update_insight_message, + model_name=self.update_insight_model, + max_token=self.update_insight_max_token, + temperature=self.update_insight_temperature, + top_k=self.update_insight_top_k, + ) + + # return if empty + if not response: + self.add_run_info( + f"update_insight insight_key={insight_key} call llm failed!" + ) + return insight_node + + profile_list = ResponseTextParser(response.message.content).parse_v1( + f"update_profile {insight_key}" + ) + if not profile_list: + self.add_run_info( + f"update_insight insight_key={insight_key} profile_list empty 1!" + ) + return insight_node + profile_list = profile_list[0] + if not profile_list: + self.add_run_info( + f"update_insight insight_key={insight_key} profile_list empty 2" + ) + return insight_node + insight_value = profile_list[0] + + if not insight_value or insight_value in self.prompt_handler.insight_value: + self.logger.info(f"insight_value={insight_value}, skip.") + return insight_node + + insight_node.meta_data[INSIGHT_VALUE] = insight_value + insight_node.obs_updated = True + return insight_node + + def _run(self): + # 获取新的obs和insight + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES) + if not new_obs_nodes: + self.logger.info("new_obs_nodes is empty, stop update sights!") + return + if not insight_nodes: + self.logger.info("insight_nodes is empty, stop update sights!") + return + + # 提交打分任务 + for node in insight_nodes: + self.submit_thread( + self.filter_obs_nodes, + sleep_time=0.1, + insight_node=node, + new_obs_nodes=new_obs_nodes, + ) + + # 选择topN + result_list = [] + for result in self.join_threads(): + insight_node, filtered_nodes, max_score = result + if not filtered_nodes: + continue + result_list.append(result) + result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True) + if len(result_sorted) > self.update_insight_max_thread: + result_sorted = result_sorted[: self.update_insight_max_thread] + + # 提交LLM update任务 + for insight_node, filtered_nodes, _ in result_sorted: + self.submit_thread( + self.update_insight, + sleep_time=1, + insight_node=insight_node, + filtered_nodes=filtered_nodes, + ) + + # 等待结果 + for result in self.join_threads(): + if result: + insight_node: MemoryNode = result + insight_key = insight_node.meta_data.get(INSIGHT_KEY, "") + insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "") + self.logger.info( + f"after_update_insight insight_key={insight_key} insight_value={insight_value}" + ) diff --git a/memory_scope/memory/worker/summary/update_profile_prompt.yaml b/memory_scope/memory/worker/summary/update_profile_prompt.yaml new file mode 100644 index 00000000..0c3f56ce --- /dev/null +++ b/memory_scope/memory/worker/summary/update_profile_prompt.yaml @@ -0,0 +1,123 @@ +update_plural_profile_system_prompt: + cn: " +从下面的句子中提取出给定类别的用户资料信息,并判断和已有信息是否重复。只输出无重复的新信息。若无法提取该类别的用户资料的新信息则回答无。 +请一步步思考,并按如下格式输出: +思考: 思考的依据和过程,150字以内。 +用户资料: <信息>或<无>, 一定加<> +" + en: + +update_plural_profile_few_shot_prompt: + cn: " +示例1: +句子:用户上周去了西溪游泳馆游泳,那个游泳馆人非常多。 +句子:用户计划每周六和朋友张三去朝阳体育馆打羽毛球。 +类别:运动(用户喜欢的运动) +已有信息:运动(用户喜欢的运动):游泳 +思考:从第一句句子可以得出游泳是用户喜欢的运动之一,但与已有信息重复。从第二句句子可以得出羽毛球是用户喜欢的运动之一,是新的信息。 +用户资料: <羽毛球> + +示例2: +句子:用户对咖啡因过敏。 +句子:用户不喜欢吃香菇。 +类别:过敏(用户的已知过敏反应) +已有信息:过敏(用户的已知过敏反应): 咖啡因 +思考:从第一句句子可以得出咖啡因是用户的已知过敏反应之一,但与已有信息重复。从第二句句子只能得出用户不喜欢香菇而非对香菇过敏,无法得出新的用户已知过敏信息。 +用户资料: <无> + +示例3: +句子:用户热衷于动作类类游戏如只狼、艾尔登法环。 +句子:用户在休闲时间经常长时间玩策略类游戏如文明6。 +句子:用户是音乐发烧友,关注各个品牌的耳机的音质和性价比。 +类别:爱好(用户的业余爱好) +已有信息:爱好(用户的业余爱好): +思考:从第一句句子可以得出动作类游戏是用户的爱好之一,是新的信息。从第二句句子可以得出策略类游戏是用户的爱好之一,是新的信息。从第三句句子可以得出音乐是用户的爱好之一,是新的信息。 +用户资料: <动作类游戏, 策略类游戏, 音乐> + +示例4: +句子:关于职场沟通你有什么具体的建议吗?最好结合一个实例。我一直听人说要加强沟通,经常和上司沟通,同步项目的进展,但是我总是感觉还有许多事情要做。 +句子:项目并没有达到一个充分的可以汇报的状态,然后准备汇报材料又很费时间,导致有时候我没有及时和上司同步项目状态。针对这个情况你有什么建议? +类别:职业(用户的职业) +已有信息:职业(用户的职业):工程师 +思考:句子中虽然提及了职场沟通等工作相关内容,但是并不能推断出用户的职位是什么,只能推知与宽泛的项目实施与管理相关。 +用户资料: <无> +" + en: + +update_plural_profile_user_query_prompt: + cn: " +{user_query} +类别:{update_profile} +已有信息:{update_profile_value} +" + en: + +update_unique_profile_system_prompt: + cn: " +从下面的句子中提取出给定类别的用户资料信息,并判断与已有信息是否矛盾。若矛盾则输出更新的信息,若不矛盾则保留已有信息,整合已有信息和新信息并输出。 +请一步步思考,并按如下格式输出: +思考: 思考的依据和过程,150字以内。 +用户资料: <信息>, 一定加<> +" + en: + + +update_unique_profile_few_shot_prompt: + cn: " +示例1: +句子:因为昨天成都下大雨,用户全身都被淋湿了。 +句子:用户关心明天成都的天气预报。 +类别:地区(用户所在地区) +已有信息:地区(用户所在地区): 杭州 +思考:从第一句句子可以得出用户在成都。第二句句子没有直接透露用户所在地信息,但与第一句句子用户在成都的信息吻合。这与已有信息(用户在杭州)矛盾,输出更新的信息。 +用户资料:<成都> + +示例2: +句子:用户女朋友下个月过生日。 +句子:用户生日在7月15日。 +类别:生日(用户的生日)。 +已有信息:生日(用户的生日):1987年7月15日。 +思考:第一句句子中提及生日,但并不是用户的生日,无法得出用户生日信息。从第二句句子可以得出用户生日在7月15日,与已有信息不矛盾,整合可以得出用户生日是1987年7月15日。 +用户资料: <1987年7月15日> + +示例3: +句子:用户在招商银行工作。 +句子:用户刚刚毕业,第一份工作是银行前台。 +句子:用户的理想工作是职业游戏选手。 +类别:职业(用户的职业) +已有信息:职业(用户的职业): +思考:整合第一和第二句句子的信息可以得出用户的现在的职业是招商银行前台。第三句句子说明了用户的理想工作但并不是现在的职业。 +用户资料:<招商银行前台> + +示例4: +句子:用户大学期间接触过优化算法的研究。 +类别:学习专业 (用户大学学习的专业) +已有信息:学习专业 (用户大学学习的专业):与人工智能相关 +思考:从句子可以得出用户大学学习的专业与优化算法相关,这与已有信息(用户大学学习的专业与人工智能相关)不矛盾,整合可以得出用户大学学习的专业与人工智能和优化算法相关。 +用户资料:<与人工智能和优化算法相关> + +示例5: +句子:今天和同学去打球了。 +句子:明天和女朋友一起去杭州旅游。 +类别:学习专业 (用户大学学习的专业) +已有信息:学习专业 (用户大学学习的专业): +思考:两个句子和学习专业都没有关联,没有新提取的信息。 +用户资料:<无> +" + en: + +update_unique_profile_user_query_prompt: + cn: " +{user_query} +类别:{update_profile} +已有信息:{update_profile_value} +" + en: + + +user_query: + cn: "句子:{content}" + en: + +update_profile_key: + cn: ["无", "重复"] \ No newline at end of file diff --git a/memory_scope/memory/worker/summary/update_profile_worker.py b/memory_scope/memory/worker/summary/update_profile_worker.py new file mode 100644 index 00000000..17edd432 --- /dev/null +++ b/memory_scope/memory/worker/summary/update_profile_worker.py @@ -0,0 +1,236 @@ +from typing import List + +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.utils.global_context import GlobalContext +from memory_scope.utils.tool_functions import prompt_to_msg +from memory_scope.constants.language_constants import COMMA_WORD, COLON_WORD + + +class UpdateProfileWorker(MemoryBaseWorker): + @property + def extra_user_attrs(self): + return GlobalContext.global_configs.get("extra_user_attrs", []) + + def filter_obs_nodes( + self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode] + ) -> (MemoryNode, List[MemoryNode], float): + max_score: float = 0 + filtered_nodes: List[MemoryNode] = [] + result = self.rank_model.call( + query=user_attr.meta_data.get("description", ""), + documents=[x.content for x in new_obs_nodes], + ) + + if not result: + self.add_run_info( + f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!" + ) + return user_attr, filtered_nodes, max_score + + # 找到大于阈值的obs node + filtered_nodes: List[MemoryNode] = [] + for index, score in result.rank_scores.items(): + node = new_obs_nodes[index] + keep_flag = "filtered" + if score >= self.update_profile_threshold: + filtered_nodes.append(node) + keep_flag = "keep" + max_score = max(max_score, score) + self.logger.info( + f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} " + f"content={node.content} score={score} keep_flag={keep_flag}" + ) + + if not filtered_nodes: + self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!") + return user_attr, filtered_nodes, max_score + + def update_user_attr( + self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode] + ) -> MemoryNode: + self.logger.info( + f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} " + f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}" + ) + + # 根据不同的参数类型是否多值,分别给出prompt + user_query_list = [] + for node in filtered_nodes: + user_query_list.append(self.prompt_handler.user_query.format(content=node.content)) + update_profile = f"{user_attr.meta_data.get("memory_key", "")}({user_attr.meta_data.get("description", "")})" + update_profile_value = update_profile + self.get_language_prompt(COLON_WORD) + self.get_language_prompt(COMMA_WORD).join(user_attr.meta_data.get("value", "")) + + if user_attr.meta_data.get("is_unique", 0) == 1: + update_profile_message = prompt_to_msg( + system_prompt=self.prompt_handler.update_unique_profile_system_prompt, + few_shot=self.prompt_handler.update_unique_profile_few_shot_prompt, + user_query=self.prompt_handler.update_unique_profile_user_query_prompt.format( + user_query="\n".join(user_query_list), + update_profile=update_profile, + update_profile_value=update_profile_value, + ), + ) + else: + update_profile_message = prompt_to_msg( + system_prompt=self.prompt_handler.update_plural_profile_system_prompt, + few_shot=self.prompt_handler.update_plural_profile_few_shot_prompt, + user_query=self.prompt_handler.update_plural_profile_user_query_prompt.format( + user_query="\n".join(user_query_list), + update_profile=update_profile, + update_profile_value=update_profile_value, + ), + ) + self.logger.info(f"update_profile_message={update_profile_message}") + + # call LLM + response: str = self.generation_model.call( + messages=update_profile_message, + model_name=self.update_profile_model, + max_token=self.update_profile_max_token, + temperature=self.update_profile_temperature, + top_k=self.update_profile_top_k, + ) + + # return if empty + if not response: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!" + ) + return user_attr + + profile_list = ResponseTextParser(response.message.content).parse_v1( + f"update_attr {user_attr.meta_data.get("memory_key", "")}" + ) + if not profile_list: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!" + ) + return user_attr + profile_list = profile_list[0] + if not profile_list: + self.add_run_info( + f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2" + ) + return user_attr + profile = profile_list[0] + + if not profile or profile in self.prompt_handler.update_profile_key: + self.logger.info(f"profile={profile}, skip.") + return user_attr + + # check 英文中午逗号 + if user_attr.meta_data.get("is_unique", 0) == 1: + user_attr.meta_data["value"] = [profile.strip()] + else: + attr_value_list = profile.replace(",", ",").split(",") + user_attr.meta_data["value"] = [ + x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list))) + ] + return user_attr + + def add_extra_user_attrs(self): + # 解析为空返回 + extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()] + if not extra_user_attr_list: + return + + for user_attr_info in extra_user_attr_list: + user_attr_split = user_attr_info.split(self.get_language_prompt(COLON_WORD)) + + # 格式不对返回 + if len(user_attr_split) < 1: + continue + user_attr_key = user_attr_split[0] + + user_attr_desc = "" + if len(user_attr_split) >= 2: + user_attr_desc = user_attr_split[1] + + user_attr_unique = 0 + if len(user_attr_split) >= 3: + user_attr_unique = int(user_attr_split[2]) + + # 已经包含返回 + if user_attr_key in self.user_profile_dict: + user_attr = self.user_profile_dict[user_attr_key] + # description为空,补充description + if not user_attr.meta_data.get("description", ""): + user_attr.meta_data["description"] = user_attr_desc + continue + + # 增加新属性 + new_attr = MemoryNode( + memory_id=self.memory_id, + meta_data={ + "memory_key": user_attr_key, + "is_unique": int(user_attr_unique), + "is_mutable": 1, + "description": user_attr_desc + }, + memory_type=MemoryTypeEnum.PROFILE, + status=1, + obs_profile_updated=true, + ) + self.user_profile_dict[user_attr_key] = new_attr + + def _run(self): + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) + if not new_obs_nodes: + self.logger.info("new_obs_nodes is empty, stop user profile!") + self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values())) + return + + # 增加环境变量配置的属性 + if self.extra_user_attrs: + self.add_extra_user_attrs() + + new_user_profile: List[MemoryNode] = [] + self.set_context(NEW_USER_PROFILE, new_user_profile) + + for user_attr_key, user_attr in self.user_profile_dict.items(): + # 不可修改直接跳过 + if user_attr.meta_data.get("is_mutable", 0) != 1: + new_user_profile.append(user_attr) + self.logger.info(f"{user_attr_key} is not mutable! continue") + continue + + self.submit_thread( + self.filter_obs_nodes, + sleep_time=0.1, + user_attr=user_attr, + new_obs_nodes=new_obs_nodes, + ) + + # 选择topN + result_list = [] + for result in self.join_threads(): + user_attr, filtered_nodes, max_score = result + if not filtered_nodes: + continue + result_list.append(result) + result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True) + if len(result_sorted) > self.update_profile_max_thread: + result_sorted = result_sorted[: self.update_profile_max_thread] + + # 提交LLM update任务 + for user_attr, filtered_nodes, _ in result_sorted: + self.submit_thread( + self.update_user_attr, + sleep_time=1, + user_attr=user_attr, + filtered_nodes=filtered_nodes, + ) + + # collect result & save + for result in self.join_threads(): + if result: + user_attribute: MemoryNode = result + self.logger.info( + f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} " + f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}" + ) + new_user_profile.append(user_attribute) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 9f3eae19..d42dcf11 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -36,7 +36,7 @@ class GetObservationWorker(MemoryBaseWorker): timestamp=message.time_created, obs_dt=dt_handler.datetime_format(), obs_reflected=False, - obs_profile_updated=False, + obs_updated=False, obs_keyword=keywords) node.gen_memory_id() return node diff --git a/memory_scope/memory/worker/write/store_memory_worker.py b/memory_scope/memory/worker/write/store_memory_worker.py index a9742ad7..1499c030 100644 --- a/memory_scope/memory/worker/write/store_memory_worker.py +++ b/memory_scope/memory/worker/write/store_memory_worker.py @@ -32,5 +32,5 @@ class StoreMemoryWorker(MemoryBaseWorker): timestamp=dt_handler.timestamp, obs_dt=dt_handler.datetime_format(), obs_reflected=False, - obs_profile_updated=False) + obs_updated=False) self.vector_store.update(node) diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index f193a124..e58433b8 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -35,10 +35,14 @@ class MemoryNode(BaseModel): obs_reflected: bool = Field(False, description="if the observation is reflected") - obs_profile_updated: bool = Field(False, description="if the observation has updated user profile") + obs_updated: bool = Field(False, description="if the observation has updated user profile or insight") obs_keyword: str = Field("", description="keywords of the content") + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.gen_memory_id() + @property def node_keys(self): return list(self.model_json_schema()["properties"].keys())