fix update insight bug & modify contra repeat prompt

This commit is contained in:
jinli.yl 2024-07-18 17:34:35 +08:00
parent d8b3f6c7e6
commit b898fb5dbc
9 changed files with 125 additions and 31 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 = []

View file

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