[dev] add contra_repeat_worker config

This commit is contained in:
jinli.yl 2024-07-01 21:13:43 +08:00
parent 4333e5b13b
commit b151da319e
14 changed files with 289 additions and 203 deletions

View file

@ -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

View file

@ -64,3 +64,8 @@ INCLUDED_WORD = {
LanguageEnum.CN: "被包含",
LanguageEnum.EN: "included"
}
COLON_WORD = {
LanguageEnum.CN: "",
LanguageEnum.EN: ":"
}

View file

@ -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

View file

@ -6,7 +6,7 @@ extract_time_prompt:
回答:
time_format_prompt:
time_string_format:
cn: |
{year}年{month}月{day}日,{year}年第{week}周,{weekday}{hour}时{minute}分{second}秒。

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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 = {}