[dev] rename get reflection system to reflection subject

This commit is contained in:
jinli.yl 2024-07-04 13:59:32 +08:00
parent 258ec08f3f
commit 9e78eb69fd
8 changed files with 61 additions and 12 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

31
tt.py Normal file
View file

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