[dev] bug fix not clear context

This commit is contained in:
jinli.yl 2024-07-09 21:58:38 +08:00
parent 8aa9c8b302
commit 907a8f1815
8 changed files with 107 additions and 26 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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