mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[dev] rename constant name & add _async_run
This commit is contained in:
parent
d8aceabbf0
commit
7751586435
6 changed files with 61 additions and 4 deletions
|
|
@ -65,4 +65,8 @@ worker:
|
|||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ CHAT_KWARGS = "chat_kwargs"
|
|||
|
||||
QUERY_WITH_TS = "query_with_ts"
|
||||
|
||||
RETRIEVE_MEMORY_NODES = "RETRIEVE_MEMORY_NODES"
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Any, Dict
|
||||
|
||||
|
|
@ -25,6 +26,13 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.continue_run: bool = True
|
||||
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])
|
||||
|
||||
return asyncio.run(async_gather())
|
||||
|
||||
@abstractmethod
|
||||
def _run(self):
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from memory_scope.utils.tool_functions import time_to_formatted_str
|
|||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
Extract_Time_PATTERN = r'-\s*(\S+):(\d+)'
|
||||
EXTRACT_TIME_PATTERN = r'-\s*(\S+):(\d+)'
|
||||
|
||||
def _run(self):
|
||||
query, query_timestamp = self.get_context(QUERY_WITH_TS)
|
||||
|
|
@ -44,7 +44,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
|
||||
# re-match time info to dict
|
||||
extract_time_dict = {}
|
||||
matches = re.findall(self.Extract_Time_PATTERN, response_text)
|
||||
matches = re.findall(self.EXTRACT_TIME_PATTERN, response_text)
|
||||
for key, value in matches:
|
||||
if key in DATATIME_KEY_MAP.keys():
|
||||
extract_time_dict[DATATIME_KEY_MAP[key]] = value
|
||||
|
|
|
|||
44
memory_scope/memory/worker/read/retrieve_store_worker.py
Normal file
44
memory_scope/memory/worker/read/retrieve_store_worker.py
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import QUERY_WITH_TS, RETRIEVE_MEMORY_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
|
||||
|
||||
|
||||
class RetrieveStoreWorker(MemoryBaseWorker):
|
||||
|
||||
async def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_id": self.user_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_id": self.user_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value],
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
memory_node_list: List[MemoryNode] = []
|
||||
fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile]
|
||||
for result in self._async_run(fn_list=fn_list, query=query):
|
||||
if result:
|
||||
memory_node_list.extend(result)
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
|
||||
|
||||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
for i, node in enumerate(memory_node_list):
|
||||
self.logger.info(f"{i}: node={node.content} score={node.score_similar} type={node.memory_type}")
|
||||
self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
|
|
@ -7,11 +7,11 @@ from memory_scope.scheme.memory_node import MemoryNode
|
|||
class BaseVectorStore(metaclass=ABCMeta):
|
||||
|
||||
@abstractmethod
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue