mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-07 08:26:06 +00:00
[dev] bug fix not clear context
This commit is contained in:
parent
8aa9c8b302
commit
907a8f1815
8 changed files with 107 additions and 26 deletions
|
|
@ -30,11 +30,11 @@ memory_service:
|
|||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 5
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 60
|
||||
# summary_memory:
|
||||
# class: memory.operation.summary_memory
|
||||
# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
# description: "summary observation memories of the user"
|
||||
# interval_time: 60
|
||||
|
||||
worker:
|
||||
dummy:
|
||||
|
|
|
|||
|
|
@ -15,8 +15,8 @@ class SummaryMemory(BaseWorkflow, BaseBackendOperation):
|
|||
self.init_workers(is_backend=True, **kwargs)
|
||||
|
||||
def _run_operation(self, **kwargs):
|
||||
self.context.clear()
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ class WriteMemory(BaseWorkflow, BaseBackendOperation):
|
|||
self.init_workers(is_backend=True, **kwargs)
|
||||
|
||||
def _run_operation(self, **kwargs):
|
||||
self.context.clear()
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
not_memorized_size = self.not_memorized_size
|
||||
if not_memorized_size < self.contextual_msg_count:
|
||||
|
|
@ -51,6 +52,5 @@ class WriteMemory(BaseWorkflow, BaseBackendOperation):
|
|||
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
self.set_memorized()
|
||||
return result
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -32,16 +32,14 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
# if self.is_multi_thread:
|
||||
# raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
self.async_task_list.append((fn, args, kwargs))
|
||||
|
||||
async def _async_gather(self):
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.async_task_list])
|
||||
|
||||
def gather_async_result(self):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
results = asyncio.run(self._async_gather())
|
||||
self.async_task_list.clear()
|
||||
return results
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
today_obs_nodes: List[MemoryNode] = self.get_memories(TODAY_NODES)
|
||||
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.timestamp, reverse=True)[:self.contra_repeat_max_count]
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from memory_scope.utils.timer import timer
|
|||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
def retrieve_not_reflected_memory(self, query: str):
|
||||
async def retrieve_not_reflected_memory(self, query: str):
|
||||
if not self.retrieve_not_reflected_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -24,13 +24,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_not_updated_memory(self, query: str):
|
||||
async def retrieve_not_updated_memory(self, query: str):
|
||||
if not self.retrieve_not_updated_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -41,13 +41,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_insight_memory(self, query: str):
|
||||
async def retrieve_insight_memory(self, query: str):
|
||||
if not self.retrieve_insight_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -57,13 +57,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_today_memory(self):
|
||||
async def retrieve_today_memory(self):
|
||||
if not self.today_obs_top_k:
|
||||
return
|
||||
|
||||
|
|
@ -80,7 +80,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt_handler.datetime_format(),
|
||||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content,
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content,
|
||||
top_k=self.today_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
|
|
@ -88,8 +88,14 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
mock_query = "-"
|
||||
self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
|
||||
self.submit_thread_task(self.retrieve_today_memory)
|
||||
self.gather_thread_result()
|
||||
# self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
|
||||
# self.submit_thread_task(self.retrieve_today_memory)
|
||||
# self.gather_thread_result()
|
||||
|
||||
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_insight_memory, query=mock_query)
|
||||
self.submit_async_task(self.retrieve_today_memory)
|
||||
self.gather_async_result()
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from memory_scope.utils.logger import Logger
|
|||
|
||||
class Timer(object):
|
||||
|
||||
def __init__(self, name: str, log_time: bool = True, use_ms: bool = False, **kwargs):
|
||||
def __init__(self, name: str, log_time: bool = True, use_ms: bool = True, **kwargs):
|
||||
self.name: str = name
|
||||
self.log_time: bool = log_time
|
||||
self.use_ms: bool = use_ms
|
||||
|
|
|
|||
76
tests/thread_test.py
Normal file
76
tests/thread_test.py
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
import sys
|
||||
|
||||
sys.path.append(".")
|
||||
|
||||
import asyncio
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
import questionary
|
||||
|
||||
|
||||
class ThreadTest(object):
|
||||
def __init__(self):
|
||||
self.task_list = []
|
||||
config = {
|
||||
"module_name": "dashscope_embedding",
|
||||
"model_name": "text-embedding-v2",
|
||||
"clazz": "models.llama_index_embedding_model",
|
||||
}
|
||||
emb = LlamaIndexEmbeddingModel(**config)
|
||||
|
||||
config = {
|
||||
"index_name": "0708_2",
|
||||
"es_url": "http://localhost:9200",
|
||||
"embedding_model": emb,
|
||||
"use_hybrid": True
|
||||
|
||||
}
|
||||
self.es_store = LlamaIndexEsMemoryStore(**config)
|
||||
self.logger = Logger.get_logger("default")
|
||||
|
||||
async def async_func(self, i: int):
|
||||
self.logger.info(f"i: {i}")
|
||||
await asyncio.sleep(i)
|
||||
self.es_store.insert(MemoryNode(
|
||||
content="xxxxxx",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="ggg567",
|
||||
meta_data={"5": "5"}
|
||||
))
|
||||
|
||||
result = self.es_store.retrieve_memories("_", top_k=10, filter_dict={"memory_id": "ggg567"})
|
||||
self.logger.info(f"result: {result}")
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
self.task_list.append((fn, args, kwargs))
|
||||
|
||||
def gather_async_result(self):
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list])
|
||||
|
||||
results = asyncio.run(async_gather())
|
||||
self.task_list.clear()
|
||||
return results
|
||||
|
||||
def run(self):
|
||||
self.submit_async_task(self.async_func, i=1)
|
||||
self.submit_async_task(self.async_func, i=2)
|
||||
self.submit_async_task(self.async_func, i=3)
|
||||
|
||||
self.gather_async_result()
|
||||
self.logger.close()
|
||||
|
||||
|
||||
executor = ThreadPoolExecutor(max_workers=5)
|
||||
t = executor.submit(ThreadTest().run)
|
||||
questionary.text(message="user:", multiline=False, qmark=">").unsafe_ask()
|
||||
|
||||
print(as_completed([t]))
|
||||
executor.shutdown()
|
||||
Loading…
Add table
Reference in a new issue