mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] add memory handler and get memories
This commit is contained in:
parent
e758cf41e6
commit
bc5278ffa0
5 changed files with 55 additions and 8 deletions
|
|
@ -4,6 +4,7 @@ from typing import List, Dict
|
|||
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
|
|
@ -31,6 +32,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._target_name: str | None = None
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
|
||||
self._contex_memory_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
@property
|
||||
def chat_messages(self) -> List[Message]:
|
||||
return self.get_context(CHAT_MESSAGES)
|
||||
|
|
@ -67,6 +70,26 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._memory_store = G_CONTEXT.memory_store
|
||||
return self._memory_store
|
||||
|
||||
def get_memories(self, key: str) -> List[MemoryNode]:
|
||||
memories: List[MemoryNode] = []
|
||||
memory_ids: List[str] = self.get_context(key)
|
||||
if memory_ids:
|
||||
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
|
||||
return memories
|
||||
|
||||
def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode):
|
||||
if not nodes:
|
||||
return
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
for node in nodes:
|
||||
if node.memory_id in self._contex_memory_dict:
|
||||
continue
|
||||
self._contex_memory_dict[node.memory_id] = node
|
||||
self.set_context(key, [n.memory_id for n in nodes])
|
||||
|
||||
@property
|
||||
def monitor(self) -> BaseMonitor:
|
||||
if self._monitor is None:
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from memory_scope.utils.datetime_handler import DatetimeHandler
|
|||
class PrintMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
|
||||
|
||||
obs_content_list: List[str] = []
|
||||
|
|
@ -29,13 +29,20 @@ class PrintMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
obs_content = "\n".join(obs_content_list)
|
||||
insight_content = "\n".join(insight_content_list)
|
||||
expired_content = "\n".join()
|
||||
result: str = f"""
|
||||
The memories of {self.user_name} about {self.target_name}.
|
||||
|
||||
observation:
|
||||
----- observation -----
|
||||
{obs_content}
|
||||
----- observation -----
|
||||
|
||||
insight:
|
||||
----- insight -----
|
||||
{insight_content}
|
||||
----- insight -----
|
||||
|
||||
----- expired -----
|
||||
{expired_content}
|
||||
----- expired -----
|
||||
""".strip()
|
||||
self.set_context(RESULT, result)
|
||||
|
|
|
|||
|
|
@ -45,4 +45,4 @@ class RetrieveStoreWorker(MemoryBaseWorker):
|
|||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}")
|
||||
self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
# query
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
self.logger.warning(f"retrieve memory nodes is empty!")
|
||||
return
|
||||
|
|
@ -33,4 +33,4 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
|
||||
self.get_context(RANKED_MEMORY_NODES, memory_node_list)
|
||||
self.set_memories(RANKED_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \
|
||||
NOT_UPDATED_NODES
|
||||
|
|
@ -9,6 +9,23 @@ from memory_scope.scheme.memory_node import MemoryNode
|
|||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
update_memories: Dict[str, MemoryNode] = {}
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if insight_nodes:
|
||||
update_memories.update({n.memory_id: n for n in insight_nodes})
|
||||
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
if not_reflected_nodes:
|
||||
update_memories.update({n.memory_id: n for n in not_reflected_nodes})
|
||||
|
||||
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
|
||||
if not_updated_nodes:
|
||||
for node in not_updated_nodes:
|
||||
if node.memory_id in update_memories:
|
||||
|
||||
|
||||
|
||||
|
||||
keys = [
|
||||
INSIGHT_NODES,
|
||||
MERGE_OBS_NODES,
|
||||
|
|
@ -20,4 +37,4 @@ class SummaryCollectWorker(MemoryBaseWorker):
|
|||
for key in keys:
|
||||
memory_nodes.extend(self.get_context(key))
|
||||
|
||||
self.memory_store.update_memories(memory_nodes)
|
||||
self.memory_store.update_memories(update_memories)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue