From b151da319ed7575e6911770d97d29940d0ae79c7 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 1 Jul 2024 21:13:43 +0800 Subject: [PATCH] =?UTF-8?q?[dev]=C2=A0add=20contra=5Frepeat=5Fworker=20con?= =?UTF-8?q?fig?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config/demo_config.yaml | 13 +- memory_scope/constants/language_constants.py | 5 + .../memory/worker/read/extract_time_worker.py | 6 +- .../worker/read/extract_time_worker.yaml | 2 +- .../worker/write/contra_repeat_worker.py | 107 ++++++++------ .../worker/write/contra_repeat_worker.yaml | 57 ++++++++ .../worker/write/es_today_obs_worker.py | 29 ---- .../write/get_observation_with_time_worker.py | 136 ++++-------------- .../get_observation_with_time_worker.yaml | 82 +++++++++++ .../worker/write/get_observation_worker.py | 20 ++- .../memory/worker/write/info_filter_worker.py | 16 ++- memory_scope/storage/base_vector_store.py | 3 + .../llama_index_elastic_search_store.py | 6 +- memory_scope/utils/datetime_handler.py | 10 +- 14 files changed, 289 insertions(+), 203 deletions(-) create mode 100644 memory_scope/memory/worker/write/contra_repeat_worker.yaml delete mode 100644 memory_scope/memory/worker/write/es_today_obs_worker.py create mode 100644 memory_scope/memory/worker/write/get_observation_with_time_worker.yaml diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 8bdc5ccd..d22f0723 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -83,6 +83,15 @@ worker: generation_model_top_k: 1 get_observation_worker: class: memory.worker.write.get_observation_worker + generation_model: dashscope_generation generation_model_top_k: 1 - - + get_observation_with_time_worker: + class: memory.worker.write.get_observation_with_time_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + contra_repeat_worker: + class: memory.worker.write.contra_repeat_worker + generation_model: dashscope_generation + generation_model_top_k: 1 + today_obs_top_k: 30 + contra_repeat_max_count: 50 diff --git a/memory_scope/constants/language_constants.py b/memory_scope/constants/language_constants.py index 7a47bab3..c91514af 100644 --- a/memory_scope/constants/language_constants.py +++ b/memory_scope/constants/language_constants.py @@ -64,3 +64,8 @@ INCLUDED_WORD = { LanguageEnum.CN: "被包含", LanguageEnum.EN: "included" } + +COLON_WORD = { + LanguageEnum.CN: ":", + LanguageEnum.EN: ":" +} diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index c0e2e8c8..25fbe6b0 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -24,9 +24,9 @@ class ExtractTimeWorker(MemoryBaseWorker): return # prepare prompt - query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_format_prompt) - extract_time_prompt: str = self.prompt_handler.extract_time_prompt - extract_time_prompt: str = extract_time_prompt.format(query=query, query_time_str=query_time_str) + query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format) + extract_time_prompt: str = self.prompt_handler.extract_time_prompt.format(query=query, + query_time_str=query_time_str) self.logger.info(f"extract_time_prompt={extract_time_prompt}") # call sft model diff --git a/memory_scope/memory/worker/read/extract_time_worker.yaml b/memory_scope/memory/worker/read/extract_time_worker.yaml index 23431946..fe2d519e 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.yaml +++ b/memory_scope/memory/worker/read/extract_time_worker.yaml @@ -6,7 +6,7 @@ extract_time_prompt: 回答: -time_format_prompt: +time_string_format: cn: | {year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。 diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 942bc8d0..1a7b7212 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -1,72 +1,92 @@ from typing import List -from common.response_text_parser import ResponseTextParser -from common.tool_functions import contains_keyword -from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \ - MODIFIED_MEMORIES -from enumeration.memory_node_status import MemoryNodeStatus -from model.memory_wrap_node import MemoryWrapNode -from worker.bailian.memory_base_worker import MemoryBaseWorker +from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES +from memory_scope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, INCLUDED_WORD +from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus +from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.scheme.message import Message +from memory_scope.utils.datetime_handler import DatetimeHandler +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.utils.timer import timer class ContraRepeatWorker(MemoryBaseWorker): + @timer + def retrieve_today_memory(self) -> List[MemoryNode]: + if not self.chat_messages: + self.logger.warning("chat_messages is empty!") + return [] + + message: Message = self.chat_messages[-1] + dt_handler = DatetimeHandler(message.time_created) + return self.vector_store.retrieve(query=message.content, + top_k=self.today_obs_top_k, + filter_dict={ + "user_id": self.user_id, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, + MemoryTypeEnum.OBS_CUSTOMIZED.value], + "obs_dt": dt_handler.datetime_format(), + }) + def _run(self): - # 合并当前的obs和今日的obs - new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES) - new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES) - today_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES) - all_obs_nodes: List[MemoryWrapNode] = [] + all_obs_nodes: List[MemoryNode] = [] + + new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES) if new_obs_nodes: all_obs_nodes.extend(new_obs_nodes) + new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES) if new_obs_with_time_nodes: all_obs_nodes.extend(new_obs_with_time_nodes) + if not all_obs_nodes: + self.logger.info("all_obs_nodes is empty!") + self.continue_run = False + return + + today_obs_nodes: List[MemoryNode] = self.retrieve_today_memory() if today_obs_nodes: all_obs_nodes.extend(today_obs_nodes) - if not all_obs_nodes: - self.add_run_info("all_obs_nodes is empty!") - return + all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.timestamp, reverse=True)[:self.contra_repeat_max_count] - # gene prompt + # build prompt user_query_list = [] - all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True) - if len(all_obs_nodes) > self.config.merge_obs_max_count: - all_obs_nodes = all_obs_nodes[:self.config.merge_obs_max_count] - for i, n in enumerate(all_obs_nodes): - user_query_list.append(f"{i + 1} {n.memory_node.content}") + user_query_list.append(f"{i + 1} {n.content}") - merge_obs_message = self.prompt_to_msg( - system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.contra_repeat_few_shot, - user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list))) - self.logger.info(f"merge_obs_message={merge_obs_message}") + system_prompt = self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list), + user_name=self.user_id) + few_shot = self.prompt_config.contra_repeat_few_shot.format(user_name=self.user_id) + user_query = self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list), + user_name=self.user_id) + contra_repeat_message = self.prompt_to_msg(system_prompt=system_prompt, + few_shot=few_shot, + user_query=user_query) + self.logger.info(f"contra_repeat_message={contra_repeat_message}") # call LLM - response_text = self.gene_client.call(messages=merge_obs_message, - model_name=self.config.merge_obs_model, - max_token=self.config.merge_obs_max_token, - temperature=self.config.merge_obs_temperature, - top_k=self.config.merge_obs_top_k) + response = self.generation_model.call(messages=contra_repeat_message, top_k=self.generation_model_top_k) # return if empty - if not response_text: - self.add_run_info("contra repeat call llm failed!") + if not response.status or not response.message.content: return + response_text = response.message.content # parse text - idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat") + idx_merge_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) 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[MemoryWrapNode] = [] + merge_obs_nodes: List[MemoryNode] = [] for obs_content_list in idx_merge_obs_list: if not obs_content_list: continue - # [6, 逃课] + # [6, skipping classes] if len(obs_content_list) != 2: self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") continue @@ -77,25 +97,24 @@ class ContraRepeatWorker(MemoryBaseWorker): self.logger.warning(f"idx={idx} is invalid!") continue - # 序号需要修正-1 + # index number needs to be corrected to -1 idx = int(idx) - 1 if idx >= len(all_obs_nodes): self.logger.warning(f"idx={idx} is invalid!") continue - if keep_flag not in ["矛盾", "被包含", "无"]: + # judge flag + if keep_flag not in self.get_language_value([NONE_WORD, CONTRADICTORY_WORD, INCLUDED_WORD]): self.logger.warning(f"keep_flag={keep_flag} is invalid!") continue - node: MemoryWrapNode = all_obs_nodes[idx] - if keep_flag != "无": + node: MemoryNode = all_obs_nodes[idx] + if keep_flag != self.get_language_value(NONE_WORD): node.memory_node.status = MemoryNodeStatus.EXPIRED.value merge_obs_nodes.append(node) # forbid keyword - if contains_keyword(text=node.memory_node.content, keywords=self.config.forbidden_key_words): - node.memory_node.status = MemoryNodeStatus.FORBIDDEN.value - self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}") + self.logger.info(f"contra_repeat stage: {node.content} {node.status}") # save context - self.set_context(MODIFIED_MEMORIES, merge_obs_nodes) + self.vector_store.update_batch(merge_obs_nodes) diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.yaml b/memory_scope/memory/worker/write/contra_repeat_worker.yaml new file mode 100644 index 00000000..cd31fc76 --- /dev/null +++ b/memory_scope/memory/worker/write/contra_repeat_worker.yaml @@ -0,0 +1,57 @@ +contra_repeat_system: + cn: | + 对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。 + 对每个句子都做一个判断,最后一共输出{num_obs}行判断。 + 请一步步思考,并按如下格式输出: + 思考:思考的依据和过程,30字以内。 + 判断:<句子序号> <矛盾,被包含,无>,一定加<> + + +contra_repeat_few_shot: + cn: | + 示例1 + 句子: + 1 {user_name}经常失眠,对安眠药的效果感兴趣,暗示可能考虑使用。 + 2 {user_name}经常失眠,寻求缓解方法。 + 3 陈伟业是{user_name}的领导 + 4 陈伟业是{user_name}的领导 + 5 陈伟业是{user_name}的领导,是银行分行行长 + + 思考:第1句不会存在与前面序号句子的矛盾或者完全重复。 + 判断:<1> <无> + 思考:第2句中所有信息都被前面序号中第1句的信息完全包含。 + 判断:<2> <被包含> + 思考:第3句信息没有在前面序号句子中出现 + 判断:<3> <无> + 思考:第4句与前面序号中第3句的信息完全重复,即被完全包含。 + 判断:<4> <被包含> + 思考:第5句中陈伟业是{user_name}的领导的信息被前面序号中第3句的信息包含,但新增了陈伟业是银行分行行长的信息,故不是被完全包含。 + 判断:<5> <无> + + 示例2 + 句子: + 1 {user_name}的孩子成绩不太好。 + 2 {user_name}的孩子在学校经常逃课。 + 3 {user_name}的父亲生日在2024年6月2日,{user_name}打算准备礼物。 + 4 {user_name}的父亲生日在2024年5月1日。 + 5 {user_name}很喜欢和同班同学打篮球。 + 6 {user_name}喜欢打篮球。 + + 思考:第1句不会存在与前面序号句子的矛盾或者完全重复。 + 判断:<1> <无> + 思考:第2句与前面序号句子既不矛盾也不重复。 + 判断:<2> <无> + 思考:第3句与前面序号句子既不矛盾也不重复。 + 判断:<3> <无> + 思考:第4句关于{user_name}父亲生日的日期信息与前面序号句子第3句矛盾了。 + 判断:<4> <矛盾> + 思考:第5句与前面序号句子既不矛盾也不重复。 + 判断:<5> <无> + 思考:第6句中所有信息都被前面序号中第5句的信息完全包含。 + 判断:<2> <被包含> + + +contra_repeat_user_query: + cn: | + 句子: + {user_query} diff --git a/memory_scope/memory/worker/write/es_today_obs_worker.py b/memory_scope/memory/worker/write/es_today_obs_worker.py deleted file mode 100644 index 904ed171..00000000 --- a/memory_scope/memory/worker/write/es_today_obs_worker.py +++ /dev/null @@ -1,29 +0,0 @@ -from typing import List - -from common.tool_functions import time_to_formatted_str -from constants.common_constants import TODAY_OBS_NODES, DT -from enumeration.memory_node_status import MemoryNodeStatus -from enumeration.memory_type_enum import MemoryTypeEnum -from model.memory_wrap_node import MemoryWrapNode -from worker.bailian.memory_base_worker import MemoryBaseWorker - - -class EsTodayObsWorker(MemoryBaseWorker): - - def _run(self): - if not self.messages: - self.logger.warning("messages is empty!") - return - msg_time_created = self.messages[-1].time_created - hits = self.es_client.exact_search_v2(size=self.config.es_today_obs_top_k, - term_filters={ - "memoryId": self.config.memory_id, - "status": MemoryNodeStatus.ACTIVE.value, - "scene": self.scene.lower(), - "memoryType": MemoryTypeEnum.OBSERVATION.value, - f"metaData.{DT}": time_to_formatted_str(msg_time_created), - }) - - today_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits] - self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}") - self.set_context(TODAY_OBS_NODES, today_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.py b/memory_scope/memory/worker/write/get_observation_with_time_worker.py index 63c13efe..e54e2a81 100644 --- a/memory_scope/memory/worker/write/get_observation_with_time_worker.py +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -1,128 +1,46 @@ -from datetime import datetime from typing import List -from common.response_text_parser import ResponseTextParser -from common.tool_functions import time_to_formatted_str, get_datetime_info_dict, extract_date_parts -from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME, KEY_WORD, DATATIME_WORD_LIST, \ - NEW_OBS_WITH_TIME_NODES -from enumeration.memory_node_status import MemoryNodeStatus -from enumeration.memory_type_enum import MemoryTypeEnum -from model.memory_wrap_node import MemoryWrapNode -from model.message import Message -from worker.bailian.memory_base_worker import MemoryBaseWorker +from memory_scope.constants.common_constants import NEW_OBS_WITH_TIME_NODES +from memory_scope.constants.language_constants import DATATIME_WORD_LIST, COLON_WORD +from memory_scope.memory.worker.write.get_observation_worker import GetObservationWorker +from memory_scope.scheme.memory_node import MemoryNode +from memory_scope.scheme.message import Message +from memory_scope.utils.datetime_handler import DatetimeHandler +from memory_scope.utils.tool_functions import prompt_to_msg -class GetObservationWithTimeWorker(MemoryBaseWorker): +class GetObservationWithTimeWorker(GetObservationWorker): - def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str): - created_dt: datetime = datetime.fromtimestamp(float(message.time_created)) - dt = time_to_formatted_str(time=created_dt) - - # 组合meta_data - meta_data = { - MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话 - REFLECTED: "0", # reflect标记 - DT: dt, # 当天标记 - NEW: "1", # summary-long标记 - MSG_TIME: message.time_created, # 对话时间 - TIME_INFER: time_infer, # 推断的时间 - KEY_WORD: keywords, # 关键词 - } - - # 事件时间 - meta_data.update({f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()}) - # 对话时间 - meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()}) - - return MemoryWrapNode.init_from_attrs(content=obs_content, - memoryId=self.config.memory_id, - timeCreated=message.time_created, - scene=self.scene, - memoryType=MemoryTypeEnum.OBSERVATION.value, - content_modified=True, # 新增的obs需要置为true - metaData=meta_data, - status=MemoryNodeStatus.ACTIVE.value, - tenantId=self.config.tenant_id) - - def _run(self): - # gene prompt + def build_prompt(self) -> List[Message]: + # build prompt user_query_list = [] i = 1 - for msg in self.messages: + for msg in self.chat_messages: match = False - for time_keyword in DATATIME_WORD_LIST: + for time_keyword in self.get_language_value(DATATIME_WORD_LIST): if time_keyword in msg.content: match = True break if match: - dt = time_to_formatted_str(time=msg.time_created, - date_format="", - string_format="{year}年{month}月{day}日{weekday}{hour}点") - user_query_list.append(f"{i} {dt} 用户:{msg.content}") + dt_handler = DatetimeHandler(dt=msg.time_created) + dt = dt_handler.string_format(self.prompt_handler.time_string_format) + user_query_list.append(f"{i} {dt} {self.user_id}{self.get_language_value(COLON_WORD)}{msg.content}") i += 1 if not user_query_list: - self.add_run_info(f"get obs with time user_query_list={user_query_list} is empty") - return + self.logger.warning(f"get obs_with_time user_query_list={user_query_list} is empty") + return [] - obtain_obs_message = self.prompt_to_msg( - system_prompt=self.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)), - few_shot=self.prompt_config.get_observation_with_time_few_shot, - user_query=self.prompt_config.get_observation_with_time_user_query.format( - user_query="\n".join(user_query_list))) + system_prompt = self.prompt_handler.get_observation_with_time_system.format(num_obs=len(user_query_list), + user_name=self.user_id) + few_shot = self.prompt_config.get_observation_with_time_few_shot.format(self.user_id) + user_query = self.prompt_config.get_observation_with_time_user_query.format( + user_query="\n".join(user_query_list), + user_name=self.user_id) + + obtain_obs_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"obtain_obs_message={obtain_obs_message}") + return obtain_obs_message - # call LLM - response_text: str = self.gene_client.call(messages=obtain_obs_message, - model_name=self.config.summary_messages_model, - max_token=self.config.summary_messages_max_token, - temperature=self.config.summary_messages_temperature, - top_k=self.config.summary_messages_top_k) - - # return if empty - if not response_text: - self.add_run_info("summary call llm failed!", continue_run=False) - return - - # parse text - idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time") - if len(idx_obs_list) <= 0: - self.add_run_info("idx_obs_list is empty!", continue_run=False) - return - - # gene new obs nodes - new_obs_nodes: List[MemoryWrapNode] = [] - for obs_content_list in idx_obs_list: - if not obs_content_list: - continue - - # [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游] - if len(obs_content_list) != 4: - self.logger.warning(f"obs_content_list={obs_content_list} is invalid!") - continue - - idx, time_infer, obs_content, keywords = obs_content_list - - if obs_content in ["无", "重复"]: - continue - - if not idx.isdigit(): - self.logger.warning(f"idx={idx} is invalid!") - continue - - if time_infer == "无": - time_infer = "" - - # 序号需要修正-1 - idx = int(idx) - 1 - if idx >= len(self.messages): - self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}") - continue - - new_obs_nodes.append(self.add_observation(message=self.messages[idx], - obs_content=obs_content, - time_infer=time_infer, - keywords=keywords)) - - # save context + def save(self, new_obs_nodes: List[MemoryNode]): self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_with_time_worker.yaml b/memory_scope/memory/worker/write/get_observation_with_time_worker.yaml new file mode 100644 index 00000000..7ca0d846 --- /dev/null +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.yaml @@ -0,0 +1,82 @@ +time_string_format: + cn: | + {year}年{month}月{day}日{weekday}{hour}点 + +get_observation_with_time_system: + cn: | + 任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。 + 每一句{user_name}句子的格式是:<序号> <对话时间> {user_name}:<句子> + 对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。 + 如果{user_name}信息涉及时间,则结合对话时间推断{user_name}信息的时间信息,没有则不输出。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。 + 对每个句子都做一次信息提取,最后一共输出{num_obs}行信息。 + 请一步步思考,并一定要按如下格式依次输出,最后的结果一定要加<>: + 信息:<句子序号> <时间信息或“无”> <明确的重要信息或“重复”或”无“> <关键词> + + +get_observation_with_time_few_shot: + cn: | + 示例1: + {user_name}句子: + 1 2022年5月1日周二3点 {user_name}:帮我写一段给同事张三女儿三岁生日的祝福语。 + 2 2022年5月2日周二17点 {user_name}:公元1400年至1550年中国历史大事表。 + 3 2022年5月3日周二18点 {user_name}:能给我整理一张如何使用大模型的技巧列表吗,要求内容尽量精简。 + 4 2022年7月3日周四12点 {user_name}:上上个月我办了游泳卡。 + + 思考:从第1句可以得知张三是{user_name}的同事,这是关于{user_name}的人际关系的重要信息。其余信息重要性不足。{user_name}信息不涉及时间。 + 信息:<1> <> <张三是{user_name}的同事。> <张三, 同事> + 思考:第2句是{user_name}提出的要求,没有明确提及{user_name}个人信息。 + 信息:<2> <> <无> <> + 思考:第3句是{user_name}提出的要求,没有明确提及{user_name}个人信息。 + 信息:<3> <> <无> <> + 思考:从第4句可以得出{user_name}上上个月办了游泳卡。{user_name}信息涉及时间,结合对话时间为2022年7月,推断{user_name}在2022年5月{user_name}办了游泳卡。 + 信息:<4> <2022年5月> <{user_name}在2022年5月办了游泳卡。> <游泳卡> + + + 示例2: + {user_name}句子: + 1 2020年1月4日周日10点 {user_name}:我花5000元买了100股海天味业。 + 2 2023年4月27日周五8点 {user_name}:明天是我和妻子的结婚纪念日,帮我推荐一家餐厅。 + 3 2020年1月4日周日10点 {user_name}:我花5000元买了100股海天味业。 + 4 2021年6月2日周四23点 {user_name}:谢啦。我中午在公司附近吃,帮我推荐一家阿里巴巴徐汇滨江园区附近的餐厅吧。 + 5 2021年7月9日周六11点 {user_name}:两个坏消息,我打羽毛球把拍子打断线了。。。然后我去我朋友家撸猫,结果我猫毛过敏,今天疯狂打喷嚏。。。 + + 思考:从第1句可以得知{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元,这是关于{user_name}的投资决策的重要信息。{user_name}信息不涉及时间。 + 信息:<1> <> <{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元。> <海天味业, 股票> + 思考:从第2句可以得知{user_name}与妻子的结婚纪念日是明天,这是关于{user_name}重要纪念日的信息。其余信息重要性不足。{user_name}信息涉及时间,结合对话时间为2023年4月27日, + 以及结婚纪念日为周期性日期,推断{user_name}与妻子的结婚纪念日是每年4月28日。 + 信息:<2> <每年4月28日> <{user_name}与妻子的结婚纪念日是每年4月28日。> <妻子, 结婚纪念日> + 思考:第3句含有的信息与第1句重复了。 + 信息:<3> <> <重复> <> + 思考:从第4句以得知{user_name}在阿里巴巴徐汇滨江园区工作,这是关于{user_name}的工作的重要信息。其余信息重要性不足。{user_name}信息不涉及时间。 + 信息:<4> <> <{user_name}在阿里巴巴徐汇滨江园区工作。> <阿里巴巴, 徐汇滨江园区, 工作> + 思考:从第5句可以得知{user_name}前天打羽毛球时把球拍打断了线,但这不是重要的信息。还可以得知{user_name}对猫毛过敏,这是关于{user_name}的健康的重要信息。{user_name}信息不涉及时间。 + 信息:<5> <> <{user_name}对猫毛过敏。> <猫毛, 过敏> + + + 示例3: + {user_name}句子: + 1 2023年6月30日周五15点 {user_name}:上个月我和家人一起去杭州旅游,景色很不错。 + 2 2023年7月2日周二10点 {user_name}:昨天是我生日,一个人过的。 + 3 2020年7月3日周四11点 {user_name}:提醒我下周一去体检。 + 4 2023年5月21日周六14点 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢 + 5 2018年3月6日周四19点 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续 + + + 思考:从第1句可以得知{user_name}和家人上个月去杭州旅游了,这是关于{user_name}的经历的重要信息。其余信息重要性不足。{user_name}信息涉及时间,结合对话时间为2023年6月推断{user_name}和家人2023年5月去杭州旅游了。 + 信息:<1> <2023年5月> <{user_name}和家人2023年5月去杭州旅游了。> <家人, 杭州, 旅游> + 思考:从第2句可以得知{user_name}的生日是昨天,这是关于{user_name}重要纪念日的信息。其余信息重要性不足。{user_name}信息涉及时间,结合对话时间为2023年7月2日, + 以及生日为周期性日期,推断{user_name}的生日是每年7月2日。 + 信息:<2> <每年7月2日> <{user_name}的生日是每年7月2日。> <生日> + 思考:从第3句可以得出{user_name}下周一去体检,这是{user_name}要求记忆的重要信息。{user_name}信息涉及时间,结合对话时间为2020年7月3日周四,推断{user_name}2020年7月6日周一去体检。 + 信息:<3> <2020年7月6日周一> <{user_name}2020年7月6日周一去体检。> <体检> + 思考:第4句是{user_name}对他人观点的讨论和疑问,没有明确提及{user_name}个人信息。 + 信息:<4> <> <无> <> + 思考:第5句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。 + 信息:<5> <> <无> <> + + +get_observation_with_time_user_query: + cn: | + {user_name}句子: + {user_query} + diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index 1d894098..cc027a08 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -2,7 +2,7 @@ from datetime import datetime from typing import List from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER -from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD +from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD, COLON_WORD from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker @@ -42,7 +42,7 @@ class GetObservationWorker(MemoryBaseWorker): node.gen_memory_id() return node - def build_prompt(self): + def build_prompt(self) -> List[Message]: # build prompt user_query_list = [] i = 1 @@ -53,12 +53,12 @@ class GetObservationWorker(MemoryBaseWorker): match = True break if not match: - user_query_list.append(f"{i} {self.user_id}:{msg.content}") + user_query_list.append(f"{i} {self.user_id}{self.get_language_value(COLON_WORD)}{msg.content}") i += 1 if not user_query_list: self.logger.warning(f"get obs user_query_list={user_query_list} is empty") - return + return [] system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list), user_name=self.user_id) @@ -70,8 +70,14 @@ class GetObservationWorker(MemoryBaseWorker): self.logger.info(f"obtain_obs_message={obtain_obs_message}") return obtain_obs_message + def save(self, new_obs_nodes: List[MemoryNode]): + self.set_context(NEW_OBS_NODES, new_obs_nodes) + def _run(self): obtain_obs_message = self.build_prompt() + if not obtain_obs_message: + self.logger.warning("get obs message is empty!") + return # call LLM response = self.generation_model.call(messages=obtain_obs_message, top_k=self.generation_model_top_k) @@ -107,6 +113,9 @@ class GetObservationWorker(MemoryBaseWorker): self.logger.warning(f"idx={idx} is invalid!") continue + if time_infer == self.get_language_value(NONE_WORD): + time_infer = "" + # index number needs to be corrected to -1 idx = int(idx) - 1 if idx >= len(self.messages): @@ -118,5 +127,4 @@ class GetObservationWorker(MemoryBaseWorker): obs_content=obs_content, keywords=keywords)) - # save context - self.set_context(NEW_OBS_NODES, new_obs_nodes) + self.save(new_obs_nodes) diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 8ea05496..3949349e 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -1,5 +1,6 @@ from typing import List +from memory_scope.constants.language_constants import COLON_WORD from memory_scope.enumeration.message_role_enum import MessageRoleEnum from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker from memory_scope.scheme.message import Message @@ -22,12 +23,19 @@ class InfoFilterWorker(MemoryBaseWorker): msg.content = msg.content[: begin_size] + msg.content[-end_size:] info_messages.append(msg) - # gene prompt - user_query = "\n".join([f"{i + 1} {self.user_id}:{msg.content}" for i, msg in enumerate(info_messages)]) + if not info_messages: + self.logger.warning("info_messages is empty!") + self.continue_run = False + return + + # generate prompt + user_query_list = [] + for i, msg in enumerate(info_messages): + user_query_list.append(f"{i + 1} {self.user_id}{self.get_language_value(COLON_WORD)}{msg.content}") system_prompt = self.prompt_handler.info_filter_system.format(batch_size=len(info_messages), user_name=self.user_id) few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.user_id) - user_query = self.prompt_handler.info_filter_user_query.format(user_query=user_query) + user_query = self.prompt_handler.info_filter_user_query.format(user_query="\n".join(user_query_list)) info_filter_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query) self.logger.info(f"info_filter_message={info_filter_message}") @@ -36,6 +44,7 @@ class InfoFilterWorker(MemoryBaseWorker): # return if empty if not response.status or not response.message.content: + self.continue_run = False return response_text = response.message.content @@ -43,6 +52,7 @@ class InfoFilterWorker(MemoryBaseWorker): info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) if len(info_score_list) != len(info_messages): self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}") + self.continue_run = False return # filter messages diff --git a/memory_scope/storage/base_vector_store.py b/memory_scope/storage/base_vector_store.py index d0ba4f6f..34be90ca 100644 --- a/memory_scope/storage/base_vector_store.py +++ b/memory_scope/storage/base_vector_store.py @@ -27,6 +27,9 @@ class BaseVectorStore(metaclass=ABCMeta): def update(self, node: MemoryNode): pass + def update_batch(self, nodes: List[MemoryNode]): + pass + def flush(self): pass diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 0881aae9..0c8f481f 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -1,6 +1,6 @@ from typing import Dict, List, Any -from llama_index.core import VectorStoreIndex, Settings +from llama_index.core import VectorStoreIndex from llama_index.core.schema import TextNode, NodeWithScore from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy @@ -123,6 +123,10 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): self.delete(node) self.insert(node) + def update_batch(self, nodes: List[MemoryNode]): + for node in nodes: + self.update(node) + def close(self): self.es_store.close() diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py index 9f3913cd..a36861ab 100644 --- a/memory_scope/utils/datetime_handler.py +++ b/memory_scope/utils/datetime_handler.py @@ -43,11 +43,11 @@ class DatetimeHandler(object): def extract_date_parts_cn(input_string: str): # Extending our pattern to handle every/每 as a possible value. patterns = { - 'year': r'(\d+|每)年', - 'month': r'(\d+|每)月', - 'day': r'(\d+|每)日', - 'weekday': r'周([一二三四五六日])', - 'hour': r'(\d+)点' + "year": r"(\d+|每)年", + "month": r"(\d+|每)月", + "day": r"(\d+|每)日", + "weekday": r"周([一二三四五六日])", + "hour": r"(\d+)点" } weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7} extracted_data = {}