Merge branch 'master' of memoryscope

This commit is contained in:
fuqingxu 2024-07-08 15:34:23 +08:00
commit 8e756fdd98
94 changed files with 941 additions and 4787 deletions

View file

@ -2,7 +2,7 @@
exclude =
scripts/*
src/agentscope/rpc/*
max-line-length = 79
max-line-length = 120
inline-quotes = "
avoid-escape = no
ignore =

View file

@ -6,8 +6,6 @@ memory_chat:
class: chat.cli_memory_chat
memory_service: memory_chat_service
generation_model: dashscope_generation
human_name: 用户
assistant_name: AI
memory_service:
memory_chat_service:
class: memory.service.chat_memory_service
@ -16,25 +14,24 @@ memory_service:
read_memory_key: read_memory
memory_operations:
read_message:
class: memory.operation.read_memory
workflow: dummy_worker
class: memory.operation.read_message
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,print_memory_worker
description: "read all memories of the user"
write_memory:
class: memory.operation.write_memory
workflow: dummy_worker
workflow: info_filter_worker,[get_observation_worker|get_observation_with_time_worker],contra_repeat_worker,store_memory_worker
description: "write observation memories of the user"
interval_time: 60
summary_memory:
class: memory.operation.summary_memory
workflow: dummy_worker
workflow: load_memory_worker,get_reflection_subject_worker,update_insight_worker,long_contra_repeat_worker,summary_collect_worker
description: "summary observation memories of the user"
interval_time: 300
worker:
@ -45,13 +42,13 @@ worker:
rank_model: dashscope_rank
set_query_worker:
class: memory.worker.read.set_query_worker
retrieve_store_worker:
class: memory.worker.read.retrieve_store_worker
retrieve_obs_top_k: 100
retrieve_ins_pf_top_k: 100
extract_time_worker:
class: memory.worker.read.extract_time_worker
generation_model_top_k: 1
retrieve_store_worker:
class: memory.worker.read.retrieve_store_worker
retrieve_obs_top_k: 5
retrieve_ins_pf_top_k: 5
semantic_rank_worker:
class: memory.worker.read.semantic_rank_worker
fuse_rerank_worker:
@ -60,12 +57,12 @@ worker:
fuse_ratio_dict:
conversation: 0.5
observation: 1
obs_customized: 1
obs_customized: 1.2
insight: 2.0
profile: 2.0
profile_customized: 2.0
fuse_time_ratio: 2.0
fuse_rerank_top_k: 10
print_memory_worker:
class: memory.worker.read.print_memory_worker
info_filter_worker:
class: memory.worker.write.info_filter_worker
generation_model: dashscope_generation
@ -85,10 +82,20 @@ worker:
generation_model_top_k: 1
retrieve_top_k: 30
contra_repeat_max_count: 50
get_reflection_worker:
class: memory.worker.summary.get_reflection_worker
store_memory_worker:
class: memory.worker.write.store_memory_worker
load_memory_worker:
class: memory.worker.summary.load_memory_worker
get_reflection_subject_worker:
class: memory.worker.summary.get_reflection_subject_worker
retrieve_top_k: 100
reflect_obs_cnt_threshold: 32
update_insight_worker:
class: memory.worker.summary.update_insight_worker
long_contra_repeat_worker:
class: memory.worker.summary.long_contra_repeat_worker
summary_collect_worker:
class: memory.worker.summary.summary_collect_worker
models:
dashscope_generation:
class: models.llama_index_generation_model
@ -102,10 +109,11 @@ models:
class: models.llama_index_rank_model
module_name: dashscope_rank
model_name: gte-rerank
vector_store:
class: storage.llama_index_elastic_search_store
memory_store:
class: storage.llama_index_es_memory_store
embedding_model: dashscope_embedding
index_name: memory_index
es_url: http://localhost:9200
use_hybrid: false
monitor:
class: storage.dummy_monitor

View file

@ -27,8 +27,8 @@ class CliMemoryChat(BaseMemoryChat):
memory_service: str,
generation_model: str,
stream: bool = True,
human_name: str = "",
assistant_name: str = "",
human_name: str = "用户",
assistant_name: str = "AI",
**kwargs):
self._memory_service: BaseMemoryService | str = memory_service

View file

@ -56,9 +56,9 @@ class CliJob(object):
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
# init vector_store
vector_store_config = self.config["vector_store"]
embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]]
G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model)
memory_store_config = self.config["memory_store"]
embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model)
# init monitor
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
@ -73,7 +73,7 @@ class CliJob(object):
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
memory_chat.run()
G_CONTEXT.vector_store.close()
G_CONTEXT.memory_store.close()
G_CONTEXT.monitor.close()

View file

@ -1,13 +0,0 @@
from enum import Enum
class MemoryMethodEnum(str, Enum):
SUMMARY = "summary"
RETRIEVE = "retrieve"
RETRIEVE_ALL = "retrieve_all"
SUMMARY_SHORT = "summary_short"
SUMMARY_LONG = "summary_long"

View file

@ -1,9 +0,0 @@
from enum import Enum
class MemoryRecallType(str, Enum):
SIMILAR = "similar"
KEYWORD = "keyword"
PROFILE = "profile"

View file

@ -2,6 +2,12 @@ from enum import Enum
class MemoryNodeStatus(str, Enum):
NEW = "new"
MODIFIED = "modified"
CONTENT_MODIFIED = "content_modified"
ACTIVE = "active"
EXPIRED = "expired"

View file

@ -8,8 +8,4 @@ class MemoryTypeEnum(str, Enum):
INSIGHT = "insight"
PROFILE = "profile"
OBS_CUSTOMIZED = "obs_customized"
PROFILE_CUSTOMIZED = "profile_customized"

View file

@ -10,6 +10,7 @@ class BaseOperation(metaclass=ABCMeta):
def __init__(self, name: str, description: str = "", **kwargs):
self.name: str = name
self.description: str = description
self.kwargs: dict = kwargs
def init_workflow(self):
pass

View file

@ -14,19 +14,17 @@ class ReadMemory(BaseWorkflow, BaseOperation):
description: str,
chat_messages: List[Message],
his_msg_count: int = 0, # supplement to the current query
contextual_msg_count: int = 0, # for the current context dialogue
**kwargs):
super().__init__(name=name, **kwargs)
BaseOperation.__init__(self, name=name, description=description)
self.chat_messages: List[Message] = chat_messages
self.his_msg_count: int = his_msg_count
self.contextual_msg_count: int = contextual_msg_count
def init_workflow(self):
self.init_workers()
def run_operation(self, **kwargs):
max_count = 1 + max(self.his_msg_count, self.contextual_msg_count)
max_count = 1 + self.his_msg_count
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
self.context[CHAT_KWARGS] = kwargs
self.run_workflow()

View file

@ -0,0 +1,21 @@
from typing import List
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
from memory_scope.scheme.message import Message
class ReadMessage(BaseOperation):
operation_type: OPERATION_TYPE = "frontend"
def __init__(self,
name: str,
description: str,
chat_messages: List[Message],
contextual_msg_count: int = 6, # for the current context dialogue
**kwargs):
super().__init__(name=name, description=description, **kwargs)
self.chat_messages: List[Message] = chat_messages
self.contextual_msg_count: int = contextual_msg_count
def run_operation(self, **kwargs):
return self.chat_messages[-self.contextual_msg_count:]

View file

@ -28,9 +28,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")
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])
@ -61,12 +66,9 @@ class BaseWorker(metaclass=ABCMeta):
def set_context(self, key: str, value: Any):
if self.is_multi_thread:
with self.context_lock:
self.context_dict[key] = value
self.context[key] = value
else:
self.context[key] = value
def has_content(self, key: str):
return key in self.context
def __getattr__(self, key: str):
return self.kwargs[key]

View file

@ -4,9 +4,10 @@ from typing import List, Dict
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
from memory_scope.memory.worker.base_worker import BaseWorker
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.scheme.message import Message
from memory_scope.storage.base_memory_store import BaseMemoryStore
from memory_scope.storage.base_monitor import BaseMonitor
from memory_scope.storage.base_vector_store import BaseVectorStore
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.prompt_handler import PromptHandler
@ -24,13 +25,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
self._generation_model: BaseModel | str = generation_model
self._rank_model: BaseModel | str = rank_model
self._vector_store: BaseVectorStore | None = None
self._memory_store: BaseMemoryStore | None = None
self._monitor: BaseMonitor | None = None
self._user_name: str | None = None
self._target_name: str | None = None
self._prompt_handler: PromptHandler | None = None
self._contex_memory_dict: Dict[str, MemoryNode] = {}
@property
def chat_messages(self) -> List[Message]:
return self.get_context(CHAT_MESSAGES)
@ -62,10 +65,30 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
return self._rank_model
@property
def vector_store(self) -> BaseVectorStore:
if self._vector_store is None:
self._vector_store = G_CONTEXT.vector_store
return self._vector_store
def memory_store(self) -> BaseMemoryStore:
if self._memory_store is None:
self._memory_store = G_CONTEXT.memory_store
return self._memory_store
def get_memories(self, key: str) -> List[MemoryNode]:
memories: List[MemoryNode] = []
memory_ids: List[str] = self.get_context(key)
if memory_ids:
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
return memories
def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode):
if not nodes:
return
if isinstance(nodes, MemoryNode):
nodes = [nodes]
for node in nodes:
if node.memory_id in self._contex_memory_dict:
continue
self._contex_memory_dict[node.memory_id] = node
self.set_context(key, [n.memory_id for n in nodes])
@property
def monitor(self) -> BaseMonitor:

View file

@ -2,9 +2,9 @@ import re
from typing import Dict
from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT
from memory_scope.constants.language_constants import DATATIME_WORD_LIST
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.utils.datetime_handler import DatetimeHandler
from memory_scope.utils.tool_functions import prompt_to_msg
class ExtractTimeWorker(MemoryBaseWorker):
@ -14,23 +14,21 @@ class ExtractTimeWorker(MemoryBaseWorker):
query, query_timestamp = self.get_context(QUERY_WITH_TS)
# find datetime keyword
contain_datetime = False
for datetime_word in self.get_language_value(DATATIME_WORD_LIST):
if datetime_word in query:
contain_datetime = True
break
contain_datetime = DatetimeHandler.has_time_word(query)
if not contain_datetime:
self.logger.info(f"contain_datetime={contain_datetime}")
return
# prepare prompt
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format)
extract_time_prompt: str = self.prompt_handler.extract_time_prompt.format(query=query,
query_time_str=query_time_str)
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
system_prompt = self.prompt_handler.extract_time_system
few_shot = self.prompt_handler.extract_time_few_shot.format(user_name=self.target_name)
user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str)
extract_time_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
self.logger.info(f"extract_time_message={extract_time_message}")
# call sft model
response = self.generation_model.call(prompt=extract_time_prompt, top_k=self.generation_model_top_k)
# call llm
response = self.generation_model.call(messages=extract_time_message, top_k=self.generation_model_top_k)
# if empty, return
if not response.status or not response.message.content:

View file

@ -1,6 +1,9 @@
extract_time_system:
cn: |
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。回答的格式严格遵照示例中的已有格式规范。
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。回答的格式严格遵照示例中的已有格式规范。若语句不涉及时间则回答无。
en: |
Instructions: From the sentences and the time when they occurred, infer and extract the time periods indicated in the content of the sentences. Answer with the most complete time periods possible. The format of the answers must strictly adhere to the specifications in the examples provided. If the sentence does not involve time, respond with "none."
extract_time_few_shot:
@ -47,6 +50,60 @@ extract_time_few_shot:
回答:
- 1995 - 月10 - 日24
示例8:
句子:我的朋友非常喜欢运动,他认为运动有助于增强身体素质。
时间2015年1月23日2015年第4周周四7时38分0秒。
回答:
en: |
Example 1:
Sentence: I remember you went to the UAE in April the year before last. Which places in the UAE are fun? Which do you prefer, Dubai or Abu Dhabi? Are the desert views spectacular?
Time: August 20, 1992, 34th week of 1992, Monday, 18:46:25.
Answer:
- Year: 1990 - Month: 4
Example 2:
Sentence: Remember to attend the meeting at 3 PM the day after tomorrow. I carefully marked this important day on my calendar to remind myself not to miss it. The meeting will be held in the company conference room, and it's an important opportunity to discuss future development directions.
Time: June 19, 2024, 25th week of 2024, Tuesday, 13:30:0.
Answer:
- Year: 2024 - Month: 6 - Day: 21 - Hour: 15
Example 3:
Sentence: Next month on the first Saturday, let's go to Hangzhou.
Time: July 15, 2005, 28th week of 2005, Saturday, 0:0:0.
Answer:
- Year: 2005 - Month: 8 - Week: 31 - Day of Week: 6
Example 4:
Sentence: The small town we visited last weekend was truly beautiful.
Time: December 2, 1999, 48th week of 1999, Tuesday, 8:40:10.
Answer:
- Year: 1999 - Week: 47 - Day of Week: 6, 7
Example 5:
Sentence: The presentation will start in half an hour, remember to prepare the materials.
Time: June 22, 2020, 25th week of 2020, Monday, 9:30:0
Answer:
- Year: 2020 - Month: 6 - Day: 22 - Hour: 10 - Minute: 0 - Second: 0
Example 6:
Sentence: The start time for the 10,000-meter race was 3 minutes and 47 seconds ago.
Time: February 17, 1987, 7th week of 1987, Wednesday, 19:54:43.
Answer:
- Year: 1987 - Month: 2 - Day: 17 - Hour: 19 - Minute: 50 - Second: 56
Example 7:
Sentence: At this time last month, we were still preparing for the concert. Every day was busy and fulfilling, and we worked hard for the successful holding of the concert. Rehearsals, set design, and program arrangements - every detail needed careful planning and preparation.
Time: November 24, 1995, 48th week of 1995, Tuesday, 17:45:0.
Answer:
- Year: 1995 - Month: 10 - Day: 24
Example 8:
Sentence: My friend loves sports very much and believes that exercise helps improve physical fitness.
Time: January 23, 2015, 4th week of 2015, Thursday, 7:38:0.
Answer:
None
extract_time_user_query:
@ -55,6 +112,10 @@ extract_time_user_query:
时间:{query_time_str}
回答:
en: |
Sentence: {query}
Time: {query_time_str}
Answer:
time_string_format:
cn: |

View file

@ -0,0 +1,72 @@
from typing import List
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.utils.datetime_handler import DatetimeHandler
from memory_scope.utils.timer import timer
class PrintMemoryWorker(MemoryBaseWorker):
@timer
def retrieve_expired_memory(self, query: str):
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.EXPIRED.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_expired_top_k,
filter_dict=filter_dict)
def _run(self):
expired_memories: List[MemoryNode] = self.retrieve_expired_memory(query="_")
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
obs_content_list: List[str] = []
insight_content_list: List[str] = []
expired_content_list: List[str] = []
i = 0
j = 0
for node in memory_node_list:
if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
i += 1
dt_handler = DatetimeHandler(node.timestamp)
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
line = f" {i} {dt} {node.content}"
obs_content_list.append(line)
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.INSIGHT, ]:
j += 1
line = f" {j} {node.content}"
insight_content_list.append(line)
for i, node in enumerate(expired_memories):
line = f" {j} {node.content}"
expired_content_list.append(line)
obs_content = "\n".join(obs_content_list)
insight_content = "\n".join(insight_content_list)
expired_content = "\n".join(expired_content_list)
result: str = f"""
The memories of {self.user_name} about {self.target_name}.
----- observation -----
{obs_content}
----- observation -----
----- insight -----
{insight_content}
----- insight -----
----- expired -----
{expired_content}
----- expired -----
""".strip()
self.set_context(RESULT, result)

View file

@ -16,26 +16,28 @@ class RetrieveStoreWorker(MemoryBaseWorker):
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return await self.vector_store.async_retrieve(query=query,
top_k=self.retrieve_obs_top_k,
filter_dict=filter_dict)
return await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_obs_top_k,
filter_dict=filter_dict)
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
filter_dict = {
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value],
"memory_type": MemoryTypeEnum.INSIGHT.value,
}
return await self.vector_store.async_retrieve(query=query,
top_k=self.retrieve_ins_pf_top_k,
filter_dict=filter_dict)
return await self.memory_store.a_retrieve_memories(query=query,
top_k=self.retrieve_ins_pf_top_k,
filter_dict=filter_dict)
def _run(self):
query, _ = self.get_context(QUERY_WITH_TS)
self.submit_async_task(self.retrieve_from_observation, query=query)
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
memory_node_list: List[MemoryNode] = []
fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile]
for result in self.async_run(fn_list=fn_list, query=query):
for result in self.gather_async_result():
if result:
memory_node_list.extend(result)
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
@ -43,4 +45,4 @@ class RetrieveStoreWorker(MemoryBaseWorker):
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
for node in memory_node_list:
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}")
self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list)
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)

View file

@ -1,7 +1,6 @@
from typing import List, Dict
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
@ -11,39 +10,27 @@ class SemanticRankWorker(MemoryBaseWorker):
def _run(self):
# query
query, _ = self.get_context(QUERY_WITH_TS)
memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES)
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
if not memory_node_list:
self.logger.warning(f"retrieve memory nodes is empty!")
return
# solve content repeat, insight & profile has higher priority
memory_node_dict: Dict[str, MemoryNode] = {}
memory_type_selected = [MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value]
for node in memory_node_list:
if node.memory_type in memory_type_selected:
return
memory_node_dict[node.content] = node
for node in memory_node_list:
if node.memory_type not in memory_type_selected:
return
memory_node_dict[node.content] = node
# drop repeated
memory_node_dict: Dict[str, MemoryNode] = {n.content: n for n in memory_node_list}
memory_node_list = list(memory_node_dict.values())
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
if not response.status or not response.rank_scores:
return
# set score
rank_memory_nodes: List[MemoryNode] = []
rank_scores = sorted(response.rank_scores.items(), key=lambda x: x[1], reverse=True)
for idx, score in rank_scores:
for idx, score in response.rank_scores.items():
if idx >= len(memory_node_list):
self.logger.warning(f"idx={idx} exceeds the maximum length of the array")
self.logger.warning(f"idx={idx} exceeds the maximum length of rank_scores!")
continue
node = memory_node_list[idx]
node.score_rank = score
memory_node_list[idx].score_rank = score
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
for node in memory_node_list:
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
self.get_context(RANKED_MEMORY_NODES, rank_memory_nodes)
self.set_memories(RANKED_MEMORY_NODES, memory_node_list)

View file

@ -1,4 +1,4 @@
get_reflect_system:
get_reflection_subject_system:
cn: |
任务:从下面的信息中提取出最重要的最多{num_questions}条{user_name}属性,要求不与已有的{user_name}属性语义重复。
要求1{user_name}属性可以是一般的{user_name}偏好,也可以是运动偏好,旅游偏好,饮食偏好等等,也可以是重要事件性质,比如最近重要的事情,也可以是一些高度概括的人生理想,价值观,人生观,性格, 也可以是和朋友的人际关系等等。
@ -6,7 +6,7 @@ get_reflect_system:
输出格式:每一行输出一个{user_name}属性,每个{user_name}属性推荐4个字如果没有信息请回答无最多输出{num_questions}条。
get_reflect_few_shot:
get_reflection_subject_few_shot:
cn: |
示例1
信息:
@ -85,7 +85,7 @@ get_reflect_few_shot:
新增{user_name}属性:
get_reflect_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()
@ -22,7 +22,7 @@ class GetReflectionWorker(MemoryBaseWorker):
meta_data=meta_data,
key=insight_key,
memory_type=MemoryTypeEnum.INSIGHT.value,
status=MemoryNodeStatus.ACTIVE.value)
status=MemoryNodeStatus.NEW.value)
def _run(self):
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
@ -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_reflect_system.format(user_name=self.target_name,
num_questions=self.reflect_num_questions)
few_shot = self.prompt_handler.get_reflect_few_shot.format(user_name=self.target_name)
user_query = self.prompt_handler.get_reflect_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))
@ -59,7 +60,7 @@ class GetReflectionWorker(MemoryBaseWorker):
return
# parse text & save
new_insight_keys = ResponseTextParser(response.message.content).parse_v2("get_reflection")
new_insight_keys = ResponseTextParser(response.message.content).parse_v2(self.__class__.__name__)
if new_insight_keys:
for insight_key in new_insight_keys:
insight_nodes.append(self.new_insight_node(insight_key))

View file

@ -19,9 +19,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"obs_reflected": False,
}
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
top_k=self.retrieve_not_reflected_top_k,
filter_dict=filter_dict)
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_context(NOT_REFLECTED_NODES, nodes)
@timer
@ -33,9 +33,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"obs_updated": False,
}
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
top_k=self.retrieve_not_updated_top_k,
filter_dict=filter_dict)
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_context(NOT_UPDATED_NODES, nodes)
@timer
@ -46,9 +46,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": MemoryTypeEnum.INSIGHT.value,
}
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
top_k=self.retrieve_insight_top_k,
filter_dict=filter_dict)
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_context(INSIGHT_NODES, nodes)
async def _run(self):

View file

@ -1,11 +1,18 @@
long_contra_repeat_system:
cn: |
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
对每个句子都做一个判断,最后一共输出{num_obs}条判断。如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
请一步步思考,并按如下格式输出:
思考思考的依据和过程30字以内。
判断:<句子序号> <矛盾,被包含,无> <修改后的内容>,一定加<>
en: |
For the following {num_obs} sentences, determine one by one if there is any contradiction with the information in any "previously numbered" sentences, or if the main information of the sentence is contained within the information of any "previously numbered" sentences. Only evaluate the relation to "previously numbered" sentences.
Make an evaluation for each sentence, resulting in a total of {num_obs} evaluations. If the sentence conflicts with a "previously numbered" sentence, the information in the earlier sentence takes precedence, and the conflicting part of the current sentence should be modified accordingly.
Please think step by step and output in the following format:
Thought: The basis and process of the thought, within 30 words.
Evaluation: <sentence number> <Contradiction, Contained or None>, <Revised content>, enclosed in <>.
long_contra_repeat_few_shot:
cn: |
@ -50,6 +57,47 @@ long_contra_repeat_few_shot:
思考第6句中所有信息都被前面序号中第5句的信息完全包含。
判断:<2> <被包含> <>
en: |
Example 1
Sentences:
1 {user_name} suffers from insomnia frequently and is interested in the effects of sleeping pills, suggesting a possible consideration of their use.
2 {user_name} suffers from insomnia frequently and seeks remedies.
3 Charles is {user_name}'s supervisor.
4 Charles is {user_name}'s supervisor.
5 Charles is {user_name}'s supervisor and the branch manager of a bank.
Thought: The first sentence does not have any contradictions or complete repetitions with the previously numbered sentences.
Evaluation: <1> <None> <>
Thought: All information in the second sentence is completely contained within the information of the first sentence.
Evaluation: <2> <Contained> <>
Thought: The information in the third sentence does not appear in the previously numbered sentences.
Evaluation: <3> <None> <>
Thought: The fourth sentence is completely repetitive of the information in the third sentence, i.e., it is completely contained.
Evaluation: <4> <Contained> <>
Thought: The information that Charles is {user_name}'s supervisor in the fifth sentence is contained within the information of the third sentence, but the new information that Charles is the branch manager of a bank is not, so it is not contained.
Evaluation: <5> <None> <>
Example 2
Sentences:
1 {user_name}'s child does not perform well academically.
2 {user_name}'s child often skips school.
3 {user_name}'s father's birthday is on June 2, 2024, and {user_name} plans to prepare a gift.
4 {user_name}'s father's birthday is on May 1, 2024.
5 {user_name} loves playing basketball with classmates.
6 {user_name} likes playing basketball.
Thought: The first sentence does not have any contradictions or complete repetitions with the previously numbered sentences.
Evaluation: <1> <None> <>
Thought: The second sentence neither contradicts nor repeats any of the previously numbered sentences.
Evaluation: <2> <None> <>
Thought: The third sentence neither contradicts nor repeats any of the previously numbered sentences.
Evaluation: <3> <None> <>
Thought: The date of {user_name}'s father's birthday in the fourth sentence contradicts the information in the third sentence.
Evaluation: <4> <Contradiction> <{user_name}'s father's birthday is on June 2, 2024.>
Thought: The fifth sentence neither contradicts nor repeats any of the previously numbered sentences.
Evaluation: <5> <None> <>
Thought: All information in the sixth sentence is completely contained within the information of the fifth sentence.
Evaluation: <6> <Contained> <>
long_contra_repeat_user_query:
cn: |
句子:

View file

@ -21,9 +21,9 @@ class LongContraRepeatWorker(MemoryBaseWorker):
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value]
}
retrieve_nodes = await self.vector_store.async_retrieve(query=node.content,
top_k=self.long_contra_repeat_top_k,
filter_dict=filter_dict)
retrieve_nodes = await self.memory_store.a_retrieve_memories(query=node.content,
top_k=self.long_contra_repeat_top_k,
filter_dict=filter_dict)
return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold]
def _run(self):

View file

@ -1,46 +1,40 @@
from typing import List, Dict
from memory_scope.constants.common_constants import (
NEW_INSIGHT_NODES,
MODIFIED_MEMORIES,
INSIGHT_NODES,
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
NEW,
NOT_REFLECTED_MERGE_NODES,
)
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \
NOT_UPDATED_NODES
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
update_memories: Dict[str, MemoryNode] = {}
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_OBS_NODES
)
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_MERGE_NODES
)
# 合并逻辑复杂务必check
all_node_dict: Dict[str, MemoryNode] = {}
if insight_nodes:
all_node_dict.update(
{n.id: n for n in insight_nodes if n.obs_updated}
)
if new_insight_nodes:
all_node_dict.update({n.content: n for n in new_insight_nodes})
if new_obs_nodes:
# 设置为非新
for n in new_obs_nodes:
n.obs_updated = "0"
all_node_dict.update({n.content: n for n in new_obs_nodes})
if not_reflected_merge_nodes and not_reflected_nodes:
# 进入reflect阶段
all_node_dict.update({n.id: n for n in not_reflected_nodes})
update_memories.update({n.memory_id: n for n in insight_nodes})
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
if not_reflected_nodes:
update_memories.update({n.memory_id: n for n in not_reflected_nodes})
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
if not_updated_nodes:
for node in not_updated_nodes:
if node.memory_id in update_memories:
keys = [
INSIGHT_NODES,
MERGE_OBS_NODES,
NOT_UPDATED_NODES,
NOT_REFLECTED_NODES,
]
memory_nodes: List[MemoryNode] = []
for key in keys:
memory_nodes.extend(self.get_context(key))
self.memory_store.update_memories(update_memories)

View file

@ -2,6 +2,7 @@ from typing import List
from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES
from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.utils.datetime_handler import DatetimeHandler
@ -48,6 +49,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
insight_node.timestamp = dt_handler.timestamp
insight_node.dt = dt_handler.datetime_format()
if insight_node.status == MemoryNodeStatus.ACTIVE.value:
insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
self.logger.info(f"after_update_{insight_node.key} value={insight_value}")
return insight_node
@ -89,6 +92,10 @@ class UpdateInsightWorker(MemoryBaseWorker):
self.logger.info(f"update_{insight_node.key} insight_value={insight_value} is invalid.")
return insight_node
if insight_node.value == insight_value:
self.logger.info(f"value={insight_value} is same!")
return insight_node
self.update_insight_node(insight_node=insight_node, insight_value=insight_value)
return insight_node
@ -102,7 +109,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
return
for node in insight_nodes:
if node.content:
if node.status == MemoryNodeStatus.ACTIVE.value:
self.submit_async_task(fn=self.filter_obs_nodes,
insight_node=node,
not_updated_nodes=not_updated_nodes)
@ -126,3 +133,6 @@ class UpdateInsightWorker(MemoryBaseWorker):
# get result
self.gather_async_result()
for node in not_updated_nodes:
node.obs_updated = True

View file

@ -29,7 +29,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"dt": dt_handler.datetime_format(),
}
return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
return self.memory_store.retrieve_memories(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
def _run(self):
all_obs_nodes: List[MemoryNode] = []

View file

@ -11,7 +11,7 @@ contra_repeat_system:
Make an evaluation for each sentence, resulting in a total of {num_obs} evaluations.
Please think step by step and output in the following format:
Thought: The basis and process of the thought, within 30 words.
Evaluation: <Contradiction, Contained, None>, enclosed in <>.
Evaluation: <sentence number> <Contradiction, Contained or None>, enclosed in <>.
contra_repeat_few_shot:

View file

@ -1,7 +1,7 @@
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_WITH_TIME_NODES
from memory_scope.constants.language_constants import DATATIME_WORD_LIST, COLON_WORD
from memory_scope.constants.language_constants import COLON_WORD
from memory_scope.memory.worker.write.get_observation_worker import GetObservationWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.scheme.message import Message
@ -16,12 +16,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
user_query_list = []
i = 1
for msg in self.chat_messages:
match = False
for time_keyword in self.get_language_value(DATATIME_WORD_LIST):
if time_keyword in msg.content:
match = True
break
if match:
if DatetimeHandler.has_time_word(query=msg.content):
dt_handler = DatetimeHandler(dt=msg.time_created)
dt = dt_handler.string_format(self.prompt_handler.time_string_format)
user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")

View file

@ -4,7 +4,7 @@ time_string_format:
get_observation_with_time_system:
cn: |
任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。
任务指令:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。
每一句{user_name}句子的格式是:<序号> <对话时间> {user_name}<句子>
对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。
如果{user_name}信息涉及时间,则结合对话时间推断{user_name}信息的时间信息,没有则不输出。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。
@ -13,6 +13,16 @@ get_observation_with_time_system:
思考思考的依据和过程50字以内。
信息:<句子序号> <时间信息或“无”> <明确的重要信息或“重复”或”无“> <关键词>
en: |
Instruction: Extract important information about {user_name}, corresponding keywords, and time information from the following {num_obs} sentences by {user_name}, one by one.
Each sentence from {user_name} is formatted as: <sentence number> <conversation time> {user_name}: <sentence>
For each sentence, only extract very clear information and make very certain inferences without any speculation. Do not extract repeated information. If all information in the sentence repeats what has already been extracted, respond with "repeat." If there is no important information, respond with "none."
If the information about {user_name} involves time, combine it with the conversation time to infer the time information of {user_name}'s information; if not, do not output time information. Be sure to distinguish information; for example, do not extract hypothetical or fictional content from {user_name} such as {user_name}'s novels or scripts.
Perform information extraction for each sentence, resulting in a total of {num_obs} pieces of information.
Please think step by step, and be sure to output in the following format, with the final output enclosed in <>:
Thought: The basis and process of the thought, within 50 words.
Information: <sentence number> <Time information or "none"> <Clear important information or "repeat" or "none"> <keywords>
get_observation_with_time_few_shot:
cn: |
@ -62,7 +72,6 @@ get_observation_with_time_few_shot:
4 2023年5月21日周六14点 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢
5 2018年3月6日周四19点 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧我听你未来的你说有一个叫难忘的朱古力232这个人他弄的设是Windows设\n这是剧本1剧本2未完待续
思考从第1句可以得知{user_name}和家人上个月去杭州旅游了,这是关于{user_name}的经历的重要信息。其余信息重要性不足。{user_name}信息涉及时间结合对话时间为2023年6月推断{user_name}和家人2023年5月去杭州旅游了。
信息:<1> <2023年5月> <{user_name}和家人2023年5月去杭州旅游了。> <家人, 杭州, 旅游>
思考从第2句可以得知{user_name}的生日是昨天,这是关于{user_name}重要纪念日的信息。其余信息重要性不足。{user_name}信息涉及时间结合对话时间为2023年7月2日
@ -75,9 +84,70 @@ get_observation_with_time_few_shot:
思考第5句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。
信息:<5> <> <无> <>
en: |
Example 1:
{user_name} sentences:
1 May 1, 2022, Tuesday, 3 PM {user_name}: Please help me write a birthday greeting for my colleague Jason's daughter who is turning three.
2 May 2, 2022, Tuesday, 5 PM {user_name}: Chronology of major events in Chinese history from 1400 to 1550 AD.
3 May 3, 2022, Tuesday, 6 PM {user_name}: Can you compile a list of tips on how to use large models for me, and try to keep the content concise?
4 July 3, 2022, Thursday, 12 PM {user_name}: I got a swimming pass two months ago.
Thought: From the first sentence, it can be inferred that Jason is {user_name}'s colleague, which is important information about {user_name}'s interpersonal relationships. The remaining information is of insufficient importance. {user_name}'s information does not involve time.
Information: <1> <> <Zhang San is {user_name}'s colleague> <Zhang San, colleague>
Thought: The second sentence is a request made by {user_name}, with no clear mention of {user_name}'s personal information.
Information: <2> <> <none> <>
Thought: The third sentence is a request made by {user_name}, with no clear mention of {user_name}'s personal information.
Information: <3> <> <none> <>
Thought: From the fourth sentence, it can be inferred that {user_name} got a swimming pass two months ago. {user_name}'s information involves time. Combining it with the conversation time of July 2022, it can be inferred that {user_name} got the swimming pass in May 2022.
Information: <4> <May 2022> <{user_name} got a swimming pass in May 2022> <swimming pass>
Example 2:
{user_name} sentences:
1 January 4, 2020, Sunday, 10 AM {user_name}: I spent $5000 to buy 100 shares of General Motors.
2 April 27, 2023, Friday, 8 AM {user_name}: Tomorrow is my wedding anniversary with my wife. Could you recommend a restaurant?
3 January 4, 2020, Sunday, 10 AM {user_name}: I spent $5000 to buy 100 shares of General Motors.
4 June 2, 2021, Thursday, 11 PM {user_name}: Thanks. I'm having lunch near the company at noon; can you recommend a restaurant near Alibaba Xuhui Riverside Campus for me?
5 July 9, 2021, Saturday, 11 AM {user_name}: Two pieces of bad news: I broke my badminton racket while playing... Then I went to my friend's house to pet the cat and ended up having an allergic reaction to the cat fur, sneezing like crazy today...
Thought: From the first sentence, it can be inferred that {user_name} bought 100 shares of General Motors stock for $5000. This is important information about {user_name}'s investment decision. {user_name}'s information does not involve time.
Information: <1> <> <{user_name} bought 100 shares of General Motors stock for $5000> <General Motors, stock>
Thought: From the second sentence, it can be inferred that {user_name}'s wedding anniversary with his wife is tomorrow, which is important information about {user_name}'s significant dates. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation date of April 27, 2023, and knowing that the anniversary is a recurring date, it can be inferred that {user_name}'s wedding anniversary is on April 28th each year.
Information: <2> <April 28 each year> <{user_name}'s wedding anniversary with his wife is on April 28 each year> <wife, wedding anniversary>
Thought: The information in the third sentence is a repeat of the first sentence.
Information: <3> <> <repeat> <>
Thought: From the fourth sentence, it can be inferred that {user_name} works at Alibaba Xuhui Riverside Campus, which is important information about {user_name}'s job. The remaining information is of insufficient importance. {user_name}'s information does not involve time.
Information: <4> <> <{user_name} works at Alibaba Xuhui Riverside Campus> <Alibaba, Xuhui Riverside Campus, job>
Thought: From the fifth sentence, it can be inferred that {user_name} broke their badminton racket the other day while playing, but this is not important information. It can also be inferred that {user_name} is allergic to cat fur, which is important information about {user_name}'s health. {user_name}'s information does not involve time.
Information: <5> <> <{user_name} is allergic to cat fur> <cat fur, allergy>
Example 3:
{user_name} sentences:
1 June 30, 2023, Friday, 3 PM {user_name}: Last month, my family and I went to San Jose for a trip. The scenery was very nice.
2 July 2, 2023, Tuesday, 10 AM {user_name}: Yesterday was my birthday. I spent it alone.
3 July 3, 2020, Thursday, 11 AM {user_name}: Remind me to go for a medical check-up next Monday.
4 May 21, 2023, Saturday, 2 PM {user_name}: Someone said that passion is the best teacher and suggested linking passion with a career, but I found that many people like playing basketball, but few make a career out of it, and even fewer make money from it. Also, how do you distinguish passion from liking?
5 March 6, 2018, Thursday, 7 PM {user_name}: Zack: This is a constellation frog setting, but I am a Virgo. My mom feels I am abnormal and doesn't let me watch it. \n The female monkey patted Zack's head, "Is that so?" \n The female monkey opened Bilibili and took a look. \n Female monkey: "Why don't you switch the setting? I heard from your future self that there's someone called 'Unforgettable Chocolate 232' who created a Windows setting." \n This is script 1; script 2 is to be continued.
Thought: From the first sentence, it can be inferred that {user_name} and their family went to San Jose for a trip last month. This is important information about {user_name}'s experience. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation time of June 2023, it can be inferred that {user_name} and their family went to San Jose for a trip in May 2023.
Information: <1> <May 2023> <{user_name} and their family went to San Jose for a trip in May 2023> <family, San Jose, trip>
Thought: From the second sentence, it can be inferred that {user_name}'s birthday was yesterday. This is important information about {user_name}'s significant dates. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation time of July 2, 2023, and knowing that the birthday is a recurring date, it can be inferred that {user_name}'s birthday is on July 2 each year.
Information: <2> <July 2 each year> <{user_name}'s birthday is on July 2 each year> <birthday>
Thought: From the third sentence, it can be inferred that {user_name} will go for a medical check-up next Monday, which is an important reminder for {user_name}. {user_name}'s information involves time. Combining it with the conversation time of July 3, 2020, Thursday, it can be inferred that {user_name} will go for a check-up on July 6, 2020, Monday.
Information: <3> <July 6, 2020, Monday> <{user_name} will go for a medical check-up on July 6, 2020, Monday> <medical check-up>
Thought: The fourth sentence is a discussion and query about other people's opinions by {user_name}, with no clear mention of {user_name}'s personal information.
Information: <4> <> <none> <>
Thought: The fifth sentence is content from a script written by {user_name}, with no extractable personal information about {user_name}.
Information: <5> <> <none> <>
get_observation_with_time_user_query:
cn: |
{user_name}句子:
{user_query}
en: |
{user_name} sentences
{user_query}

View file

@ -1,7 +1,7 @@
from typing import List
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD, COLON_WORD
from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
@ -33,7 +33,7 @@ class GetObservationWorker(MemoryBaseWorker):
meta_data=meta_data,
content=obs_content,
memory_type=MemoryTypeEnum.OBSERVATION.value,
status=MemoryNodeStatus.ACTIVE.value,
status=MemoryNodeStatus.NEW.value,
timestamp=message.time_created,
obs_reflected=False,
obs_updated=False)
@ -43,12 +43,7 @@ class GetObservationWorker(MemoryBaseWorker):
user_query_list = []
i = 1
for msg in self.chat_messages:
match = False
for time_keyword in self.get_language_value(DATATIME_WORD_LIST):
if time_keyword in msg.content:
match = True
break
if not match:
if not DatetimeHandler.has_time_word(query=msg.content):
user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
i += 1

View file

@ -13,7 +13,7 @@ get_observation_system:
Perform information extraction for each sentence, resulting in a total of {num_obs} pieces of information.
Please output the results in the following format, with the final output enclosed in <>:
Thought: The basis and process of the thought, within 50 words.
Information: <> <Clear important information or “Repeat” or “None”> <keywords>
Information: <sentence number> <> <Clear important information or “Repeat” or “None”> <keywords>
get_observation_few_shot:

View file

@ -14,7 +14,7 @@ class StoreMemoryWorker(MemoryBaseWorker):
if self.has_content(store_key):
memory_nodes: List[MemoryNode] = self.get_context(store_key)
self.vector_store.update_batch(memory_nodes)
self.memory_store.update_memories(memory_nodes)
elif store_key in self.chat_kwargs:
query = self.chat_kwargs[store_key]
@ -27,8 +27,8 @@ class StoreMemoryWorker(MemoryBaseWorker):
target_name=self.target_name,
content=query,
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
status=MemoryNodeStatus.ACTIVE.value,
status=MemoryNodeStatus.NEW.value,
timestamp=dt_handler.timestamp,
obs_reflected=False,
obs_updated=False)
self.vector_store.update(node)
self.memory_store.update_memories(node)

View file

@ -1,13 +1,12 @@
import datetime
from typing import Dict, List
from uuid import uuid4
from pydantic import Field, BaseModel
from memory_scope.utils.tool_functions import md5_hash
class MemoryNode(BaseModel):
memory_id: str = Field("", description="unique id for memory")
memory_id: str = Field(uuid4(), description="unique id for memory")
user_name: str = Field("", description="the user who owns the memory")
@ -43,7 +42,6 @@ class MemoryNode(BaseModel):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}"
self.dt = datetime.datetime.fromtimestamp(self.timestamp).strftime("%Y%m%d")
@property
@ -52,4 +50,3 @@ class MemoryNode(BaseModel):
def __getitem__(self, key: str):
return self.model_dump().get(key)

View file

@ -0,0 +1,33 @@
from abc import ABCMeta, abstractmethod
from typing import Dict, List
from memory_scope.scheme.memory_node import MemoryNode
class BaseMemoryStore(metaclass=ABCMeta):
@abstractmethod
def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
@abstractmethod
async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
@abstractmethod
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
"""
status:
1. new: emb & insert
2. modified: update
3. content_modified: emb & update
4. active: do nothing
5. expired: update
"""
def flush(self):
pass
@abstractmethod
def close(self):
pass

View file

@ -1,37 +0,0 @@
from abc import ABCMeta, abstractmethod
from typing import Dict, List
from memory_scope.scheme.memory_node import MemoryNode
class BaseVectorStore(metaclass=ABCMeta):
@abstractmethod
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
@abstractmethod
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
@abstractmethod
def insert(self, node: MemoryNode):
pass
def insert_batch(self, nodes: List[MemoryNode]):
pass
def delete(self, node: MemoryNode):
pass
def update(self, node: MemoryNode):
pass
def update_batch(self, nodes: List[MemoryNode]):
pass
def flush(self):
pass
def close(self):
pass

View file

@ -0,0 +1,24 @@
from typing import Dict, List
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_memory_store import BaseMemoryStore
class DummyMemoryStore(BaseMemoryStore):
def __init__(self, embedding_model: BaseModel, **kwargs):
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 update_memories(self, nodes: MemoryNode | List[MemoryNode]):
pass
def close(self):
pass

View file

@ -1,21 +0,0 @@
from typing import Dict, List
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_vector_store import BaseVectorStore
class DummyVectorStore(BaseVectorStore):
def __init__(self, embedding_model: BaseModel, **kwargs):
self.embedding_model: BaseModel = embedding_model
self.kwargs = kwargs
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
pass
def insert(self, node: MemoryNode):
pass

View file

@ -1,143 +0,0 @@
from typing import Dict, List, Any
from llama_index.core import VectorStoreIndex
from llama_index.core.schema import TextNode, NodeWithScore
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_vector_store import BaseVectorStore
from memory_scope.utils.logger import Logger
class _ElasticsearchStore(ElasticsearchStore):
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
"""
Async delete node from Elasticsearch index.
Args:
ref_doc_id: ID of the node to delete.
delete_kwargs: Optional. Additional arguments to
pass to AsyncElasticsearch delete_by_query.
Raises:
Exception: If AsyncElasticsearch delete_by_query fails.
"""
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
"""
Convert standard filters to Elasticsearch filter.
Args:
standard_filters: Standard Llama-index filters.
Returns:
Elasticsearch filter.
"""
result = {
"bool": {}
}
for key, value in standard_filters.items():
if isinstance(value, list):
operands = []
for v in value:
operands.append(
{
"term":
{
f"metadata.{key}.keyword": {"value": v}
}
}
)
result['bool'].update({"should": operands})
result['bool'].update({"minimum_should_match": 1})
else:
operand = [{
"term": {
f"metadata.{key}.keyword": {
"value": value,
}
}
}]
if "must" in result['bool']:
result['bool']['must'].extend(operand)
else:
result['bool'].update({"must": operand})
return result
class LlamaIndexElasticSearchStore(BaseVectorStore):
def __init__(self,
embedding_model: BaseModel,
index_name: str,
es_url: str,
use_hybrid: bool = True,
**kwargs):
self.embedding_model: BaseModel = embedding_model
self.es_store = _ElasticsearchStore(index_name=index_name,
es_url=es_url,
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid),
**kwargs)
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
embed_model=self.embedding_model.model)
self.logger = Logger.get_logger()
def retrieve(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k)
text_nodes = retriever.retrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
async def async_retrieve(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}")
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(
vector_store_kwargs={"es_filter": es_filter},
similarity_top_k=top_k)
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
def insert(self, node: MemoryNode):
self.index.insert_nodes([self._memory_node_2_text_node(node)])
def delete(self, node: MemoryNode):
memory_id = node.memory_id
return self.es_store.delete(memory_id)
def update(self, node: MemoryNode):
self.delete(node)
self.insert(node)
def update_batch(self, nodes: List[MemoryNode]):
for node in nodes:
self.update(node)
def close(self):
self.es_store.close()
@staticmethod
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
return TextNode(id_=memory_node.memory_id,
text=memory_node.content,
metadata=memory_node.model_dump(exclude={"content"}))
@staticmethod
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
return MemoryNode(content=text_node.text, **text_node.metadata)

View file

@ -0,0 +1,252 @@
from typing import Dict, List, Any, Optional, cast
from llama_index.core import VectorStoreIndex
from llama_index.core.schema import TextNode, NodeWithScore
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_memory_store import BaseMemoryStore
from memory_scope.utils.logger import Logger
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
def _hybrid(
self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int,
) -> Dict[str, Any]:
# Add a query to the knn query.
# RRF is used to even the score from the knn query and text query
# RRF has two optional parameters: {'rank_constant':int, 'window_size':int}
# https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html
query_body = {
"knn": knn,
"query": {
"bool": {
"must": [
{
"match": {
self.text_field: {
"query": query,
}
}
}
],
"filter": filter,
}
},
}
if isinstance(self.rrf, Dict):
query_body["rank"] = {"rrf": self.rrf}
elif isinstance(self.rrf, bool) and self.rrf is True:
query_body["rank"] = {"rrf": {"window_size": top_k}}
return query_body
def es_query(
self,
*,
query: Optional[str],
query_vector: Optional[List[float]],
text_field: str,
vector_field: str,
k: int,
num_candidates: int,
filter: List[Dict[str, Any]] = [],
) -> Dict[str, Any]:
knn = {
"filter": filter,
"field": vector_field,
"k": k,
"num_candidates": num_candidates,
}
if query_vector is not None:
knn["query_vector"] = query_vector
else:
# Inference in Elasticsearch. When initializing we make sure to always have
# a model_id if don't have an embedding_service.
knn["query_vector_builder"] = {
"text_embedding": {
"model_id": self.model_id,
"model_text": query,
}
}
if self.hybrid:
return self._hybrid(query=cast(str, query), knn=knn, filter=filter, top_k=k)
return {"knn": knn}
class _ElasticsearchStore(ElasticsearchStore):
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
"""
Async delete node from Elasticsearch index.
Args:
ref_doc_id: ID of the node to delete.
delete_kwargs: Optional. Additional arguments to
pass to AsyncElasticsearch delete_by_query.
Raises:
Exception: If AsyncElasticsearch delete_by_query fails.
"""
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
"""
Convert standard filters to Elasticsearch filter.
Args:
standard_filters: Standard Llama-index filters.
Returns:
Elasticsearch filter.
"""
result = {
"bool": {}
}
for key, value in standard_filters.items():
if isinstance(value, list):
operands = []
for v in value:
operands.append(
{
"term":
{
f"metadata.{key}.keyword": {"value": v}
}
}
)
result['bool'].update({"should": operands})
result['bool'].update({"minimum_should_match": 1})
else:
operand = [{
"term": {
f"metadata.{key}.keyword": {
"value": value,
}
}
}]
if "must" in result['bool']:
result['bool']['must'].extend(operand)
else:
result['bool'].update({"must": operand})
return result
class LlamaIndexEsMemoryStore(BaseMemoryStore):
def __init__(self,
embedding_model: BaseModel,
index_name: str,
es_url: str,
use_hybrid: bool = True,
**kwargs):
self.embedding_model: BaseModel = embedding_model
self.es_store = _ElasticsearchStore(index_name=index_name,
es_url=es_url,
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
**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(text="text")])
self.logger = Logger.get_logger()
def retrieve_memories(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k, sparse_top_k=top_k)
text_nodes = retriever.retrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
async def a_retrieve_memories(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}")
if filter_dict is None:
filter_dict = {}
es_filter = _to_elasticsearch_filter(filter_dict)
retriever = self.index.as_retriever(
vector_store_kwargs={"es_filter": es_filter},
similarity_top_k=top_k)
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
return [self._text_node_2_memory_node(n) for n in text_nodes]
def insert(self, node: MemoryNode):
self.index.insert_nodes([self._memory_node_2_text_node(node)])
def delete(self, node: MemoryNode):
memory_id = node.memory_id
return self.es_store.delete(memory_id)
def update(self, node: MemoryNode):
self.delete(node)
self.insert(node)
def update_batch(self, nodes: List[MemoryNode]):
for node in nodes:
self.update(node)
def close(self):
self.es_store.close()
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
if not nodes:
self.logger.warning("empty nodes!")
return
if isinstance(nodes, MemoryNode):
nodes = [nodes]
# emb & insert new memories
# TODO batch insert
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW]
if new_memories:
for n in new_memories:
n.status = MemoryNodeStatus.ACTIVE.value
self.insert(n)
# emb & update new memories
# TODO insert overwrite
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED]
if c_modified_memories:
for n in c_modified_memories:
n.status = MemoryNodeStatus.ACTIVE.value
self.delete(n)
self.insert(n)
# update new memories
# TODO no emb
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED]
if modified_memories:
for n in modified_memories:
n.status = MemoryNodeStatus.ACTIVE.value
self.delete(n)
self.insert(n)
# set memories expired
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED]
if expired_memories:
for n in expired_memories:
n.status = MemoryNodeStatus.ACTIVE.value
self.delete(n)
self.insert(n)
@staticmethod
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
return TextNode(id_=memory_node.memory_id,
text=memory_node.content,
metadata=memory_node.model_dump(exclude={"content"}))
@staticmethod
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
return MemoryNode(content=text_node.text, **text_node.metadata)

View file

@ -2,7 +2,7 @@ import datetime
import re
from typing import Dict
from memory_scope.constants.language_constants import WEEKDAYS
from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.logger import Logger
@ -65,6 +65,82 @@ class DatetimeHandler(object):
extracted_data[key] = int(match.group(1))
return extracted_data
@classmethod
def extract_date_parts_en(cls, input_string: str) -> dict:
date_info = {
"year": -1,
"month": -1,
"day": -1,
"hour": -1,
"minute": -1,
"second": -1,
"weekday": -1
}
# Patterns to extract the parts of the date/time
patterns = {
"year": r"\b(\d{4})\b",
"month": r"\b(January|February|March|April|May|June|July|August|September|October|November|December)\b",
"day_month_year": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October"
r"|November|December) (?P<day>\d{1,2}),? (?P<year>\d{4})\b",
"day_month": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October|November"
r"|December) (?P<day>\d{1,2})\b",
"hour_12": r"\b(\d{1,2})\s*(AM|PM|am|pm)\b",
"hour_24": r"\b(\d{1,2}):(\d{2}):(\d{2})\b"
}
month_mapping = {
"January": 1, "February": 2, "March": 3, "April": 4, "May": 5, "June": 6, "July": 7, "August": 8,
"September": 9, "October": 10, "November": 11, "December": 12
}
weekday_mapping = {
"Monday": 1, "Tuesday": 2, "Wednesday": 3, "Thursday": 4, "Friday": 5, "Saturday": 6, "Sunday": 7
}
day_month_year_match = re.search(patterns["day_month_year"], input_string)
if day_month_year_match:
date_info["year"] = int(day_month_year_match.group("year"))
date_info["month"] = month_mapping[day_month_year_match.group("month")]
date_info["day"] = int(day_month_year_match.group("day"))
# Extract month and day without year
elif date_info["year"] == -1:
day_month_match = re.search(patterns["day_month"], input_string)
if day_month_match:
date_info["month"] = month_mapping[day_month_match.group("month")]
date_info["day"] = int(day_month_match.group("day"))
# Extract year
if date_info["year"] == -1:
year_match = re.search(patterns["year"], input_string)
if year_match:
date_info["year"] = int(year_match.group(0))
# Extract month
if date_info["month"] == -1:
month_match = re.search(patterns["month"], input_string)
if month_match:
date_info["month"] = month_mapping[month_match.group(0)]
# Extract 12-hour format time
hour_12_match = re.search(patterns["hour_12"], input_string)
if hour_12_match:
hour, period = int(hour_12_match.group(1)), hour_12_match.group(2).lower()
if period == 'pm' and hour != 12:
hour += 12
elif period == 'am' and hour == 12:
hour = 0
date_info["hour"] = hour
# Extract weekday
for week_day, value in weekday_mapping.items():
if week_day in input_string:
date_info["weekday"] = value
break
return date_info
@classmethod
def extract_date_parts(cls, input_string: str) -> dict:
func_name = f"extract_date_parts_{G_CONTEXT.language}"
@ -96,6 +172,16 @@ class DatetimeHandler(object):
return ""
return getattr(cls, func_name)(extract_time_dict, meta_data)
@classmethod
def has_time_word(cls, query: str) -> bool:
contain_datetime = False
# TODO use re
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
if datetime_word in query:
contain_datetime = True
break
return contain_datetime
def datetime_format(self, dt_format: str = "%Y%m%d"):
return self._dt.strftime(dt_format)

View file

@ -6,7 +6,7 @@ from memory_scope.enumeration.language_enum import LanguageEnum
from memory_scope.memory.service.base_memory_service import BaseMemoryService
from memory_scope.models.base_model import BaseModel
from memory_scope.storage.base_monitor import BaseMonitor
from memory_scope.storage.base_vector_store import BaseVectorStore
from memory_scope.storage.base_memory_store import BaseMemoryStore
class GlobalContext(object):
@ -18,7 +18,7 @@ class GlobalContext(object):
self.model_dict: Dict[str, BaseModel] = {}
self.memory_chat_dict: Dict[str, BaseMemoryChat] = {}
self.vector_store: BaseVectorStore | None = None
self.memory_store: BaseMemoryStore | None = None
self.monitor: BaseMonitor | None = None
self.thread_pool: ThreadPoolExecutor | None = None
self.language: LanguageEnum = LanguageEnum.EN

View file

@ -50,7 +50,7 @@ def init_instance_by_config(config: dict,
**kwargs: Additional keyword arguments to pass to the class constructor.
Returns:
object: An instance of the class initialized with the provided config and kwargs.
instance: An instance of the class initialized with the provided config and kwargs.
"""
config_copy = deepcopy(config)

View file

View file

@ -1,124 +0,0 @@
import json
import time
from http import HTTPStatus
import requests
from utils.logger import Logger
from utils.timer import Timer
from enumeration.env_type import EnvType
class DashClient(object):
def __init__(self,
request_id: str,
dash_scope_uid: str,
authorization: str,
workspace: str,
model_name: str,
env_type: EnvType | str = EnvType.DAILY,
timeout: int = None,
max_retry_count: int = 2,
retry_sleep_time: float = 1.0,
**kwargs):
self.model_name: str = model_name
self.env_type: EnvType = EnvType(env_type)
self.timeout: int = timeout
self.max_retry_count: int = max_retry_count
self.retry_sleep_time: float = retry_sleep_time
self.kwargs: dict = kwargs
# 20240506 update by 泉雨
# if authorization:
# workspace = ""
# dash_scope_uid = ""
self.headers = {
'Content-Type': 'application/json',
'Authorization': authorization,
'X-Request-Id': request_id,
'X-DashScope-Uid': dash_scope_uid,
'X-DashScope-WorkSpace': workspace,
}
self.url: str = ""
self.data = {}
self.logger = Logger.get_logger()
def before_call(self, model_name: str = None, **kwargs):
pass
def after_call(self, response_obj, **kwargs):
pass
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
if model_name is None:
model_name = self.model_name
self.before_call(model_name=model_name, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
self.logger.debug(f"url={self.url} header={self.headers} data={self.data} timeout={self.timeout}")
response = requests.post(url=self.url,
headers=self.headers,
data=json.dumps(self.data),
timeout=self.timeout)
if response.status_code == HTTPStatus.OK:
response_obj = json.loads(response.text)
self.logger.info(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
f"call model={model_name} success! retry_cnt={retry_cnt}",
stacklevel=3)
return self.after_call(response_obj, **kwargs), True
else:
self.logger.warning(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
f"call model={model_name} failed! retry_cnt={retry_cnt} details={response.text}",
stacklevel=3)
return None, False
def call(self, model_name: str = None, **kwargs):
for i in range(self.max_retry_count):
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
if flag:
return result
else:
time.sleep(self.retry_sleep_time)
return None
class LLIClient(object):
def __init__(self,
model_name: str,
timeout: int = None,
max_retry_count: int = 2,
retry_sleep_time: float = 1.0,
**kwargs):
self.model_name: str = model_name
self.timeout: int = timeout
self.max_retry_count: int = max_retry_count
self.retry_sleep_time: float = retry_sleep_time
self.kwargs: dict = kwargs
self.data = {}
self.logger = Logger.get_logger()
def before_call(self, **kwargs):
pass
def after_call(self, **kwargs):
pass
def call_once(self, **kwargs):
pass
def call(self, **kwargs):
pass

View file

@ -1,103 +0,0 @@
from typing import List, Dict
import dashscope
import time
from models import EMB
from models.dash_client import DashClient, LLIClient
from typing import List, Dict
from utils.registry import build_from_cfg
from utils.timer import Timer
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
from enumeration.dash_api_enum import DashApiEnum
class DashEmbeddingClient(DashClient):
"""
url: https://help.aliyun.com/document_detail/2782232.html?spm=a2c4g.2782227.0.0.76195b1d9UeBAk#a6a39590fegqx
"""
def __init__(self, model_name: str = dashscope.TextEmbedding.Models.text_embedding_v2, **kwargs):
super(DashEmbeddingClient, self).__init__(model_name=model_name, **kwargs)
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.EMBEDDING)
def before_call(self, model_name: str = None, **kwargs):
text: str | List[str] = kwargs.pop("text", "")
# text_type: query or document
text_type: str = kwargs.pop("text_type", "query")
if isinstance(text, str):
text = [text]
self.kwargs["text_type"] = text_type
self.data = {
"model": model_name,
"input": {
"texts": text,
},
"parameters": {**kwargs, **self.kwargs},
}
def after_call(self, response_obj, **kwargs) -> Dict[int, List[float]] | List[float]:
embedding_results = {}
for emb in response_obj["output"]["embeddings"]:
embedding_results[emb["text_index"]] = emb["embedding"]
if len(embedding_results) == 1:
embedding_results = list(embedding_results.values())[0]
return embedding_results
class LLIEmbedding(LLIClient):
def __init__(self, method, model_name, **kwargs):
super(LLIEmbedding, self).__init__(model_name, **kwargs)
self.config = {
"method": method,
"model_name": model_name,
**kwargs}
self.embedder = build_from_cfg(self.config, EMB)
def before_call(self, **kwargs):
text: str | List[str] = kwargs.pop("text", "")
if isinstance(text, str):
text = [text]
self.data = dict(texts=text)
def after_call(self, emb: Dict[int, List[float]], **kwargs) -> Dict[int, List[float]] | List[float]:
embedding_results = {}
for idx, e in enumerate(emb):
embedding_results[idx] = e
if len(embedding_results) == 1:
embedding_results = list(embedding_results.values())[0]
return embedding_results
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
if model_name is None:
model_name = self.model_name
self.before_call(model_name=model_name, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
self.logger.debug(f"data={self.data} timeout={self.timeout}")
try:
results = self.embedder.get_text_embedding_batch(**self.data)
results = self.after_call(results)
return results, True
except Exception as e:
self.logger.debug(f"Get Error in Embedding: {e}")
return None, False
def call(self, model_name: str = None, **kwargs):
for i in range(self.max_retry_count):
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
if flag:
return result
else:
time.sleep(self.retry_sleep_time)
return None

View file

@ -1,129 +0,0 @@
from typing import List, Dict
import dashscope
from models.dash_client import DashClient, LLIClient
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
from enumeration.dash_api_enum import DashApiEnum
import time
from typing import List, Dict
from utils.timer import Timer
from models import LLM
from utils.registry import build_from_cfg
from llama_index.core.base.llms.types import ChatMessage
from llama_index.core.base.llms.types import (
ChatResponse,
CompletionResponse,
)
class DashGenerateClient(DashClient):
"""
url: https://help.aliyun.com/document_detail/2712576.html
"""
def __init__(self, model_name: str = dashscope.Generation.Models.qwen_max, **kwargs):
super(DashGenerateClient, self).__init__(model_name=model_name, **kwargs)
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.GENERATION)
def before_call(self, model_name: str = None, **kwargs):
prompt: str = kwargs.pop("prompt", "")
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
input_text = {}
if prompt:
input_text["prompt"] = prompt
elif messages:
input_text["messages"] = messages
else:
raise RuntimeError("prompt and messages is both empty!")
self.data = {
"model": model_name,
"input": input_text,
"parameters": {**kwargs, **self.kwargs},
}
def after_call(self, response_obj, **kwargs):
self.logger.debug(f"response_obj={response_obj}")
output = response_obj["output"]
if "text" in output:
return output["text"]
elif "choices" in output:
return output["choices"][0]["message"]["content"]
else:
raise NotImplementedError
class LLILLM(LLIClient):
def __init__(self, method, model_name: str, **kwargs):
super(LLILLM, self).__init__(model_name, **kwargs)
self.config = {
"method": method,
"model_name": model_name,
**kwargs}
self.llm = build_from_cfg(self.config, LLM)
def before_call(self, model_name: str = None, **kwargs):
prompt: str = kwargs.pop("prompt", "")
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
if prompt:
input_text = prompt
input_type = 'prompt'
llama_input = input_text
elif messages:
input_text = messages
input_type = 'messages'
llama_input = [ChatMessage(
role=x['role'], content=x['content']
) for x in input_text]
else:
raise RuntimeError("prompt and messages is both empty!")
self.data = {
input_type: llama_input,
}
def after_call(self, response_obj: ChatResponse | CompletionResponse, **kwargs) -> str:
self.logger.debug(f"response_obj={response_obj}")
if isinstance(response_obj, CompletionResponse):
return response_obj.text
elif isinstance(response_obj, ChatResponse):
return response_obj.message.content
else:
raise NotImplementedError
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
if model_name is None:
model_name = self.model_name
self.before_call(model_name=model_name, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
self.logger.debug(f"data={self.data} timeout={self.timeout}")
if True:
# try:
if 'prompt' in self.data:
results = self.llm.complete(**self.data)
else:
results = self.llm.chat(**self.data)
results = self.after_call(results)
return results, True
# except:
# return None, False
def call(self, model_name: str = None, **kwargs):
for i in range(self.max_retry_count):
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
print("dashscope llm results:",result)
if flag:
return result
else:
time.sleep(self.retry_sleep_time)
return None

View file

@ -1,119 +0,0 @@
from typing import List
import dashscope
from models.dash_client import DashClient, LLIClient
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
from enumeration.dash_api_enum import DashApiEnum
import time
from typing import List
from models import RERANKER
from utils.timer import Timer
from utils.registry import build_from_cfg
from llama_index.core.data_structs import Node
from llama_index.core.schema import NodeWithScore # type: ignore
class DashReRankClient(DashClient):
"""
url: https://help.aliyun.com/document_detail/2780059.html
"""
def __init__(self, model_name: str = dashscope.TextReRank.Models.gte_rerank, **kwargs):
super(DashReRankClient, self).__init__(model_name=model_name, **kwargs)
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.RERANK)
def before_call(self, model_name: str = None, **kwargs):
query: str = kwargs.pop("query", "")
documents: List[str] = kwargs.pop("documents", [])
top_n: int | None = kwargs.pop("top_n", None)
return_documents: bool = kwargs.pop("return_documents", False)
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
if top_n is None:
top_n = len(documents)
self.kwargs.update({
"top_n": top_n,
"return_documents": return_documents,
})
self.data = {
"model": model_name,
"input": {
"query": query,
"documents": documents,
},
"parameters": {**kwargs, **self.kwargs},
}
def after_call(self, response_obj, **kwargs):
return response_obj["output"]["results"]
class LLIReRank(LLIClient):
def __init__(self, method, model_name, **kwargs):
super(LLIReRank, self).__init__(model_name, **kwargs)
self.config = {
"method": method,
"model_name": model_name,
**kwargs}
self.reranker = build_from_cfg(self.config, RERANKER)
def before_call(self, model_name: str = None, **kwargs):
query: str = kwargs.pop("query", "")
documents: List[str] = kwargs.pop("documents", [])
top_n: int | None = kwargs.pop("top_n", None)
return_documents: bool = kwargs.pop("return_documents", False)
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
if top_n is None:
top_n = len(documents)
nodes = [NodeWithScore(node=Node(text=text), score=1.0) for text in documents]
self.data = {
"nodes": nodes,
"query_str": query,
}
def after_call(self, nodes, **kwargs):
results = []
for idx, node in enumerate(nodes):
results.append(dict(index=idx,
relevance_score=node.score,
document=node.node.text))
return results
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
if model_name is None:
model_name = self.model_name
self.before_call(model_name=model_name, **kwargs)
with Timer(self.__class__.__name__, log_time=False) as t:
self.logger.debug(f"data={self.data} timeout={self.timeout}")
try:
results = self.reranker.postprocess_nodes(**self.data)
results = self.after_call(results)
return results, True
except Exception as e:
self.logger.debug(f"Rerank falls, data={self.data}")
# return None, False
def call(self, model_name: str = None, **kwargs):
for i in range(self.max_retry_count):
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
if flag:
return result
else:
time.sleep(self.retry_sleep_time)
return None

View file

@ -1,419 +0,0 @@
from elasticsearch import Elasticsearch
from elasticsearch.helpers import bulk
from models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
from common.dash_embedding_client import DashEmbeddingClient
from common.logger import Logger
from constants.common_constants import ES_ENV_URL_DICT
from enumeration.env_type import EnvType
from utils.logger import Logger
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
from llama_index.vector_stores.elasticsearch import ElasticsearchStore
from llama_index.core.schema import TextNode
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
class ElasticSearchClient(object):
def __init__(self,
es_user_name: str,
es_password: str,
es_index_name: str,
embedding_client: DashEmbeddingClient | None = None,
env_type: EnvType | str = EnvType.DAILY,
content_key: str = "content",
vector_key: str = "vector",
**kwargs):
self.es_index_name: str = es_index_name
self.embedding_client: DashEmbeddingClient = embedding_client
self.content_key: str = content_key
self.vector_key: str = vector_key
self.es_client = Elasticsearch(
hosts=[ES_ENV_URL_DICT.get(EnvType(env_type))],
basic_auth=(es_user_name, es_password),
**kwargs)
self.logger = Logger.get_logger()
self.logger.debug(f"connect es_client info={self.es_client.info()}")
def log_index_info(self):
index_info = self.es_client.indices.get(index=self.es_index_name)
self.logger.info(f"index={self.es_index_name} exists. index_info={index_info}")
def insert(self, _id: str, body: dict):
assert body and self.content_key in body, f"body={body} is illegal!"
# text_type: document
content = body[self.content_key]
vector = self.embedding_client.call(text=content, text_type="document")
if not vector:
self.logger.warning(f"embedding_client call failed, stop es insert!")
return
body[self.vector_key] = vector
response = self.es_client.index(id=_id, index=self.es_index_name, body=body)
self.logger.info(f"insert response={response}")
def insert_batch(self, doc_list: list):
"""
doc_list = [
{
"_id": 2,
"_source": {
"author": "john",
"text": "Elasticsearch: cool.",
"timestamp": "2023-03-23T10:00:00"
}
},
{
"_id": 3,
"_source": {
"author": "jane",
"text": "Elasticsearch: very cool.",
"timestamp": "2023-03-23T11:00:00"
}
}
]
"""
text_list = []
for doc in doc_list:
assert "_id" in doc and "_source" in doc
content = doc["_source"][self.content_key]
text_list.append(content)
vector_dict = self.embedding_client.call(text=text_list, text_type="document")
if not vector_dict:
self.logger.warning(f"embedding_client call failed, stop es insert!")
return
# add _index
for i, doc in enumerate(doc_list):
doc["_index"] = self.es_index_name
vector = vector_dict[i]
doc["_source"][self.vector_key] = vector
# 执行批量插入
responses = bulk(self.es_client, doc_list)
# 输出批量插入的响应
for response in responses[1]:
self.logger.info(f"insert_batch response={response}")
def print_hits(self, hits: list):
for hit in hits:
print_kwargs = {
"_id": hit['_id'],
"_score": hit['_score'],
}
for k, v in hit['_source'].items():
# 不打印vector
if k == self.vector_key:
v = len(v)
print_kwargs[k] = v
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
def exact_search(self,
size: int,
exact_filters: dict = None,
wildcard_filters: dict = None,
print_hits: bool = False,
exclude_vector: bool = True):
"""
{
"match": {
"category": "electronics" # 一级字段过滤
}
},
{
"match": {
"product.name": "laptop" # 二级字段过滤
}
}
{
"terms": {
"product.keyA": ["a", "b", "c"] # 二级字段keyA的精确值必须为a、b、c中的一
}
}
"""
must_list = []
for key, value in exact_filters.items():
if not key:
continue
if isinstance(value, str):
must_list.append({"match": {key: value}})
elif isinstance(value, list):
must_list.append({"terms": {key: value}})
query = {
"size": size,
"query": {
"bool": {
"must": must_list
}
},
# 添加_source配置以排除vector字段
"_source": {
"excludes": [self.vector_key] if exclude_vector else []
}
}
if wildcard_filters:
should_list = []
for key, value in wildcard_filters.items():
if not key:
continue
if isinstance(value, str):
should_list.append({"wildcard": {key: f"*{value}*"}})
elif isinstance(value, list):
for v in value:
should_list.append({"wildcard": {key: f"*{v}*"}})
query["query"]["bool"].update({
"should": should_list,
"minimum_should_match": 1,
})
self.logger.info(f"query={query}")
response = self.es_client.search(index=self.es_index_name, body=query)
hits = response['hits']['hits']
# 耗时log
self.logger.info(f"exact_search cost={response['took']}ms "
f"size={len(hits)} "
f"timed_out={response['timed_out']} "
f"shards={response['_shards']} "
f"exact_filters={exact_filters}", stacklevel=2)
# 每一条结果log一次
if print_hits:
self.print_hits(hits)
return hits
def exact_search_v2(self,
size: int,
term_filters: dict = None,
match_filters: dict = None,
print_hits: bool = False,
exclude_vector: bool = True):
"""
"bool": {
"must": [
{"term": {"field1": "固定值"}}, # 一级目录关键字过滤(等于某个值)
{"terms": {"field2": ["a", "b", "c"]}} # 二级目录关键字过滤(等于三个中的任意一个)
],
"should": [ # 至少匹配其中之一
{"match": {"key": "ccc"}}, # key包含"ccc"
{"match": {"key": "bbb"}} # 或者key包含"bbb"
],
"minimum_should_match": 1 # 至少有一个`should`条件匹配
}
"""
query = {
"size": size,
"query": {
"bool": {
}
},
# 添加_source配置以排除vector字段
"_source": {
"excludes": [self.vector_key] if exclude_vector else []
}
}
if term_filters:
must_list = []
for k, v in term_filters.items():
if isinstance(v, list):
must_list.append({"terms": {k: v}})
elif isinstance(v, str):
must_list.append({"term": {k: v}})
else:
raise NotImplemented
query["query"]["bool"]["must"] = must_list
if match_filters:
match_list = []
for k, v in match_filters.items():
if isinstance(v, list):
for v_sub in v:
match_list.append({"match": {k: v_sub}})
elif isinstance(v, str):
match_list.append({"match": {k: v}})
else:
raise NotImplemented
query["query"]["bool"]["should"] = match_list
query["query"]["bool"]["minimum_should_match"] = 1
self.logger.info(query)
response = self.es_client.search(index=self.es_index_name, body=query)
hits = response['hits']['hits']
# 耗时log
self.logger.info(f"exact_search cost={response['took']}ms "
f"size={len(hits)} "
f"timed_out={response['timed_out']} "
f"shards={response['_shards']}", stacklevel=2)
# 每一条结果log一次
if print_hits:
self.print_hits(hits)
return hits
def similar_search(self,
text: str,
size: int,
exact_filters: dict = None,
print_hits: bool = False,
exclude_vector: bool = True):
if exact_filters is None:
exact_filters = {}
# 过滤or
or_filters = {}
for k in list(exact_filters.keys()):
v = exact_filters[k]
if isinstance(v, list):
exact_filters.pop(k)
or_filters[k] = v
vector = self.embedding_client.call(text=text)
if not vector:
self.logger.warning(f"embedding_client call failed, stop select from es!")
return
query = {
# 返回最相似的top_k个文档
"size": size,
"query": {
"bool": {
"must": {
"script_score": {
# 对所有文档执行
"query": {
"match_all": {}
},
"script": {
# 使用余弦相似度+1,es不能返回负数
"source": f"cosineSimilarity(params.query_vector, '{self.vector_key}') + 1.0",
"params": {"query_vector": vector}
}
}
},
"filter": [
{"term": {k: v}} for k, v in exact_filters.items()
],
}
},
# 添加_source配置以排除vector字段
"_source": {
"excludes": [self.vector_key] if exclude_vector else []
}
}
if or_filters:
k_v_pair = []
for k, v_list in or_filters.items():
for v in v_list:
k_v_pair.append((k, v))
query["query"]["bool"]["should"] = [{"term": {k: v}} for k, v in k_v_pair]
query["query"]["bool"]["minimum_should_match"] = 1
response = self.es_client.search(index=self.es_index_name, body=query)
hits = response['hits']['hits']
# 耗时log
self.logger.info(f"similar_search cost={response['took']}ms "
f"size={len(hits)} "
f"timed_out={response['timed_out']} "
f"shards={response['_shards']} "
f"text={text} "
f"exact_filters={exact_filters}", stacklevel=2)
# 还原打分
for hit in hits:
hit['_score'] -= 1
# 每一条结果log一次
if print_hits:
self.print_hits(hits)
return hits
class LLIElasticSearch(object):
def __init__(self,
es_index_name: str,
embedding_client: LLIEmbedding | None = None,
retrieve_topk: int = 3,
content_key: str = "text",
):
self.es_index_name = es_index_name
self.content_key = content_key
self.embedding_client: LLIEmbedding = embedding_client
# using local es for debug convenient
self.es_client = ElasticsearchStore(index_name="my_index",
es_url="http://localhost:9200",
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True))
self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_client, llm=None)
self.storage_context = StorageContext.from_defaults(vector_store=self.es_client)
self.index = VectorStoreIndex(storage_context=self.storage_context,
service_context=self.service_context)
self.retriever = self.index.as_retriever(similarity_top_k=retrieve_topk)
self.logger = Logger.get_logger()
def log_index_info(self, ):
pass
def print_hits(self, hits: list):
for hit in hits:
print_kwargs = {
"_id": hit['_id'],
"_score": hit['_score'],
}
for k, v in hit['_source'].items():
# 不打印vector
if k == self.vector_key:
v = len(v)
print_kwargs[k] = v
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
def similar_search(self,
text: str,
size: int, ):
ret_nodes = self.retriever.retrieve(text)
return ret_nodes
def insert_batch(self, doc_list:list[str]):
node_list = []
for doc in doc_list:
assert "_id" in doc and "_source" in doc
content = doc["_source"]["text"]
doc["_source"].pop("text")
meta = doc["_source"]
node = TextNode(text=content, metadata=meta)
node.node_id(doc['_id'])
node_list.append(node)
self.index.insert_nodes(node_list)
def insert(self, _id: str, body: dict):
assert body and self.content_key in body, f"body={body} is illegal!"
content = body[self.content_key]
body.pop(self.content_key)
meta = body
node = TextNode(text=content, metadata=meta)
self.index.insert_nodes([node])

View file

@ -1,71 +0,0 @@
import re
from typing import Dict, List
from pydantic import Field, BaseModel
class MemoryNode(BaseModel):
"""
除了 content_modified其他均和数据库字段保持统一
根据code判断如果code是空则为新增的memoryNode如果有值则为更新
if content_modified is true则需要调用embedding服务
"""
id: str = Field("", description="唯一主键 uuid64")
code: str = Field("", description="和id保持一致为空则是新增")
# 0520新增
timeCreated: str = Field("", description="Memory创建时间算法不关注")
# 0520新增
timeModified: str = Field("", description="Memory更新时间算法不关注")
content: str = Field("", description="记忆内容")
memoryId: str = Field("", description="记忆 id检索区分字段")
# 0520新增
scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT")
# 0520新增
# NOTE 百炼服务端只召回observation, insight, profile, obs_customized, profile_customized
memoryType: str = Field("", description="conversation, observation, insight, "
"profile, obs_customized, profile_customized")
# 0520新增但不是数据库字段
content_modified: bool = Field(False, description="content是否被更新if true则需要调用embedding服务")
# reflected: 1 is reflected before, 0 has not reflected, 如果是用户自定义,写入空值"".
metaData: Dict[str, str] = Field({}, description="元信息: infoScore, algoVersion, datetime, reflected")
status: str = Field("active", description="active or expired")
tenantId: str = Field("", description="request id")
vector: List[float] = Field([], description="content embedding result, return empty")
def get_time_info(self, time_format: str):
pattern = re.compile(r'\{([^}]*)}')
keys = pattern.findall(time_format)
match_flag = True
kv_dict = {}
for k in keys:
if k not in self.metaData:
match_flag = False
break
v = self.metaData[k]
if not v:
match_flag = False
break
kv_dict[k] = v
if match_flag:
return time_format.format(**kv_dict)
return ""
def to_dict(self):
res = {"content": self.content, "memoryId": self.memoryId, "memoryType": self.memoryType,
"status": self.status, "metaData": self.metaData}
return res

View file

@ -1,33 +0,0 @@
from pydantic import Field, BaseModel
from scheme.memory_node import MemoryNode
class MemoryNode(BaseModel):
id: str = Field("", description="uuid64")
score_similar: float = Field(0, description="相似度打分")
score_rank: float = Field(0, description="排序打分")
score_rerank: float = Field(0, description="重排打分")
memory_node: MemoryNode = Field(None, description="memory node 核心,返回给上游的结构")
@classmethod
def init_from_es(cls, hit: dict):
memory_node = MemoryNode(**hit['_source'])
return cls(id=hit['_id'], score_similar=hit['_score'], memory_node=memory_node)
@classmethod
def init_from_attrs(cls, **kwargs):
_id: str = kwargs.get("_id", "")
score_similar: float = kwargs.pop("score_similar", 0)
score_rank: float = kwargs.pop("score_rank", 0)
score_rerank: float = kwargs.pop("score_rerank", 0)
memory_node = MemoryNode(**kwargs)
return cls(id=_id,
score_similar=score_similar,
score_rank=score_rank,
score_rerank=score_rerank,
memory_node=memory_node)

View file

@ -1,140 +0,0 @@
from datetime import datetime
from typing import List
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERGE_NODES, NEW_INSIGHT_KEYS, INSIGHT_KEY, \
INSIGHT_VALUE, REFLECTED
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class GetInsightWorker(MemoryBaseWorker):
def __init__(self, insight_obs_max_cnt, es_insight_similar_top_k, get_insight_model, get_insight_max_token, get_insight_temperature, get_insight_top_k, **kwargs):
super(GetInsightWorker,self).__init__(*args,**kwargs)
self.insight_obs_max_cnt = insight_obs_max_cnt
self.get_insight_model = get_insight_model
self.get_insight_max_token = get_insight_max_token
self.get_insight_temperature = get_insight_temperature
self.get_insight_top_k = get_insight_top_k
self.es_insight_similar_top_k = es_insight_similar_top_k
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
created_dt = datetime.now()
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
DT: dt,
INSIGHT_KEY: insight_key,
INSIGHT_VALUE: insight_value,
}
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
content = f"用户的{insight_key}{insight_value}"
return MemoryNode.init_from_attrs(content=content,
memoryId=self.memory_id,
scene=self.scene,
memoryType=MemoryTypeEnum.INSIGHT.value,
content_modified=True, # 新增的insight需要置为true
metaData=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
tenantId=self.tenant_id)
def reflect_new_insight_key(self,
insight_key: str,
not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None:
# 检索历史memory
hits = self.es_client.similar_search(text=insight_key,
size=self.es_insight_similar_top_k,
exact_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value],
})
# 转化成 MemoryNodeWrap 合并新增nodes
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
related_nodes.extend(not_reflected_merge_nodes)
# content去重
related_node_dict = {n.memory_node.content: n for n in related_nodes}
related_nodes = sorted(list(related_node_dict.values()), key=lambda x: x.memory_node.id)
documents = [n.memory_node.content for n in related_nodes]
# 重排所有记忆
result = self.rerank_client.call(query=insight_key, documents=documents)
if not result:
self.add_run_info(f"reflect insight_key={insight_key} call rerank client failed!")
return
# 根据打分过滤
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
related_nodes[index].score_rank = score
related_nodes_sorted = sorted(related_nodes, key=lambda x: x.score_rank, reverse=True)[
:self.insight_obs_max_cnt]
# 生成prompt
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
get_insight_message = self.prompt_to_msg(
system_prompt=self.prompt_config.get_insight_system,
few_shot=self.prompt_config.get_insight_few_shot,
user_query=self.prompt_config.get_insight_user_query.format(
insight_key=insight_key, user_query="\n".join(user_query_list)))
self.logger.info(f"get_insight_message={get_insight_message}")
# call LLM, 提取insight
response_text = self.gene_client.call(messages=get_insight_message,
model_name=self.get_insight_model,
max_token=self.get_insight_max_token,
temperature=self.get_insight_temperature,
top_k=self.get_insight_top_k)
# return if empty
if not response_text:
self.add_run_info("reflect_upon_user_attr call llm failed!")
return
response_text = response_text.strip()
if response_text in [""]:
return
return self.new_insight_node(insight_key=insight_key, insight_value=response_text)
def _run(self):
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
if not new_insight_keys:
self.add_run_info("new_insight_keys is empty! stop insight.")
return
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
if not not_reflected_merge_nodes:
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
return
# submit insight task
for insight_key in new_insight_keys:
self.submit_thread(self.reflect_new_insight_key,
sleep_time=1,
insight_key=insight_key,
not_reflected_merge_nodes=not_reflected_merge_nodes)
# save output
new_insight_nodes: List[MemoryNode] = []
for result in self.join_threads():
if result:
new_insight_nodes.append(result)
assert isinstance(result, MemoryNode)
insight_key = result.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = result.memory_node.metaData.get(INSIGHT_VALUE, "")
self.logger.info(f"after_get_insight insight_key={insight_key} insight_value={insight_value}")
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
# set REFLECTED
for node in not_reflected_merge_nodes:
scheme.memory_node.metaData[REFLECTED] = "1"

View file

@ -1,80 +0,0 @@
from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, NOT_REFLECTED_OBS_NODES, REFLECTED, INSIGHT_NODES, INSIGHT_KEY, \
NEW_INSIGHT_KEYS, NOT_REFLECTED_MERGE_NODES
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class GetReflectionWorker(MemoryBaseWorker):
def __init__(self, reflect_obs_cnt_threshold, reflect_num_questions, reflect_obs_model, reflect_obs_max_token, reflect_obs_temperature, reflect_obs_top_k, *args, **kwargs):
super(GetReflectionWorker,self).__init__(*args, **kwargs)
self.reflect_obs_cnt_threshold = reflect_obs_cnt_threshold
self.reflect_num_questions = reflect_num_questions
self.reflect_obs_model = reflect_obs_model
self.reflect_obs_max_token = reflect_obs_max_token
self.reflect_obs_temperature = reflect_obs_temperature
self.reflect_obs_top_k = reflect_obs_top_k
def _run(self):
# 过滤得到 not_reflected_merge_nodes
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
not_reflected_merge_nodes: List[MemoryNode] = []
if new_obs_nodes:
not_reflected_merge_nodes.extend(new_obs_nodes)
if not_reflected_nodes:
not_reflected_merge_nodes.extend(not_reflected_nodes)
not_reflected_merge_nodes = [node for node in not_reflected_merge_nodes
if scheme.memory_node.metaData.get(REFLECTED, "") == "0"]
# count
not_reflected_count = len(not_reflected_merge_nodes)
if not_reflected_count <= self.reflect_obs_cnt_threshold:
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop reflect.")
return
# save context
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
# get profile_keys
exist_keys: List[str] = []
profile_keys: List[str] = list(self.user_profile_dict.keys())
exist_keys.extend(profile_keys)
self.logger.info(f"profile_keys={profile_keys}")
# get insight_keys
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if insight_nodes:
insight_keys = [n.memory_node.metaData.get(INSIGHT_KEY) for n in insight_nodes]
insight_keys = [x.strip() for x in insight_keys if x]
exist_keys.extend(insight_keys)
self.logger.info(f"insight_keys={insight_keys}")
# gen reflect prompt
user_query_list = [n.memory_node.content for n in not_reflected_merge_nodes]
reflect_message = self.prompt_to_msg(
system_prompt=self.prompt_config.get_reflect_system.format(
num_questions=self.reflect_num_questions),
few_shot=self.prompt_config.get_reflect_few_shot,
user_query=self.prompt_config.get_reflect_user_query.format(exist_keys="".join(exist_keys),
user_query="\n".join(user_query_list)))
self.logger.info(f"reflect_message={reflect_message}")
# # call LLM
response_text = self.gene_client.call(messages=reflect_message,
model_name=self.reflect_obs_model,
max_token=self.reflect_obs_max_token,
temperature=self.reflect_obs_temperature,
top_k=self.reflect_obs_top_k)
# return if empty
if not response_text:
self.add_run_info("reflect_obs_questions call llm failed!")
return
# parse text & save
new_insight_keys = ResponseTextParser(response_text).parse_v2("get_insight_keys")
if new_insight_keys:
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)

View file

@ -1,118 +0,0 @@
from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
MODIFIED_MEMORIES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class LongContraRepeatWorker(MemoryBaseWorker):
def __init__(es_contra_repeat_similar_top_k, long_contra_repeat_threshold, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
super(LongContraRepeatWorker, self).__init__(*args, **kwargs)
self.es_contra_repeat_similar_top_k = es_contra_repeat_similar_top_k
self.merge_obs_model = merge_obs_model
self.merge_obs_max_token = merge_obs_max_token
self.merge_obs_temperature = merge_obs_temperature
self.merge_obs_top_k = merge_obs_top_k
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
# new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
# oday_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
for new_obs_node in new_obs_nodes:
text = new_obs_scheme.memory_node.content
hits = self.es_client.similar_search(text=text,
size=self.es_contra_repeat_similar_top_k,
exact_filters={
"memoryId": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"scene": self.scene.lower(),
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value],
})
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
has_match = False
for related_node in related_nodes:
if related_node.score_similar < self.long_contra_repeat_threshold:
continue
else:
has_match = True
all_obs_nodes.append(related_node)
if has_match:
all_obs_nodes.append(new_obs_node)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.memory_node.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.prompt_config.long_contra_repeat_system.format(num_obs=len(user_query_list)),
few_shot=self.prompt_config.long_contra_repeat_few_shot,
user_query=self.prompt_config.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.gene_client.call(messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -1,33 +0,0 @@
from typing import List, Dict
from constants.common_constants import NEW_INSIGHT_NODES, MODIFIED_MEMORIES, INSIGHT_NODES, NEW_OBS_NODES, \
NOT_REFLECTED_OBS_NODES, NEW, NOT_REFLECTED_MERGE_NODES
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
# 合并逻辑复杂务必check
all_node_dict: Dict[str, MemoryNode] = {}
if insight_nodes:
all_node_dict.update({n.id: n for n in insight_nodes if n.memory_node.content_modified})
if new_insight_nodes:
all_node_dict.update({n.memory_node.content: n for n in new_insight_nodes})
if new_obs_nodes:
# 设置为非新
for n in new_obs_nodes:
n.memory_node.metaData[NEW] = "0"
all_node_dict.update({n.memory_node.content: n for n in new_obs_nodes})
if not_reflected_merge_nodes and not_reflected_nodes:
# 进入reflect阶段
all_node_dict.update({n.id: n for n in not_reflected_nodes})
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))

View file

@ -1,151 +0,0 @@
from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class UpdateInsightWorker(MemoryBaseWorker):
def __init__(update_insight_threshold, update_insight_max_thread, update_insight_model, update_insight_max_token, update_insight_temperature, update_insight_top_k,*args, **kwargs):
super(UpdateInsightWorker, self).__init__(*args, **kwargs)
self.update_insight_threshold = update_insight_threshold
self.update_insight_max_thread = update_insight_max_thread
self.update_insight_model = update_insight_model
self.update_insight_max_token = update_insight_max_token
self.update_insight_temperature = update_insight_temperature
self.update_insight_top_k = update_insight_top_k
def filter_obs_nodes(self,
insight_node: MemoryNode,
new_obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
if not insight_key or not insight_value:
self.logger.warning(f"insight_key={insight_key} insight_value={insight_value} is empty!")
return insight_node, filtered_nodes, max_score
result = self.rerank_client.call(query=insight_key,
documents=[x.memory_node.content for x in new_obs_nodes])
if not result:
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
return insight_node, filtered_nodes, max_score
# 找到大于阈值的obs node
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_insight_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(f"insight_key={insight_key} insight_value={insight_value} "
f"score={score} keep_flag={keep_flag}")
if not filtered_nodes:
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
return insight_node, filtered_nodes, max_score
def update_insight(self,
insight_node: MemoryNode,
filtered_nodes: List[MemoryNode]) -> MemoryNode:
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
self.logger.info(f"update_insight insight_key={insight_key} insight_value={insight_value} "
f"doc.size={len(filtered_nodes)}")
# gen prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{scheme.memory_node.content}")
update_insight_message = self.prompt_to_msg(
system_prompt=self.prompt_config.update_insight_system,
few_shot=self.prompt_config.update_insight_few_shot,
user_query=self.prompt_config.update_insight_user_query.format(
user_query="\n".join(user_query_list),
insight_key=insight_key,
insight_key_value=insight_key + "" + insight_value))
self.logger.info(f"update_insight_message={update_insight_message}")
# call LLM
response_text: str = self.gene_client.call(messages=update_insight_message,
model_name=self.update_insight_model,
max_token=self.update_insight_max_token,
temperature=self.update_insight_temperature,
top_k=self.update_insight_top_k)
# return if empty
if not response_text:
self.add_run_info(f"update_insight insight_key={insight_key} call llm failed!")
return insight_node
profile_list = ResponseTextParser(response_text).parse_v1(f"update_profile {insight_key}")
if not profile_list:
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 1!")
return insight_node
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 2")
return insight_node
insight_value = profile_list[0]
if not insight_value or insight_value in ["", "重复"]:
self.logger.info(f"insight_value={insight_value}, skip.")
return insight_node
insight_scheme.memory_node.metaData[INSIGHT_VALUE] = insight_value
insight_scheme.memory_node.content_modified = True
return insight_node
def _run(self):
# 获取新的obs和insight
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop update sights!")
return
if not insight_nodes:
self.logger.info("insight_nodes is empty, stop update sights!")
return
# 提交打分任务
for node in insight_nodes:
self.submit_thread(self.filter_obs_nodes,
sleep_time=0.1,
insight_node=node,
new_obs_nodes=new_obs_nodes)
# 选择topN
result_list = []
for result in self.join_threads():
insight_node, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_insight_max_thread:
result_sorted = result_sorted[:update_insight_max_thread]
# 提交LLM update任务
for insight_node, filtered_nodes, _ in result_sorted:
self.submit_thread(self.update_insight,
sleep_time=1,
insight_node=insight_node,
filtered_nodes=filtered_nodes)
# 等待结果
for result in self.join_threads():
if result:
insight_node: MemoryNode = result
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
self.logger.info(f"after_update_insight insight_key={insight_key} insight_value={insight_value}")

View file

@ -1,210 +0,0 @@
from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from node.user_attribute import UserAttribute
from worker.memory_base_worker import MemoryBaseWorker
class UpdateProfileWorker(MemoryBaseWorker):
def __init__(update_profile_max_thread, update_profile_threshold, extra_user_attrs, update_profile_model, update_profile_max_token, update_profile_temperature, update_profile_top_k, *args, **kwargs):
super(UpdateProfileWorker,self).__init__(*args, **kwargs)
self.update_profile_max_thread = update_profile_max_thread
self.extra_user_attrs = extra_user_attrs
self.update_profile_threshold = update_profile_threshold
self.update_profile_model = update_profile_model
self.update_profile_max_token = update_profile_max_token
self.update_profile_temperature = update_profile_temperature
self.update_profile_top_k = update_profile_top_k
# @property
# def extra_user_attrs(self):
# return self.request.extra_user_attrs
def filter_obs_nodes(self,
user_attr: UserAttribute,
new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
result = self.rerank_client.call(query=user_attr.description,
documents=[x.memory_node.content for x in new_obs_nodes])
if not result:
self.add_run_info(f"update_user_attr={user_attr.memory_key} call rerank failed!")
return user_attr, filtered_nodes, max_score
# 找到大于阈值的obs node
filtered_nodes: List[MemoryNode] = []
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_profile_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(f"key={user_attr.memory_key} desc={user_attr.description} "
f"content={scheme.memory_node.content} score={score} keep_flag={keep_flag}")
if not filtered_nodes:
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
return user_attr, filtered_nodes, max_score
def update_user_attr(self, user_attr: UserAttribute, filtered_nodes: List[MemoryNode]) -> UserAttribute:
self.logger.info(f"update_user_attr memory_key={user_attr.memory_key} desc={user_attr.description} "
f"value={user_attr.value} doc.size={len(filtered_nodes)}")
# 根据不同的参数类型是否多值分别给出prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{scheme.memory_node.content}")
update_profile = f"{user_attr.memory_key}{user_attr.description}"
update_profile_value = update_profile + "" + "".join(user_attr.value)
if user_attr.is_unique == 1:
update_profile_message = self.prompt_to_msg(
system_prompt=self.prompt_config.update_unique_profile_system,
few_shot=self.prompt_config.update_unique_profile_few_shot,
user_query=self.prompt_config.update_unique_profile_user_query.format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value))
else:
update_profile_message = self.prompt_to_msg(
system_prompt=self.prompt_config.update_plural_profile_system,
few_shot=self.prompt_config.update_plural_profile_few_shot,
user_query=self.prompt_config.update_plural_profile_user_query.format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value))
self.logger.info(f"update_profile_message={update_profile_message}")
# call LLM
response_text: str = self.gene_client.call(messages=update_profile_message,
model_name=self.update_profile_model,
max_token=self.update_profile_max_token,
temperature=self.update_profile_temperature,
top_k=self.update_profile_top_k)
# return if empty
if not response_text:
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} call llm failed!")
return user_attr
profile_list = ResponseTextParser(response_text).parse_v1(f"update_attr {user_attr.memory_key}")
if not profile_list:
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 1!")
return user_attr
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 2")
return user_attr
profile = profile_list[0]
if not profile or profile in ["", "重复"]:
self.logger.info(f"profile={profile}, skip.")
return user_attr
# check 英文中午逗号
if user_attr.is_unique == 1:
user_attr.value = [profile.strip()]
else:
attr_value_list = profile.replace("", ",").split(",")
user_attr.value = [x.strip() for x in sorted(list(set(user_attr.value + attr_value_list)))]
return user_attr
def add_extra_user_attrs(self):
# 解析为空返回
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
if not extra_user_attr_list:
return
for user_attr_info in extra_user_attr_list:
user_attr_split = user_attr_info.split(":")
# 格式不对返回
if len(user_attr_split) < 1:
continue
user_attr_key = user_attr_split[0]
user_attr_desc = ""
if len(user_attr_split) >= 2:
user_attr_desc = user_attr_split[1]
user_attr_unique = 0
if len(user_attr_split) >= 3:
user_attr_unique = int(user_attr_split[2])
# 已经包含返回
if user_attr_key in self.user_profile_dict:
user_attr = self.user_profile_dict[user_attr_key]
# description为空补充description
if not user_attr.description:
user_attr.description = user_attr_desc
continue
# 增加新属性
new_attr = UserAttribute(memory_id=self.config.memory_id,
scene=self.scene,
memory_key=user_attr_key,
is_unique=int(user_attr_unique),
is_mutable=1,
memory_type=MemoryTypeEnum.PROFILE,
description=user_attr_desc,
status=1)
self.user_profile_dict[user_attr_key] = new_attr
def _run(self):
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop user profile!")
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
return
# 增加环境变量配置的属性
if self.extra_user_attrs:
self.add_extra_user_attrs()
new_user_profile: List[UserAttribute] = []
self.set_context(NEW_USER_PROFILE, new_user_profile)
for user_attr_key, user_attr in self.user_profile_dict.items():
# 不可修改直接跳过
if user_attr.is_mutable != 1:
new_user_profile.append(user_attr)
self.logger.info(f"{user_attr_key} is not mutable! continue")
continue
self.submit_thread(self.filter_obs_nodes,
sleep_time=0.1,
user_attr=user_attr,
new_obs_nodes=new_obs_nodes)
# 选择topN
result_list = []
for result in self.join_threads():
user_attr, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_profile_max_thread:
result_sorted = result_sorted[:self.update_profile_max_thread]
# 提交LLM update任务
for user_attr, filtered_nodes, _ in result_sorted:
self.submit_thread(self.update_user_attr,
sleep_time=1,
user_attr=user_attr,
filtered_nodes=filtered_nodes)
# collect result & save
for result in self.join_threads():
if result:
user_attribute: UserAttribute = result
self.logger.info(f"after_update_profile memory_key={user_attribute.memory_key} "
f"desc={user_attribute.description} value={user_attribute.value}")
new_user_profile.append(user_attribute)

View file

@ -1,98 +0,0 @@
from typing import List
from common.response_text_parser import ResponseTextParser
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
MODIFIED_MEMORIES
from enumeration.memory_status_enum import MemoryNodeStatus
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class ContraRepeatWorker(MemoryBaseWorker):
def __init__(self, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
super(ContraRepeatWorker, self).__init__(*args, **kwargs)
self.merge_obs_model = merge_obs_model
self.merge_obs_max_token = merge_obs_max_token
self.merge_obs_temperature = merge_obs_temperature
self.merge_obs_top_k = merge_obs_top_k
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
if new_obs_nodes:
all_obs_nodes.extend(new_obs_nodes)
if new_obs_with_time_nodes:
all_obs_nodes.extend(new_obs_with_time_nodes)
if today_obs_nodes:
all_obs_nodes.extend(today_obs_nodes)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.memory_node.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)),
few_shot=self.prompt_config.contra_repeat_few_shot,
user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.gene_client.call(messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -1,134 +0,0 @@
from datetime import datetime
from typing import List
from common.response_text_parser import ResponseTextParser
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict, extract_date_parts
from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME, KEY_WORD, DATATIME_WORD_LIST, \
NEW_OBS_WITH_TIME_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
class GetObservationWithTimeWorker(MemoryBaseWorker):
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
super(GetObservationWithTimeWorker, self).__init__(*args, **kwargs)
self.summary_messages_model = summary_messages_model
self.summary_messages_max_token = summary_messages_max_token
self.summary_messages_temperature = summary_messages_temperature
self.summary_messages_top_k = summary_messages_top_k
def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: time_infer, # 推断的时间
KEY_WORD: keywords, # 关键词
}
# 事件时间
meta_data.update({f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()})
# 对话时间
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
return MemoryNode.init_from_attrs(content=obs_content,
memoryId=self.memory_id,
timeCreated=message.time_created,
scene=self.scene,
memoryType=MemoryTypeEnum.OBSERVATION.value,
content_modified=True, # 新增的obs需要置为true
metaData=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
tenantId=self.tenant_id)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if match:
dt = time_to_formatted_str(time=msg.time_created,
date_format="",
string_format="{year}{month}{day}{weekday}{hour}")
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(f"get obs with time user_query_list={user_query_list} is empty")
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)),
few_shot=self.prompt_config.get_observation_with_time_few_shot,
user_query=self.prompt_config.get_observation_with_time_user_query.format(
user_query="\n".join(user_query_list)))
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.gene_client.call(messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
if time_infer == "":
time_infer = ""
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
continue
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
obs_content=obs_content,
time_infer=time_infer,
keywords=keywords))
# save context
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)

View file

@ -1,122 +0,0 @@
from datetime import datetime
from typing import List
from common.response_text_parser import ResponseTextParser
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER, NEW, MSG_TIME, KEY_WORD, \
DATATIME_WORD_LIST
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from node.message import Message
from worker.memory_base_worker import MemoryBaseWorker
class GetObservationWorker(MemoryBaseWorker):
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
super(GetObservationWorker, self).__init__(*args,**kwargs)
self.summary_messages_model = summary_messages_model
self.summary_messages_max_token = summary_messages_max_token
self.summary_messages_temperature = summary_messages_temperature
self.summary_messages_top_k = summary_messages_top_k
def add_observation(self, message: Message, obs_content: str, keywords: str):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: "", # 推断的时间
KEY_WORD: keywords, # 关键词
}
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
return MemoryNode.init_from_attrs(content=obs_content,
memoryId=self.memory_id,
timeCreated=message.time_created,
scene=self.scene,
memoryType=MemoryTypeEnum.OBSERVATION.value,
content_modified=True, # 新增的obs需要置为true
metaData=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
tenantId=self.tenant_id)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if not match:
user_query_list.append(f"{i} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.prompt_config.get_observation_system.format(num_obs=len(user_query_list)),
few_shot=self.prompt_config.get_observation_few_shot,
user_query=self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list)))
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.gene_client.call(messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
continue
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
obs_content=obs_content,
keywords=keywords))
# save context
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -1,64 +0,0 @@
from common.response_text_parser import ResponseTextParser
from enumeration.message_role_enum import MessageRoleEnum
from worker.memory_base_worker import MemoryBaseWorker
class InfoFilterWorker(MemoryBaseWorker):
def __init__(self, info_filter_msg_max_size, info_filter_model, info_filter_max_token, info_filter_temperature, info_filter_top_k, *args, **kwargs):
super(InfoFilterWorker,self).__init__(*args,**kwargs)
self.info_filter_msg_max_size
self.info_filter_model = info_filter_model
self.info_filter_max_token = info_filter_max_token
self.info_filter_temperature = info_filter_temperature
self.info_filter_top_k = info_filter_top_k
def _run(self):
# filter user msg
info_messages = []
for msg in self.messages:
if msg.role != MessageRoleEnum.USER.value:
continue
if len(msg.content) >= self.info_filter_msg_max_size:
continue
info_messages.append(msg)
# gene prompt
user_query = "\n".join([f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)])
info_filter_message = self.prompt_to_msg(
system_prompt=self.prompt_config.info_filter_system.format(batch_size=len(info_messages)),
few_shot=self.prompt_config.info_filter_few_shot,
user_query=self.prompt_config.info_filter_user_query.format(user_query=user_query))
self.logger.info(f"info_filter_message={info_filter_message}")
# call llm
response_text = self.gene_client.call(messages=info_filter_message,
model_name=self.info_filter_model,
max_token=self.info_filter_max_token,
temperature=self.info_filter_temperature,
top_k=self.info_filter_top_k)
# return if empty
if not response_text:
self.add_run_info("info score call llm failed!", continue_run=False)
return
# parse text
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
if len(info_score_list) != len(info_messages):
self.add_run_info(f"info_score_size != info_messages_size, "
f"{len(info_score_list)} vs {len(info_messages)}", continue_run=False)
return
# 过滤value=0的messages
filtered_messages = []
for msg, info_score in zip(info_messages, info_score_list):
if not info_score:
continue
score = info_score[0]
# if score in ("1", "2",):
if score in ("2",):
msg.info_score = score
filtered_messages.append(msg)
# 后续不会关注为0的msg直接丢弃
self.messages = filtered_messages

View file

@ -1,198 +0,0 @@
import re
from datetime import datetime
from importlib import import_module
from typing import Dict, List
from constants.common_constants import WEEKDAYS
from enumeration.message_role_enum import MessageRoleEnum
def under_line_to_hump(underline_str):
sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str)
return sub[0:1].upper() + sub[1:]
def parse_response_text_v1(response_text: str) -> dict:
"""
parse text like:
<1> <AAA>
<2> <BBB> ddd
<4> <CCC> dddd<555>
result = {1: "AAA", 2: "BBB", 4: "CCC"}
"""
result_dict: Dict[int, str] = {}
# 确保第一个数字后面是string
matches = re.findall(r'<(\d+)>\s*<([^>]+)>', response_text.strip())
# matches 为空返回
for key, value in matches:
result_dict[int(key)] = value
return result_dict
def parse_response_text_v2(response_text: str) -> Dict[int, List[str]]:
"""
parse text like:
XXX
<1> <AAA> <222>
<2> <BBB>
<4,5> <CCC>
result = {1: ["AAA", "222"], 2: "BBB", 4: "CCC"}
"""
result_dict: Dict[int, List[str]] = {}
for line in response_text.strip().split("\n"):
if "> <" not in line:
continue
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
idx: str = ll[0]
values: List[str] = ll[1:]
if idx.isdigit():
idx_int = int(idx)
else:
idx_split = idx.split(",")
if len(idx_split) == 0:
continue
idx = idx_split[0]
if idx.isdigit():
idx_int = int(idx)
else:
continue
if values:
result_dict[idx_int] = values
return result_dict
def parse_response_text_v3(response_text: str) -> List[List[str]]:
"""
parse text like:
XXX
<1> <AAA>
<2c> <BBB>
<41> <CCC> <BBB>
result = [["1", "AAA"], ["2c", "BBB"], ["41", "CCC", "BBB"]]
"""
result_list: List[List[str]] = []
for line in response_text.strip().split("\n"):
if "> <" not in line:
continue
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
result_list.append(ll)
return result_list
def get_datetime_info_dict(parse_dt: datetime):
return {
"year": parse_dt.year,
"month": parse_dt.month,
"day": parse_dt.day,
"hour": parse_dt.hour,
"minute": parse_dt.minute,
"second": parse_dt.second,
"week": parse_dt.isocalendar().week,
"weekday": WEEKDAYS[parse_dt.isocalendar().weekday - 1],
}
def extract_date_parts(input_string: str):
# Extending our pattern to handle "每" (every) as a possible value.
patterns = {
'year': r'(\d+|每)年',
'month': r'(\d+|每)月',
'day': r'(\d+|每)日',
'weekday': r'周([一二三四五六日])?',
'hour': r'(\d+)点'
}
weekday_dict = {"": 1, "": 2, "": 3, "": 4, "": 5, "": 6, "": 7}
extracted_data = {}
# Search for patterns in the input string and populate the dictionary
for key, pattern in patterns.items():
match = re.search(pattern, input_string)
if match: # If there is a match, include it in the output dictionary
if match.group(1) == "":
extracted_data[key] = -1
elif match.group(1) in weekday_dict.keys():
extracted_data[key] = weekday_dict[match.group(1)]
else:
extracted_data[key] = int(match.group(1))
return extracted_data
def time_to_formatted_str(time: datetime | str | int | float = None,
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
string_format: str = "") -> str:
if isinstance(time, str | int | float):
if isinstance(time, str):
time = float(time)
current_dt = datetime.fromtimestamp(time)
elif isinstance(time, datetime):
current_dt = time
else:
current_dt = datetime.now()
return_str = ""
if date_format:
return_str = current_dt.strftime(date_format)
elif string_format:
return_str = string_format.format(**get_datetime_info_dict(current_dt))
return return_str
def init_instance_by_config(config: dict|object, default_module_path: str = None, try_kwargs: dict = {}, accept_types: type = None):
if isinstance(config, accept_types):
return config
import_module(config.pop("path", default_module_path))
clazz = getattr(module, config.pop("name"))
try:
return clazz(**config, **try_kwargs)
except:
return clazz(**config)
def init_instance_by_config_v2(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs):
clazz_path = config.pop("clazz")
if not clazz_path:
raise RuntimeError("empty clazz_path!")
clazz_name_split = clazz_path.split(".")
clazz_name: str = clazz_name_split[-1]
if suffix_name and not clazz_name.endswith(suffix_name):
clazz_name = f"{clazz_name}_{suffix_name}"
# 构造path
clazz_paths = []
if default_clazz_path:
clazz_paths.append(default_clazz_path)
clazz_paths.extend(clazz_name_split[:-1])
clazz_paths.append(clazz_name)
module = import_module(".".join(clazz_paths))
cls_name = under_line_to_hump(clazz_name)
return getattr(module, cls_name)(**config, **kwargs)
def complete_config_name(config_name: str, suffix: str = ".json"):
if not config_name.endswith(suffix):
config_name += suffix
return config_name
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
return [
{
"role": MessageRoleEnum.SYSTEM.value,
"content": system_prompt.strip(),
},
{
"role": MessageRoleEnum.USER.value,
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
},
]

View file

@ -1,33 +0,0 @@
from typing import Dict, List
from pydantic import Field, BaseModel
class UserAttribute(BaseModel):
"""
用户画像的一条属性和数据库保持一致只会选择status为1的属性透传过来
status会透传过来
如果code为空则为新增否则是更新
确保请求是10条返回是原始10条+加上新增的条数如果可以新增只会对正确的请求操作数据库
"""
id: str = Field("", description="唯一主键")
memory_id: str = Field("", description="memory id")
# 从key改成memory_key
memory_key: str = Field("", description="memory key")
value: List[str] = Field([], description="value")
is_unique: int = Field(1, description="属性是否唯一if 1 value只有一个if 0, value 可以很多个")
is_mutable: int = Field(1, description="是否可变if 1value可变if 1不可变用户定义")
memory_type: str = Field("", description="profile, profile_customized")
description: str = Field("", description="memory id")
status: int = Field(1,
description="0为删除1为active状态算法不感知只为了保存用户删除的画像给算法传status为valid的用户画像")
ext_info: Dict[str, str] = Field({}, description="占位符字典")

View file

@ -1,102 +0,0 @@
import json
from typing import List, Dict
from enumeration.memory_status_enum import MemoryNodeStatus
from scheme.memory_node import MemoryNode
from node.user_attribute import UserAttribute
class UserProfileHandler(object):
@classmethod
def format_content(cls, key: str, description: str, value: str | List[str] = None):
if not key.startswith("用户"):
key = f"用户的{key}"
if not description.startswith("用户"):
description = f"用户{description}"
content = f"{key}{description}"
if value:
if isinstance(value, list):
value = "".join(value)
content = f"{content}{value}"
return content
"""
提供UserAttribute MemoryNode 的相互转化
"""
@classmethod
def to_nodes(cls,
user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None,
split_value: bool = False) -> List[MemoryNode]:
user_profile_dict: Dict[str, UserAttribute] = {}
if user_profile:
if isinstance(user_profile, list):
for user_attr in user_profile:
user_profile_dict[user_attr.memory_key] = user_attr
elif isinstance(user_profile, dict):
user_profile_dict.update(user_profile)
user_profile_nodes: List[MemoryNode] = []
for _, user_attr in user_profile_dict.items():
# 获取id
_id = user_attr.code
if not _id:
_id = f"{user_attr.memory_id}_{user_attr.scene}_profile_{user_attr.memory_key}"
attr_node = MemoryNode.init_from_attrs(id=_id,
code=_id,
content="",
memoryId=user_attr.memory_id,
scene=user_attr.scene,
memoryType=user_attr.memory_type,
content_modified=True,
metaData={
"memory_key": user_attr.memory_key,
"value": json.dumps(user_attr.value, ensure_ascii=False),
"is_unique": str(user_attr.is_unique),
"is_mutable": str(user_attr.is_mutable),
"description": user_attr.description,
"status": MemoryNodeStatus.ACTIVE.value,
"ext_info": json.dumps(user_attr.ext_info,
ensure_ascii=False),
},
status=MemoryNodeStatus.ACTIVE.value)
if split_value:
for value in user_attr.value:
content = cls.format_content(user_attr.memory_key, user_attr.description, value)
attr_node_copy = attr_node.copy(deep=True)
attr_node_copy.memory_node.content = content
user_profile_nodes.append(attr_node_copy)
else:
content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value)
attr_scheme.memory_node.content = content
user_profile_nodes.append(attr_node)
return user_profile_nodes
@classmethod
def to_user_attr(cls, user_profile_nodes: List[MemoryNode]) -> Dict[str, UserAttribute]:
user_profile_dict: Dict[str, UserAttribute] = {}
for node in user_profile_nodes:
user_attr = UserAttribute(
code=node.id,
memory_id=scheme.memory_node.memoryId,
scene=scheme.memory_node.scene,
memory_key=scheme.memory_node.metaData["memory_key"],
value=json.loads(scheme.memory_node.metaData["value"]),
is_unique=int(scheme.memory_node.metaData["is_unique"]),
is_mutable=int(scheme.memory_node.metaData["is_mutable"]),
memory_type=scheme.memory_node.memoryType,
description=scheme.memory_node.metaData["description"],
status=1 if scheme.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
ext_info=json.loads(scheme.memory_node.metaData["ext_info"]),
)
user_profile_dict[user_attr.memory_key] = user_attr
return user_profile_dict

View file

View file

@ -1,73 +0,0 @@
from typing import Any, Dict
from ..utils.logger import Logger
from ..utils.timer import Timer
class BaseWorker(object):
def __init__(self, raise_exception: bool = True, **kwargs):
super(BaseWorker, self).__init__(**kwargs)
# 异常是否继续执行
self.raise_exception: bool = raise_exception
# True 为正常运行False会结束整个pipeline
self.continue_run: bool = True
# 短name
self._name_simple: str = ""
# 是否多线程环境
self.is_multi_thread: bool = False
# pipeline 上下文
self.context_dict: Dict[str, Any] | None = None
self.context_lock = None
# 日志
self.logger: Logger = Logger.get_logger()
# worker 参数
self.kwargs: dict = kwargs
def _run(self):
raise NotImplementedError
def run(self):
self.logger.info(f"----- Begin {self.name_simple} -----")
with Timer(self.name_simple, log_time=False) as t:
if self.raise_exception:
self._run()
else:
try:
self._run()
except Exception as e:
self.logger.exception(f"run {self.name_simple} failed! args={e.args}")
self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----")
def set_context_dict(self, context_dict: dict, context_lock=None):
self.context_dict = context_dict
if context_lock is not None:
self.context_lock = context_lock
self.is_multi_thread = True
def get_context(self, key: str, default=None):
return self.context_dict.get(key, default)
def set_context(self, key: str, value: Any):
if self.is_multi_thread:
# add lock to multi thread
with self.context_lock:
self.context_dict[key] = value
else:
self.context_dict[key] = value
def __getattr__(self, key):
return self.kwargs[key]
@property
def name_simple(self) -> str:
if not self._name_simple:
self._name_simple = self.__class__.__name__.replace("Worker", "")
return self._name_simple

View file

@ -1,6 +0,0 @@
from memory_base_worker import MemoryBaseWorker
class DummyWorker(MemoryBaseWorker):
def _run(self):
pass

View file

@ -1,22 +0,0 @@
from typing import List
from constants.common_constants import INSIGHT_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
from cli import GLOBAL_CONTEXT
class EsInsightWorker(MemoryBaseWorker):
def _run(self):
insight_nodes = self.vector_store.retrieve(
size=self.kwargs.es_insight_top_k,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": MemoryTypeEnum.INSIGHT.value,
},
)
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
self.set_context(INSIGHT_NODES, insight_nodes)

View file

@ -1,22 +0,0 @@
from typing import List
from constants.common_constants import NEW, NEW_OBS_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNewObsWorker(MemoryBaseWorker):
def _run(self):
new_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": MemoryTypeEnum.OBSERVATION.value,
f"meta_data.{NEW}": "1",
},
)
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -1,29 +0,0 @@
from typing import List
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsNotReflectedWorker(MemoryBaseWorker):
def _run(self):
not_reflected_obs_nodes = self.vector_store.retrieve(
size=self.kwargs.es_new_obs_top_k,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
f"meta_data.{REFLECTED}": "0",
},
)
self.logger.info(
f"retrieve_not_reflected_obs.size={len(not_reflected_obs_nodes)}"
)
self.set_context(NOT_REFLECTED_OBS_NODES, not_reflected_obs_nodes)

View file

@ -1,37 +0,0 @@
from typing import List
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_recall_type import MemoryRecallType
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsSimilarWorker(MemoryBaseWorker):
def __init__(self, es_similar_top_k, *args, **kwargs):
super(EsSimilarWorker, self).__init__(*args, **kwargs)
self.es_similar_top_k = es_similar_top_k
def _run(self):
query = self.messages[-1].content
similar_obs_nodes = self.vector_store.retrieve(
text=query,
size=self.es_similar_top_k,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.INSIGHT.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
for node in similar_obs_nodes:
node.meta_data[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
for node in similar_obs_nodes:
self.logger.info(f"node={node.content} score_similar={node.score_similar}")
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)

View file

@ -1,32 +0,0 @@
from typing import List
from utils.tool_functions import time_to_formatted_str
from constants.common_constants import TODAY_OBS_NODES, DT
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class EsTodayObsWorker(MemoryBaseWorker):
def __init__(self, es_today_obs_top_k, *args, **kwargs):
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
self.es_today_obs_top_k = es_today_obs_top_k
def _run(self):
if not self.messages:
self.logger.warning("messages is empty!")
return
msg_time_created = self.messages[-1].time_created
today_obs_nodes = self.vector_store.retrieve(
size=self.es_today_obs_top_k,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": MemoryTypeEnum.OBSERVATION.value,
f"meta_Data.{DT}": time_to_formatted_str(msg_time_created),
},
)
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
self.set_context(TODAY_OBS_NODES, today_obs_nodes)

View file

@ -1,25 +0,0 @@
from typing import List, Dict
from constants import common_constants
from enumeration.memory_status_enum import MemoryNodeStatus
from enumeration.memory_type_enum import MemoryTypeEnum
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class LoadProfileWorker(MemoryBaseWorker):
def _run(self):
user_profile_node = self.vector_store(
size=10000,
filter_dict={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.PROFILE.value,
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
],
},
)
self.set_context(common_constants.USER_PROFILE, user_profile_node)
self.logger.info(f"retrieve_user_profile.size={len(user_profile_node)}")

View file

@ -1,74 +0,0 @@
import re
from utils.tool_functions import time_to_formatted_str
from constants.common_constants import (
DATATIME_WORD_LIST,
DATATIME_KEY_MAP,
EXTRACT_TIME_DICT,
)
from worker.memory_base_worker import MemoryBaseWorker
class ExtractTimeWorker(MemoryBaseWorker):
# TODO add en version
@staticmethod
def get_parse_time_prompt(query: str, query_time_str: str):
return f"""
任务指令从语句与语句发生的时间推断并提取语句内容中指向的时间段回答尽可能完整的时间段
语句{query}
时间{query_time_str}
回答
""".strip()
def _run(self):
# save to context
extract_time_dict = {}
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)
# get query & time_created_dt
query = self.messages[-1].content
time_created = self.messages[-1].time_created
# find datetime keyword
contain_datetime = False
for datetime_word in DATATIME_WORD_LIST:
if datetime_word in query:
contain_datetime = True
break
if not contain_datetime:
self.logger.info(f"contain_datetime={contain_datetime}")
return
# prepare prompt
# TODO add en version
time_format = "{year}{month}{day}日,{year}年第{week}周,{weekday}{hour}{minute}{second}秒。"
query_time_str = time_to_formatted_str(
time=time_created, date_format="", string_format=time_format
)
extract_time_prompt = self.get_parse_time_prompt(
query=query, query_time_str=query_time_str
)
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
# call sft model
response_text = self.generation_model.call(
prompt=extract_time_prompt,
model_name=self.parse_time_model,
max_token=self.parse_time_max_token,
temperature=self.parse_time_temperature,
top_k=self.parse_time_top_k,
)
# if empty, return
if not response_text:
return
# re-match time info to dict
pattern = r"-\s*(\S+)(\d+)"
matches = re.findall(pattern, response_text)
for key, value in matches:
if key in DATATIME_KEY_MAP.keys():
extract_time_dict[DATATIME_KEY_MAP[key]] = value
self.logger.info(f"response_text={response_text} filters={extract_time_dict}")

View file

@ -1,120 +0,0 @@
from typing import Dict, List
from constants.common_constants import (
RELATED_MEMORIES,
EXTRACT_TIME_DICT,
ALL_ONLINE_NODES,
TIME_MATCHED,
)
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class FuseRerankWorker(MemoryBaseWorker):
@staticmethod
def format_time_infer(
time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]
):
if time_infer:
return time_infer
time_infer = ""
if "year" in extract_time_dict:
value = meta_data.get("msg_year")
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += "每年"
if "month" in extract_time_dict:
value = meta_data.get("msg_month")
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += "每月"
if "day" in extract_time_dict:
value = meta_data.get("msg_day")
if value:
time_infer += f"{value}"
elif value == "-1":
time_infer += "每日"
if "weekday" in extract_time_dict:
value = meta_data.get("msg_weekday")
if value:
time_infer += value
return time_infer
def _run(self):
# 解析时间meta信息
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
all_online_nodes: List[MemoryNode] = self.get_context(ALL_ONLINE_NODES)
if not all_online_nodes:
self.add_run_info("all_online_nodes is empty, stop")
return
filtered_nodes = []
for node in all_online_nodes:
if node.score_rank < self.fuse_score_threshold:
continue
# 根据类型给ratio
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
# 时间系数,完全匹配才行
fuse_time_ratio: float = 1.0
match_event_flag = False
match_msg_flag = False
if extract_time_dict:
match_event_flag = True
for k, v in extract_time_dict.items():
event_value = node.meta_data.get(f"event_{k}", "")
if event_value in ["-1", v]:
continue
else:
match_event_flag = False
break
match_msg_flag = True
for k, v in extract_time_dict.items():
msg_value = node.meta_data.get(f"msg_{k}", "")
if msg_value == v:
continue
else:
match_msg_flag = False
break
if match_event_flag or match_msg_flag:
fuse_time_ratio = self.fuse_time_ratio
node.meta_data[TIME_MATCHED] = "1"
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
self.logger.info(
f"content={node.content} f_event={int(match_event_flag)} "
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}"
)
filtered_nodes.append(node)
# get output & save context
filtered_nodes = sorted(
filtered_nodes, key=lambda x: x.score_rerank, reverse=True
)
filtered_nodes = filtered_nodes[: self.output_max_count]
related_memories: List[str] = []
for node in filtered_nodes:
content = node.content
# 如果命中时间逻辑
if node.meta_data.get(TIME_MATCHED, "") == "1":
time_infer = self.format_time_infer(
time_infer="",
extract_time_dict=extract_time_dict,
meta_data=node.meta_data,
)
content = f"{time_infer}: {content}"
related_memories.append(content)
self.set_context(RELATED_MEMORIES, related_memories)

View file

@ -1,35 +0,0 @@
from typing import List
from utils.user_profile_handler import UserProfileHandler
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class MemoryStoreWorker(MemoryBaseWorker):
def _run(self):
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(
MODIFIED_MEMORIES
)
if modified_memories:
if isinstance(modified_memories[0], MemoryNode):
modified_memories = [n.memory_node for n in modified_memories]
for n in modified_memories:
if not n.id:
n.id = f"{n.memory_id}_content_{n.content}"
n.code = n.id
# TODO add batch insert
n.meta_data.pop(CONTENT_MODIFIED)
self.vector_store.insert(n)
else:
self.logger.warning("modified_memories is empty!")
new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE)
if new_user_profile:
for n in new_user_profile:
n.meta_data.pop(CONTENT_MODIFIED)
self.vector_store.insert(n)
else:
self.logger.warning("new_user_profile is empty!")

View file

@ -1,62 +0,0 @@
from typing import List, Dict
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
QUERY_KEYWORDS
from enumeration.memory_recall_enum import MemoryRecallType
from scheme.memory_node import MemoryNode
from worker.memory_base_worker import MemoryBaseWorker
class SemanticRankWorker(MemoryBaseWorker):
def user_profile_to_nodes(self) -> List[MemoryNode]:
user_profile_nodes: List[MemoryNode] = self.user_profile_dict
for node in user_profile_nodes:
# 从画像侧召回
node.meta_data[RECALL_TYPE] = MemoryRecallType.PROFILE
self.logger.info(f"user profile node={node.content}")
return user_profile_nodes
def _run(self):
all_node_dict: Dict[str, MemoryNode] = {}
# 优先级: similar_obs_nodes < profile_nodes
similar_obs_nodes: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES)
if similar_obs_nodes:
for node in similar_obs_nodes:
all_node_dict[node.content] = node
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
if profile_nodes:
for node in profile_nodes:
all_node_dict[node.content] = node
if not all_node_dict:
self.add_run_info("all_node_dict is empty!", continue_run=False)
return
# call recall model
query_keywords = self.get_context(QUERY_KEYWORDS)
# TODO 根据效果更改
# query: str = "用户:" + self.messages[-1].content
query: str = self.messages[-1].content
if query_keywords:
query_keyword_join = "".join(query_keywords)
query = f"{query} 用户的{query_keyword_join}"
documents = list(all_node_dict.keys())
result = self.rank_model.call(query=query, documents=documents)
if not result:
self.add_run_info("semantic call recall model failed!")
return
# set score
for index, score in result.rank_scores.items():
content = documents[index]
node = all_node_dict[content]
node.score_rank = score
self.logger.info(f"query={query} content={node.content} score_rank={node.score_rank}")
# save context
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())
self.set_context(ALL_ONLINE_NODES, all_online_nodes)

View file

@ -1,166 +0,0 @@
from datetime import datetime
from typing import List
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
from ...constants.common_constants import (
NEW_INSIGHT_NODES,
DT,
NOT_REFLECTED_MERGE_NODES,
NEW_INSIGHT_KEYS,
INSIGHT_KEY,
INSIGHT_VALUE,
REFLECTED,
CONTENT_MODIFIED
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_insight_prompt import (
GET_INSIGHT_FEW_SHOT_PROMPT,
GET_INSIGHT_SYSTEM_PROMPT,
GET_INSIGHT_USER_QUERY_PROMPT
)
class GetInsightWorker(MemoryBaseWorker):
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
created_dt = datetime.now()
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
DT: dt,
INSIGHT_KEY: insight_key,
INSIGHT_VALUE: insight_value,
CONTENT_MODIFIED: True, # 新增的insight需要置为true
}
meta_data.update(
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
content = f"用户的{insight_key}{insight_value}"
return MemoryNode(
content=content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.INSIGHT.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def reflect_new_insight_key(
self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode]
) -> MemoryNode | None:
# 检索历史memory
hits = self.vector_store.similar_search(
text=insight_key,
size=self.es_insight_similar_top_k,
exact_filters={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
# 转化成 MemoryNodeWrap 合并新增nodes
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
related_nodes.extend(not_reflected_merge_nodes)
# content去重
related_node_dict = {n.memory_node.content: n for n in related_nodes}
related_nodes = sorted(
list(related_node_dict.values()), key=lambda x: x.memory_node.id
)
documents = [n.memory_node.content for n in related_nodes]
# 重排所有记忆
result = self.rank_model.call(query=insight_key, documents=documents)
if not result:
self.add_run_info(
f"reflect insight_key={insight_key} call rerank client failed!"
)
return
# 根据打分过滤
for rank_node in result:
index = rank_node["index"]
score = rank_node["relevance_score"]
related_nodes[index].score_rank = score
related_nodes_sorted = sorted(
related_nodes, key=lambda x: x.score_rank, reverse=True
)[: self.insight_obs_max_cnt]
# 生成prompt
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
get_insight_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_INSIGHT_SYSTEM_PROMPT),
few_shot=self.get_prompt(GET_INSIGHT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_INSIGHT_USER_QUERY_PROMPT).format(
insight_key=insight_key, user_query="\n".join(user_query_list)
),
)
self.logger.info(f"get_insight_message={get_insight_message}")
# call LLM, 提取insight
response_text = self.generation_model.call(
messages=get_insight_message,
model_name=self.get_insight_model,
max_token=self.get_insight_max_token,
temperature=self.get_insight_temperature,
top_k=self.get_insight_top_k,
)
# return if empty
if not response_text:
self.add_run_info("reflect_upon_user_attr call llm failed!")
return
response_text = response_text.strip()
if response_text in [""]:
return
return self.new_insight_node(
insight_key=insight_key, insight_value=response_text
)
def _run(self):
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
if not new_insight_keys:
self.add_run_info("new_insight_keys is empty! stop insight.")
return
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_MERGE_NODES
)
if not not_reflected_merge_nodes:
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
return
# submit insight task
for insight_key in new_insight_keys:
self.submit_thread(
self.reflect_new_insight_key,
sleep_time=1,
insight_key=insight_key,
not_reflected_merge_nodes=not_reflected_merge_nodes,
)
# save output
new_insight_nodes: List[MemoryNode] = []
for result in self.join_threads():
if result:
new_insight_nodes.append(result)
assert isinstance(result, MemoryNode)
insight_key = result.meta_data.get(INSIGHT_KEY, "")
insight_value = result.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"after_get_insight insight_key={insight_key} insight_value={insight_value}"
)
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
# set REFLECTED
for node in not_reflected_merge_nodes:
node.meta_data[REFLECTED] = "1"

View file

@ -1,99 +0,0 @@
from typing import List
from ...utilsresponse_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
REFLECTED,
INSIGHT_NODES,
INSIGHT_KEY,
NEW_INSIGHT_KEYS,
NOT_REFLECTED_MERGE_NODES,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_reflection_prompt import (
GET_REFLECTION_FEW_SHOT_PROMPT,
GET_REFLECTION_SYSTEM_PROMPT,
GET_REFLECTION_USER_QUERY_PROMPT
)
class GetReflectionWorker(MemoryBaseWorker):
def _run(self):
# 过滤得到 not_reflected_merge_nodes
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_OBS_NODES
)
not_reflected_merge_nodes: List[MemoryNode] = []
if new_obs_nodes:
not_reflected_merge_nodes.extend(new_obs_nodes)
if not_reflected_nodes:
not_reflected_merge_nodes.extend(not_reflected_nodes)
not_reflected_merge_nodes = [
node
for node in not_reflected_merge_nodes
if node.meta_data.get(REFLECTED, "") == "0"
]
# count
not_reflected_count = len(not_reflected_merge_nodes)
if not_reflected_count <= self.reflect_obs_cnt_threshold:
self.logger.info(
f"not_reflected_count={not_reflected_count} is not enough, stop reflect."
)
return
# save context
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
# get profile_keys
exist_keys: List[str] = []
profile_keys: List[str] = list(self.user_profile_dict.keys())
exist_keys.extend(profile_keys)
self.logger.info(f"profile_keys={profile_keys}")
# get insight_keys
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if insight_nodes:
insight_keys = [
n.meta_data.get(INSIGHT_KEY) for n in insight_nodes
]
insight_keys = [x.strip() for x in insight_keys if x]
exist_keys.extend(insight_keys)
self.logger.info(f"insight_keys={insight_keys}")
# gen reflect prompt
user_query_list = [n.content for n in not_reflected_merge_nodes]
reflect_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_REFLECTION_SYSTEM_PROMPT).format(
num_questions=self.reflect_num_questions
),
few_shot=self.get_prompt(GET_REFLECTION_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_REFLECTION_USER_QUERY_PROMPT).format(
exist_keys="".join(exist_keys), user_query="\n".join(user_query_list)
),
)
self.logger.info(f"reflect_message={reflect_message}")
# # call LLM
response_text = self.generation_model.call(
messages=reflect_message,
model_name=self.reflect_obs_model,
max_token=self.reflect_obs_max_token,
temperature=self.reflect_obs_temperature,
top_k=self.reflect_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("reflect_obs_questions call llm failed!")
return
# parse text & save
new_insight_keys = ResponseTextParser(response_text).parse_v2(
"get_insight_keys"
)
if new_insight_keys:
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)

View file

@ -1,129 +0,0 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
MSG_TIME,
MODIFIED_MEMORIES,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.long_contra_repeat_prompt import (
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT,
LONG_CONTRA_REPEAT_SYSTEM_PROMPT,
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT,
)
class LongContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
for new_obs_node in new_obs_nodes:
text = new_obs_node.content
related_nodes = self.vector_store.similar_search(
text=text,
size=self.es_contra_repeat_similar_top_k,
exact_filters={
"memory_id": self.memory_id,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
has_match = False
for related_node in related_nodes:
if related_node.score_similar < self.long_contra_repeat_threshold:
continue
else:
has_match = True
all_obs_nodes.append(related_node)
if has_match:
all_obs_nodes.append(new_obs_node)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(
all_obs_nodes,
key=lambda x: x.meta_data.get(MSG_TIME, ""),
reverse=True,
)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(LONG_CONTRA_REPEAT_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(LONG_CONTRA_REPEAT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.generation_model.call(
messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(f"after contra repeat: {node.content} {node.status}")
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -1,47 +0,0 @@
from typing import List, Dict
from ...constants.common_constants import (
NEW_INSIGHT_NODES,
MODIFIED_MEMORIES,
INSIGHT_NODES,
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
NEW,
NOT_REFLECTED_MERGE_NODES,
CONTENT_MODIFIED,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
class SummaryCollectWorker(MemoryBaseWorker):
def _run(self):
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
not_reflected_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_OBS_NODES
)
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
NOT_REFLECTED_MERGE_NODES
)
# 合并逻辑复杂务必check
all_node_dict: Dict[str, MemoryNode] = {}
if insight_nodes:
all_node_dict.update(
{n.id: n for n in insight_nodes if n.meta_data.get(CONTENT_MODIFIED, False)}
)
if new_insight_nodes:
all_node_dict.update({n.content: n for n in new_insight_nodes})
if new_obs_nodes:
# 设置为非新
for n in new_obs_nodes:
n.meta_data[NEW] = "0"
all_node_dict.update({n.content: n for n in new_obs_nodes})
if not_reflected_merge_nodes and not_reflected_nodes:
# 进入reflect阶段
all_node_dict.update({n.id: n for n in not_reflected_nodes})
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))

View file

@ -1,177 +0,0 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
INSIGHT_NODES,
NEW_OBS_NODES,
INSIGHT_KEY,
INSIGHT_VALUE,
CONTENT_MODIFIED,
)
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.update_insight_prompt import (
UPDATE_INSIGHT_FEW_SHOT_PROMPT,
UPDATE_INSIGHT_SYSTEM_PROMPT,
UPDATE_INSIGHT_USER_QUERY_PROMPT,
)
class UpdateInsightWorker(MemoryBaseWorker):
def filter_obs_nodes(
self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode]
) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
if not insight_key or not insight_value:
self.logger.warning(
f"insight_key={insight_key} insight_value={insight_value} is empty!"
)
return insight_node, filtered_nodes, max_score
result = self.rank_model.call(
query=insight_key, documents=[x.content for x in new_obs_nodes]
)
if not result:
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
return insight_node, filtered_nodes, max_score
# 找到大于阈值的obs node
for index, score in result.rank_scores.items():
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_insight_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(
f"insight_key={insight_key} insight_value={insight_value} "
f"score={score} keep_flag={keep_flag}"
)
if not filtered_nodes:
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
return insight_node, filtered_nodes, max_score
def update_insight(
self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]
) -> MemoryNode:
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"update_insight insight_key={insight_key} insight_value={insight_value} "
f"doc.size={len(filtered_nodes)}"
)
# gen prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{node.content}")
update_insight_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_INSIGHT_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_INSIGHT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_INSIGHT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
insight_key=insight_key,
insight_key_value=insight_key + "" + insight_value,
),
)
self.logger.info(f"update_insight_message={update_insight_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=update_insight_message,
model_name=self.update_insight_model,
max_token=self.update_insight_max_token,
temperature=self.update_insight_temperature,
top_k=self.update_insight_top_k,
)
# return if empty
if not response_text:
self.add_run_info(
f"update_insight insight_key={insight_key} call llm failed!"
)
return insight_node
profile_list = ResponseTextParser(response_text).parse_v1(
f"update_profile {insight_key}"
)
if not profile_list:
self.add_run_info(
f"update_insight insight_key={insight_key} profile_list empty 1!"
)
return insight_node
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(
f"update_insight insight_key={insight_key} profile_list empty 2"
)
return insight_node
insight_value = profile_list[0]
if not insight_value or insight_value in ["", "重复"]:
self.logger.info(f"insight_value={insight_value}, skip.")
return insight_node
insight_node.meta_data[INSIGHT_VALUE] = insight_value
insight_node.meta_data[CONTENT_MODIFIED] = True
return insight_node
def _run(self):
# 获取新的obs和insight
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop update sights!")
return
if not insight_nodes:
self.logger.info("insight_nodes is empty, stop update sights!")
return
# 提交打分任务
for node in insight_nodes:
self.submit_thread(
self.filter_obs_nodes,
sleep_time=0.1,
insight_node=node,
new_obs_nodes=new_obs_nodes,
)
# 选择topN
result_list = []
for result in self.join_threads():
insight_node, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_insight_max_thread:
result_sorted = result_sorted[: self.update_insight_max_thread]
# 提交LLM update任务
for insight_node, filtered_nodes, _ in result_sorted:
self.submit_thread(
self.update_insight,
sleep_time=1,
insight_node=insight_node,
filtered_nodes=filtered_nodes,
)
# 等待结果
for result in self.join_threads():
if result:
insight_node: MemoryNode = result
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
self.logger.info(
f"after_update_insight insight_key={insight_key} insight_value={insight_value}"
)

View file

@ -1,241 +0,0 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ....memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.update_profile_prompt import (
UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT,
UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT,
UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT,
UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT,
UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT,
UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT
)
from ...chat.global_context import GlobalContext
class UpdateProfileWorker(MemoryBaseWorker):
@property
def extra_user_attrs(self):
return GlobalContext.global_configs.get("extra_user_attrs", [])
def filter_obs_nodes(
self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode]
) -> (MemoryNode, List[MemoryNode], float):
max_score: float = 0
filtered_nodes: List[MemoryNode] = []
result = self.rank_model.call(
query=user_attr.meta_data.get("description", ""),
documents=[x.content for x in new_obs_nodes],
)
if not result:
self.add_run_info(
f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!"
)
return user_attr, filtered_nodes, max_score
# 找到大于阈值的obs node
filtered_nodes: List[MemoryNode] = []
for index, score in result.rank_scores.items():
node = new_obs_nodes[index]
keep_flag = "filtered"
if score >= self.update_profile_threshold:
filtered_nodes.append(node)
keep_flag = "keep"
max_score = max(max_score, score)
self.logger.info(
f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
f"content={node.content} score={score} keep_flag={keep_flag}"
)
if not filtered_nodes:
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
return user_attr, filtered_nodes, max_score
def update_user_attr(
self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode]
) -> MemoryNode:
self.logger.info(
f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}"
)
# 根据不同的参数类型是否多值分别给出prompt
user_query_list = []
for node in filtered_nodes:
user_query_list.append(f"句子:{node.content}")
update_profile = f"{user_attr.meta_data.get("memory_key", "")}{user_attr.meta_data.get("description", "")}"
update_profile_value = update_profile + "" + "".join(user_attr.meta_data.get("value", ""))
if user_attr.meta_data.get("is_unique", 0) == 1:
update_profile_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value,
),
)
else:
update_profile_message = self.prompt_to_msg(
system_prompt=self.get_prompt(UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT),
few_shot=self.get_prompt(UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT),
user_query=self.get_prompt(UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list),
update_profile=update_profile,
update_profile_value=update_profile_value,
),
)
self.logger.info(f"update_profile_message={update_profile_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=update_profile_message,
model_name=self.update_profile_model,
max_token=self.update_profile_max_token,
temperature=self.update_profile_temperature,
top_k=self.update_profile_top_k,
)
# return if empty
if not response_text:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!"
)
return user_attr
profile_list = ResponseTextParser(response_text).parse_v1(
f"update_attr {user_attr.meta_data.get("memory_key", "")}"
)
if not profile_list:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!"
)
return user_attr
profile_list = profile_list[0]
if not profile_list:
self.add_run_info(
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2"
)
return user_attr
profile = profile_list[0]
if not profile or profile in ["", "重复"]:
self.logger.info(f"profile={profile}, skip.")
return user_attr
# check 英文中午逗号
if user_attr.meta_data.get("is_unique", 0) == 1:
user_attr.meta_data["value"] = [profile.strip()]
else:
attr_value_list = profile.replace("", ",").split(",")
user_attr.meta_data["value"] = [
x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list)))
]
return user_attr
def add_extra_user_attrs(self):
# 解析为空返回
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
if not extra_user_attr_list:
return
for user_attr_info in extra_user_attr_list:
user_attr_split = user_attr_info.split(":")
# 格式不对返回
if len(user_attr_split) < 1:
continue
user_attr_key = user_attr_split[0]
user_attr_desc = ""
if len(user_attr_split) >= 2:
user_attr_desc = user_attr_split[1]
user_attr_unique = 0
if len(user_attr_split) >= 3:
user_attr_unique = int(user_attr_split[2])
# 已经包含返回
if user_attr_key in self.user_profile_dict:
user_attr = self.user_profile_dict[user_attr_key]
# description为空补充description
if not user_attr.meta_data.get("description", ""):
user_attr.meta_data["description"] = user_attr_desc
continue
# 增加新属性
new_attr = MemoryNode(
memory_id=self.memory_id,
meta_data={
"memory_key": user_attr_key,
"is_unique": int(user_attr_unique),
"is_mutable": 1,
"description": user_attr_desc
},
memory_type=MemoryTypeEnum.PROFILE,
status=1,
)
self.user_profile_dict[user_attr_key] = new_attr
def _run(self):
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
if not new_obs_nodes:
self.logger.info("new_obs_nodes is empty, stop user profile!")
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
return
# 增加环境变量配置的属性
if self.extra_user_attrs:
self.add_extra_user_attrs()
new_user_profile: List[MemoryNode] = []
self.set_context(NEW_USER_PROFILE, new_user_profile)
for user_attr_key, user_attr in self.user_profile_dict.items():
# 不可修改直接跳过
if user_attr.meta_data.get("is_mutable", 0) != 1:
new_user_profile.append(user_attr)
self.logger.info(f"{user_attr_key} is not mutable! continue")
continue
self.submit_thread(
self.filter_obs_nodes,
sleep_time=0.1,
user_attr=user_attr,
new_obs_nodes=new_obs_nodes,
)
# 选择topN
result_list = []
for result in self.join_threads():
user_attr, filtered_nodes, max_score = result
if not filtered_nodes:
continue
result_list.append(result)
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
if len(result_sorted) > self.update_profile_max_thread:
result_sorted = result_sorted[: self.update_profile_max_thread]
# 提交LLM update任务
for user_attr, filtered_nodes, _ in result_sorted:
self.submit_thread(
self.update_user_attr,
sleep_time=1,
user_attr=user_attr,
filtered_nodes=filtered_nodes,
)
# collect result & save
for result in self.join_threads():
if result:
user_attribute: MemoryNode = result
self.logger.info(
f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} "
f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}"
)
new_user_profile.append(user_attribute)

View file

@ -1,117 +0,0 @@
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...constants.common_constants import (
NEW_OBS_NODES,
TODAY_OBS_NODES,
MSG_TIME,
NEW_OBS_WITH_TIME_NODES,
MODIFIED_MEMORIES,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...scheme.memory_node import MemoryNode
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.contra_repeat_prompt import (
CONTRA_REPEAT_FEW_SHOT_PROMPT,
CONTRA_REPEAT_SYSTEM_PROMPT,
CONTRA_REPEAT_USER_QUERY_PROMPT,
)
class ContraRepeatWorker(MemoryBaseWorker):
def _run(self):
# 合并当前的obs和今日的obs
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(
NEW_OBS_WITH_TIME_NODES
)
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
all_obs_nodes: List[MemoryNode] = []
if new_obs_nodes:
all_obs_nodes.extend(new_obs_nodes)
if new_obs_with_time_nodes:
all_obs_nodes.extend(new_obs_with_time_nodes)
if today_obs_nodes:
all_obs_nodes.extend(today_obs_nodes)
if not all_obs_nodes:
self.add_run_info("all_obs_nodes is empty!")
return
# gene prompt
user_query_list = []
all_obs_nodes = sorted(
all_obs_nodes,
key=lambda x: x.meta_data.get(MSG_TIME, ""),
reverse=True,
)
for i, n in enumerate(all_obs_nodes):
user_query_list.append(f"{i + 1} {n.content}")
merge_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(CONTRA_REPEAT_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(CONTRA_REPEAT_FEW_SHOT_PROMPT),
user_query=self.get_prompt(CONTRA_REPEAT_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response_text = self.generation_model.call(
messages=merge_obs_message,
model_name=self.merge_obs_model,
max_token=self.merge_obs_max_token,
temperature=self.merge_obs_temperature,
top_k=self.merge_obs_top_k,
)
# return if empty
if not response_text:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
if len(idx_merge_obs_list) <= 0:
self.add_run_info("idx_merge_obs_list is empty!")
return
# add merged obs
merge_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_merge_obs_list:
if not obs_content_list:
continue
# [6, 逃课]
if len(obs_content_list) != 2:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, keep_flag = obs_content_list
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(all_obs_nodes):
self.logger.warning(f"idx={idx} is invalid!")
continue
if keep_flag not in ["矛盾", "被包含", ""]:
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != "":
node.status = MemoryNodeStatus.EXPIRED.value
merge_obs_nodes.append(node)
self.logger.info(
f"after contra repeat: {node.content} {node.status}"
)
# save context
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)

View file

@ -1,167 +0,0 @@
from datetime import datetime
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...utils.tool_functions import (
time_to_formatted_str,
get_datetime_info_dict,
extract_date_parts,
)
from ...constants.common_constants import (
REFLECTED,
DT,
TIME_INFER,
NEW,
MSG_TIME,
KEY_WORD,
DATATIME_WORD_LIST,
NEW_OBS_WITH_TIME_NODES,
CONTENT_MODIFIED,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ...scheme.message import Message
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_observation_with_time_prompt import (
GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT,
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT,
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT,
)
class GetObservationWithTimeWorker(MemoryBaseWorker):
def add_observation(
self, message: Message, obs_content: str, time_infer: str, keywords: str
):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: time_infer, # 推断的时间
KEY_WORD: keywords, # 关键词
CONTENT_MODIFIED: True, # 新增的obs需要置为true
}
# 事件时间
meta_data.update(
{f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()}
)
# 对话时间
meta_data.update(
{f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
return MemoryNode.init_from_attrs(
content=obs_content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.OBSERVATION.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if match:
dt = time_to_formatted_str(
time=msg.time_created,
date_format="",
string_format="{year}{month}{day}{weekday}{hour}",
)
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(
f"get obs with time user_query_list={user_query_list} is empty"
)
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT
).format(num_obs=len(user_query_list)),
few_shot=self.get_prompt(GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT),
user_query=self.get_prompt(
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT
).format(user_query="\n".join(user_query_list)),
)
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k,
)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
if time_infer == "":
time_infer = ""
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(
f"idx={idx} is invalid! messages.size={len(self.messages)}"
)
continue
new_obs_nodes.append(
self.add_observation(
message=self.messages[idx],
obs_content=obs_content,
time_infer=time_infer,
keywords=keywords,
)
)
# save context
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)

View file

@ -1,144 +0,0 @@
from datetime import datetime
from typing import List
from ...utils.response_text_parser import ResponseTextParser
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
from ...constants.common_constants import (
REFLECTED,
DT,
NEW_OBS_NODES,
TIME_INFER,
NEW,
MSG_TIME,
KEY_WORD,
DATATIME_WORD_LIST,
CONTENT_MODIFIED,
)
from ...enumeration.memory_status_enum import MemoryNodeStatus
from ...enumeration.memory_type_enum import MemoryTypeEnum
from ...scheme.memory_node import MemoryNode
from ...scheme.message import Message
from ..memory_base_worker import MemoryBaseWorker
from ...prompts.get_observation_prompt import (
GET_OBSERVATION_FEW_SHOT_PROMPT,
GET_OBSERVATION_SYSTEM_PROMPT,
GET_OBSERVATION_USER_QUERY_PROMPT,
)
class GetObservationWorker(MemoryBaseWorker):
def add_observation(self, message: Message, obs_content: str, keywords: str):
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
REFLECTED: "0", # reflect标记
DT: dt, # 当天标记
NEW: "1", # summary-long标记
MSG_TIME: message.time_created, # 对话时间
TIME_INFER: "", # 推断的时间
KEY_WORD: keywords, # 关键词
CONTENT_MODIFIED: True, # 新增的obs需要置为true
}
meta_data.update(
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
)
return MemoryNode(
content=obs_content,
memory_id=self.memory_id,
memory_type=MemoryTypeEnum.OBSERVATION.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
)
def _run(self):
# gene prompt
user_query_list = []
i = 1
for msg in self.messages:
match = False
for time_keyword in DATATIME_WORD_LIST:
if time_keyword in msg.content:
match = True
break
if not match:
user_query_list.append(f"{i} 用户:{msg.content}")
i += 1
if not user_query_list:
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
return
obtain_obs_message = self.prompt_to_msg(
system_prompt=self.get_prompt(GET_OBSERVATION_SYSTEM_PROMPT).format(
num_obs=len(user_query_list)
),
few_shot=self.get_prompt(GET_OBSERVATION_FEW_SHOT_PROMPT),
user_query=self.get_prompt(GET_OBSERVATION_USER_QUERY_PROMPT).format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
# call LLM
response_text: str = self.generation_model.call(
messages=obtain_obs_message,
model_name=self.summary_messages_model,
max_token=self.summary_messages_max_token,
temperature=self.summary_messages_temperature,
top_k=self.summary_messages_top_k,
)
# return if empty
if not response_text:
self.add_run_info("summary call llm failed!", continue_run=False)
return
# parse text
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
if len(idx_obs_list) <= 0:
self.add_run_info("idx_obs_list is empty!", continue_run=False)
return
# gene new obs nodes
new_obs_nodes: List[MemoryNode] = []
for obs_content_list in idx_obs_list:
if not obs_content_list:
continue
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
if len(obs_content_list) != 4:
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
continue
idx, time_infer, obs_content, keywords = obs_content_list
if obs_content in ["", "重复"]:
continue
if not idx.isdigit():
self.logger.warning(f"idx={idx} is invalid!")
continue
# 序号需要修正-1
idx = int(idx) - 1
if idx >= len(self.messages):
self.logger.warning(
f"idx={idx} is invalid! messages.size={len(self.messages)}"
)
continue
new_obs_nodes.append(
self.add_observation(
message=self.messages[idx],
obs_content=obs_content,
keywords=keywords,
)
)
# save context
self.set_context(NEW_OBS_NODES, new_obs_nodes)

View file

@ -1,70 +0,0 @@
from ...utils.response_text_parser import ResponseTextParser
from enumeration.message_role_enum import MessageRoleEnum
from worker.memory_base_worker import MemoryBaseWorker
from ...chat.global_context import GlobalContext
from ...prompts.info_filter_prompt import INFO_FILTER_FEW_SHOT_PROMPT, INFO_FILTER_SYSTEM_PROMPT, INFO_FILTER_USER_QUERY_PROMPT
class InfoFilterWorker(MemoryBaseWorker):
def _run(self):
# filter user msg
info_messages = []
for msg in self.messages:
if msg.role != MessageRoleEnum.USER.value:
continue
if len(msg.content) >= self.info_filter_msg_max_size:
continue
info_messages.append(msg)
# gene prompt
user_query = "\n".join(
[f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)]
)
info_filter_message = self.prompt_to_msg(
system_prompt=self.get_prompt(INFO_FILTER_SYSTEM_PROMPT).format(
batch_size=len(info_messages)
),
few_shot=self.get_prompt(INFO_FILTER_FEW_SHOT_PROMPT),
user_query=self.get_prompt(INFO_FILTER_USER_QUERY_PROMPT).format(
user_query=user_query
),
)
self.logger.info(f"info_filter_message={info_filter_message}")
# call llm
response_text = self.generation_model.call(
messages=info_filter_message,
model_name=self.info_filter_model,
max_token=self.info_filter_max_token,
temperature=self.info_filter_temperature,
top_k=self.info_filter_top_k,
)
# return if empty
if not response_text:
self.add_run_info("info score call llm failed!", continue_run=False)
return
# parse text
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
if len(info_score_list) != len(info_messages):
self.add_run_info(
f"info_score_size != info_messages_size, "
f"{len(info_score_list)} vs {len(info_messages)}",
continue_run=False,
)
return
# 过滤value=0的messages
filtered_messages = []
for msg, info_score in zip(info_messages, info_score_list):
if not info_score:
continue
score = info_score[0]
# if score in ("1", "2",):
if score in ("2",):
msg.info_score = score
filtered_messages.append(msg)
# 后续不会关注为0的msg直接丢弃
self.messages = filtered_messages

View file

@ -2,7 +2,7 @@ import unittest
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
@ -12,16 +12,18 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
config = {
"module_name": "dashscope_embedding",
"model_name": "text-embedding-v2",
"clazz": "models.llama_index_embedding_model"
"clazz": "models.llama_index_embedding_model",
}
emb = LlamaIndexEmbeddingModel(**config)
config = {
"index_name": "0626_1",
"index_name": "0708_2",
"es_url": "http://localhost:9200",
"embedding_model": emb,
"use_hybrid": True
}
self.es_store = LlamaIndexElasticSearchStore(**config)
self.es_store = LlamaIndexEsMemoryStore(**config)
self.data = [
MemoryNode(
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
@ -100,15 +102,53 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
meta_data={"5": "5"}
),
MemoryNode(
content="ggggggggg",
memory_type="profile",
user_id="6",
status="valid",
memory_id="hhh234",
meta_data={"5": "5"}
),
MemoryNode(
content="ggggggggg",
memory_type="profile",
user_id="6",
status="valid",
memory_id="iii234",
meta_data={"5": "5"}
),
MemoryNode(
content="ggggggggg",
memory_type="profile",
user_id="6",
status="valid",
memory_id="jjj234",
meta_data={"5": "5"}
),
MemoryNode(
content="ggggggggg",
memory_type="profile",
user_id="6",
status="valid",
memory_id="kkk234",
meta_data={"5": "5"}
),
]
def test_retrieve(self):
filter_dict = {
"user_id": "6",
}
# filter_dict = {
# "user_id": "6",
# }
filter_dict = {}
for node in self.data:
self.es_store.insert(node)
self.es_store.insert(MemoryNode(
content="xxxxxx",
memory_type="profile",
@ -117,9 +157,10 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
memory_id="ggg567",
meta_data={"5": "5"}
))
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
print(len(res))
print(res)
self.es_store.update(MemoryNode(
content="test update",
@ -128,7 +169,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
status="invalid",
memory_id="ggg567"
))
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
print(len(res))
print(res)
@ -140,7 +181,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
memory_id="ggg567"
))
import asyncio
res = asyncio.run(self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10))
res = asyncio.run(self.es_store.a_retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15))
#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()