cn version fix

This commit is contained in:
jinli.yl 2024-07-18 18:07:53 +08:00
commit b971a10807
31 changed files with 991 additions and 506 deletions

View file

@ -30,7 +30,7 @@ memory_service:
delete_memory:
class: memory.operation.frontend_operation
workflow: set_query,retrieve_all_memory,delete_query
workflow: set_query,retrieve_all_memory,delete_memory
description: "delete a single long-term memory"
delete_all:
@ -72,12 +72,14 @@ worker:
extract_time:
class: memory.worker.frontend.extract_time_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
semantic_rank:
class: memory.worker.frontend.semantic_rank_worker
rank_model: dashscope_rank
fuse_rerank:
class: memory.worker.frontend.fuse_rerank_worker
fuse_score_threshold: 0.1
fuse_score_threshold: 0.05
fuse_ratio_dict:
conversation: 0.5
observation: 1
@ -94,12 +96,12 @@ worker:
class: memory.worker.frontend.print_memory_worker
retrieve_all_memory:
class: memory.worker.frontend.retrieve_memory_worker
retrieve_obs_top_k: 10000
retrieve_ins_top_k: 10000
retrieve_expired_top_k: 10000
delete_query:
retrieve_obs_top_k: 100
retrieve_ins_top_k: 100
retrieve_expired_top_k: 100
delete_memory:
class: memory.worker.write.update_memory_worker
method: delete_query
method: delete_memory
delete_all:
class: memory.worker.write.update_memory_worker
method: delete_all
@ -109,18 +111,26 @@ worker:
info_filter:
class: memory.worker.write.info_filter_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
load_today_memory:
class: memory.worker.write.load_memory_worker
retrieve_today_top_k: 100
get_observation:
class: memory.worker.write.get_observation_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
get_observation_with_time:
class: memory.worker.write.get_observation_with_time_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
contra_repeat:
class: memory.worker.write.contra_repeat_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
store_memory:
class: memory.worker.write.update_memory_worker
method: from_memory_key
@ -132,16 +142,27 @@ worker:
retrieve_insight_top_k: 100
get_reflection_subject:
class: memory.worker.summary.get_reflection_subject_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
update_insight:
class: memory.worker.summary.update_insight_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
rank_model: dashscope_rank
long_contra_repeat:
class: memory.worker.summary.long_contra_repeat_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
models:
dashscope_generation:
class: models.llama_index_generation_model
module_name: dashscope_generation
model_name: qwen-max
max_tokens: 2000
dashscope_embedding:
class: models.llama_index_embedding_model
module_name: dashscope_embedding
@ -160,7 +181,7 @@ memory_store:
embedding_model: dashscope_embedding
index_name: memory_index
es_url: http://localhost:9200
use_hybrid: false
use_hybrid: true
monitor:
class: storage.dummy_monitor

View file

@ -30,7 +30,7 @@ memory_service:
delete_memory:
class: memory.operation.frontend_operation
workflow: set_query,retrieve_all_memory,delete_query
workflow: set_query,retrieve_all_memory,delete_memory
description: "delete a single long-term memory"
delete_all:
@ -72,12 +72,14 @@ worker:
extract_time:
class: memory.worker.frontend.extract_time_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
semantic_rank:
class: memory.worker.frontend.semantic_rank_worker
rank_model: dashscope_rank
fuse_rerank:
class: memory.worker.frontend.fuse_rerank_worker
fuse_score_threshold: 0.1
fuse_score_threshold: 0.05
fuse_ratio_dict:
conversation: 0.5
observation: 1
@ -94,12 +96,12 @@ worker:
class: memory.worker.frontend.print_memory_worker
retrieve_all_memory:
class: memory.worker.frontend.retrieve_memory_worker
retrieve_obs_top_k: 10000
retrieve_ins_top_k: 10000
retrieve_expired_top_k: 10000
delete_query:
retrieve_obs_top_k: 100
retrieve_ins_top_k: 100
retrieve_expired_top_k: 100
delete_memory:
class: memory.worker.write.update_memory_worker
method: delete_query
method: delete_memory
delete_all:
class: memory.worker.write.update_memory_worker
method: delete_all
@ -109,18 +111,26 @@ worker:
info_filter:
class: memory.worker.write.info_filter_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
load_today_memory:
class: memory.worker.write.load_memory_worker
retrieve_today_top_k: 100
get_observation:
class: memory.worker.write.get_observation_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
get_observation_with_time:
class: memory.worker.write.get_observation_with_time_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
contra_repeat:
class: memory.worker.write.contra_repeat_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
store_memory:
class: memory.worker.write.update_memory_worker
method: from_memory_key
@ -132,16 +142,27 @@ worker:
retrieve_insight_top_k: 100
get_reflection_subject:
class: memory.worker.summary.get_reflection_subject_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
update_insight:
class: memory.worker.summary.update_insight_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
rank_model: dashscope_rank
long_contra_repeat:
class: memory.worker.summary.long_contra_repeat_worker
generation_model: dashscope_generation
generation_model_kwargs:
top_k: 1
models:
dashscope_generation:
class: models.llama_index_generation_model
module_name: dashscope_generation
model_name: qwen-max
max_tokens: 2000
dashscope_embedding:
class: models.llama_index_embedding_model
module_name: dashscope_embedding
@ -160,7 +181,7 @@ memory_store:
embedding_model: dashscope_embedding
index_name: memory_index
es_url: http://localhost:9200
use_hybrid: false
use_hybrid: true
monitor:
class: storage.dummy_monitor

View file

@ -53,6 +53,8 @@ class CliMemoryChat(BaseMemoryChat):
"""
self._memory_service: BaseMemoryService | str = memory_service
self._generation_model: BaseModel | str = generation_model
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
self.stream: bool = stream
self.human_name: str = human_name
self.assistant_name: str = assistant_name
@ -174,7 +176,7 @@ class CliMemoryChat(BaseMemoryChat):
self.logger.info(f"messages={messages}")
# Invoke the Language Model with the constructed message context, respecting streaming setting
generated = self.generation_model.call(messages=messages, stream=self.stream)
generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs)
# In non-streaming interactions, explicitly save the AI's reply to memory if instructed
if remember_response:

View file

@ -6,7 +6,6 @@ system_prompt:
memory_prompt:
cn: |
如果用户问题和以下信息没有关联,请忘记这些信息;如果用户问题和以下信息有关联,请记住这些信息,他们可以帮助更好地理解用户的问题
如果用户问题和以下信息没有关联,请忘记这些信息;如果用户问题和以下信息有关联,请记住这些信息
en: |
If the user's question is not related to the following information, please disregard it; if the user's question is related to the following information, please retain it, as it can help better understand the user's query.
If the user's question is not related to the information provided below, please disregard this information; if the user's question is related to the information provided below, please remember this information.

View file

@ -112,6 +112,37 @@ WEEKDAYS = {
]
}
MONTH_DICT = {
LanguageEnum.CN: [
"1月",
"2月",
"3月",
"4月",
"5月",
"6月",
"7月",
"8月",
"9月",
"10月",
"11月",
"12月",
],
LanguageEnum.EN: [
"January",
"February",
"March",
"April",
"May",
"June",
"July",
"August",
"September",
"October",
"November",
"December",
]
}
# Constants for the word 'none' in different languages
NONE_WORD = {
LanguageEnum.CN: "",

View file

@ -19,7 +19,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def _parse_params(self, **kwargs):
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
def _run(self):
"""
@ -47,7 +47,7 @@ class ExtractTimeWorker(MemoryBaseWorker):
self.logger.info(f"extract_time_message={extract_time_message}")
# Invoke the LLM to generate a response
response = self.generation_model.call(messages=extract_time_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=extract_time_message, **self.generation_model_kwargs)
# Handle empty or unsuccessful responses
if not response.status or not response.message.content:

View file

@ -1,107 +1,113 @@
time_string_format:
cn: |
{year}年{month}{day}日,{year}年第{week}周,{weekday}{hour}时。
en: |
{month} {day}, {year}, {week}th week of {year}, {weekday}, at {hour}.
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."
Task: 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:
cn: |
示例1:
句子:我记得你前年四月份去了阿联酋,阿联酋有哪些好玩的地方?迪拜和阿布扎比你更喜欢哪个?沙漠的景色壮观吗?
时间1992年8月20日1992年第34周周一18时46分25秒
时间1992年8月20日1992年第34周周一18时
回答:
- 1990 - 月4
- 1990 - 月4
示例2:
句子:后天下午三点的会议记得参加。我在日历上仔细标注了这个重要的日子,提醒自己不要错过。会议将在公司会议室举行,这是一个讨论未来发展方向的重要机会。
时间2024年6月19日2024年第25周周二13时30分0秒
时间2024年6月19日2024年第25周周二13时
回答:
- 2024 - 月6 - 日21 - 时15
- 2024 - 月6 - 日21 - 时15
示例3:
句子:下个月第一个周六去杭州玩。
时间2005年7月15日2005年第28周周六0时0分0秒
时间2005年7月15日2005年第28周周六0时
回答:
- 2005 - 月8 - 周31 - 星期几6
- 2005 - 月8月 - 周31 - 星期几:周六
示例4:
句子:上周末我们去的那个小镇真是太美了。
时间1999年12月2日1999年第48周周二8时40分10秒
时间1999年12月2日1999年第48周周二8时
回答:
- 1999 - 周47 - 星期几:6,7
- 1999 - 周47 - 星期几:周六,周日
示例5:
句子:再过半小时就要宣讲了,记得准备材料。
时间2020年6月22日2020年第25周周一9时30分0秒
时间2020年6月22日2020年第25周周一9时
回答:
- 2020 - 月6 - 日22 - 时10 - 分0 - 秒0
- 2020 - 月6 - 日22 - 时10 - 分0 - 秒0
示例6:
句子10000米长跑比赛的开始时间是3分47秒前。
时间1987年2月17日1987年第7周周三19时54分43秒
时间1987年2月17日1987年第7周周三19时
回答:
- 1987 - 月2 - 日17 - 时19 - 分50 - 秒56
示例7:
句子:上个月的这个时候我们还在筹备音乐会。每天都是忙碌而充实的日子,我们为音乐会的顺利举办而努力奋斗着。彩排、布景、节目安排,每一个细节都需要精心安排和准备。
时间1995年11月24日1995年第48周周二17时45分0秒
时间1995年11月24日1995年第48周周二17时
回答:
- 1995 - 月10 - 日24
示例8:
句子:我的朋友非常喜欢运动,他认为运动有助于增强身体素质。
时间2015年1月23日2015年第4周周四7时38分0秒
时间2015年1月23日2015年第4周周四7时
回答:
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.
Time: August 20, 1992, 34th week of 1992, Monday, at 18.
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.
Time: June 19, 2024, 25th week of 2024, Tuesday, at 13.
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.
Time: July 15, 2005, 28th week of 2005, Saturday, at 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.
Time: December 2, 1999, 48th week of 1999, Tuesday, at 8.
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
Time: June 22, 2020, 25th week of 2020, Monday, at 9.
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.
Time: February 17, 1987, 7th week of 1987, Wednesday, at 19.
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.
Time: November 24, 1995, 48th week of 1995, Tuesday, at 17.
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.
Time: January 23, 2015, 4th week of 2015, Thursday, at 7.
Answer:
None
@ -117,7 +123,5 @@ extract_time_user_query:
Time: {query_time_str}
Answer:
time_string_format:
cn: |
{year}年{month}月{day}日,{year}年第{week}周,{weekday}{hour}时{minute}分{second}秒。

View file

@ -21,7 +21,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
def _parse_params(self, **kwargs):
self.reflect_obs_cnt_threshold: int = kwargs.get("reflect_obs_cnt_threshold", 10)
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
self.reflect_num_questions: int = kwargs.get("reflect_num_questions", 5)
def new_insight_node(self, insight_key: str) -> MemoryNode:
@ -88,7 +88,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
self.logger.info(f"reflect_message={reflect_message}")
# Invoke Language Model for new insights
response = self.generation_model.call(messages=reflect_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=reflect_message, **self.generation_model_kwargs)
# Handle empty response
if not response.status or not response.message.content:
@ -98,7 +98,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
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))
self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
# Mark unaudited nodes as reflected
for node in not_reflected_nodes:

View file

@ -1,7 +1,7 @@
get_reflection_subject_system:
cn: |
任务:从下面的信息中提取出最重要的最多{num_questions}条{user_name}属性,要求不与已有的{user_name}属性语义重复。
要求1{user_name}属性可以是一般的{user_name}偏好,也可以是运动偏好,旅游偏好,饮食偏好等等,也可以是重要事件性质,比如最近重要的事情,也可以是一些高度概括的人生理想,价值观,人生观,性格, 也可以是和朋友的人际关系等等。
要求1{user_name}属性可以是基本信息,基础画像,也可以是运动偏好,旅游偏好,饮食偏好等等兴趣偏好,也可以是重要事件性质,比如最近重要的事情,也可以是一些高度概括的人生理想,价值观,人生观,性格,也可以是和朋友的人际关系等等。
要求2根据{user_name}属性,我们可以生成“{user_name}的<{user_name}属性>是什么?”的问题,以此可以从下面的信息中提取{user_name}属性对应的值。
输出格式:每一行输出一个{user_name}属性,每个{user_name}属性推荐4个字如果没有信息请回答无最多输出{num_questions}条。
@ -81,9 +81,9 @@ get_reflection_subject_few_shot:
{user_name}年龄为28岁。
{user_name}体重为70kg。
{user_name}是男性。
已有{user_name}属性:性别,年龄,体重,当前学习进展
已有{user_name}属性:性别,体重,当前学习进展
新增{user_name}属性:
年龄
get_reflection_subject_user_query:
cn: |

View file

