mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] add contra_repeat_worker config
This commit is contained in:
parent
4333e5b13b
commit
b151da319e
14 changed files with 289 additions and 203 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -64,3 +64,8 @@ INCLUDED_WORD = {
|
|||
LanguageEnum.CN: "被包含",
|
||||
LanguageEnum.EN: "included"
|
||||
}
|
||||
|
||||
COLON_WORD = {
|
||||
LanguageEnum.CN: ":",
|
||||
LanguageEnum.EN: ":"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ extract_time_prompt:
|
|||
回答:
|
||||
|
||||
|
||||
time_format_prompt:
|
||||
time_string_format:
|
||||
cn: |
|
||||
{year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
57
memory_scope/memory/worker/write/contra_repeat_worker.yaml
Normal file
57
memory_scope/memory/worker/write/contra_repeat_worker.yaml
Normal 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}
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue