[dev] modify status from MemoryNodeStatus to ActionStatusEnum

This commit is contained in:
jinli.yl 2024-07-13 20:42:43 +08:00
parent 44bf421487
commit 4c91212b16
10 changed files with 60 additions and 39 deletions

View file

@ -112,7 +112,7 @@ worker:
retrieve_not_reflected_top_k: 0
retrieve_not_updated_top_k: 0
retrieve_insight_top_k: 0
today_obs_top_k: 100
retrieve_today_top_k: 100
get_observation:
class: memory.worker.write.get_observation_worker
generation_model: dashscope_generation
@ -135,7 +135,7 @@ worker:
retrieve_not_reflected_top_k: 100
retrieve_not_updated_top_k: 100
retrieve_insight_top_k: 100
today_obs_top_k: 0
retrieve_today_top_k: 0
get_reflection_subject:
class: memory.worker.summary.get_reflection_subject_worker
retrieve_top_k: 100

View file

@ -2,7 +2,7 @@ from typing import List
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES
from memory_scope.constants.language_constants import COMMA_WORD
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
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
@ -30,14 +30,15 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
MemoryNode: A new MemoryNode instance representing the insight, marked as new and of type INSIGHT.
"""
dt_handler = DatetimeHandler()
meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()} # ⭐ Prepare metadata with current datetime info
# Prepare metadata with current datetime info
meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()}
return MemoryNode(user_name=self.user_name, # ⭐ Populate MemoryNode attributes
return MemoryNode(user_name=self.user_name,
target_name=self.target_name,
meta_data=meta_data,
key=insight_key,
memory_type=MemoryTypeEnum.INSIGHT.value,
status=MemoryNodeStatus.NEW.value)
action_status=ActionStatusEnum.NEW.value)
def _run(self):
"""
@ -58,7 +59,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
# Count unaudited nodes
not_reflected_count = len(not_reflected_nodes)
if not_reflected_count <= self.reflect_obs_cnt_threshold:
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stopping process.")
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, skip.")
self.continue_run = False
return

View file

@ -2,8 +2,9 @@ from typing import List, Dict
from memory_scope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_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.action_status_enum import ActionStatusEnum
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
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
@ -34,7 +35,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.ACTIVE.value,
"store_status": StoreStatusEnum.VALID.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value]
}
# Retrieve memories similar to the node's content, limited by top_k and filtered by filter_dict
@ -129,13 +130,13 @@ class LongContraRepeatWorker(MemoryBaseWorker):
node: MemoryNode = all_obs_nodes[idx]
if status == self.get_language_value(CONTRADICTORY_WORD):
if not content:
node.status = MemoryNodeStatus.EXPIRED.value
node.store_status = StoreStatusEnum.EXPIRED.value
else:
node.content = content
node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
node.action_status = ActionStatusEnum.CONTENT_MODIFIED.value
elif status == self.get_language_value(INCLUDED_WORD):
node.status = MemoryNodeStatus.EXPIRED.value
node.store_status = StoreStatusEnum.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after_long_contra_repeat: {node.content} {node.status}")

View file

@ -2,7 +2,7 @@ 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.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
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
@ -76,8 +76,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
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()
if insight_node.status == MemoryNodeStatus.ACTIVE.value:
insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
if insight_node.action_status == ActionStatusEnum.NONE.value:
insight_node.action_status = ActionStatusEnum.CONTENT_MODIFIED.value
self.logger.info(f"after_update_{insight_node.key} value={insight_value}")
return insight_node
@ -162,7 +162,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
# Process active insight nodes with corresponding not updated nodes
for node in insight_nodes:
if node.status == MemoryNodeStatus.ACTIVE.value:
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)

View file

@ -1,7 +1,8 @@
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_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.store_status_enum import StoreStatusEnum
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
@ -106,8 +107,8 @@ class ContraRepeatWorker(MemoryBaseWorker):
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != self.get_language_value(NONE_WORD):
node.status = MemoryNodeStatus.EXPIRED.value
self.logger.info(f"contra_repeat stage: {node.content} {node.status}")
node.store_status = StoreStatusEnum.EXPIRED.value
self.logger.info(f"contra_repeat stage: {node.content} {node.store_status} {node.action_status}")
merge_obs_nodes.append(node)
# save context

View file

@ -2,7 +2,7 @@ from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
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
@ -38,10 +38,8 @@ class GetObservationWorker(MemoryBaseWorker):
meta_data=meta_data,
content=obs_content,
memory_type=MemoryTypeEnum.OBSERVATION.value,
status=MemoryNodeStatus.NEW.value,
timestamp=message.time_created,
obs_reflected=False,
obs_updated=False)
action_status=ActionStatusEnum.NEW.value,
timestamp=message.time_created)
def filter_messages(self) -> List[Message]:
filter_messages = []
@ -134,6 +132,7 @@ class GetObservationWorker(MemoryBaseWorker):
idx, time_infer, obs_content, keywords = obs_content_list
# Skips processing if content indicates no meaningful observation
obs_content = obs_content.lower()
if obs_content in self.get_language_value([NONE_WORD, REPEATED_WORD]):
continue
@ -142,7 +141,7 @@ class GetObservationWorker(MemoryBaseWorker):
self.logger.warning(f"idx={idx} is invalid!")
continue
if time_infer == self.get_language_value(NONE_WORD):
if time_infer.lower() == self.get_language_value(NONE_WORD):
time_infer = ""
# Adjusts index to zero-based and checks validity against filtered messages

View file

@ -33,16 +33,18 @@ class InfoFilterWorker(MemoryBaseWorker):
# filter user msg
info_messages: List[Message] = []
for msg in self.chat_messages:
if msg.memorized:
continue
# TODO: add memory for assistant
if msg.role != MessageRoleEnum.USER.value:
continue
if len(msg.content) >= self.info_filter_msg_max_size:
half_size = int(self.info_filter_msg_max_size * 0.5 + 0.5)
msg.content = msg.content[: half_size] + msg.content[-half_size:]
info_messages.append(msg)
self.logger.warning(info_messages)
if not info_messages:
self.logger.warning("info_messages is empty!")
self.continue_run = False
@ -79,12 +81,25 @@ class InfoFilterWorker(MemoryBaseWorker):
# filter messages
filtered_messages: List[Message] = []
for msg, info_score in zip(info_messages, info_score_list):
for info_score in info_score_list:
if not info_score:
continue
score = info_score[0]
if len(info_score) != 2:
self.logger.warning(f"info_score={info_score} is invalid!")
continue
idx, score = info_score
idx = int(idx) - 1
if idx >= len(info_messages):
self.logger.warning(f"idx={idx} is invalid! info_messages.size={len(info_messages)}")
continue
message = info_messages[idx]
if score in self.preserved_scores:
msg.meta_data["info_score"] = score
filtered_messages.append(msg)
message.meta_data["info_score"] = score
filtered_messages.append(message)
self.logger.info(f"info filter stage: keep {message.content}")
self.chat_messages = filtered_messages

View file

@ -81,7 +81,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
"dt": dt_handler.datetime_format(),
}
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content,
top_k=self.today_obs_top_k,
top_k=self.retrieve_today_top_k,
filter_dict=filter_dict)
self.memory_handler.set_memories(TODAY_NODES, nodes)

View file

@ -1,6 +1,7 @@
from typing import Dict, List, Set
from memory_scope.enumeration.action_status_enum import ActionStatusEnum
from memory_scope.enumeration.store_status_enum import StoreStatusEnum
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_memory_store import BaseMemoryStore
from memory_scope.utils.global_context import G_CONTEXT
@ -123,6 +124,11 @@ class MemoryHandler(object):
if not nodes:
return
for node in nodes:
# Non-deleted expired memory nodes need to be changed to a modified state.
if node.store_status == StoreStatusEnum.EXPIRED.value and node.action_status != ActionStatusEnum.DELETE:
node.action_status = ActionStatusEnum.MODIFIED
# emb & insert new memories
new_memories = [n for n in nodes if n.action_status == ActionStatusEnum.NEW.value]
if new_memories:

View file

@ -1,6 +1,8 @@
import re
from memory_scope.utils.logger import Logger
from memory_scope.constants.language_constants import NONE_WORD
from memory_scope.utils.global_context import G_CONTEXT
class ResponseTextParser(object):
@ -30,19 +32,15 @@ class ResponseTextParser(object):
matches = [match.group(1) for match in self.pattern_v1.finditer(line)]
if matches:
result.append(matches)
self.logger.info(
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
)
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
return result
def parse_v2(self, prefix: str = ""):
result = []
for line in self.response_text.split("\n"):
line = line.strip()
if not line or line == "":
if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language):
continue
result.append(line)
self.logger.info(
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
)
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
return result