fix es interface params

This commit is contained in:
jinli.yl 2024-07-24 21:19:35 +08:00
parent c0eae3f802
commit 6a5e830d46
18 changed files with 47 additions and 50 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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():

View file

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

View file

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

View file

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