From 4333e5b13b4d1bb67bc0e11d6e9e6dccf7acd820 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 1 Jul 2024 19:47:43 +0800 Subject: [PATCH] [dev] add info filter system prompt --- config/demo_config.yaml | 9 ++ ...lter_prompt.py => info_filter_worker.yaml} | 9 +- memory_scope/chat/cli_memory_chat.py | 3 +- .../chat}/cli_memory_chat.yaml | 0 memory_scope/constants/common_constants.py | 24 ---- memory_scope/constants/language_constants.py | 35 +++++ memory_scope/memory/worker/dummy_worker.py | 3 +- .../memory/worker/memory_base_worker.py | 9 +- .../memory/worker/read/extract_time_worker.py | 8 +- .../worker/read}/extract_time_worker.yaml | 0 .../worker/write/contra_repeat_worker.py | 101 ++++++++++++++ .../worker/write/es_today_obs_worker.py | 29 ++++ .../write/get_observation_with_time_worker.py | 128 ++++++++++++++++++ .../worker/write/get_observation_worker.py | 122 +++++++++++++++++ .../worker/write/get_observation_worker.yaml | 70 ++++++++++ .../memory/worker/write/info_filter_worker.py | 58 ++++++++ .../worker/write/info_filter_worker.yaml | 63 +++++++++ memory_scope/scheme/memory_node.py | 13 +- memory_scope/storage/dummy_vector_store.py | 1 + memory_scope/utils/datetime_handler.py | 78 +++++++++++ memory_scope/utils/prompt_handler.py | 37 +++-- memory_scope/utils/tool_functions.py | 52 ++----- 22 files changed, 753 insertions(+), 99 deletions(-) rename config/prompts/{info_filter_prompt.py => info_filter_worker.yaml} (82%) rename {config/prompts => memory_scope/chat}/cli_memory_chat.yaml (100%) rename {config/prompts => memory_scope/memory/worker/read}/extract_time_worker.yaml (100%) create mode 100644 memory_scope/memory/worker/write/contra_repeat_worker.py create 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.py create mode 100644 memory_scope/memory/worker/write/get_observation_worker.py create mode 100644 memory_scope/memory/worker/write/get_observation_worker.yaml create mode 100644 memory_scope/memory/worker/write/info_filter_worker.py create mode 100644 memory_scope/memory/worker/write/info_filter_worker.yaml create mode 100644 memory_scope/utils/datetime_handler.py diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 9b7fc090..8bdc5ccd 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -76,4 +76,13 @@ worker: observation: 1 fuse_time_ratio: 2.0 fuse_rerank_top_k: 10 + info_filter_worker: + class: memory.worker.write.info_filter_worker + generation_model: dashscope_generation + info_filter_msg_max_size: 200 + generation_model_top_k: 1 + get_observation_worker: + class: memory.worker.write.get_observation_worker + generation_model_top_k: 1 + diff --git a/config/prompts/info_filter_prompt.py b/config/prompts/info_filter_worker.yaml similarity index 82% rename from config/prompts/info_filter_prompt.py rename to config/prompts/info_filter_worker.yaml index 4b4dbaaf..593cdd9d 100644 --- a/config/prompts/info_filter_prompt.py +++ b/config/prompts/info_filter_worker.yaml @@ -1,5 +1,10 @@ -from memory_scope.enumeration.language_enum import LanguageEnum - +info_filter_system_prompt: + cn: | + 任务指令:对所给{batch_size}个句子中所含有的关于用户的信息打分,分数为0,1,2或3。 + 注意:其中0表示不包含用户信息,1表示句子中只包含用户假设的信息或者用户虚构的内容,2表示包含用户的一般信息,时效性信息或者需要猜测才能得到的用户信息,3表示明确含有或者可以确定推断出关于用户的重要信息,或者用户要求记录。 + 按如下格式输出, 每一行输出一个打分,一定加<>,一共输出{batch_size}个分数: + 结果: + <分数:0或1或2或3> INFO_FILTER_SYSTEM_PROMPT = { LanguageEnum.CN: """ diff --git a/memory_scope/chat/cli_memory_chat.py b/memory_scope/chat/cli_memory_chat.py index d882aec5..b9990453 100644 --- a/memory_scope/chat/cli_memory_chat.py +++ b/memory_scope/chat/cli_memory_chat.py @@ -46,8 +46,7 @@ class CliMemoryChat(BaseMemoryChat): @property def prompt_handler(self) -> PromptHandler: if self._prompt_handler is None: - self._prompt_handler = PromptHandler() - self._prompt_handler.add_file_prompts(self.__class__.__name__) + self._prompt_handler = PromptHandler(__file__, **self.kwargs) return self._prompt_handler def print_logo(self): diff --git a/config/prompts/cli_memory_chat.yaml b/memory_scope/chat/cli_memory_chat.yaml similarity index 100% rename from config/prompts/cli_memory_chat.yaml rename to memory_scope/chat/cli_memory_chat.yaml diff --git a/memory_scope/constants/common_constants.py b/memory_scope/constants/common_constants.py index 915c7536..a0f4e9e4 100644 --- a/memory_scope/constants/common_constants.py +++ b/memory_scope/constants/common_constants.py @@ -83,31 +83,7 @@ TIME_MATCHED = "time_matched" QUERY_KEYWORDS = "query_keywords" -WEEKDAYS = ["周一", "周二", "周三", "周四", "周五", "周六", "周日"] -DATATIME_WORD_LIST = [ - "天", - "周", - "月", - "年", - "星期", - "点", - "分钟", - "小时", - "秒", - "上午", - "下午", - "早上", - "早晨", - "晚上", - "中午", - "日", - "夜", - "清晨", - "傍晚", - "凌晨", - "岁", -] TIME_FORMAT_V1 = "{year}年{month}月{day}日{weekday}{hour}点" diff --git a/memory_scope/constants/language_constants.py b/memory_scope/constants/language_constants.py index 921d6719..7a47bab3 100644 --- a/memory_scope/constants/language_constants.py +++ b/memory_scope/constants/language_constants.py @@ -29,3 +29,38 @@ DATATIME_WORD_LIST = { ] } + +WEEKDAYS = { + LanguageEnum.CN: [ + "周一", + "周二", + "周三", + "周四", + "周五", + "周六", + "周日" + ], + LanguageEnum.EN: [ + + ] +} + +NONE_WORD = { + LanguageEnum.CN: "无", + LanguageEnum.EN: "none" +} + +REPEATED_WORD = { + LanguageEnum.CN: "重复", + LanguageEnum.EN: "repeated" +} + +CONTRADICTORY_WORD = { + LanguageEnum.CN: "矛盾", + LanguageEnum.EN: "contradictory" +} + +INCLUDED_WORD = { + LanguageEnum.CN: "被包含", + LanguageEnum.EN: "included" +} diff --git a/memory_scope/memory/worker/dummy_worker.py b/memory_scope/memory/worker/dummy_worker.py index 239a422c..2b488b7e 100644 --- a/memory_scope/memory/worker/dummy_worker.py +++ b/memory_scope/memory/worker/dummy_worker.py @@ -10,4 +10,5 @@ class DummyWorker(BaseWorker): chat_kwargs = self.get_context(CHAT_KWARGS) self.logger.info(f"enter workflow={workflow_name}.dummy_worker!") ts = int(datetime.datetime.now().timestamp()) - self.set_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} \nts={ts}") + file_path = __file__ + self.set_context(RESULT, f"test {workflow_name} kwargs={chat_kwargs} file_path={file_path} \nts={ts}") diff --git a/memory_scope/memory/worker/memory_base_worker.py b/memory_scope/memory/worker/memory_base_worker.py index 3f973ed0..29cb09b6 100644 --- a/memory_scope/memory/worker/memory_base_worker.py +++ b/memory_scope/memory/worker/memory_base_worker.py @@ -83,13 +83,14 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta): @property def prompt_handler(self) -> PromptHandler: if self._prompt_handler is None: - self._prompt_handler = PromptHandler() - self._prompt_handler.add_file_prompts(self.__class__.__name__) + self._prompt_handler = PromptHandler(__file__, **self.kwargs) return self._prompt_handler def __getattr__(self, key: str): return self.kwargs[key] @staticmethod - def get_language_prompt(prompt: dict) -> str: - return prompt[G_CONTEXT.language] + def get_language_value(languages: dict | list) -> str | list[str]: + if isinstance(languages, list): + return [x[G_CONTEXT.language] for x in languages] + return languages[G_CONTEXT.language] diff --git a/memory_scope/memory/worker/read/extract_time_worker.py b/memory_scope/memory/worker/read/extract_time_worker.py index 3ade4b05..c0e2e8c8 100644 --- a/memory_scope/memory/worker/read/extract_time_worker.py +++ b/memory_scope/memory/worker/read/extract_time_worker.py @@ -4,7 +4,7 @@ from typing import Dict from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT from memory_scope.constants.language_constants import DATATIME_WORD_LIST from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker -from memory_scope.utils.tool_functions import time_to_formatted_str +from memory_scope.utils.datetime_handler import DatetimeHandler class ExtractTimeWorker(MemoryBaseWorker): @@ -15,7 +15,7 @@ class ExtractTimeWorker(MemoryBaseWorker): # find datetime keyword contain_datetime = False - for datetime_word in self.get_language_prompt(DATATIME_WORD_LIST): + for datetime_word in self.get_language_value(DATATIME_WORD_LIST): if datetime_word in query: contain_datetime = True break @@ -24,9 +24,7 @@ class ExtractTimeWorker(MemoryBaseWorker): return # prepare prompt - query_time_str = time_to_formatted_str(dt=query_timestamp, - date_format="", - string_format=self.prompt_handler.time_format_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) self.logger.info(f"extract_time_prompt={extract_time_prompt}") diff --git a/config/prompts/extract_time_worker.yaml b/memory_scope/memory/worker/read/extract_time_worker.yaml similarity index 100% rename from config/prompts/extract_time_worker.yaml rename to memory_scope/memory/worker/read/extract_time_worker.yaml diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py new file mode 100644 index 00000000..942bc8d0 --- /dev/null +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -0,0 +1,101 @@ +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 + + +class ContraRepeatWorker(MemoryBaseWorker): + + 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] = [] + if new_obs_nodes: + all_obs_nodes.extend(new_obs_nodes) + if new_obs_with_time_nodes: + all_obs_nodes.extend(new_obs_with_time_nodes) + 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 + + # gene 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}") + + 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}") + + # 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) + + # return if empty + if not response_text: + self.add_run_info("contra repeat call llm failed!") + return + + # parse text + idx_merge_obs_list = ResponseTextParser(response_text).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[MemoryWrapNode] = [] + 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 + + if keep_flag not in ["矛盾", "被包含", "无"]: + self.logger.warning(f"keep_flag={keep_flag} is invalid!") + continue + + node: MemoryWrapNode = all_obs_nodes[idx] + if keep_flag != "无": + 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}") + + # save context + self.set_context(MODIFIED_MEMORIES, merge_obs_nodes) diff --git a/memory_scope/memory/worker/write/es_today_obs_worker.py b/memory_scope/memory/worker/write/es_today_obs_worker.py new file mode 100644 index 00000000..904ed171 --- /dev/null +++ b/memory_scope/memory/worker/write/es_today_obs_worker.py @@ -0,0 +1,29 @@ +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 new file mode 100644 index 00000000..63c13efe --- /dev/null +++ b/memory_scope/memory/worker/write/get_observation_with_time_worker.py @@ -0,0 +1,128 @@ +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 + + +class GetObservationWithTimeWorker(MemoryBaseWorker): + + 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 + user_query_list = [] + i = 1 + for msg in self.messages: + match = False + for time_keyword in 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}") + 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 + + 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))) + self.logger.info(f"obtain_obs_message={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 + self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py new file mode 100644 index 00000000..1d894098 --- /dev/null +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -0,0 +1,122 @@ +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.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.tool_functions import prompt_to_msg + + +class GetObservationWorker(MemoryBaseWorker): + def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str): + created_dt: datetime = datetime.fromtimestamp(float(message.time_created)) + dt_handler = DatetimeHandler(dt=created_dt) + + # 组合meta_data + meta_data = { + MemoryTypeEnum.CONVERSATION.value: message.content, + TIME_INFER: time_infer, + **{f"msg_{k}": str(v) for k, v in dt_handler.dt_info_dict.items()}, + } + + if time_infer: + dt_infer_handler = DatetimeHandler(dt=time_infer) + meta_data.update({f"event_{k}": str(v) for k, v in dt_infer_handler.dt_info_dict.items()}) + + node = MemoryNode(user_id=self.user_id, + meta_data=meta_data, + content=obs_content, + memoryType=MemoryTypeEnum.OBSERVATION.value, + status=MemoryNodeStatus.ACTIVE.value, + timestamp=message.time_created, + obs_dt=dt_handler.datetime_format(), + obs_reflected=False, + obs_profile_updated=False, + keywords=keywords) + node.gen_memory_id() + return node + + def build_prompt(self): + # build prompt + user_query_list = [] + i = 1 + for msg in self.chat_messages: + match = False + for time_keyword in self.get_language_value(DATATIME_WORD_LIST): + if time_keyword in msg.content: + match = True + break + if not match: + user_query_list.append(f"{i} {self.user_id}:{msg.content}") + i += 1 + + if not user_query_list: + self.logger.warning(f"get obs user_query_list={user_query_list} is empty") + return + + system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list), + user_name=self.user_id) + few_shot = self.prompt_config.get_observation_few_shot.format(self.user_id) + user_query = self.prompt_config.get_observation_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 + + def _run(self): + obtain_obs_message = self.build_prompt() + + # call LLM + response = self.generation_model.call(messages=obtain_obs_message, top_k=self.generation_model_top_k) + + # return if empty + if not response.status or not response.message.content: + return + response_text = response.message.content + + # parse text + idx_obs_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__) + if len(idx_obs_list) <= 0: + self.logger.warning("idx_obs_list is empty!") + return + + # gene new obs nodes + new_obs_nodes: List[MemoryNode] = [] + for obs_content_list in idx_obs_list: + if not obs_content_list: + continue + + # [1, In June 2022, the user will travel to Hangzhou for tourism, tourism] + 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 self.get_language_value([NONE_WORD, REPEATED_WORD]): + continue + + if not idx.isdigit(): + self.logger.warning(f"idx={idx} is invalid!") + continue + + # index number needs to be corrected to -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], + time_infer=time_infer, + obs_content=obs_content, + keywords=keywords)) + + # save context + self.set_context(NEW_OBS_NODES, new_obs_nodes) diff --git a/memory_scope/memory/worker/write/get_observation_worker.yaml b/memory_scope/memory/worker/write/get_observation_worker.yaml new file mode 100644 index 00000000..4b840721 --- /dev/null +++ b/memory_scope/memory/worker/write/get_observation_worker.yaml @@ -0,0 +1,70 @@ +get_observation_system: + cn: | + 任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,与相应的关键词。最多提取{num_obs}条信息。对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。 + 不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。 + 对每个句子都做一次信息提取,最后一共输出{num_obs}行信息。 + 请一定要按如下格式依次输出,最后的结果一定要加<>: + 信息:<句子序号> <> <明确的重要信息或“重复”或”无“> <关键词> + + +get_observation_few_shot: + cn: | + 示例1: + {user_name}句子: + 1 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办 + 2 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢 + 3 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办 + 4 {user_name}:我是一个刚毕业的学生,对社会,行业不了解,给我介绍一下社会系统和行业格局 + 思考:从第1句可以得知{user_name}现在没有工作,负债几万,这是关于{user_name}工作与经济状况的重要信息。 + 信息:<1> <> <{user_name}当前无工作且负债几万> <无工作, 负债几万> + 思考:第2句是{user_name}对他人观点的讨论和疑问,没有明确提及{user_name}个人信息。 + 信息:<2> <> <无> <> + 思考:第3句含有的信息与第1句重复了。 + 信息:<3> <> <重复> <> + 思考:从第4句可以得知{user_name}是一个刚毕业的学生,这是关于{user_name}身份背景状况的重要信息。其余信息重要性不足。 + 信息:<4> <> <{user_name}是一名刚毕业的学生。> <刚毕业, 学生> + + 示例2: + {user_name}句子: + 1 {user_name}:帮我写一段给同事张三女儿三岁生日的祝福语。 + 2 {user_name}:能给我整理一张如何使用大模型的技巧列表吗,要求内容尽量精简。 + 3 {user_name}:两个坏消息,我打羽毛球把拍子打断线了。。。然后我去我朋友家撸猫,结果我猫毛过敏,今天疯狂打喷嚏。。。 + 4 {user_name}:公元1400年至1550年中国历史大事表。 + 5 {user_name}:谢啦。我中午在公司附近吃,帮我推荐一家阿里巴巴徐汇滨江园区附近的餐厅吧。 + 思考:从第1句可以得知张三是{user_name}的同事,这是关于{user_name}的人际关系的重要信息。其余信息重要性不足。 + 信息:<1> <> <张三是{user_name}的同事。> <张三, 同事> + 思考:第2句是{user_name}提出的要求,没有明确提及{user_name}个人信息。 + 信息:<2> <> <无> <> + 思考:从第3句可以得知{user_name}前天打羽毛球时把球拍打断了线,但这不是重要的信息。还可以得知{user_name}对猫毛过敏,这是关于{user_name}的健康的重要信息。 + 信息:<3> <> <{user_name}对猫毛过敏。> <猫毛, 过敏> + 思考:从第4句是{user_name}提出的要求,没有明确提及{user_name}个人信息。 + 信息:<4> <> <无> <> + 思考:从第5句可以得知{user_name}在阿里巴巴徐汇滨江园区工作,这是关于{user_name}的工作地点的重要信息。 + 信息:<5> <> <{user_name}在阿里巴巴徐汇滨江园区工作。> <阿里巴巴, 徐汇滨江园区, 工作> + + 示例3: + {user_name}句子: + 1 {user_name}:我想买辆新能源汽车,有什么推荐吗? + 2 {user_name}:我在上海,想买辆新能源汽车,有什么推荐吗? + 3 {user_name}:案外人异议审查期间,人民法院不得对执行标的进行处分,不就是中止执行的意思吗? + 4 {user_name}:请写两句藏头诗分别以“胜”和“利”开头。 + 5 {user_name}:我花5000元买了100股海天味业。 + 6 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续 + 思考:从第1句可以得知{user_name}寻求购买新能源汽车的建议或推荐,这是这是关于{user_name}的大宗消费的重要的信息。 + 信息:<1> <> <{user_name}寻求购买新能源汽车的建议或推荐。> <购买, 新能源汽车> + 思考:从第2句可以得知{user_name}当前所在城市为上海,这是关于{user_name}的生活地区的重要信息。其余信息与第1句重复了。 + 信息:<2> <> <{user_name}所在的城市是上海。> <上海> + 思考:第3句是{user_name}对某个观点的讨论和疑问,没有明确提及{user_name}个人信息。 + 信息:<3> <> <无> <> + 思考:第4句是{user_name}提出的要求,没有明确提及{user_name}个人信息。 + 信息:<4> <> <无> <> + 思考:从第5句可以得知{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元,这是关于{user_name}的投资决策的重要信息。 + 信息:<5> <> <{user_name}购买了海天味业股票,购买数量为100股,购买金额为5000元。> <海天味业, 股票> + 思考:第6句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。 + 信息:<6> <> <无> <> + + +get_observation_user_query: + cn: | + {user_name}句子: + {user_query} diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py new file mode 100644 index 00000000..8ea05496 --- /dev/null +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -0,0 +1,58 @@ +from typing import List + +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 +from memory_scope.utils.response_text_parser import ResponseTextParser +from memory_scope.utils.tool_functions import prompt_to_msg + + +class InfoFilterWorker(MemoryBaseWorker): + + def _run(self): + # filter user msg + info_messages: List[Message] = [] + for msg in self.chat_messages: + # TODO: add memory for assistant + if msg.role != MessageRoleEnum.USER.value: + continue + if len(msg.content) >= self.info_filter_msg_max_size: + begin_size = int(self.info_filter_msg_max_size * 0.75 + 0.5) + end_size = int(self.info_filter_msg_max_size * 0.25 + 0.5) + 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)]) + 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) + 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}") + + # call llm + response = self.generation_model.call(messages=info_filter_message, top_k=self.generation_model_top_k) + + # return if empty + if not response.status or not response.message.content: + return + response_text = response.message.content + + # parse text + 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)}") + return + + # filter messages + filtered_messages: List[Message] = [] + for msg, info_score in zip(info_messages, info_score_list): + if not info_score: + continue + + score = info_score[0] + if score in ("3",): + msg.meta_data["info_score"] = score + filtered_messages.append(msg) + self.chat_messages = filtered_messages diff --git a/memory_scope/memory/worker/write/info_filter_worker.yaml b/memory_scope/memory/worker/write/info_filter_worker.yaml new file mode 100644 index 00000000..eed496fa --- /dev/null +++ b/memory_scope/memory/worker/write/info_filter_worker.yaml @@ -0,0 +1,63 @@ +info_filter_system: + cn: | + 任务指令:对所给{batch_size}个句子中所含有的关于{user_name}的信息打分,分数为0,1,2或3。 + 注意:其中0表示不包含{user_name}信息,1表示句子中只包含{user_name}假设的信息或者{user_name}虚构的内容,2表示包含{user_name}的一般信息,时效性信息或者需要猜测才能得到的{user_name}信息,3表示明确含有或者可以确定推断出关于{user_name}的重要信息,或者{user_name}要求记录。 + 按如下格式输出, 每一行输出一个打分,一定加<>,一共输出{batch_size}个分数: + 结果: + <分数:0或1或2或3> + +info_filter_few_shot: + cn: | + 示例1 + 句子: + 1 {user_name}:帮我写一段给同事张三女儿三岁生日的祝福语。 + 2 {user_name}:公元1400年至1550年中国历史大事表。 + 3 {user_name}:你吃午饭了吗? + 4 {user_name}:我今天心情不好,可以安慰我一下吗? + 5 {user_name}:能给我整理一张如何使用大模型的技巧列表吗,要求内容尽量精简。 + 6 {user_name}:记一下,明天下午3点提醒我去拿一下文件。 + 结果: + <3> + <0> + <0> + <2> + <2> + <3> + + 示例2 + 句子: + 1 {user_name}:我刚刚入职了阿里巴巴。 + 2 {user_name}:露天睡觉蚊子多,咋搞。 + 3 {user_name}:创造力和外倾性有关? + 4 {user_name}:一个区县的所有的事业人员的档案审核、修改和规范,应该是县委组织部下属的干部档案中心负责还是县人社局负责? + 5 {user_name}:假如我要和一个女人准备要孩子,我作为男人,怎么保护女人和孩子以及怎么备孕确保精子质量高对后代好 + 6 {user_name}:我和你一起出去玩,你会感觉开心吗? + 7 {user_name}:林浅,一位对未来充满好奇的年轻女孩,偶然间发现了这家能寄信给未来的邮局。出于对逝去祖父的怀念,她决定写下一封信,寄给五年后的自己,希望能收到祖父生前未说完的故事。五年期限将至,当她几乎忘记这段往事时,一封泛黄的回信悄然降临,不仅带来了祖父未完的冒险故事,还藏着一段关于勇气、爱与自我发现的深刻启示。续写成3000字小说。 + 结果: + <3> + <2> + <0> + <0> + <3> + <1> + <1> + + 示例3 + 句子: + 1 {user_name}:你的妈妈患有焦虑症,怎么安慰和开导她? + 2 {user_name}:肾脏严重亏空 + 3 {user_name}:我很喜欢打篮球,所以我身体很好 + 4 {user_name}:篮球明星有哪些? + 5 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续 + 结果: + <1> + <1> + <3> + <0> + <1> + +info_filter_user_query: + cn: | + 句子: + {user_query} + 结果: diff --git a/memory_scope/scheme/memory_node.py b/memory_scope/scheme/memory_node.py index 8572bec0..76a88370 100644 --- a/memory_scope/scheme/memory_node.py +++ b/memory_scope/scheme/memory_node.py @@ -11,6 +11,8 @@ class MemoryNode(BaseModel): user_id: str = Field("", description="unique memory id for user") + meta_data: Dict[str, str] = Field({}, description="other data infos") + content: str = Field("", description="memory content") score_similar: float = Field(0, description="es similar score") @@ -21,14 +23,21 @@ class MemoryNode(BaseModel): memory_type: str = Field("", description="conversation/observation/insight...") - meta_data: Dict[str, str] = Field({}, description="other data infos") - status: str = Field("active", description="active or expired") vector: List[float] = Field([], description="content embedding result, return empty") timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node") + obs_dt: str = Field("", description="dt of the observation") + + 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") + + keyword: str = Field("", description="keywords of the content") + + @property def node_keys(self): return list(self.model_json_schema()["properties"].keys()) diff --git a/memory_scope/storage/dummy_vector_store.py b/memory_scope/storage/dummy_vector_store.py index 290fe970..a766dccb 100644 --- a/memory_scope/storage/dummy_vector_store.py +++ b/memory_scope/storage/dummy_vector_store.py @@ -9,6 +9,7 @@ class DummyVectorStore(BaseVectorStore): def __init__(self, embedding_model: BaseModel, **kwargs): self.embedding_model: BaseModel = embedding_model + self.kwargs = kwargs def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]: pass diff --git a/memory_scope/utils/datetime_handler.py b/memory_scope/utils/datetime_handler.py new file mode 100644 index 00000000..9f3913cd --- /dev/null +++ b/memory_scope/utils/datetime_handler.py @@ -0,0 +1,78 @@ +import datetime +import re + +from memory_scope.constants.language_constants import WEEKDAYS +from memory_scope.utils.global_context import G_CONTEXT +from memory_scope.utils.logger import Logger + + +class DatetimeHandler(object): + + def __init__(self, dt: datetime.datetime | str | int | float = None): + if isinstance(dt, str | int | float): + if isinstance(dt, str): + dt = float(dt) + self._dt: datetime.datetime = datetime.datetime.fromtimestamp(dt) + elif isinstance(dt, datetime.datetime): + self._dt: datetime.datetime = dt + else: + self._dt: datetime.datetime = datetime.datetime.now() + + self._dt_info_dict: dict | None = None + self.logger = Logger.get_logger() + + def _parse_dt_info(self): + return { + "year": self._dt.year, + "month": self._dt.month, + "day": self._dt.day, + "hour": self._dt.hour, + "minute": self._dt.minute, + "second": self._dt.second, + "week": self._dt.isocalendar().week, + "weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1], + } + + @property + def dt_info_dict(self): + if self._dt_info_dict is None: + self._dt_info_dict = self._parse_dt_info() + return self._dt_info_dict + + @staticmethod + 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+)点' + } + weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7} + extracted_data = {} + + # Search for patterns in the input string and populate the dictionary + for key, pattern in patterns.items(): + match = re.search(pattern, input_string) + if match: # If there is a match, include it in the output dictionary + if match.group(1) == "每": + extracted_data[key] = -1 + elif match.group(1) in weekday_dict.keys(): + extracted_data[key] = weekday_dict[match.group(1)] + else: + extracted_data[key] = int(match.group(1)) + return extracted_data + + def extract_date_parts(self): + func_name = f"extract_date_parts_{G_CONTEXT.language}" + if not hasattr(self, func_name): + self.logger.warning(f"language={G_CONTEXT.language} needs to complete extract_date_parts function!") + return {} + return getattr(self, func_name)() + + def datetime_format(self, dt_format: str = "%Y%m%d"): + return self._dt.strftime(dt_format) + + def string_format(self, string_format: str): + return string_format.format(**self.dt_info_dict) diff --git a/memory_scope/utils/prompt_handler.py b/memory_scope/utils/prompt_handler.py index 490bc2f8..986022cc 100644 --- a/memory_scope/utils/prompt_handler.py +++ b/memory_scope/utils/prompt_handler.py @@ -5,30 +5,37 @@ from typing import Dict import yaml from memory_scope.utils.global_context import G_CONTEXT -from memory_scope.utils.tool_functions import camelcase_to_underscore class PromptHandler(object): - def __init__(self, default_prompt_dir: str = "config/prompts"): - self._default_prompt_dir: str = default_prompt_dir + def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs): + self._class_path: str = class_path self._prompt_dict: Dict[str, str] = {} - def add_file_prompts(self, name: str, to_underscore: bool = True): - if to_underscore: - name: str = camelcase_to_underscore(name) + file_path = self._class_path.strip(".py") + self.add_prompt_file(file_path) - class_path = os.path.join(self._default_prompt_dir, name) - if os.path.exists(f"{class_path}.yaml"): - with open(f"{class_path}.yaml") as f: - prompt_language_dict = yaml.load(f, yaml.FullLoader) - elif os.path.exists(f"{class_path}.json"): - with open(f"{class_path}.json") as f: - prompt_language_dict = json.load(f) + if prompt_file: + self.add_prompt_file(prompt_file) + + if prompt_dict: + self.add_prompt_dict(prompt_dict) + + def add_prompt_file(self, file_path: str): + if os.path.exists(f"{file_path}.yaml"): + with open(f"{file_path}.yaml") as f: + prompt_dict = yaml.load(f, yaml.FullLoader) + elif os.path.exists(f"{file_path}.json"): + with open(f"{file_path}.json") as f: + prompt_dict = json.load(f) else: - raise RuntimeError(f"{class_path}.yaml/json is not exists!") + raise RuntimeError(f"{file_path}.yaml/json is not exists!") - for key, language_dict in prompt_language_dict.items(): + self.add_prompt_dict(prompt_dict) + + def add_prompt_dict(self, prompt_dict: dict): + for key, language_dict in prompt_dict.items(): prompts = language_dict.get(G_CONTEXT.language) if not prompts: raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!") diff --git a/memory_scope/utils/tool_functions.py b/memory_scope/utils/tool_functions.py index 60c42e03..f4ca8b3d 100644 --- a/memory_scope/utils/tool_functions.py +++ b/memory_scope/utils/tool_functions.py @@ -9,8 +9,9 @@ from importlib import import_module import pyfiglet from termcolor import colored, COLORS -from memory_scope.constants.common_constants import WEEKDAYS +from memory_scope.constants.language_constants import WEEKDAYS from memory_scope.enumeration.message_role_enum import MessageRoleEnum +from memory_scope.utils.global_context import G_CONTEXT def underscore_to_camelcase(name: str, is_first_title: bool = True): @@ -25,11 +26,15 @@ def camelcase_to_underscore(name: str): return re.sub(r'(? "20240528", add %H:%M:%S - string_format: str = "") -> str: - - if isinstance(dt, str | int | float): - if isinstance(dt, str): - dt = float(dt) - current_dt = datetime.fromtimestamp(dt) - elif isinstance(dt, datetime): - current_dt = dt - else: - current_dt = datetime.now() - - return_str = "" - if date_format: - return_str = current_dt.strftime(date_format) - elif string_format: - return_str = string_format.format(**get_datetime_info_dict(current_dt)) - - return return_str - - def char_logo(words: str, seed: int = time.time_ns(), color=None): font = pyfiglet.Figlet() rendered_text = font.renderText(words)