From 654ba05ab988b62f0817034a30bd0f4b3a162dee Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 1 Jul 2024 23:17:39 +0800 Subject: [PATCH] [dev] fix es retrieve params --- .../memory/worker/write/contra_repeat_worker.py | 17 ++++++++--------- .../write/get_observation_with_time_worker.py | 8 ++++---- .../worker/write/get_observation_worker.py | 10 +++++----- .../memory/worker/write/info_filter_worker.py | 14 +++++++------- 4 files changed, 24 insertions(+), 25 deletions(-) diff --git a/memory_scope/memory/worker/write/contra_repeat_worker.py b/memory_scope/memory/worker/write/contra_repeat_worker.py index 1a7b7212..aec0d71e 100644 --- a/memory_scope/memory/worker/write/contra_repeat_worker.py +++ b/memory_scope/memory/worker/write/contra_repeat_worker.py @@ -22,15 +22,14 @@ class ContraRepeatWorker(MemoryBaseWorker): 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(), - }) + filter_dict = { + "user_name": self.user_name, + "target_name": self.target_name, + "status": MemoryNodeStatus.ACTIVE.value, + "memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value], + "obs_dt": dt_handler.datetime_format(), + } + return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict) def _run(self): all_obs_nodes: List[MemoryNode] = [] 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 e54e2a81..71c3692a 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 @@ -24,7 +24,7 @@ class GetObservationWithTimeWorker(GetObservationWorker): if match: 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}") + user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") i += 1 if not user_query_list: @@ -32,11 +32,11 @@ class GetObservationWithTimeWorker(GetObservationWorker): return [] 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_name=self.target_name) + few_shot = self.prompt_config.get_observation_with_time_few_shot.format(user_name=self.target_name) user_query = self.prompt_config.get_observation_with_time_user_query.format( user_query="\n".join(user_query_list), - user_name=self.user_id) + user_name=self.target_name) 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}") diff --git a/memory_scope/memory/worker/write/get_observation_worker.py b/memory_scope/memory/worker/write/get_observation_worker.py index b9897a4c..fc6f9e44 100644 --- a/memory_scope/memory/worker/write/get_observation_worker.py +++ b/memory_scope/memory/worker/write/get_observation_worker.py @@ -20,7 +20,7 @@ class GetObservationWorker(MemoryBaseWorker): meta_data = { MemoryTypeEnum.CONVERSATION.value: message.content, TIME_INFER: time_infer, - **dt_handler.dt_info_dict.items(), + **dt_handler.dt_info_dict, } if time_infer: @@ -52,7 +52,7 @@ class GetObservationWorker(MemoryBaseWorker): match = True break if not match: - user_query_list.append(f"{i} {self.user_id}{self.get_language_value(COLON_WORD)}{msg.content}") + user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}") i += 1 if not user_query_list: @@ -60,10 +60,10 @@ class GetObservationWorker(MemoryBaseWorker): 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_name=self.target_name) + few_shot = self.prompt_config.get_observation_few_shot.format(user_name=self.target_name) user_query = self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list), - user_name=self.user_id) + user_name=self.target_name) 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}") diff --git a/memory_scope/memory/worker/write/info_filter_worker.py b/memory_scope/memory/worker/write/info_filter_worker.py index 3949349e..a085cd48 100644 --- a/memory_scope/memory/worker/write/info_filter_worker.py +++ b/memory_scope/memory/worker/write/info_filter_worker.py @@ -18,9 +18,8 @@ class InfoFilterWorker(MemoryBaseWorker): 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:] + half_size = int(self.info_filter_msg_max_size * 0.5 + 0.5) + msg.content = msg.content[: half_size] + msg.content[-half_size:] info_messages.append(msg) if not info_messages: @@ -31,11 +30,12 @@ class InfoFilterWorker(MemoryBaseWorker): # 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}") + user_query_list.append(f"{i + 1} {self.target_name}{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="\n".join(user_query_list)) + user_name=self.target_name) + few_shot = self.prompt_handler.info_filter_few_shot.format(user_name=self.target_name) + user_query = self.prompt_handler.info_filter_user_query.format(user_query="\n".join(user_query_list), + user_name=self.target_name) 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}")