@ -21,9 +21,10 @@ class LongContraRepeatWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def _parse_params(self, **kwargs):
self.unit_test_flag = False
self.long_contra_repeat_top_k: int = kwargs.get("long_contra_repeat_top_k", 2)
self.long_contra_repeat_threshold: float = kwargs.get("long_contra_repeat_threshold", 0.1)
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
def retrieve_similar_content(self, node: MemoryNode) -> (MemoryNode, List[MemoryNode]):
"""
@ -66,16 +67,19 @@ class LongContraRepeatWorker(MemoryBaseWorker):
for node in not_updated_nodes:
self.submit_thread_task(fn=self.retrieve_similar_content, node=node)
obs_node_dict: Dict[str, MemoryNode] = {}
for origin_node, retrieve_nodes in self.gather_thread_result():
if not retrieve_nodes:
continue
obs_node_dict[origin_node.memory_id] = origin_node
for node in retrieve_nodes:
if node.memory_id in obs_node_dict:
if self.unit_test_flag:
all_obs_nodes: List[MemoryNode] = not_updated_nodes
else:
obs_node_dict: Dict[str, MemoryNode] = {}
for origin_node, retrieve_nodes in self.gather_thread_result():
if not retrieve_nodes:
continue
obs_node_dict[node.memory_id] = node
all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True)
obs_node_dict[origin_node.memory_id] = origin_node
for node in retrieve_nodes:
if node.memory_id in obs_node_dict:
continue
obs_node_dict[node.memory_id] = node
all_obs_nodes: List[MemoryNode] = sorted(obs_node_dict.values(), key=lambda x: x.timestamp, reverse=True)
if not all_obs_nodes:
self.logger.warning("all_obs_nodes is empty, stop.")
@ -96,7 +100,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
self.logger.info(f"long_contra_repeat_message={long_contra_repeat_message}")
# Invokes the language model for processing the constructed prompt
response = self.generation_model.call(messages=long_contra_repeat_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=long_contra_repeat_message, **self.generation_model_kwargs)
# Handles the case where the model's response is empty
if not response or not response.message.content:

View file

@ -1,6 +1,8 @@
long_contra_repeat_system:
cn: |
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。
注意:只判断与“前面序号”的句子的关系,不要判断“后面序号”。
其中矛盾的形式可以有很多种,可以是逻辑上的矛盾,可以是属性上的变化导致的矛盾,比如不能同时在两个地方工作,同一个时刻不能在两个地点,同一个时刻不能干两件事情等等。
对每个句子都做一个判断,最后一共输出{num_obs}条判断。如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
请一步步思考,并按如下格式输出:
思考思考的依据和过程30字以内。
@ -102,3 +104,6 @@ long_contra_repeat_user_query:
cn: |
句子:
{user_query}
en: |
Sentences:
{user_query}

View file

@ -21,7 +21,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
def _parse_params(self, **kwargs):
self.update_insight_threshold: float = kwargs.get("update_insight_threshold", 0.1)
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
self.update_insight_max_count: int = kwargs.get("update_insight_max_count", 10)
def filter_obs_nodes(self,
@ -44,8 +44,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
filtered_nodes: List[MemoryNode] = []
# Check if insight node key or value is empty and log a warning
if not insight_node.key or not insight_node.value:
self.logger.warning(f"insight_key={insight_node.key} insight_value={insight_node.value} is empty!")
if not insight_node.key:
self.logger.warning(f"insight_key={insight_node.key} is empty!")
return insight_node, filtered_nodes, max_score
# Call the ranking model to get scores for each observed node's content against the insight key
@ -62,7 +62,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
filtered_nodes.append(node)
max_score = max(max_score, score)
# Log information about each node's processing
self.logger.info(f"insight_key={insight_node.key} insight_value={insight_node.value} "
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
f"score={score} keep_flag={keep_flag}")
# Warn if no nodes were filtered
@ -74,8 +74,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
def update_insight_node(self, insight_node: MemoryNode, insight_value: str):
dt_handler = DatetimeHandler()
content = (f"{self.user_name}{self.get_language_value(COLON_WORD)}{insight_node.key}"
f"{self.get_language_value(COLON_WORD)}{insight_value}")
key = self.prompt_handler.insight_string_format.format(name=self.target_name, key=insight_node.key)
content = f"{key}{self.get_language_value(COLON_WORD)}{insight_value}"
insight_node.content = content
insight_node.value = insight_value
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
@ -97,24 +97,23 @@ class UpdateInsightWorker(MemoryBaseWorker):
Returns:
MemoryNode: The updated MemoryNode with potentially revised insight value.
"""
self.logger.info(f"Updating insight for key={insight_node.key}, value={insight_node.value}, "
self.logger.info(f"Updating insight for key={insight_node.key}, old_value={insight_node.value}, "
f"with {len(filtered_nodes)} documents considered.")
# Generate the prompt for updating insight
user_query_list = [n.content for n in filtered_nodes]
system_prompt = self.prompt_handler.update_insight_system.foramt(user_name=self.target_name)
few_shot = self.prompt_handler.update_insight_few_shot.foramt(user_name=self.target_name)
user_query = self.prompt_handler.update_insight_user_query.foramt(
system_prompt = self.prompt_handler.update_insight_system.format(user_name=self.target_name)
few_shot = self.prompt_handler.update_insight_few_shot.format(user_name=self.target_name)
user_query = self.prompt_handler.update_insight_user_query.format(
user_query="\n".join(user_query_list),
insight_key=insight_node.key,
insight_key_value=insight_node.key + self.get_language_value(COLON_WORD) + insight_node.value)
# Construct the message for LLM interaction
update_insight_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
self.logger.info(f"Generated insight update message: {update_insight_message}")
# Call the Language Model for insight update
response = self.generation_model.call(messages=update_insight_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=update_insight_message, **self.generation_model_kwargs)
# Handle empty or invalid responses
if not response.status or not response.message.content:
@ -170,11 +169,11 @@ class UpdateInsightWorker(MemoryBaseWorker):
if node.action_status == ActionStatusEnum.NONE.value:
self.submit_thread_task(fn=self.filter_obs_nodes,
insight_node=node,
not_updated_nodes=not_updated_nodes)
obs_nodes=not_updated_nodes)
else:
self.submit_thread_task(fn=self.filter_obs_nodes,
insight_node=node,
not_updated_nodes=not_reflected_nodes)
obs_nodes=not_reflected_nodes)
# select top n
result_list = []
@ -190,7 +189,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
self.submit_thread_task(fn=self.update_insight, insight_node=insight_node, filtered_nodes=filtered_nodes)
# Gather the final results from all update tasks
self.gather_thread_result()
for _ in self.gather_thread_result():
pass
for node in not_updated_nodes:
node.obs_updated = 1

View file

@ -105,3 +105,9 @@ update_insight_user_query:
{user_query}
Category: {insight_key}
Existing information: {insight_key_value}
insight_string_format:
cn: |
{name}的{key}
en: |
The {key} of {name}

View file

@ -24,7 +24,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
FILE_PATH: str = __file__
def _parse_params(self, **kwargs):
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
self.retrieve_top_k: int = kwargs.get("retrieve_top_k", 30)
self.contra_repeat_max_count: int = kwargs.get("contra_repeat_max_count", 50)
@ -69,7 +69,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
self.logger.info(f"contra_repeat_message={contra_repeat_message}")
# call LLM
response = self.generation_model.call(messages=contra_repeat_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=contra_repeat_message, **self.generation_model_kwargs)
# return if empty
if not response.status or not response.message.content:

