From 9e78eb69fd7330e6050b39b77a3bd06182c92d77 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 4 Jul 2024 13:59:32 +0800 Subject: [PATCH] [dev] rename get reflection system to reflection subject --- config/demo_config.yaml | 4 +-- memory_scope/memory/worker/base_worker.py | 4 +++ .../memory/worker/read/read_all_memory.py | 14 +++++++++ ...aml => get_reflection_subject_prompt.yaml} | 6 ++-- ...er.py => get_reflection_subject_worker.py} | 11 ++++--- .../llama_index_elastic_search_store.py | 2 +- tests/storages/test_storages_lli_es.py | 1 - tt.py | 31 +++++++++++++++++++ 8 files changed, 61 insertions(+), 12 deletions(-) create mode 100644 memory_scope/memory/worker/read/read_all_memory.py rename memory_scope/memory/worker/summary/{get_reflection_prompt.yaml => get_reflection_subject_prompt.yaml} (98%) rename memory_scope/memory/worker/summary/{get_reflection_worker.py => get_reflection_subject_worker.py} (86%) create mode 100644 tt.py diff --git a/config/demo_config.yaml b/config/demo_config.yaml index 9d4513c5..a63380dc 100644 --- a/config/demo_config.yaml +++ b/config/demo_config.yaml @@ -21,11 +21,11 @@ memory_service: description: "read session messages of the user" read_memory: class: memory.operation.read_memory - workflow: set_query_worker,[extract_time_worker|retrieve_store_worker,semantic_rank_worker],fuse_rerank_worker + workflow: set_query_worker,retrieve_store_worker,[extract_time_worker|semantic_rank_worker],fuse_rerank_worker description: "read related memories of the user" list_memory: class: memory.operation.read_memory - workflow: dummy_worker + workflow: set_query_worker,retrieve_store_worker, description: "read all memories of the user" write_memory: class: memory.operation.write_memory diff --git a/memory_scope/memory/worker/base_worker.py b/memory_scope/memory/worker/base_worker.py index d2da7d26..6d565aaa 100644 --- a/memory_scope/memory/worker/base_worker.py +++ b/memory_scope/memory/worker/base_worker.py @@ -28,9 +28,13 @@ 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") self.task_list.append((fn, args, kwargs)) def gather_async_result(self): + if self.is_multi_thread: + raise RuntimeError(f"async_task is not allowed in multi_thread condition") async def async_gather(): return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list]) diff --git a/memory_scope/memory/worker/read/read_all_memory.py b/memory_scope/memory/worker/read/read_all_memory.py new file mode 100644 index 00000000..e204b4e9 --- /dev/null +++ b/memory_scope/memory/worker/read/read_all_memory.py @@ -0,0 +1,14 @@ +from typing import List + +from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES +from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker +from memory_scope.scheme.memory_node import MemoryNode + + +class ReadAllMemory(MemoryBaseWorker): + + def _run(self): + memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES) + memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True) + obs_content_list: List[str] = [] + diff --git a/memory_scope/memory/worker/summary/get_reflection_prompt.yaml b/memory_scope/memory/worker/summary/get_reflection_subject_prompt.yaml similarity index 98% rename from memory_scope/memory/worker/summary/get_reflection_prompt.yaml rename to memory_scope/memory/worker/summary/get_reflection_subject_prompt.yaml index 0dcca2d2..32274f1e 100644 --- a/memory_scope/memory/worker/summary/get_reflection_prompt.yaml +++ b/memory_scope/memory/worker/summary/get_reflection_subject_prompt.yaml @@ -1,4 +1,4 @@ -get_reflection_system: +get_reflection_subject_system: cn: | 任务:从下面的信息中提取出最重要的最多{num_questions}条{user_name}属性,要求不与已有的{user_name}属性语义重复。 要求1:{user_name}属性可以是一般的{user_name}偏好,也可以是运动偏好,旅游偏好,饮食偏好等等,也可以是重要事件性质,比如最近重要的事情,也可以是一些高度概括的人生理想,价值观,人生观,性格, 也可以是和朋友的人际关系等等。 @@ -6,7 +6,7 @@ get_reflection_system: 输出格式:每一行输出一个{user_name}属性,每个{user_name}属性推荐4个字,如果没有信息请回答无,最多输出{num_questions}条。 -get_reflection_few_shot: +get_reflection_subject_few_shot: cn: | 示例1 信息: @@ -85,7 +85,7 @@ get_reflection_few_shot: 新增{user_name}属性: 无 -get_reflection_user_query: +get_reflection_subject_user_query: cn: | 信息: {user_query} diff --git a/memory_scope/memory/worker/summary/get_reflection_worker.py b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py similarity index 86% rename from memory_scope/memory/worker/summary/get_reflection_worker.py rename to memory_scope/memory/worker/summary/get_reflection_subject_worker.py index 92158378..10c08942 100644 --- a/memory_scope/memory/worker/summary/get_reflection_worker.py +++ b/memory_scope/memory/worker/summary/get_reflection_subject_worker.py @@ -11,7 +11,7 @@ from memory_scope.utils.response_text_parser import ResponseTextParser from memory_scope.utils.tool_functions import prompt_to_msg -class GetReflectionWorker(MemoryBaseWorker): +class GetReflectionSubjectWorker(MemoryBaseWorker): def new_insight_node(self, insight_key: str) -> MemoryNode: dt_handler = DatetimeHandler() @@ -40,10 +40,11 @@ class GetReflectionWorker(MemoryBaseWorker): # gen reflect prompt user_query_list = [n.content for n in not_reflected_nodes] - system_prompt = self.prompt_handler.get_reflection_system.format(user_name=self.target_name, - num_questions=self.reflect_num_questions) - few_shot = self.prompt_handler.get_reflection_few_shot.format(user_name=self.target_name) - user_query = self.prompt_handler.get_reflection_user_query.format( + system_prompt = self.prompt_handler.get_reflection_subject_system.format( + user_name=self.target_name, + num_questions=self.reflect_num_questions) + few_shot = self.prompt_handler.get_reflection_subject_few_shot.format(user_name=self.target_name) + user_query = self.prompt_handler.get_reflection_subject_user_query.format( user_name=self.target_name, exist_keys=self.get_language_value(COMMA_WORD).join(exist_keys), user_query="\n".join(user_query_list)) diff --git a/memory_scope/storage/llama_index_elastic_search_store.py b/memory_scope/storage/llama_index_elastic_search_store.py index 61537f0c..d36ac5f8 100644 --- a/memory_scope/storage/llama_index_elastic_search_store.py +++ b/memory_scope/storage/llama_index_elastic_search_store.py @@ -84,7 +84,7 @@ class LlamaIndexElasticSearchStore(BaseVectorStore): **kwargs) self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store, embed_model=self.embedding_model.model) - self.index.build_index_from_nodes([TextNode()]) + self.index.build_index_from_nodes([TextNode(text="text")]) self.logger = Logger.get_logger() def retrieve(self, diff --git a/tests/storages/test_storages_lli_es.py b/tests/storages/test_storages_lli_es.py index 34c66735..0fbe0e41 100644 --- a/tests/storages/test_storages_lli_es.py +++ b/tests/storages/test_storages_lli_es.py @@ -141,7 +141,6 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase): )) import asyncio res = asyncio.run(self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10)) - #res = self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10) print(len(res)) print(res) diff --git a/tt.py b/tt.py new file mode 100644 index 00000000..8b315e87 --- /dev/null +++ b/tt.py @@ -0,0 +1,31 @@ +import asyncio + + +class TT(object): + def __init__(self): + self.task_list = [] + + async def async_func(self, i: int): + await asyncio.sleep(i) # 模拟异步操作 + print(f"函数{i}的结果") + + 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() + + +TT().run()