mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-07 08:26:06 +00:00
cn version fix
This commit is contained in:
commit
b971a10807
31 changed files with 991 additions and 506 deletions
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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: "无",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}秒。
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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: |
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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}信息不涉及时间。
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ class BaseModel(metaclass=ABCMeta):
|
|||
else:
|
||||
kwargs = self.kwargs
|
||||
self._model = obj_cls(**kwargs)
|
||||
|
||||
return self._model
|
||||
|
||||
@abstractmethod
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
352
tests/worker/test_workers_cn.py
Normal file
352
tests/worker/test_workers_cn.py
Normal 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}")
|
||||
352
tests/worker/test_workers_en.py
Normal file
352
tests/worker/test_workers_en.py
Normal 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}")
|
||||
Loading…
Add table
Reference in a new issue