mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
move get reflection worker retrieve logic to another worker
This commit is contained in:
parent
d53cb3580c
commit
dadc2ef771
8 changed files with 294 additions and 282 deletions
|
|
@ -22,6 +22,8 @@ DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
|
|||
|
||||
NOT_REFLECTED_NODES = "not_reflected_nodes"
|
||||
|
||||
NOT_UPDATED_NODES = "not_updated_nodes"
|
||||
|
||||
|
||||
MODIFIED_MEMORIES = "modified_memories"
|
||||
|
||||
|
|
|
|||
|
|
@ -24,14 +24,19 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.kwargs: dict = kwargs
|
||||
|
||||
self.continue_run: bool = True
|
||||
self.task_list: list = []
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
@staticmethod
|
||||
def async_run(fn_list, *args, **kwargs):
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn in fn_list])
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
self.task_list.append((fn, args, kwargs))
|
||||
|
||||
return asyncio.run(async_gather())
|
||||
def gather_async_result(self):
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list])
|
||||
|
||||
results = asyncio.run(async_gather())
|
||||
self.task_list.clear()
|
||||
return results
|
||||
|
||||
@abstractmethod
|
||||
def _run(self):
|
||||
|
|
|
|||
|
|
@ -8,31 +8,10 @@ from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
|||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.timer import timer
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
@timer
|
||||
def retrieve_not_reflected_memory(self, query: str) -> List[MemoryNode]:
|
||||
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_reflected": False,
|
||||
}
|
||||
return self.vector_store.retrieve(query=query, top_k=self.retrieve_not_reflected_top_k, filter_dict=filter_dict)
|
||||
|
||||
@timer
|
||||
def retrieve_insight_memory(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
return self.vector_store.retrieve(query=query, top_k=self.retrieve_insight_top_k, filter_dict=filter_dict)
|
||||
|
||||
def new_insight_node(self, insight_key: str) -> MemoryNode:
|
||||
dt_handler = DatetimeHandler()
|
||||
|
|
@ -46,20 +25,8 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
status=MemoryNodeStatus.ACTIVE.value)
|
||||
|
||||
def _run(self):
|
||||
not_reflected_nodes: List[MemoryNode] = []
|
||||
insight_nodes: List[MemoryNode] = []
|
||||
|
||||
fn_list = [self.retrieve_not_reflected_memory, self.retrieve_insight_memory]
|
||||
for nodes in self.async_run(fn_list=fn_list, query="_"):
|
||||
if not nodes:
|
||||
continue
|
||||
for node in nodes:
|
||||
if node.memory_type == MemoryTypeEnum.INSIGHT.value:
|
||||
insight_nodes.append(node)
|
||||
else:
|
||||
not_reflected_nodes.append(node)
|
||||
|
||||
self.set_context(INSIGHT_NODES, insight_nodes)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_nodes)
|
||||
|
|
@ -67,9 +34,6 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop.")
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_NODES, not_reflected_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = [n.key for n in insight_nodes]
|
||||
self.logger.info(f"exist_keys={exist_keys}")
|
||||
|
|
@ -99,3 +63,6 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
if new_insight_keys:
|
||||
for insight_key in new_insight_keys:
|
||||
insight_nodes.append(self.new_insight_node(insight_key))
|
||||
|
||||
for node in not_reflected_nodes:
|
||||
node.obs_reflected = True
|
||||
|
|
|
|||
60
memory_scope/memory/worker/summary/load_memory_worker.py
Normal file
60
memory_scope/memory/worker/summary/load_memory_worker.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES
|
||||
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.utils.timer import timer
|
||||
|
||||
|
||||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
async def retrieve_not_reflected_memory(self, query: str):
|
||||
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_reflected": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
async def retrieve_not_updated_memory(self, query: str):
|
||||
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_updated": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
async def retrieve_insight_memory(self, query: str):
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(INSIGHT_NODES, nodes)
|
||||
|
||||
async def _run(self):
|
||||
mock_query = "-"
|
||||
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_insight_memory, query=mock_query)
|
||||
|
||||
self.gather_async_result()
|
||||
|
|
@ -1,12 +1,56 @@
|
|||
systemp_prompt:
|
||||
cn:
|
||||
en:
|
||||
long_contra_repeat_system:
|
||||
cn: |
|
||||
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
|
||||
如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
|
||||
请一步步思考,并按如下格式输出:
|
||||
思考:思考的依据和过程,30字以内。
|
||||
判断:<句子序号> <矛盾,被包含,无> <修改后的内容>,一定加<>
|
||||
|
||||
few_shot_prompt:
|
||||
cn:
|
||||
en:
|
||||
|
||||
user_query_prompt:
|
||||
cn:
|
||||
en:
|
||||
long_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> <矛盾> <{user_name}的父亲生日在2024年6月2日>
|
||||
思考:第5句与前面序号句子既不矛盾也不重复。
|
||||
判断:<5> <无> <>
|
||||
思考:第6句中所有信息都被前面序号中第5句的信息完全包含。
|
||||
判断:<2> <被包含> <>
|
||||
|
||||
long_contra_repeat_user_query:
|
||||
cn: |
|
||||
句子:
|
||||
{user_query}
|
||||
|
|
|
|||
|
|
@ -1,87 +1,68 @@
|
|||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
MSG_TIME,
|
||||
MODIFIED_MEMORIES,
|
||||
MODIFIED_MEMORIES, NOT_UPDATED_NODES,
|
||||
)
|
||||
from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_node.content
|
||||
related_nodes = self.vector_store.similar_search(
|
||||
text=text,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
async def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]):
|
||||
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]
|
||||
}
|
||||
retrieve_nodes = await self.vector_store.async_retrieve(query=node.content,
|
||||
top_k=self.long_contra_repeat_top_k,
|
||||
filter_dict=filter_dict)
|
||||
return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold]
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
def _run(self):
|
||||
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
|
||||
for node in not_updated_nodes:
|
||||
self.submit_async_task(fn=self.retrieve_similar_content, node=node)
|
||||
|
||||
obs_node_dict: Dict[str, MemoryNode] = {}
|
||||
for origin_node, retrieve_nodes in self.gather_async_result():
|
||||
if not retrieve_nodes:
|
||||
continue
|
||||
obs_node_dict[origin_node.memory_id] = origin_node
|
||||
for node in retrieve_nodes:
|
||||
if node.memory_id in obs_node_dict:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
all_obs_nodes.append(related_node)
|
||||
if has_match:
|
||||
all_obs_nodes.append(new_obs_node)
|
||||
obs_node_dict[node.memory_id] = node
|
||||
all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True)
|
||||
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
self.logger.warning("all_obs_nodes is empty, stop.")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(
|
||||
all_obs_nodes,
|
||||
key=lambda x: x.meta_data.get(MSG_TIME, ""),
|
||||
reverse=True,
|
||||
)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
merge_obs_message = prompt_to_msg(
|
||||
system_prompt=self.prompt_handler.system_prompt.format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.prompt_handler.few_shot_prompt,
|
||||
user_query=self.prompt_handler.user_query_prompt.format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
system_prompt = self.prompt_handler.long_contra_repeat_system.format(num_obs=len(user_query_list),
|
||||
user_name=self.target_name)
|
||||
few_shot = self.prompt_handler.long_contra_repeat_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list))
|
||||
|
||||
long_contra_repeat_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"long_contra_repeat_message={long_contra_repeat_message}")
|
||||
|
||||
# call LLM
|
||||
response = self.generation_model.call(
|
||||
messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k,
|
||||
)
|
||||
response = self.generation_model.call(messages=long_contra_repeat_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
|
|
|
|||
|
|
@ -1,64 +1,56 @@
|
|||
system_prompt:
|
||||
cn: "从下面的句子中提取出给定类别的用户资料信息,并判断与已有信息是否矛盾,若矛盾以新信息为准。整合已有信息和新信息并输出。若已有信息为空则直接输出提取的
|
||||
信息。若无法提取出给定类别的用户资料信息则输出“无”。
|
||||
请一步步思考,并按如下格式输出:
|
||||
思考: 思考的依据和过程,150字以内。
|
||||
用户资料: <信息或无>, 一定加<>"
|
||||
en:
|
||||
update_insight_system:
|
||||
cn: |
|
||||
从下面的句子中提取出给定类别的{user_name}的资料信息,并判断与已有信息是否矛盾,若矛盾以新信息为准。整合已有信息和新信息并输出。
|
||||
请一步步思考,并按如下格式输出:
|
||||
思考: 思考的依据和过程,150字以内。
|
||||
{user_name}的资料: <信息>, 一定加<>
|
||||
|
||||
few_shot_promp:
|
||||
cn: "
|
||||
示例1:
|
||||
句子:因为昨天成都下大雨,用户全身都被淋湿了。
|
||||
句子:用户关心明天成都的天气预报。
|
||||
类别:用户所在地区
|
||||
已有信息:用户所在地区:
|
||||
思考:从第一句句子可以得出用户在成都。第二句句子没有直接透露用户所在地信息,但与第一句句子用户在成都的信息吻合。已有信息为空,直接输出得出的信息。
|
||||
用户资料:<成都>
|
||||
|
||||
示例2:
|
||||
句子:用户最近养好了肠胃。
|
||||
句子:用户关注中医养生。
|
||||
类别:用户健康状况
|
||||
已有信息:用户健康状况: 肠胃不好,高血压
|
||||
思考:从第一句句子可以得出用户最近养好了肠胃,与已有信息矛盾,以新信息为准。第二句句子与用户健康状况无关。整合已有信息和新信息得到用户健康状况是肠胃健康,高血压。
|
||||
用户资料:<肠胃健康,高血压>
|
||||
update_insight_few_shot:
|
||||
cn: |
|
||||
示例1:
|
||||
因为昨天成都下大雨,{user_name}全身都被淋湿了。
|
||||
{user_name}关心明天成都的天气预报。
|
||||
类别:{user_name}所在地区
|
||||
已有信息:{user_name}所在地区: 杭州
|
||||
思考:从第一句句子可以得出{user_name}在成都。第二句句子没有直接透露{user_name}所在地信息,但与第一句句子{user_name}在成都的信息吻合。这与已有信息({user_name}在杭州)矛盾,输出更新的信息。
|
||||
{user_name}的资料:<成都>
|
||||
|
||||
示例2:
|
||||
{user_name}最近养好了肠胃。
|
||||
{user_name}关注中医养生。
|
||||
类别:{user_name}健康状况
|
||||
已有信息:{user_name}健康状况: 肠胃不好,高血压
|
||||
思考:从第一句句子可以得出{user_name}最近养好了肠胃,与已有信息矛盾,以新信息为准。第二句句子与{user_name}健康状况无关。整合已有信息和新信息得到{user_name}健康状况是肠胃健康,高血压。
|
||||
{user_name}的资料:<肠胃健康,高血压>
|
||||
|
||||
示例3:
|
||||
{user_name}刚刚毕业,第一份工作是银行前台。
|
||||
{user_name}的理想工作是职业游戏选手。
|
||||
类别:{user_name}职业
|
||||
已有信息:{user_name}职业:在招商银行工作
|
||||
思考:整合已有信息和第一句句子的信息可以得出{user_name}的现在的职业是招商银行前台。第二句句子说明了{user_name}的理想工作但并不是现在的职业。
|
||||
{user_name}的资料:<招商银行前台>
|
||||
|
||||
示例4:
|
||||
{user_name}大学期间接触过优化算法的研究。
|
||||
类别:{user_name}学习专业
|
||||
已有信息:{user_name}学习专业:与人工智能相关
|
||||
思考:从句子可以得出{user_name}大学学习的专业与优化算法相关,这与已有信息({user_name}学习专业与人工智能相关)不矛盾,整合可以得出{user_name}大学学习的专业与人工智能和优化算法相关。
|
||||
{user_name}的资料:<与人工智能和优化算法相关>
|
||||
|
||||
示例5:
|
||||
{user_name}单身。
|
||||
{user_name}受到一名18岁男生的追求,但不想接受又不想伤害他。
|
||||
{user_name}喜欢成熟且情绪稳定的男生。
|
||||
类别:{user_name}情感状况
|
||||
已有信息:{user_name}情感状况:有男朋友
|
||||
思考:从第一句句子可以得出{user_name}现在单身,与已有信息矛盾,以新信息为准。从第二句句子得出{user_name}受到一名18岁男生的追求但并不喜欢他。第三句话表达了{user_name}理想的伴侣类型但与{user_name}
|
||||
情感状况无关。整合得出{user_name}情感状况为单身,受到一名18岁男生的追求但并不喜欢他。
|
||||
{user_name}的资料:<单身,受到一名18岁男生的追求但并不喜欢他。>
|
||||
|
||||
示例3:
|
||||
句子:用户刚刚毕业,第一份工作是银行前台。
|
||||
句子:用户的理想工作是职业游戏选手。
|
||||
类别:用户职业
|
||||
已有信息:用户职业:在招商银行工作
|
||||
思考:整合已有信息和第一句句子的信息可以得出用户的现在的职业是招商银行前台。第二句句子说明了用户的理想工作但并不是现在的职业。
|
||||
用户资料:<招商银行前台>
|
||||
|
||||
示例4:
|
||||
句子:用户大学期间接触过优化算法的研究。
|
||||
类别:用户学习专业
|
||||
已有信息:用户学习专业:与人工智能相关
|
||||
思考:从句子可以得出用户大学学习的专业与优化算法相关,这与已有信息(用户学习专业与人工智能相关)不矛盾,整合可以得出用户大学学习的专业与人工智能和优化算法相关。
|
||||
用户资料:<与人工智能和优化算法相关>
|
||||
|
||||
示例5:
|
||||
句子:用户单身。
|
||||
句子:用户受到一名18岁男生的追求,但不想接受又不想伤害他。
|
||||
句子:用户喜欢成熟且情绪稳定的男生。
|
||||
类别:用户情感状况
|
||||
已有信息:用户情感状况:有男朋友
|
||||
思考:从第一句句子可以得出用户现在单身,与已有信息矛盾,以新信息为准。从第二句句子得出用户受到一名18岁男生的追求但并不喜欢他。第三句话表达了用户理想的伴侣类型但与用户
|
||||
情感状况无关。整合得出用户情感状况为单身,受到一名18岁男生的追求但并不喜欢他。
|
||||
用户资料:<单身,受到一名18岁男生的追求但并不喜欢他。>
|
||||
"
|
||||
en:
|
||||
|
||||
user_query_prompt:
|
||||
cn: "
|
||||
{user_query}
|
||||
类别:{insight_key}
|
||||
已有信息:{insight_key_value}
|
||||
"
|
||||
en:
|
||||
|
||||
user_query:
|
||||
cn: "句子:{content}"
|
||||
|
||||
update_insight_user_query:
|
||||
cn: |
|
||||
{user_query}
|
||||
类别:{insight_key}
|
||||
已有信息:{insight_key_value}
|
||||
|
|
|
|||
|
|
@ -1,167 +1,128 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES
|
||||
from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.response_text_parser import ResponseTextParser
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
||||
def filter_obs_nodes(
|
||||
self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode]
|
||||
) -> (MemoryNode, List[MemoryNode], float):
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_node.insight_key
|
||||
insight_value = insight_node.insight_value
|
||||
if not insight_key or not insight_value:
|
||||
self.logger.warning(
|
||||
f"insight_key={insight_key} insight_value={insight_value} is empty!"
|
||||
)
|
||||
if not insight_node.key or not insight_node.value:
|
||||
self.logger.warning(f"insight_key={insight_node.key} insight_value={insight_node.value} is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
result = self.rank_model.call(
|
||||
query=insight_key, documents=[x.content for x in new_obs_nodes]
|
||||
)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
|
||||
response = self.rank_model.call(query=insight_node.key, documents=[x.content for x in obs_nodes])
|
||||
if not response.status:
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
|
||||
for index, score in result.rank_scores.items():
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_insight_threshold:
|
||||
# find nodes related to query
|
||||
for index, score in response.rank_scores.items():
|
||||
node = obs_nodes[index]
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(
|
||||
f"insight_key={insight_key} insight_value={insight_value} "
|
||||
f"score={score} keep_flag={keep_flag}"
|
||||
)
|
||||
self.logger.info(f"insight_key={insight_node.key} insight_value={insight_node.value} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
|
||||
|
||||
self.logger.warning(f"update_insight={insight_node.key} filtered_nodes is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(
|
||||
self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]
|
||||
) -> MemoryNode:
|
||||
def update_insight_node(self, insight_node: MemoryNode, insight_value: str):
|
||||
dt_handler = DatetimeHandler()
|
||||
content = (f"{self.user_name}{self.get_language_value(COLON_WORD)}{insight_node.key}"
|
||||
f"{self.get_language_value(COLON_WORD)}{insight_value}")
|
||||
insight_node.content = content
|
||||
insight_node.value = insight_value
|
||||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
|
||||
insight_node.timestamp = dt_handler.timestamp
|
||||
insight_node.dt = dt_handler.datetime_format()
|
||||
self.logger.info(f"after_update_{insight_node.key} value={insight_value}")
|
||||
return insight_node
|
||||
|
||||
insight_key = insight_node.insight_key
|
||||
insight_value = insight_node.insight_value
|
||||
self.logger.info(
|
||||
f"update_insight insight_key={insight_key} insight_value={insight_value} "
|
||||
f"doc.size={len(filtered_nodes)}"
|
||||
)
|
||||
def update_insight(self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]) -> MemoryNode:
|
||||
self.logger.info(f"update_insight insight_key={insight_node.key} insight_value={insight_node.value} "
|
||||
f"doc.size={len(filtered_nodes)}")
|
||||
|
||||
# gen prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(self.prompt_handler.user_query.format(content=node.content))
|
||||
update_insight_message = prompt_to_msg(
|
||||
system_prompt=self.prompt_handler.system_prompt,
|
||||
few_shot=self.prompt_handler.few_shot_prompt,
|
||||
user_query=self.prompt_handler.user_query_prompt.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key + self.get_language_value(COLON_WORD) + insight_value,
|
||||
),
|
||||
)
|
||||
user_query_list = [n.content for n in filtered_nodes]
|
||||
system_prompt = self.prompt_handler.update_insight_system.foramt(user_name=self.target_name)
|
||||
few_shot = self.prompt_handler.update_insight_few_shot.foramt(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.update_insight_user_query.foramt(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_node.key,
|
||||
insight_key_value=insight_node.key + self.get_language_value(COLON_WORD) + insight_node.value)
|
||||
|
||||
update_insight_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"update_insight_message={update_insight_message}")
|
||||
|
||||
# call LLM
|
||||
response: str = self.generation_model.call(
|
||||
messages=update_insight_message,
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k,
|
||||
)
|
||||
response = self.generation_model.call(messages=update_insight_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} call llm failed!"
|
||||
)
|
||||
if not response.status or not response.message.content:
|
||||
return insight_node
|
||||
|
||||
profile_list = ResponseTextParser(response.message.content).parse_v1(
|
||||
f"update_profile {insight_key}"
|
||||
)
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 1!"
|
||||
)
|
||||
insight_value_list = ResponseTextParser(response.message.content).parse_v1(f"update_{insight_node.key}")
|
||||
if not insight_value_list:
|
||||
self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!")
|
||||
return insight_node
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 2"
|
||||
)
|
||||
return insight_node
|
||||
insight_value = profile_list[0]
|
||||
|
||||
insight_value_list = insight_value_list[0]
|
||||
if not insight_value_list:
|
||||
self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!")
|
||||
return insight_node
|
||||
|
||||
insight_value = insight_value_list[0]
|
||||
if not insight_value or insight_value in self.get_language_value([NONE_WORD, REPEATED_WORD]):
|
||||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
self.logger.info(f"update_{insight_node.key} insight_value={insight_value} is invalid.")
|
||||
return insight_node
|
||||
|
||||
insight_node.insight_value = insight_value
|
||||
insight_node.obs_updated = True
|
||||
self.update_insight_node(insight_node=insight_node, insight_value=insight_value)
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop update sights!")
|
||||
return
|
||||
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
|
||||
if not insight_nodes:
|
||||
self.logger.info("insight_nodes is empty, stop update sights!")
|
||||
self.logger.warning("insight_nodes is empty, stop.")
|
||||
return
|
||||
|
||||
# 提交打分任务
|
||||
for node in insight_nodes:
|
||||
self.submit_thread(
|
||||
self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
insight_node=node,
|
||||
new_obs_nodes=new_obs_nodes,
|
||||
)
|
||||
if node.content:
|
||||
self.submit_async_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_updated_nodes)
|
||||
else:
|
||||
self.submit_async_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_reflected_nodes)
|
||||
|
||||
# 选择topN
|
||||
# select top n
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
for result in self.gather_async_result():
|
||||
insight_node, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[: self.update_insight_max_thread]
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)[: self.update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
# submit llm update task
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(
|
||||
self.update_insight,
|
||||
sleep_time=1,
|
||||
insight_node=insight_node,
|
||||
filtered_nodes=filtered_nodes,
|
||||
)
|
||||
self.submit_async_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes)
|
||||
|
||||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryNode = result
|
||||
insight_key = insight_node.insight_key
|
||||
insight_value = insight_node.insight_value
|
||||
self.logger.info(
|
||||
f"after_update_insight insight_key={insight_key} insight_value={insight_value}"
|
||||
)
|
||||
# get result
|
||||
self.gather_async_result()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue