mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-07 08:26:06 +00:00
fix es interface params
This commit is contained in:
parent
c0eae3f802
commit
6a5e830d46
18 changed files with 47 additions and 50 deletions
|
|
@ -103,43 +103,43 @@ worker:
|
|||
retrieve_ins_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
delete_memory:
|
||||
class: memory.worker.write.update_memory_worker
|
||||
class: memory.worker.summary.update_memory_worker
|
||||
method: delete_memory
|
||||
delete_all:
|
||||
class: memory.worker.write.update_memory_worker
|
||||
class: memory.worker.summary.update_memory_worker
|
||||
method: delete_all
|
||||
add_memory:
|
||||
class: memory.worker.write.update_memory_worker
|
||||
class: memory.worker.summary.update_memory_worker
|
||||
method: from_query
|
||||
info_filter:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
class: memory.worker.summary.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
load_today_memory:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
class: memory.worker.summary.load_memory_worker
|
||||
retrieve_today_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.write.get_observation_worker
|
||||
class: memory.worker.summary.get_observation_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
get_observation_with_time:
|
||||
class: memory.worker.write.get_observation_with_time_worker
|
||||
class: memory.worker.summary.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
contra_repeat:
|
||||
class: memory.worker.write.contra_repeat_worker
|
||||
class: memory.worker.summary.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
store_memory:
|
||||
class: memory.worker.write.update_memory_worker
|
||||
class: memory.worker.summary.update_memory_worker
|
||||
method: from_memory_key
|
||||
memory_key: all
|
||||
load_obs_and_insight:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
class: memory.worker.summary.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 100
|
||||
retrieve_not_updated_top_k: 100
|
||||
retrieve_insight_top_k: 100
|
||||
|
|
|
|||
|
|
@ -43,8 +43,8 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
# ⭐ Retrieve memories matching the query, filtered by the specified conditions,
|
||||
# limited to a certain number, and sorted by relevance.
|
||||
# Retrieve memories matching the query, filtered by the specified conditions,
|
||||
# limited to a certain number, and sorted by relevance.
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class SetQueryWorker(MemoryBaseWorker):
|
|||
Otherwise, the content of the last message in `self.chat_messages` is used as the query,
|
||||
along with its creation timestamp.
|
||||
"""
|
||||
query = "_" # Default query value
|
||||
query = "" # Default query value
|
||||
query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
|
||||
|
||||
if "query" in self.chat_kwargs:
|
||||
|
|
|
|||
|
|
@ -194,7 +194,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
return self.get_context(MEMORY_HANDLER)
|
||||
|
||||
@staticmethod
|
||||
def get_language_value(languages: dict | list[dict]) -> Any | list[Any]:
|
||||
def get_language_value(languages: dict | List[dict]) -> Any | List[Any]:
|
||||
"""
|
||||
Retrieves the value(s) corresponding to the current language context.
|
||||
|
||||
|
|
|
|||
|
|
@ -17,12 +17,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
self.retrieve_today_top_k: int = kwargs.get("retrieve_today_top_k", 0)
|
||||
|
||||
@timer
|
||||
def retrieve_not_reflected_memory(self, query: str):
|
||||
def retrieve_not_reflected_memory(self):
|
||||
"""
|
||||
Retrieves top-K not reflected memories based on the query and stores them in the memory handler.
|
||||
|
||||
Args:
|
||||
query (str): The search query for retrieving memories.
|
||||
"""
|
||||
if not self.retrieve_not_reflected_top_k:
|
||||
return
|
||||
|
|
@ -34,18 +31,14 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": 0,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_not_updated_memory(self, query: str):
|
||||
def retrieve_not_updated_memory(self):
|
||||
"""
|
||||
Retrieves top-K not updated memories based on the query and stores them in the memory handler.
|
||||
|
||||
Args:
|
||||
query (str): The search query for retrieving memories.
|
||||
"""
|
||||
if not self.retrieve_not_updated_top_k:
|
||||
return
|
||||
|
|
@ -57,8 +50,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": 0,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
|
|
@ -66,9 +58,6 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
def retrieve_insight_memory(self, query: str):
|
||||
"""
|
||||
Retrieves top-K insight memories based on the query and stores them in the memory handler.
|
||||
|
||||
Args:
|
||||
query (str): The search query for retrieving memories.
|
||||
"""
|
||||
if not self.retrieve_insight_top_k:
|
||||
return
|
||||
|
|
@ -79,18 +68,16 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"store_status": StoreStatusEnum.VALID.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_today_memory(self, query: str, dt: str):
|
||||
def retrieve_today_memory(self, dt: str):
|
||||
"""
|
||||
Retrieves top-K memories from today based on the query and stores them in the memory handler.
|
||||
|
||||
Args:
|
||||
query (str): The search query for retrieving memories.
|
||||
dt (str): The date string to filter today's memories.
|
||||
"""
|
||||
if not self.retrieve_today_top_k:
|
||||
|
|
@ -103,8 +90,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_today_top_k,
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
self.memory_handler.set_memories(TODAY_NODES, nodes)
|
||||
|
|
@ -120,12 +106,11 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"""
|
||||
|
||||
# Placeholder query
|
||||
query = "-"
|
||||
dt = DatetimeHandler().datetime_format()
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_insight_memory, query=query)
|
||||
self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt)
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory)
|
||||
self.submit_thread_task(self.retrieve_insight_memory)
|
||||
self.submit_thread_task(self.retrieve_today_memory, dt=dt)
|
||||
|
||||
# Waits for all submitted tasks to complete
|
||||
for _ in self.gather_thread_result():
|
||||
|
|
@ -11,7 +11,10 @@ class BaseMemoryStore(metaclass=ABCMeta):
|
|||
"""
|
||||
|
||||
@abstractmethod
|
||||
def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
def retrieve_memories(self,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]:
|
||||
"""
|
||||
Retrieves a list of MemoryNode objects that are most relevant to the query,
|
||||
considering a filter dictionary for additional constraints. The number of nodes returned
|
||||
|
|
@ -30,7 +33,10 @@ class BaseMemoryStore(metaclass=ABCMeta):
|
|||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
async def a_retrieve_memories(self,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]:
|
||||
"""
|
||||
Asynchronously retrieves a list of MemoryNode objects that best match the query,
|
||||
respecting a filter dictionary, with the result size capped at top_k.
|
||||
|
|
|
|||
|
|
@ -12,6 +12,18 @@ class DummyMemoryStore(BaseMemoryStore):
|
|||
semantic retrieval. Actual storage operations are not implemented.
|
||||
"""
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
async def a_retrieve_memories(self,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
"""
|
||||
Initializes the DummyMemoryStore with an embedding model and additional keyword arguments.
|
||||
|
|
@ -23,12 +35,6 @@ class DummyMemoryStore(BaseMemoryStore):
|
|||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
def batch_insert(self, nodes: List[MemoryNode]):
|
||||
pass
|
||||
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: Optional[str] = None,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
# if index is not created, return []
|
||||
|
|
@ -59,7 +59,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
async def a_retrieve_memories(self,
|
||||
query: Optional[str] = None,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
# if index is not created, return []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue