[dev] fix es retrieve params

This commit is contained in:
jinli.yl 2024-07-01 23:17:39 +08:00
parent c2406e2eba
commit 654ba05ab9
4 changed files with 24 additions and 25 deletions

View file

@ -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] = []

View file

@ -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}")

View file

@ -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}")

View file

@ -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}")