View file

@ -1,6 +1,8 @@
contra_repeat_system:
cn: |
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
任务:对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。
注意:只判断与“前面序号”的句子的关系,不要判断“后面序号”。
其中矛盾的形式可以有很多种,可以是逻辑上的矛盾,可以是属性上的变化导致的矛盾,比如不能同时在两个地方工作,同一个时刻不能在两个地点,同一个时刻不能干两件事情等等。
对每个句子都做一个判断,最后一共输出{num_obs}条判断。
请一步步思考,并按如下格式输出:
思考思考的依据和过程30字以内。
@ -107,4 +109,4 @@ contra_repeat_user_query:
en: |
Sentences:
{user_query}
{user_query}

View file

@ -50,7 +50,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
dt_handler = DatetimeHandler(dt=msg.time_created)
dt = dt_handler.string_format(self.prompt_handler.time_string_format)
# Append formatted timestamp-query pairs to the user_query_list
user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
user_query_list.append(f"{i + 1} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
# Construct the system prompt with the count of observations
system_prompt = self.prompt_handler.get_observation_with_time_system.format(num_obs=len(user_query_list),

View file

@ -1,17 +1,20 @@
time_string_format:
cn: |
{year}年{month}月{day}日{weekday}{hour}点
{year}年{month}{day}日{weekday}{hour}点
en: |
{month} {day}, {year}, {weekday}, at {hour}
get_observation_with_time_system:
cn: |
任务指令:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。
任务:从下面的{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}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,回答“无”。
如果{user_name}信息涉及时间,则结合对话时间推断{user_name}信息的时间信息,没有则不输出。
对每个句子都做一次信息提取,最后一共输出{num_obs}条信息。
请一步步思考,并一定要按如下格式依次输出,最后的结果一定要加<>:
思考思考的依据和过程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.
@ -47,7 +50,7 @@ get_observation_with_time_few_shot:
{user_name}句子:
1 2020年1月4日周日10点 {user_name}我花5000元买了100股海天味业。
2 2023年4月27日周五8点 {user_name}:明天是我和妻子的结婚纪念日,帮我推荐一家餐厅。
3 2020年1月4日周日10点 {user_name}我花5000元买了100股海天味业
3 2020年1月4日周日10点 {user_name}我花50000元买了100股茅台
4 2021年6月2日周四23点 {user_name}:谢啦。我中午在公司附近吃,帮我推荐一家阿里巴巴徐汇滨江园区附近的餐厅吧。
5 2021年7月9日周六11点 {user_name}:两个坏消息,我打羽毛球把拍子打断线了。。。然后我去我朋友家撸猫,结果我猫毛过敏,今天疯狂打喷嚏。。。
@ -56,8 +59,8 @@ get_observation_with_time_few_shot:
思考从第2句可以得知{user_name}与妻子的结婚纪念日是明天,这是关于{user_name}重要纪念日的信息。其余信息重要性不足。{user_name}信息涉及时间结合对话时间为2023年4月27日
以及结婚纪念日为周期性日期,推断{user_name}与妻子的结婚纪念日是每年4月28日。
信息:<2> <每年4月28日> <{user_name}与妻子的结婚纪念日是每年4月28日。> <妻子, 结婚纪念日>
思考第3句含有的信息与第1句重复了
信息:<3> <> <重复> <>
思考第3句含有的信息与第1句相似,但是不重复,可以得知{user_name}购买了茅台股票
信息:<3> <> <{user_name}购买了茅台股票购买数量为100股购买金额为50000元。> <茅台, 股票>
思考从第4句以得知{user_name}在阿里巴巴徐汇滨江园区工作,这是关于{user_name}的工作的重要信息。其余信息重要性不足。{user_name}信息不涉及时间。
信息:<4> <> <{user_name}在阿里巴巴徐汇滨江园区工作。> <阿里巴巴, 徐汇滨江园区, 工作>
思考从第5句可以得知{user_name}前天打羽毛球时把球拍打断了线,但这不是重要的信息。还可以得知{user_name}对猫毛过敏,这是关于{user_name}的健康的重要信息。{user_name}信息不涉及时间。

View file

@ -17,7 +17,7 @@ class GetObservationWorker(MemoryBaseWorker):
OBS_STORE_KEY: str = NEW_OBS_NODES
def _parse_params(self, **kwargs):
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
def add_observation(self, message: Message, time_infer: str, obs_content: str, keywords: str):
dt_handler = DatetimeHandler(dt=message.time_created)
@ -65,7 +65,7 @@ class GetObservationWorker(MemoryBaseWorker):
user_query_list = []
for i, msg in enumerate(filter_messages):
# Construct each user query item with index, target name, and message content
user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
user_query_list.append(f"{i + 1} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
# Format the system prompt with the number of observations and target name
system_prompt = self.prompt_handler.get_observation_system.format(num_obs=len(user_query_list),
@ -109,7 +109,7 @@ class GetObservationWorker(MemoryBaseWorker):
obtain_obs_message = self.build_message(filter_messages)
# Generates observations using the language model
response = self.generation_model.call(messages=obtain_obs_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=obtain_obs_message, **self.generation_model_kwargs)
if not response.status or not response.message.content:
return

View file

@ -1,37 +1,45 @@
get_observation_system:
cn: |
任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,与相应的关键词。最多提取{num_obs}条信息。对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。
不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。
任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,与相应的关键词。如果没有重要信息则回答“无”,最多提取{num_obs}条信息。
{user_name}的重要信息可以包含用户基本信息,用户画像信息,用户兴趣偏好信息,用户性格,用户价值观,用户人际关系,用户重大事件转折点等等重要信息。
如果句子中只包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,回答“无”。
对每个句子都做一次信息提取,最后一共输出{num_obs}条信息。
请一定要按如下格式依次输出,最后的结果一定要加<>:
请一步步思考,并一定要按如下格式依次输出,最后的结果一定要加<>:
思考思考的依据和过程50字以内。
信息:<句子序号> <> <明确的重要信息或“重复”或”无“> <关键词>
信息:<句子序号> <> <明确的重要信息或“无”> <关键词>
en: |
Task: Extract important information and corresponding keywords from the following {num_obs} sentences about {user_name}. Extract up to {num_obs} pieces of information. 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." 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.
Task: Extract important information, interests and corresponding keywords from the following {num_obs} sentences about {user_name} in sequence. Extract up to {num_obs} pieces of information.
If all the information in a sentence is completely repetitive of what has already been extracted, reply "repetitive," and if there is no important information, reply "none."
The user information may include basic user information, user profiles, user interests and preferences, personality, values, significant life events, turning points, and other important 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 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: <sentence number> <> <Clear important information or “Repeat” or “None”> <keywords>
get_observation_few_shot:
cn: |
示例1
{user_name}句子:
1 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办
2 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢
3 {user_name}:我现在处境很糟,没有工作,负债几万,怎么办
3 {user_name}:我现在心情很糟糕
4 {user_name}:我是一个刚毕业的学生,对社会,行业不了解,给我介绍一下社会系统和行业格局
5 {user_name}我花5000元买了100股海天味业。
6 {user_name}我花50000元买了100股茅台。
思考从第1句可以得知{user_name}现在没有工作,负债几万,这是关于{user_name}工作与经济状况的重要信息。
信息:<1> <> <{user_name}当前无工作且负债几万> <无工作, 负债几万>
思考第2句是{user_name}对他人观点的讨论和疑问,没有明确提及{user_name}个人信息。
信息:<2> <> <无> <>
思考:第3句含有的信息与第1句重复了
信息:<3> <> <重复> <>
思考:从第3句可以得知{user_name}当前心情不好
信息:<3> <> <{user_name}当前心情不好> <心情>
思考从第4句可以得知{user_name}是一个刚毕业的学生,这是关于{user_name}身份背景状况的重要信息。其余信息重要性不足。
信息:<4> <> <{user_name}是一名刚毕业的学生。> <刚毕业, 学生>
思考从第5句可以得知{user_name}购买了海天味业股票购买数量为100股购买金额为5000元这是关于{user_name}的投资决策的重要信息。
信息:<5> <> <{user_name}购买了海天味业股票购买数量为100股购买金额为5000元。> <海天味业, 股票>
思考第6句含有的信息与第1句相似可以得知{user_name}购买了茅台股票。
信息:<6> <> <{user_name}购买了茅台股票购买数量为100股购买金额为50000元。> <茅台, 股票>
示例2
{user_name}句子:
@ -61,7 +69,7 @@ get_observation_few_shot:
6 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧我听你未来的你说有一个叫难忘的朱古力232这个人他弄的设是Windows设\n这是剧本1剧本2未完待续
思考从第1句可以得知{user_name}寻求购买新能源汽车的建议或推荐,这是这是关于{user_name}的大宗消费的重要的信息。
信息:<1> <> <{user_name}寻求购买新能源汽车的建议或推荐。> <购买, 新能源汽车>
思考从第2句可以得知{user_name}当前所在城市为上海,这是关于{user_name}的生活地区的重要信息。其余信息与第1句重复了。
思考从第2句可以得知{user_name}当前所在城市为上海,这是关于{user_name}的生活地区的重要信息。
信息:<2> <> <{user_name}所在的城市是上海。> <上海>
思考第3句是{user_name}对某个观点的讨论和疑问,没有明确提及{user_name}个人信息。
信息:<3> <> <无> <>
@ -71,6 +79,16 @@ get_observation_few_shot:
信息:<5> <> <{user_name}购买了海天味业股票购买数量为100股购买金额为5000元。> <海天味业, 股票>
思考第6句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。
信息:<6> <> <无> <>
示例4
{user_name}句子:
1 {user_name}:李子好酸啊,我不太喜欢吃。
2 {user_name}:桃子上的毛太多了,我不爱吃他。
思考从第1句可以得知{user_name}不太喜欢吃李子。
信息:<1> <> <{user_name}不喜欢吃李子。> <李子>
思考从第2句可以得知{user_name}不喜欢吃桃子,和上一句相似都是对某一种水果不喜欢,但是表达了不同的信息。
信息:<2> <> <{user_name}不喜欢吃桃子。> <西瓜>
en: |
Example 1:

View file

@ -20,7 +20,7 @@ class InfoFilterWorker(MemoryBaseWorker):
def _parse_params(self, **kwargs):
self.preserved_scores: str = kwargs.get("preserved_scores", "2,3")
self.info_filter_msg_max_size: int = kwargs.get("info_filter_msg_max_size", 200)
self.generation_model_top_k: int = kwargs.get("generation_model_top_k", 1)
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
def _run(self):
"""
@ -69,7 +69,7 @@ class InfoFilterWorker(MemoryBaseWorker):
self.logger.info(f"info_filter_message={info_filter_message}")
# call llm
response = self.generation_model.call(messages=info_filter_message, top_k=self.generation_model_top_k)
response = self.generation_model.call(messages=info_filter_message, **self.generation_model_kwargs)
# return if empty
if not response.status or not response.message.content:
@ -81,7 +81,6 @@ class InfoFilterWorker(MemoryBaseWorker):
info_score_list = ResponseTextParser(response_text).parse_v1(self.__class__.__name__)
if len(info_score_list) != len(info_messages):
self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}")
return
# filter messages
filtered_messages: List[Message] = []

View file

@ -1,15 +1,17 @@
info_filter_system:
cn: |
任务指令:对所给{batch_size}个句子中所含有的关于{user_name}的信息打分分数为0,1,2或3。
注意其中0表示不包含{user_name}信息1表示句子中只包含{user_name}假设的信息或者{user_name}虚构的内容2表示包含{user_name}的一般信息,时效性信息或者需要猜测才能得到的{user_name}信息3表示明确含有或者可以确定推断出关于{user_name}的重要信息,或者{user_name}要求记录。
对每个句子都做一次信息打分,一共输出{batch_size}个分数。
任务:对所给{batch_size}个句子中所含有的关于{user_name}的信息打分分数为0,1,2或3。
注意其中0表示不包含用户信息1表示句子中只包含用户假设的信息或者用户虚构的内容比如用户创作的小说或剧本2表示包含用户的一般信息时效性信息或者需要猜测才能得到的用户信息3表示明确含有或者可以确定推断出关于用户的重要信息或者用户要求记录。
{user_name}的重要信息可以包含用户基本信息,用户画像信息,用户兴趣偏好信息,用户性格,用户价值观,用户人际关系,用户重大事件转折点等等重要信息。
对每个句子都做一次信息打分,一共输出{batch_size}个分数,不需要写最终结果。
请一定要按如下格式依次输出,最后的结果一定要加<>:
思考思考的依据和过程30字以内。
结果:<句子序号> <分数:0或1或2或3>
en: |
Instruction: Evaluate the given {batch_size} sentences for information about {user_name}, with a score of 0, 1, 2, or 3.
Note: A score of 0 indicates no information about {user_name}, 1 indicates hypothetical or fictional content about or from {user_name}, 2 indicates general, timely, or speculative information about {user_name}, and 3 indicates significant and verifiable information about {user_name} or information that {user_name} requested to be recorded.
Task: Evaluate the given {batch_size} sentences for information about {user_name}, with a score of 0, 1, 2, or 3.
Note: A score of 0 indicates no information about {user_name}, 1 indicates the sentence contains hypothetical information about the user or content concocted by the user, such as novels or scripts they have authored, 2 indicates general, timely, or speculative information about {user_name}, and 3 indicates significant and verifiable information about {user_name} or information that {user_name} requested to be recorded.
The user information may include basic user details, user profile data, user preferences, user personality, user values, user critical life events, turning points, and other important information.
Perform information scoring for each sentence, and output a total of {batch_size} scores.
Please ensure to output in the following format, and the final result must be enclosed in <>:
Thought: The basis and process of thinking, within 30 words.

View file

@ -103,4 +103,5 @@ class LoadMemoryWorker(MemoryBaseWorker):
self.submit_thread_task(self.retrieve_today_memory, query=query, dt=dt)
# Waits for all submitted tasks to complete
self.gather_thread_result()
for _ in self.gather_thread_result():
pass

View file

@ -43,22 +43,36 @@ class UpdateMemoryWorker(MemoryBaseWorker):
self.logger.info(f"delete_all.size={len(nodes)}")
return nodes
def delete_query(self):
if "query" not in self.chat_kwargs:
return
def delete_memory(self):
if "query" in self.chat_kwargs:
query = self.chat_kwargs["query"].strip()
if not query:
return
query = self.chat_kwargs["query"].strip()
if not query:
return
i = 0
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
for node in nodes:
if node.content == query:
i += 1
node.action_status = ActionStatusEnum.DELETE.value
self.logger.info(f"delete_memory.query.size={len(nodes)}")
return nodes
i = 0
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
for node in nodes:
if node.content == query:
i += 1
node.action_status = ActionStatusEnum.DELETE.value
self.logger.info(f"delete_query.size={len(nodes)}")
return nodes
elif "memory_id" in self.chat_kwargs:
memory_id = self.chat_kwargs["memory_id"].strip()
if not memory_id:
return
i = 0
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
for node in nodes:
if node.memory_id == memory_id:
i += 1
node.action_status = ActionStatusEnum.DELETE.value
self.logger.info(f"delete_memory.memory_id.size={len(nodes)}")
return nodes
return []
def _run(self):
method = self.method.strip()

View file

@ -52,6 +52,7 @@ class BaseModel(metaclass=ABCMeta):
else:
kwargs = self.kwargs
self._model = obj_cls(**kwargs)
return self._model
@abstractmethod

View file

@ -49,6 +49,7 @@ class LlamaIndexGenerationModel(BaseModel):
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
else:
raise RuntimeError("prompt and messages are both empty!")
self.data.update(**kwargs)
def after_call(self,
model_response: ModelResponse,

View file

@ -216,16 +216,14 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
# Connecting to ElasticsearchStore locally
es_local = ElasticsearchStore(
index_name=index_name,
es_url=es_url,
)
es_url=es_url)
# Connecting to Elastic Cloud with username and password
es_cloud_user_pass = ElasticsearchStore(
index_name=index_name,
es_cloud_id=es_cloud_id,
es_user=es_user,
es_password=es_password,
)
es_password=es_password)
# Connecting to Elastic Cloud with API Key
es_cloud_api_key = ElasticsearchStore(

View file

@ -1,7 +1,7 @@
import datetime
import re
from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST
from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.logger import Logger
@ -52,7 +52,7 @@ class DatetimeHandler(object):
"""
return {
"year": self._dt.year,
"month": self._dt.month,
"month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1],
"day": self._dt.day,
"hour": self._dt.hour,
"minute": self._dt.minute,

View file

@ -37,6 +37,23 @@ class MemoryHandler(object):
self._id_memory_dict.clear()
self._key_id_dict.clear()
def add_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
if key not in self._key_id_dict:
return self.set_memories(key, nodes, log_repeat)
if isinstance(nodes, MemoryNode):
nodes = [nodes]
for node in nodes:
_id = node.memory_id
if _id not in self._key_id_dict[key]:
self._key_id_dict[key].append(_id)
if _id not in self._id_memory_dict:
self._id_memory_dict[_id] = node
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content or node.key} "
f"store_status={node.store_status} action_status={node.action_status}")
def set_memories(self, key: str, nodes: MemoryNode | List[MemoryNode], log_repeat: bool = True):
if nodes is None:
nodes = []

View file

@ -1,368 +0,0 @@
from typing import Dict, List, Any, Optional, cast
import ray
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_memory_store import BaseMemoryStore
from memory_scope.utils.logger import Logger
ray.init(ignore_reinit_error=True)
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
"""
Custom asynchronous dense vector strategy extending LlamaIndex's ElasticsearchStore's strategy.
This strategy enables hybrid search combining KNN queries with text queries and supports customizable ranking functions.
"""
def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]:
"""
Constructs a hybrid query body combining KNN search with a text query, and applies filters.
Args:
query (str): The text query to be combined with the KNN results.
knn (Dict[str, Any]): The KNN query part specifying the vector search parameters.
filter (List[Dict[str, Any]]): A list of filters to apply to the search.
top_k (int): The number of top results to retrieve.
Returns:
Dict[str, Any]: The constructed query body for Elasticsearch to perform the hybrid search.
"""
# Combines KNN query with a text query and applies optional RRF ranking for result balancing
query_body = {
"knn": knn,
"query": {
"bool": {
"must": [
{
"match": {
self.text_field: {
"query": query,
}
}
}
],
"filter": filter,
}
},
}
# Configures Rank-Risk Function (RRF) if enabled or specified, to balance scores between KNN and text matches
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]] = None,
) -> Dict[str, Any]:
if filter is None:
filter = []
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 we 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]:
"""
Converts the provided standard Llama-index filters into an Elasticsearch compatible filter format.
This function processes each key-value pair in the input dictionary. If the value is a list,
it constructs a 'should' clause with multiple 'term' sub-clauses for each item in the list,
requiring at least one to match. If the value is not a list, it forms a 'must' clause with a single 'term'
sub-clause. The resulting structure is nested within a 'bool' clause which is the standard way to combine
boolean logic in Elasticsearch queries.
Args:
standard_filters (Dict[str, List[str]]): A dictionary where keys represent filter fields and values are
either single values or lists of values to filter by.
Returns:
Dict[str, Any]: An Elasticsearch query filter dictionary ready to be used in a query.
"""
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
# The following decorator '@ray.remote' is used to define a function or class that should be executed remotely
# by Ray. This facilitates parallel and distributed computation. However, due to the instruction constraints,
# no modification or additional explanation is provided for this part.
@ray.remote
class _LlamaIndexEsMemoryStore(BaseMemoryStore):
def __init__(self,
embedding_model: dict,
index_name: str,
es_url: str,
use_hybrid: bool = True,
**kwargs):
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
embedding_model = LlamaIndexEmbeddingModel(**embedding_model)
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]]) -> List[MemoryNode]:
pass
def batch_insert(self, nodes: List[MemoryNode]):
pass
def batch_update(self, nodes: List[MemoryNode], update_embedding: bool = True):
pass
def batch_delete(self, nodes: List[MemoryNode]):
pass
def insert(self, node: MemoryNode):
"""
Inserts a MemoryNode into the Elasticsearch store by converting it to aTextNode.
Args:
node (MemoryNode): The MemoryNode to be inserted into the store.
"""
self.index.insert_nodes([self._memory_node_2_text_node(node)])
def delete(self, node: MemoryNode):
"""
Deletes a MemoryNode from the Elasticsearch store based on its memory_id.
Args:
node (MemoryNode): The MemoryNode to be deleted, identified by its memory_id.
Returns:
bool: The result of the deletion operation, typically True if successful.
"""
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):
"""
Closes the Elasticsearch store, releasing any resources associated with it.
This method ensures that the connection to the Elasticsearch instance is properly closed,
which is a good practice to prevent resource leaks when you're done interacting with the store.
"""
self.es_store.close()
@staticmethod
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
"""
Converts a MemoryNode object into a TextNode object.
Args:
memory_node (MemoryNode): The MemoryNode to be converted.
Returns:
TextNode: The converted TextNode object with the content and metadata from the MemoryNode.
"""
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:
"""
Converts a NodeWithScore object into a MemoryNode object.
Args:
text_node (NodeWithScore): The NodeWithScore to be converted.
Returns:
MemoryNode: The converted MemoryNode object with the text and metadata from the NodeWithScore.
"""
return MemoryNode(content=text_node.text, **text_node.metadata)
class LlamaIndexEsMemoryStore():
def __init__(self,
embedding_model: BaseModel,
index_name: str,
es_url: str,
use_hybrid: bool = True,
**kwargs):
if 'embedding_model' in kwargs: kwargs.pop('embedding_model')
self.proxy_obj = _LlamaIndexEsMemoryStore.remote(embedding_model.kwargs, index_name, es_url, use_hybrid,
**kwargs)
def retrieve_memories(self,
query: str,
top_k: int,
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
return ray.get(self.proxy_obj.retrieve_memories.remote(query, top_k, filter_dict))
def insert(self, node: MemoryNode):
return ray.get(self.proxy_obj.insert.remote(node))
def delete(self, node: MemoryNode):
return ray.get(self.proxy_obj.delete.remote(node))
def update(self, node: MemoryNode):
return ray.get(self.proxy_obj.update.remote(node))
def update_batch(self, nodes: List[MemoryNode]) -> Any:
"""
Updates a batch of memory nodes asynchronously using Ray.
Args:
nodes (List[MemoryNode]): A list of MemoryNode objects to be updated.
Returns:
Any: The result from the remote task once completed.
"""
return ray.get(self.proxy_obj.update_batch.remote(nodes))
def close(self) -> Any:
"""
Closes the Elasticsearch memory store asynchronously using Ray.
Returns:
Any: The result from the remote task once completed.
"""
return ray.get(self.proxy_obj.close.remote())
def update_memories(self, nodes: MemoryNode | List[MemoryNode]) -> Any:
"""
Updates one or more memory nodes asynchronously using Ray.
Args:
nodes (MemoryNode | List[MemoryNode]): A single MemoryNode or a list of MemoryNode objects to be updated.
Returns:
Any: The result from the remote task once completed.
"""
return ray.get(self.proxy_obj.update_memories.remote(nodes))
@staticmethod
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
"""
Converts a MemoryNode object into a TextNode object.
Args:
memory_node (MemoryNode): The MemoryNode to convert.
Returns:
TextNode: The converted TextNode object with content and metadata.
"""
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:
"""
Converts a TextNode (with score) into a MemoryNode object.
Args:
text_node (NodeWithScore): The TextNode to convert, which includes a 'score' attribute.
Returns:
MemoryNode: The converted MemoryNode object with content and metadata.
"""
return MemoryNode(content=text_node.text, **text_node.metadata)

View file

@ -0,0 +1,352 @@
import datetime
import unittest
from memory_scope.cli import MemoryScope
from memory_scope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.scheme.message import Message
from memory_scope.utils import Logger
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.tool_functions import init_instance_by_config
class TestWorkersCn(unittest.TestCase):
"""Tests for LLIEmbedding"""
def setUp(self):
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
ms = MemoryScope()
ms.load_config("config/demo_config_cn.yaml")
ms.init_global_content_by_config()
def tearDown(self):
self.logger.close()
@unittest.skip
def test_extract_time(self):
name = "extract_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
query = "明天我去上海出差"
query_timestamp = int(datetime.datetime.now().timestamp())
worker.set_context(QUERY_WITH_TS, (query, query_timestamp))
worker.run()
result = worker.get_context(EXTRACT_TIME_DICT)
worker.logger.info(f"result={result}")
@unittest.skip
def test_info_filter(self):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃苹果"),
Message(role=MessageRoleEnum.USER.value, content="明天我要去高考"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [msg.content for msg in worker.chat_messages]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_info_filter2(self):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"),
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"),
Message(role=MessageRoleEnum.USER.value, content="听说篮球运动对身体很好,是真的吗?"),
Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"),
Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"),
Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"),
Message(role=MessageRoleEnum.USER.value, content="听你这么说,我感觉挺有信心的,谢了!"),
Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"),
Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"),
Message(role=MessageRoleEnum.USER.value, content="听说打篮球可以长高,这是真的吗?"),
Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"),
Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"),
Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [msg.content for msg in worker.chat_messages]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation(self):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃苹果"),
Message(role=MessageRoleEnum.USER.value, content="我准备去高考"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃西瓜"),
Message(role=MessageRoleEnum.USER.value, content="我不爱吃西瓜"),
Message(role=MessageRoleEnum.USER.value, content="我在一家叫京东的公司干活"),
]
# chat_messages = [
# Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃西瓜"),
# Message(role=MessageRoleEnum.USER.value, content="我在一家叫京东的公司干活"),
# ]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation2(self):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"),
Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"),
Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"),
Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"),
Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"),
Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"),
Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"),
Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"),
Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation_with_time(self):
name = "get_observation_with_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"),
Message(role=MessageRoleEnum.USER.value, content="上个月我去了杭州旅游"),
Message(role=MessageRoleEnum.USER.value, content="下周我要去高考"),
Message(role=MessageRoleEnum.USER.value, content="明天我去北京出差"),
Message(role=MessageRoleEnum.USER.value, content="前天我把苹果扔掉了"),
Message(role=MessageRoleEnum.USER.value, content="前天我把苹果扔掉了,我不喜欢吃"),
Message(role=MessageRoleEnum.USER.value, content="明天是我生日"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_contra_repeat(self):
name = "contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"),
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"),
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"),
MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴干活"),
MemoryNode(user_name="AI", target_name="用户", content="我不爱吃西瓜"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
worker.logger.info(f"result1={result1}")
worker.logger.info(f"result2={result2}")
worker.logger.info(f"result3={result3}")
worker.logger.info(f"result4={result4}")
@unittest.skip
def test_get_reflection_subject(self):
name = "get_reflection_subject"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
MemoryNode(content="用户有要好朋友,常一起外出就餐。"),
MemoryNode(content="用户打算换工作,关心北京的工作机会分布。"),
MemoryNode(content="用户喜爱尝试新美食,求美食应用推荐。"),
MemoryNode(content="用户喜欢在家做饭,寻求海鲜菜谱。"),
MemoryNode(content="用户在北京阿里云园区工作。"),
MemoryNode(content="用户是阿里云百炼的工程师。"),
MemoryNode(content="用户目前的工作是大语言模型的应用开发"),
MemoryNode(content="用户想知道维持广泛社交关系的方法。"),
]
worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_handler.set_memories(INSIGHT_NODES, [])
worker.run()
result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.get_reflection={result}")
return worker
@unittest.skip
def test_update_insight_worker(self):
reflection_worker = self.test_get_reflection_subject.__wrapped__(self)
name = "update_insight"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context=reflection_worker.context,
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户喜欢打王者荣耀"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.update_insight={result}")
# @unittest.skip
def test_long_contra_repeat_worker(self):
name = "long_contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
MemoryNode(content="用户在上海工作。"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.unit_test_flag = True
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")

View file

@ -0,0 +1,352 @@
import datetime
import unittest
from memory_scope.cli import MemoryScope
from memory_scope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.scheme.message import Message
from memory_scope.utils import Logger
from memory_scope.utils.global_context import G_CONTEXT
from memory_scope.utils.tool_functions import init_instance_by_config
class TestWorkersEn(unittest.TestCase):
"""Tests for LLIEmbedding"""
def setUp(self):
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
ms = MemoryScope()
ms.load_config("config/demo_config_en.yaml")
ms.init_global_content_by_config()
def tearDown(self):
self.logger.close()
@unittest.skip
def test_extract_time(self):
name = "extract_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
query = "明天我去上海出差"
query_timestamp = int(datetime.datetime.now().timestamp())
worker.set_context(QUERY_WITH_TS, (query, query_timestamp))
worker.run()
result = worker.get_context(EXTRACT_TIME_DICT)
worker.logger.info(f"result={result}")
@unittest.skip
def test_info_filter(self):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃苹果"),
Message(role=MessageRoleEnum.USER.value, content="明天我要去高考"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [msg.content for msg in worker.chat_messages]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_info_filter2(self):
name = "info_filter"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"),
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"),
Message(role=MessageRoleEnum.USER.value, content="听说篮球运动对身体很好,是真的吗?"),
Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"),
Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"),
Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"),
Message(role=MessageRoleEnum.USER.value, content="听你这么说,我感觉挺有信心的,谢了!"),
Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"),
Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"),
Message(role=MessageRoleEnum.USER.value, content="听说打篮球可以长高,这是真的吗?"),
Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"),
Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"),
Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [msg.content for msg in worker.chat_messages]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation(self):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃苹果"),
Message(role=MessageRoleEnum.USER.value, content="我准备去高考"),
Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃西瓜"),
Message(role=MessageRoleEnum.USER.value, content="我不爱吃西瓜"),
Message(role=MessageRoleEnum.USER.value, content="我在一家叫京东的公司干活"),
]
# chat_messages = [
# Message(role=MessageRoleEnum.USER.value, content="我不喜欢吃西瓜"),
# Message(role=MessageRoleEnum.USER.value, content="我在一家叫京东的公司干活"),
# ]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation2(self):
name = "get_observation"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"),
Message(role=MessageRoleEnum.USER.value, content="最近在北京的工作压力太大,有什么放松的建议吗?"),
Message(role=MessageRoleEnum.USER.value, content="说到朋友,我确实有几位很要好的朋友,我们经常一起出去吃饭。"),
Message(role=MessageRoleEnum.USER.value, content="对了,最近想换工作,你觉得北京的哪个区工作机会更多?"),
Message(role=MessageRoleEnum.USER.value, content="我很喜欢尝试新的美食,有没有推荐的美食应用?"),
Message(role=MessageRoleEnum.USER.value, content="我有时也喜欢自己在家做饭,你有没有好的海鲜菜谱推荐?"),
Message(role=MessageRoleEnum.USER.value, content="我在北京阿里云园区工作"),
Message(role=MessageRoleEnum.USER.value, content="我是阿里云百炼的工程师"),
Message(role=MessageRoleEnum.USER.value, content="最后一个问题,你知道怎么才能维持广泛的社交关系吗?"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_get_observation_with_time(self):
name = "get_observation_with_time"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
chat_messages = [
Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"),
Message(role=MessageRoleEnum.USER.value, content="上个月我去了杭州旅游"),
Message(role=MessageRoleEnum.USER.value, content="下周我要去高考"),
Message(role=MessageRoleEnum.USER.value, content="明天我去北京出差"),
Message(role=MessageRoleEnum.USER.value, content="前天我把苹果扔掉了"),
Message(role=MessageRoleEnum.USER.value, content="前天我把苹果扔掉了,我不喜欢吃"),
Message(role=MessageRoleEnum.USER.value, content="明天是我生日"),
]
worker.set_context(CHAT_MESSAGES, chat_messages)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(NEW_OBS_WITH_TIME_NODES)]
result = "\n".join(result)
worker.logger.info(f"result={result}")
@unittest.skip
def test_contra_repeat(self):
name = "contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴工作"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result1 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"),
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result2 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="用户在京东工作"),
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result3 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
nodes = [
MemoryNode(user_name="AI", target_name="用户", content="我喜欢吃西瓜"),
MemoryNode(user_name="AI", target_name="用户", content="用户在阿里巴巴干活"),
MemoryNode(user_name="AI", target_name="用户", content="我不爱吃西瓜"),
]
worker.memory_handler.set_memories(NEW_OBS_NODES, nodes)
worker.run()
result4 = "\n".join([" ".join([node.content, node.store_status, node.action_status])
for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)])
worker.logger.info(f"result1={result1}")
worker.logger.info(f"result2={result2}")
worker.logger.info(f"result3={result3}")
worker.logger.info(f"result4={result4}")
@unittest.skip
def test_get_reflection_subject(self):
name = "get_reflection_subject"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
MemoryNode(content="用户有要好朋友,常一起外出就餐。"),
MemoryNode(content="用户打算换工作,关心北京的工作机会分布。"),
MemoryNode(content="用户喜爱尝试新美食,求美食应用推荐。"),
MemoryNode(content="用户喜欢在家做饭,寻求海鲜菜谱。"),
MemoryNode(content="用户在北京阿里云园区工作。"),
MemoryNode(content="用户是阿里云百炼的工程师。"),
MemoryNode(content="用户目前的工作是大语言模型的应用开发"),
MemoryNode(content="用户想知道维持广泛社交关系的方法。"),
]
worker.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
worker.memory_handler.set_memories(INSIGHT_NODES, [])
worker.run()
result = [node.key for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.get_reflection={result}")
return worker
@unittest.skip
def test_update_insight_worker(self):
reflection_worker = self.test_get_reflection_subject.__wrapped__(self)
name = "update_insight"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context=reflection_worker.context,
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户喜欢打王者荣耀"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(INSIGHT_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.update_insight={result}")
# @unittest.skip
def test_long_contra_repeat_worker(self):
name = "long_contra_repeat"
worker: MemoryBaseWorker = init_instance_by_config(
config=G_CONTEXT.worker_config[name],
suffix_name="worker",
name=name,
is_multi_thread=False,
context={},
context_lock=None,
thread_pool=G_CONTEXT.thread_pool)
nodes = [
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
MemoryNode(content="用户在北京工作,感到压力大,寻求放松方式。"),
MemoryNode(content="用户在上海工作。"),
]
worker.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
worker.unit_test_flag = True
worker.run()
result = [node.content for node in worker.memory_handler.get_memories(MERGE_OBS_NODES)]
result = "\n".join(result)
worker.logger.info(f"result.long_contra_repeat={result}")