mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
fix update insight bug & modify contra repeat prompt
This commit is contained in:
parent
d8b3f6c7e6
commit
b898fb5dbc
9 changed files with 125 additions and 31 deletions
|
|
@ -148,6 +148,7 @@ worker:
|
|||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
rank_model: dashscope_rank
|
||||
long_contra_repeat:
|
||||
class: memory.worker.summary.long_contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.reflect_obs_cnt_threshold: int = kwargs.get("reflect_obs_cnt_threshold", 10)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
|
||||
self.reflect_num_questions: int = kwargs.get("reflect_num_questions", 5)
|
||||
|
||||
def new_insight_node(self, insight_key: str) -> MemoryNode:
|
||||
|
|
@ -88,7 +88,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# Invoke Language Model for new insights
|
||||
response = self.generation_model.call(messages=reflect_message, top_k=self.generation_model_top_k)
|
||||
response = self.generation_model.call(messages=reflect_message, **self.generation_model_kwargs)
|
||||
|
||||
# Handle empty response
|
||||
if not response.status or not response.message.content:
|
||||
|
|
@ -98,7 +98,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
new_insight_keys = ResponseTextParser(response.message.content).parse_v2(self.__class__.__name__)
|
||||
if new_insight_keys:
|
||||
for insight_key in new_insight_keys:
|
||||
insight_nodes.append(self.new_insight_node(insight_key))
|
||||
self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
|
||||
|
||||
# Mark unaudited nodes as reflected
|
||||
for node in not_reflected_nodes:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
FILE_PATH: str = __file__
|
||||
|
||||
def _parse_params(self, **kwargs):
|
||||
self.unit_test_flag = False
|
||||
self.long_contra_repeat_top_k: int = kwargs.get("long_contra_repeat_top_k", 2)
|
||||
self.long_contra_repeat_threshold: float = kwargs.get("long_contra_repeat_threshold", 0.1)
|
||||
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
|
||||
|
|
@ -66,16 +67,19 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
for node in not_updated_nodes:
|
||||
self.submit_thread_task(fn=self.retrieve_similar_content, node=node)
|
||||
|
||||
obs_node_dict: Dict[str, MemoryNode] = {}
|
||||
for origin_node, retrieve_nodes in self.gather_thread_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:
|
||||
if self.unit_test_flag:
|
||||
all_obs_nodes: List[MemoryNode] = not_updated_nodes
|
||||
else:
|
||||
obs_node_dict: Dict[str, MemoryNode] = {}
|
||||
for origin_node, retrieve_nodes in self.gather_thread_result():
|
||||
if not retrieve_nodes:
|
||||
continue
|
||||
obs_node_dict[node.memory_id] = node
|
||||
all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True)
|
||||
obs_node_dict[origin_node.memory_id] = origin_node
|
||||
for node in retrieve_nodes:
|
||||
if node.memory_id in obs_node_dict:
|
||||
continue
|
||||
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.logger.warning("all_obs_nodes is empty, stop.")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
long_contra_repeat_system:
|
||||
cn: |
|
||||
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
|
||||
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。
|
||||
注意:只判断与“前面序号”的句子的关系,不要判断“后面序号”。
|
||||
其中矛盾的形式可以有很多种,可以是逻辑上的矛盾,可以是属性上的变化导致的矛盾,比如不能同时在两个地方工作,同一个时刻不能在两个地点,同一个时刻不能干两件事情等等。
|
||||
对每个句子都做一个判断,最后一共输出{num_obs}条判断。如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
|
||||
请一步步思考,并按如下格式输出:
|
||||
思考:思考的依据和过程,30字以内。
|
||||
|
|
@ -102,3 +104,6 @@ long_contra_repeat_user_query:
|
|||
cn: |
|
||||
句子:
|
||||
{user_query}
|
||||
en: |
|
||||
Sentences:
|
||||
{user_query}
|
||||
|
|
|
|||
|
|
@ -44,8 +44,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
# Check if insight node key or value is empty and log a warning
|
||||
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!")
|
||||
if not insight_node.key:
|
||||
self.logger.warning(f"insight_key={insight_node.key} is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Call the ranking model to get scores for each observed node's content against the insight key
|
||||
|
|
@ -62,7 +62,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} insight_value={insight_node.value} "
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
# Warn if no nodes were filtered
|
||||
|
|
@ -74,8 +74,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
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}")
|
||||
key = self.prompt_handler.insight_string_format.format(name=self.target_name, key=insight_node.key)
|
||||
content = f"{key}{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()})
|
||||
|
|
@ -97,18 +97,17 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
Returns:
|
||||
MemoryNode: The updated MemoryNode with potentially revised insight value.
|
||||
"""
|
||||
self.logger.info(f"Updating insight for key={insight_node.key}, value={insight_node.value}, "
|
||||
self.logger.info(f"Updating insight for key={insight_node.key}, old_value={insight_node.value}, "
|
||||
f"with {len(filtered_nodes)} documents considered.")
|
||||
|
||||
# Generate the prompt for updating insight
|
||||
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(
|
||||
system_prompt = self.prompt_handler.update_insight_system.format(user_name=self.target_name)
|
||||
few_shot = self.prompt_handler.update_insight_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.update_insight_user_query.format(
|
||||
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)
|
||||
|
||||
# Construct the message for LLM interaction
|
||||
update_insight_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"Generated insight update message: {update_insight_message}")
|
||||
|
|
@ -170,11 +169,11 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
if node.action_status == ActionStatusEnum.NONE.value:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_updated_nodes)
|
||||
obs_nodes=not_updated_nodes)
|
||||
else:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_reflected_nodes)
|
||||
obs_nodes=not_reflected_nodes)
|
||||
|
||||
# select top n
|
||||
result_list = []
|
||||
|
|
@ -190,7 +189,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.submit_thread_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes)
|
||||
|
||||
# Gather the final results from all update tasks
|
||||
self.gather_thread_result()
|
||||
for _ in self.gather_thread_result():
|
||||
pass
|
||||
|
||||
for node in not_updated_nodes:
|
||||
node.obs_updated = 1
|
||||
|
|
|
|||
|
|
@ -105,3 +105,9 @@ update_insight_user_query:
|
|||
{user_query}
|
||||
Category: {insight_key}
|
||||
Existing information: {insight_key_value}
|
||||
|
||||
insight_string_format:
|
||||
cn: |
|
||||
{name}的{key}
|
||||
en: |
|
||||
The {key} of {name}
|
||||
|
|
|
|||
|
|
@ -103,4 +103,5 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt)
|
||||
|
||||
# Waits for all submitted tasks to complete
|
||||
self.gather_thread_result()
|
||||
for _ in self.gather_thread_result():
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -37,6 +37,23 @@ class MemoryHandler(object):
|
|||
self._id_memory_dict.clear()
|
||||
self._key_id_dict.clear()
|
||||
|
||||
def add_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
|
||||
if key not in self._key_id_dict:
|
||||
return self.set_memories(key, nodes, log_repeat)
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
for node in nodes:
|
||||
_id = node.memory_id
|
||||
if _id not in self._key_id_dict[key]:
|
||||
self._key_id_dict[key].append(_id)
|
||||
|
||||
if _id not in self._id_memory_dict:
|
||||
self._id_memory_dict[_id] = node
|
||||
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content or node.key} "
|
||||
f"store_status={node.store_status} action_status={node.action_status}")
|
||||
|
||||
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
|
||||
if nodes is None:
|
||||
nodes = []
|
||||
|
|
|
|||
|
|
@ -3,11 +3,12 @@ import unittest
|
|||
|
||||
from memory_scope.cli import MemoryScope
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
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 import Logger
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
|
@ -16,9 +17,16 @@ class TestWorkersCn(unittest.TestCase):
|
|||
"""Tests for LLIEmbedding"""
|
||||
|
||||
def setUp(self):
|
||||
ms = MemoryScope().load_config("config/demo_config_cn.yaml")
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
|
||||
|
||||
ms = MemoryScope()
|
||||
ms.load_config("config/demo_config_cn.yaml")
|
||||
ms.init_global_content_by_config()
|
||||
|
||||
def tearDown(self):
|
||||
self.logger.close()
|
||||
|
||||
@unittest.skip
|
||||
def test_extract_time(self):
|
||||
name = "extract_time"
|
||||
|
|
@ -257,7 +265,7 @@ class TestWorkersCn(unittest.TestCase):
|
|||
worker.logger.info(f"result3={result3}")
|
||||
worker.logger.info(f"result4={result4}")
|
||||
|
||||
# @unittest.skip
|
||||
@unittest.skip
|
||||
def test_get_reflection_subject(self):
|
||||
name = "get_reflection_subject"
|
||||
|
||||
|
|
@ -287,6 +295,58 @@ class TestWorkersCn(unittest.TestCase):
|
|||
worker.memory_handler.set_memories(INSIGHT_NODES, [])
|
||||
worker.run()
|
||||
|
||||
result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
|
||||
result = "\n".join(result)
|
||||
worker.logger.info(f"result.get_reflection={result}")
|
||||
return worker
|
||||
|
||||
@unittest.skip
|
||||
def test_update_insight_worker(self):
|
||||
reflection_worker = self.test_get_reflection_subject.__wrapped__(self)
|
||||
|
||||
name = "update_insight"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context=reflection_worker.context,
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="用户喜欢打王者荣耀"),
|
||||
]
|
||||
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
worker.run()
|
||||
|
||||
result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
|
||||
result = "\n".join(result)
|
||||
worker.logger.info(f"result={result}")
|
||||
worker.logger.info(f"result.update_insight={result}")
|
||||
|
||||
# @unittest.skip
|
||||
def test_long_contra_repeat_worker(self):
|
||||
name = "long_contra_repeat"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
|
||||
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
|
||||
MemoryNode(content="用户在上海工作。"),
|
||||
]
|
||||
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
worker.unit_test_flag = True
|
||||
worker.run()
|
||||
|
||||
result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]
|
||||
result = "\n".join(result)
|
||||
worker.logger.info(f"result.long_contra_repeat={result}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue