[dev] rename constant name & add _async_run

This commit is contained in:
jinli.yl 2024-07-01 11:45:34 +08:00
parent d8aceabbf0
commit 7751586435
6 changed files with 61 additions and 4 deletions

View file

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

View file

@ -8,6 +8,7 @@ CHAT_KWARGS = "chat_kwargs"
QUERY_WITH_TS = "query_with_ts"
RETRIEVE_MEMORY_NODES = "RETRIEVE_MEMORY_NODES"

View file

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

View file

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

View 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)

View file

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