mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] fix es retrieve params
This commit is contained in:
parent
c2406e2eba
commit
654ba05ab9
4 changed files with 24 additions and 25 deletions
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue