summay_long

This commit is contained in:
hs 2024-07-02 12:30:38 +08:00
parent 344685b308
commit d590f98abd
16 changed files with 1101 additions and 41 deletions

View file

@ -12,11 +12,6 @@ RETRIEVE_MEMORY_NODES = "retrieve_memory_nodes"
RANKED_MEMORY_NODES = "ranked_memory_nodes"
PIPELINE = "pipeline"
WORKER = "worker"
@ -58,8 +53,6 @@ INSIGHT_KEY = "insight_key"
INSIGHT_VALUE = "insight_value"
DT = "dt"
MSG_TIME = "msg_time"
NEW = "new"
@ -68,8 +61,6 @@ TIME_INFER = "time_infer"
KEY_WORD = "key_word"
REFLECTED = "reflected"
NEW_USER_PROFILE = "new_user_profile"
RECALL_TYPE = "recall_type"
@ -82,9 +73,6 @@ TIME_MATCHED = "time_matched"
QUERY_KEYWORDS = "query_keywords"
TIME_FORMAT_V1 = "{year}{month}{day}{weekday}{hour}"
DATATIME_KEY_MAP = {
@ -94,5 +82,3 @@ DATATIME_KEY_MAP = {
"": "week",
"星期几": "weekday",
}
CONTENT_MODIFIED = "content_modified"

View file

@ -1,30 +1,29 @@
from memory_scope.enumeration.language_enum import LanguageEnum
DATATIME_WORD_LIST = {
LanguageEnum.CN:
[
"",
"",
"",
"",
"星期",
"",
"分钟",
"小时",
"",
"上午",
"下午",
"早上",
"早晨",
"晚上",
"中午",
"",
"",
"清晨",
"傍晚",
"凌晨",
"",
],
LanguageEnum.CN: [
"",
"",
"",
"",
"星期",
"",
"分钟",
"小时",
"",
"上午",
"下午",
"早上",
"早晨",
"晚上",
"中午",
"",
"",
"清晨",
"傍晚",
"凌晨",
"",
],
LanguageEnum.EN: [
]
@ -69,3 +68,9 @@ COLON_WORD = {
LanguageEnum.CN: "",
LanguageEnum.EN: ":"
}
COMMA_WORD = {
LanguageEnum.CN: "",
LanguageEnum.EN: ","
}

View file

@ -0,0 +1,11 @@
system_prompt:
cn:
few_shot_prompt:
cn:
user_query_prompt:
cn:
content_format:
cn: "用户的{insight_key}{insight_value}"

View file

@ -0,0 +1,160 @@
from datetime import datetime
from typing import List
from memory_scope.utils.tool_functions import time_to_formatted_str, get_datetime_info_dict, prompt_to_msg
from memory_scope.constants.common_constants import (
NEW_INSIGHT_NODES,
DT,
NOT_REFLECTED_MERGE_NODES,
NEW_INSIGHT_KEYS,
INSIGHT_KEY,
INSIGHT_VALUE,
)
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.constants.language_constants import NONE_WORD
class GetInsightWorker(MemoryBaseWorker):
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
created_dt = datetime.now()
obs_dt = time_to_formatted_str(time=created_dt)
# 组合meta_data
meta_data = {
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 = self.prompt_handler.content_format.format(insight_key=insight_key, insight_value=insight_value)
return MemoryNode(
content=content,
user_name=self.user_name,
target_name=self.target_name,
memory_type=MemoryTypeEnum.INSIGHT.value,
meta_data=meta_data,
status=MemoryNodeStatus.ACTIVE.value,
obs_dt=obs_dt,
obs_updated=True,
)
def reflect_new_insight_key(
self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode]
) -> MemoryNode | None:
# 检索历史memory
related_nodes = self.vector_store.similar_search(
text=insight_key,
size=self.es_insight_similar_top_k,
exact_filters={
"user_name": self.user_name,
"target_name": self.target_name,
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": [
MemoryTypeEnum.OBSERVATION.value,
MemoryTypeEnum.OBS_CUSTOMIZED.value,
],
},
)
# 合并新增nodes
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.memory_id # memory_id or use_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 index, score in result.rank_scores.items():
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]
# prepare prompt
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
get_insight_message = prompt_to_msg(
system_prompt=self.prompt_handler.system_prompt,
few_shot=self.prompt_handler.few_shot_prompt,
user_query=self.prompt_handler.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 = 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:
self.add_run_info("reflect_upon_user_attr call llm failed!")
return
response_text = response.message.content.strip()
if response_text in [self.get_language_value(NONE_WORD)]:
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.obs_reflected = True

View file

@ -0,0 +1,9 @@
system_prompt:
cn:
few_shot_prompt:
cn:
user_query_prompt:
cn:

View file

@ -0,0 +1,95 @@
from typing import List
from memory_scope.utils.response_text_parser import ResponseTextParser
from memory_scope.constants.common_constants import (
NEW_OBS_NODES,
NOT_REFLECTED_OBS_NODES,
INSIGHT_NODES,
INSIGHT_KEY,
NEW_INSIGHT_KEYS,
NOT_REFLECTED_MERGE_NODES,
)
from memory_scope.constants.language_constants import COLON_WORD, COMMA_WORD
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.utils.tool_functions import prompt_to_msg
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 not node.obs_reflected
]
# 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 = prompt_to_msg(
system_prompt=self.prompt_handler.system_prompt.format(
num_questions=self.reflect_num_questions
),
few_shot=self.prompt_handler.few_shot_prompt,
user_query=self.prompt_handler.user_query_prompt.format(
exist_keys=self.get_language_value(COMMA_WORD).join(exist_keys), user_query="\n".join(user_query_list)
),
)
self.logger.info(f"reflect_message={reflect_message}")
# # call LLM
response = 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:
self.add_run_info("reflect_obs_questions call llm failed!")
return
# parse text & save
new_insight_keys = ResponseTextParser(response.message.content).parse_v2(
"get_insight_keys"
)
if new_insight_keys:
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)

View file

@ -0,0 +1,12 @@
systemp_prompt:
cn:
en:
few_shot_prompt:
cn:
en:
user_query_prompt:
cn:
en:

View file

@ -0,0 +1,133 @@
from typing import List
from memory_scope.utils.response_text_parser import ResponseTextParser
from memory_scope.constants.common_constants import (
NEW_OBS_NODES,
MSG_TIME,
MODIFIED_MEMORIES,
)
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.utils.tool_functions import prompt_to_msg
from memory_scope.constants.language_constants import NONE_WORD, INCLUDED_WORD, CONTRADICTORY_WORD
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={
"user_name": self.user_name,
"target_name": self.target_name,
"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 = prompt_to_msg(
system_prompt=self.prompt_handler.system_prompt.format(
num_obs=len(user_query_list)
),
few_shot=self.prompt_handler.few_shot_prompt,
user_query=self.prompt_handler.user_query_prompt.format(
user_query="\n".join(user_query_list)
),
)
self.logger.info(f"merge_obs_message={merge_obs_message}")
# call LLM
response = 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:
self.add_run_info("contra repeat call llm failed!")
return
# parse text
idx_merge_obs_list = ResponseTextParser(response.message.content).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
long_contra_repeat_keep_flag = [
self.get_language_value(CONTRADICTORY_WORD),
self.get_language_value(INCLUDED_WORD),
self.get_language_value(NONE_WORD),
]
if keep_flag not in long_contra_repeat_keep_flag.values():
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
continue
node: MemoryNode = all_obs_nodes[idx]
if keep_flag != self.get_language_value(NONE_WORD):
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

@ -0,0 +1,46 @@
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.memory.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.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})
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))

View file

@ -0,0 +1,67 @@
system_prompt:
cn: "从下面的句子中提取出给定类别的用户资料信息,并判断与已有信息是否矛盾,若矛盾以新信息为准。整合已有信息和新信息并输出。若已有信息为空则直接输出提取的
信息。若无法提取出给定类别的用户资料信息则输出“无”。
请一步步思考,并按如下格式输出:
思考: 思考的依据和过程150字以内。
用户资料: <信息或无>, 一定加<>"
en:
few_shot_promp:
cn: "
示例1:
句子:因为昨天成都下大雨,用户全身都被淋湿了。
句子:用户关心明天成都的天气预报。
类别:用户所在地区
已有信息:用户所在地区:
思考:从第一句句子可以得出用户在成都。第二句句子没有直接透露用户所在地信息,但与第一句句子用户在成都的信息吻合。已有信息为空,直接输出得出的信息。
用户资料:<成都>
示例2:
句子:用户最近养好了肠胃。
句子:用户关注中医养生。
类别:用户健康状况
已有信息:用户健康状况: 肠胃不好,高血压
思考:从第一句句子可以得出用户最近养好了肠胃,与已有信息矛盾,以新信息为准。第二句句子与用户健康状况无关。整合已有信息和新信息得到用户健康状况是肠胃健康,高血压。
用户资料:<肠胃健康,高血压>
示例3:
句子:用户刚刚毕业,第一份工作是银行前台。
句子:用户的理想工作是职业游戏选手。
类别:用户职业
已有信息:用户职业:在招商银行工作
思考:整合已有信息和第一句句子的信息可以得出用户的现在的职业是招商银行前台。第二句句子说明了用户的理想工作但并不是现在的职业。
用户资料:<招商银行前台>
示例4:
句子:用户大学期间接触过优化算法的研究。
类别:用户学习专业
已有信息:用户学习专业:与人工智能相关
思考:从句子可以得出用户大学学习的专业与优化算法相关,这与已有信息(用户学习专业与人工智能相关)不矛盾,整合可以得出用户大学学习的专业与人工智能和优化算法相关。
用户资料:<与人工智能和优化算法相关>
示例5:
句子:用户单身。
句子用户受到一名18岁男生的追求但不想接受又不想伤害他。
句子:用户喜欢成熟且情绪稳定的男生。
类别:用户情感状况
已有信息:用户情感状况:有男朋友
思考从第一句句子可以得出用户现在单身与已有信息矛盾以新信息为准。从第二句句子得出用户受到一名18岁男生的追求但并不喜欢他。第三句话表达了用户理想的伴侣类型但与用户
情感状况无关。整合得出用户情感状况为单身受到一名18岁男生的追求但并不喜欢他。
用户资料:<单身受到一名18岁男生的追求但并不喜欢他。>
"
en:
user_query_prompt:
cn: "
{user_query}
类别:{insight_key}
已有信息:{insight_key_value}
"
en:
user_query:
cn: "句子:{content}"
insight_value:
cn: ["无", "重复"]
en:

View file

@ -0,0 +1,173 @@
from typing import List
from memory_scope.utils.response_text_parser import ResponseTextParser
from memory_scope.constants.common_constants import (
INSIGHT_NODES,
NEW_OBS_NODES,
INSIGHT_KEY,
INSIGHT_VALUE,
)
from memory_scope.utils.tool_functions import prompt_to_msg
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.constants.language_constants import COMMA_WORD, COLON_WORD
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(self.prompt_handler.user_query.format(content=node.content))
update_insight_message = prompt_to_msg(
system_prompt=self.prompt_handler.system_prompt,
few_shot=self.prompt_handler.few_shot_prompt,
user_query=self.prompt_handler.user_query_prompt.format(
user_query="\n".join(user_query_list),
insight_key=insight_key,
insight_key_value=insight_key + self.get_language_value(COLON_WORD) + insight_value,
),
)
self.logger.info(f"update_insight_message={update_insight_message}")
# call LLM
response: 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:
self.add_run_info(
f"update_insight insight_key={insight_key} call llm failed!"
)
return insight_node
profile_list = ResponseTextParser(response.message.content).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.prompt_handler.insight_value:
self.logger.info(f"insight_value={insight_value}, skip.")
return insight_node
insight_node.meta_data[INSIGHT_VALUE] = insight_value
insight_node.obs_updated = 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

@ -0,0 +1,123 @@
update_plural_profile_system_prompt:
cn: "
从下面的句子中提取出给定类别的用户资料信息,并判断和已有信息是否重复。只输出无重复的新信息。若无法提取该类别的用户资料的新信息则回答无。
请一步步思考,并按如下格式输出:
思考: 思考的依据和过程150字以内。
用户资料: <信息>或<无>, 一定加<>
"
en:
update_plural_profile_few_shot_prompt:
cn: "
示例1:
句子:用户上周去了西溪游泳馆游泳,那个游泳馆人非常多。
句子:用户计划每周六和朋友张三去朝阳体育馆打羽毛球。
类别:运动(用户喜欢的运动)
已有信息:运动(用户喜欢的运动):游泳
思考:从第一句句子可以得出游泳是用户喜欢的运动之一,但与已有信息重复。从第二句句子可以得出羽毛球是用户喜欢的运动之一,是新的信息。
用户资料: <羽毛球>
示例2:
句子:用户对咖啡因过敏。
句子:用户不喜欢吃香菇。
类别:过敏(用户的已知过敏反应)
已有信息:过敏(用户的已知过敏反应): 咖啡因
思考:从第一句句子可以得出咖啡因是用户的已知过敏反应之一,但与已有信息重复。从第二句句子只能得出用户不喜欢香菇而非对香菇过敏,无法得出新的用户已知过敏信息。
用户资料: <无>
示例3:
句子:用户热衷于动作类类游戏如只狼、艾尔登法环。
句子用户在休闲时间经常长时间玩策略类游戏如文明6。
句子:用户是音乐发烧友,关注各个品牌的耳机的音质和性价比。
类别:爱好(用户的业余爱好)
已有信息:爱好(用户的业余爱好):
思考:从第一句句子可以得出动作类游戏是用户的爱好之一,是新的信息。从第二句句子可以得出策略类游戏是用户的爱好之一,是新的信息。从第三句句子可以得出音乐是用户的爱好之一,是新的信息。
用户资料: <动作类游戏, 策略类游戏, 音乐>
示例4:
句子:关于职场沟通你有什么具体的建议吗?最好结合一个实例。我一直听人说要加强沟通,经常和上司沟通,同步项目的进展,但是我总是感觉还有许多事情要做。
句子:项目并没有达到一个充分的可以汇报的状态,然后准备汇报材料又很费时间,导致有时候我没有及时和上司同步项目状态。针对这个情况你有什么建议?
类别:职业(用户的职业)
已有信息:职业(用户的职业):工程师
思考:句子中虽然提及了职场沟通等工作相关内容,但是并不能推断出用户的职位是什么,只能推知与宽泛的项目实施与管理相关。
用户资料: <无>
"
en:
update_plural_profile_user_query_prompt:
cn: "
{user_query}
类别:{update_profile}
已有信息:{update_profile_value}
"
en:
update_unique_profile_system_prompt:
cn: "
从下面的句子中提取出给定类别的用户资料信息,并判断与已有信息是否矛盾。若矛盾则输出更新的信息,若不矛盾则保留已有信息,整合已有信息和新信息并输出。
请一步步思考,并按如下格式输出:
思考: 思考的依据和过程150字以内。
用户资料: <信息>, 一定加<>
"
en:
update_unique_profile_few_shot_prompt:
cn: "
示例1:
句子:因为昨天成都下大雨,用户全身都被淋湿了。
句子:用户关心明天成都的天气预报。
类别:地区(用户所在地区)
已有信息:地区(用户所在地区): 杭州
思考:从第一句句子可以得出用户在成都。第二句句子没有直接透露用户所在地信息,但与第一句句子用户在成都的信息吻合。这与已有信息(用户在杭州)矛盾,输出更新的信息。
用户资料:<成都>
示例2:
句子:用户女朋友下个月过生日。
句子用户生日在7月15日。
类别:生日(用户的生日)。
已有信息生日用户的生日1987年7月15日。
思考第一句句子中提及生日但并不是用户的生日无法得出用户生日信息。从第二句句子可以得出用户生日在7月15日与已有信息不矛盾整合可以得出用户生日是1987年7月15日。
用户资料: <1987年7月15日>
示例3:
句子:用户在招商银行工作。
句子:用户刚刚毕业,第一份工作是银行前台。
句子:用户的理想工作是职业游戏选手。
类别:职业(用户的职业)
已有信息:职业(用户的职业):
思考:整合第一和第二句句子的信息可以得出用户的现在的职业是招商银行前台。第三句句子说明了用户的理想工作但并不是现在的职业。
用户资料:<招商银行前台>
示例4:
句子:用户大学期间接触过优化算法的研究。
类别:学习专业 (用户大学学习的专业)
已有信息:学习专业 (用户大学学习的专业):与人工智能相关
思考:从句子可以得出用户大学学习的专业与优化算法相关,这与已有信息(用户大学学习的专业与人工智能相关)不矛盾,整合可以得出用户大学学习的专业与人工智能和优化算法相关。
用户资料:<与人工智能和优化算法相关>
示例5:
句子:今天和同学去打球了。
句子:明天和女朋友一起去杭州旅游。
类别:学习专业 (用户大学学习的专业)
已有信息:学习专业 (用户大学学习的专业)
思考:两个句子和学习专业都没有关联,没有新提取的信息。
用户资料:<无>
"
en:
update_unique_profile_user_query_prompt:
cn: "
{user_query}
类别:{update_profile}
已有信息:{update_profile_value}
"
en:
user_query:
cn: "句子:{content}"
en:
update_profile_key:
cn: ["无", "重复"]

View file

@ -0,0 +1,236 @@
from typing import List
from memory_scope.utils.response_text_parser import ResponseTextParser
from memory_scope.constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.utils.global_context import GlobalContext
from memory_scope.utils.tool_functions import prompt_to_msg
from memory_scope.constants.language_constants import COMMA_WORD, COLON_WORD
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(self.prompt_handler.user_query.format(content=node.content))
update_profile = f"{user_attr.meta_data.get("memory_key", "")}({user_attr.meta_data.get("description", "")})"
update_profile_value = update_profile + self.get_language_prompt(COLON_WORD) + self.get_language_prompt(COMMA_WORD).join(user_attr.meta_data.get("value", ""))
if user_attr.meta_data.get("is_unique", 0) == 1:
update_profile_message = prompt_to_msg(
system_prompt=self.prompt_handler.update_unique_profile_system_prompt,
few_shot=self.prompt_handler.update_unique_profile_few_shot_prompt,
user_query=self.prompt_handler.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 = prompt_to_msg(
system_prompt=self.prompt_handler.update_plural_profile_system_prompt,
few_shot=self.prompt_handler.update_plural_profile_few_shot_prompt,
user_query=self.prompt_handler.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: 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:
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.message.content).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.prompt_handler.update_profile_key:
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(self.get_language_prompt(COLON_WORD))
# 格式不对返回
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,
obs_profile_updated=true,
)
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

@ -36,7 +36,7 @@ class GetObservationWorker(MemoryBaseWorker):
timestamp=message.time_created,
obs_dt=dt_handler.datetime_format(),
obs_reflected=False,
obs_profile_updated=False,
obs_updated=False,
obs_keyword=keywords)
node.gen_memory_id()
return node

View file

@ -32,5 +32,5 @@ class StoreMemoryWorker(MemoryBaseWorker):
timestamp=dt_handler.timestamp,
obs_dt=dt_handler.datetime_format(),
obs_reflected=False,
obs_profile_updated=False)
obs_updated=False)
self.vector_store.update(node)

View file

@ -35,10 +35,14 @@ class MemoryNode(BaseModel):
obs_reflected: bool = Field(False, description="if the observation is reflected")
obs_profile_updated: bool = Field(False, description="if the observation has updated user profile")
obs_updated: bool = Field(False, description="if the observation has updated user profile or insight")
obs_keyword: str = Field("", description="keywords of the content")
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.gen_memory_id()
@property
def node_keys(self):
return list(self.model_json_schema()["properties"].keys())