mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
Merge branch 'master' of memoryscope
This commit is contained in:
commit
8e756fdd98
94 changed files with 941 additions and 4787 deletions
2
.flake8
2
.flake8
|
|
@ -2,7 +2,7 @@
|
|||
exclude =
|
||||
scripts/*
|
||||
src/agentscope/rpc/*
|
||||
max-line-length = 79
|
||||
max-line-length = 120
|
||||
inline-quotes = "
|
||||
avoid-escape = no
|
||||
ignore =
|
||||
|
|
|
|||
|
|
@ -6,8 +6,6 @@ memory_chat:
|
|||
class: chat.cli_memory_chat
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
human_name: 用户
|
||||
assistant_name: AI
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.service.chat_memory_service
|
||||
|
|
@ -16,25 +14,24 @@ memory_service:
|
|||
read_memory_key: read_memory
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
class: memory.operation.read_message
|
||||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: set_query_worker,[extract_time_worker|retrieve_store_worker,semantic_rank_worker],fuse_rerank_worker
|
||||
workflow: set_query_worker,retrieve_store_worker,[extract_time_worker|semantic_rank_worker],fuse_rerank_worker
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
workflow: set_query_worker,retrieve_store_worker,print_memory_worker
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: dummy_worker
|
||||
workflow: info_filter_worker,[get_observation_worker|get_observation_with_time_worker],contra_repeat_worker,store_memory_worker
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 60
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: dummy_worker
|
||||
workflow: load_memory_worker,get_reflection_subject_worker,update_insight_worker,long_contra_repeat_worker,summary_collect_worker
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 300
|
||||
worker:
|
||||
|
|
@ -45,13 +42,13 @@ worker:
|
|||
rank_model: dashscope_rank
|
||||
set_query_worker:
|
||||
class: memory.worker.read.set_query_worker
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
extract_time_worker:
|
||||
class: memory.worker.read.extract_time_worker
|
||||
generation_model_top_k: 1
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 5
|
||||
retrieve_ins_pf_top_k: 5
|
||||
semantic_rank_worker:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
fuse_rerank_worker:
|
||||
|
|
@ -60,12 +57,12 @@ worker:
|
|||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
observation: 1
|
||||
obs_customized: 1
|
||||
obs_customized: 1.2
|
||||
insight: 2.0
|
||||
profile: 2.0
|
||||
profile_customized: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
print_memory_worker:
|
||||
class: memory.worker.read.print_memory_worker
|
||||
info_filter_worker:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
|
|
@ -85,10 +82,20 @@ worker:
|
|||
generation_model_top_k: 1
|
||||
retrieve_top_k: 30
|
||||
contra_repeat_max_count: 50
|
||||
get_reflection_worker:
|
||||
class: memory.worker.summary.get_reflection_worker
|
||||
store_memory_worker:
|
||||
class: memory.worker.write.store_memory_worker
|
||||
load_memory_worker:
|
||||
class: memory.worker.summary.load_memory_worker
|
||||
get_reflection_subject_worker:
|
||||
class: memory.worker.summary.get_reflection_subject_worker
|
||||
retrieve_top_k: 100
|
||||
reflect_obs_cnt_threshold: 32
|
||||
update_insight_worker:
|
||||
class: memory.worker.summary.update_insight_worker
|
||||
long_contra_repeat_worker:
|
||||
class: memory.worker.summary.long_contra_repeat_worker
|
||||
summary_collect_worker:
|
||||
class: memory.worker.summary.summary_collect_worker
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model
|
||||
|
|
@ -102,10 +109,11 @@ models:
|
|||
class: models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
vector_store:
|
||||
class: storage.llama_index_elastic_search_store
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
use_hybrid: false
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
|
|
@ -27,8 +27,8 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
memory_service: str,
|
||||
generation_model: str,
|
||||
stream: bool = True,
|
||||
human_name: str = "",
|
||||
assistant_name: str = "",
|
||||
human_name: str = "用户",
|
||||
assistant_name: str = "AI",
|
||||
**kwargs):
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
|
|
|
|||
|
|
@ -56,9 +56,9 @@ class CliJob(object):
|
|||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
vector_store_config = self.config["vector_store"]
|
||||
embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model)
|
||||
memory_store_config = self.config["memory_store"]
|
||||
embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
|
@ -73,7 +73,7 @@ class CliJob(object):
|
|||
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
memory_chat.run()
|
||||
|
||||
G_CONTEXT.vector_store.close()
|
||||
G_CONTEXT.memory_store.close()
|
||||
G_CONTEXT.monitor.close()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,13 +0,0 @@
|
|||
from enum import Enum
|
||||
|
||||
|
||||
class MemoryMethodEnum(str, Enum):
|
||||
SUMMARY = "summary"
|
||||
|
||||
RETRIEVE = "retrieve"
|
||||
|
||||
RETRIEVE_ALL = "retrieve_all"
|
||||
|
||||
SUMMARY_SHORT = "summary_short"
|
||||
|
||||
SUMMARY_LONG = "summary_long"
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
from enum import Enum
|
||||
|
||||
|
||||
class MemoryRecallType(str, Enum):
|
||||
SIMILAR = "similar"
|
||||
|
||||
KEYWORD = "keyword"
|
||||
|
||||
PROFILE = "profile"
|
||||
|
|
@ -2,6 +2,12 @@ from enum import Enum
|
|||
|
||||
|
||||
class MemoryNodeStatus(str, Enum):
|
||||
NEW = "new"
|
||||
|
||||
MODIFIED = "modified"
|
||||
|
||||
CONTENT_MODIFIED = "content_modified"
|
||||
|
||||
ACTIVE = "active"
|
||||
|
||||
EXPIRED = "expired"
|
||||
|
|
|
|||
|
|
@ -8,8 +8,4 @@ class MemoryTypeEnum(str, Enum):
|
|||
|
||||
INSIGHT = "insight"
|
||||
|
||||
PROFILE = "profile"
|
||||
|
||||
OBS_CUSTOMIZED = "obs_customized"
|
||||
|
||||
PROFILE_CUSTOMIZED = "profile_customized"
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
def __init__(self, name: str, description: str = "", **kwargs):
|
||||
self.name: str = name
|
||||
self.description: str = description
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
def init_workflow(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -14,19 +14,17 @@ class ReadMemory(BaseWorkflow, BaseOperation):
|
|||
description: str,
|
||||
chat_messages: List[Message],
|
||||
his_msg_count: int = 0, # supplement to the current query
|
||||
contextual_msg_count: int = 0, # for the current context dialogue
|
||||
**kwargs):
|
||||
super().__init__(name=name, **kwargs)
|
||||
BaseOperation.__init__(self, name=name, description=description)
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
self.his_msg_count: int = his_msg_count
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
|
||||
def init_workflow(self):
|
||||
self.init_workers()
|
||||
|
||||
def run_operation(self, **kwargs):
|
||||
max_count = 1 + max(self.his_msg_count, self.contextual_msg_count)
|
||||
max_count = 1 + self.his_msg_count
|
||||
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
self.run_workflow()
|
||||
|
|
|
|||
21
memory_scope/memory/operation/read_message.py
Normal file
21
memory_scope/memory/operation/read_message.py
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
|
||||
class ReadMessage(BaseOperation):
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
description: str,
|
||||
chat_messages: List[Message],
|
||||
contextual_msg_count: int = 6, # for the current context dialogue
|
||||
**kwargs):
|
||||
super().__init__(name=name, description=description, **kwargs)
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
|
||||
def run_operation(self, **kwargs):
|
||||
return self.chat_messages[-self.contextual_msg_count:]
|
||||
|
|
@ -28,9 +28,14 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
self.task_list.append((fn, args, kwargs))
|
||||
|
||||
def gather_async_result(self):
|
||||
if self.is_multi_thread:
|
||||
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
|
||||
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list])
|
||||
|
||||
|
|
@ -61,12 +66,9 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
self.context[key] = value
|
||||
else:
|
||||
self.context[key] = value
|
||||
|
||||
def has_content(self, key: str):
|
||||
return key in self.context
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
|
|
|
|||
|
|
@ -4,9 +4,10 @@ from typing import List, Dict
|
|||
from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.prompt_handler import PromptHandler
|
||||
|
||||
|
|
@ -24,13 +25,15 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
self._generation_model: BaseModel | str = generation_model
|
||||
self._rank_model: BaseModel | str = rank_model
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._memory_store: BaseMemoryStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
self._user_name: str | None = None
|
||||
self._target_name: str | None = None
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
|
||||
self._contex_memory_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
@property
|
||||
def chat_messages(self) -> List[Message]:
|
||||
return self.get_context(CHAT_MESSAGES)
|
||||
|
|
@ -62,10 +65,30 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self) -> BaseVectorStore:
|
||||
if self._vector_store is None:
|
||||
self._vector_store = G_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
def memory_store(self) -> BaseMemoryStore:
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
return self._memory_store
|
||||
|
||||
def get_memories(self, key: str) -> List[MemoryNode]:
|
||||
memories: List[MemoryNode] = []
|
||||
memory_ids: List[str] = self.get_context(key)
|
||||
if memory_ids:
|
||||
memories.extend([self._contex_memory_dict[x] for x in memory_ids])
|
||||
return memories
|
||||
|
||||
def set_memories(self, key: str, nodes: List[MemoryNode] | MemoryNode):
|
||||
if not nodes:
|
||||
return
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
for node in nodes:
|
||||
if node.memory_id in self._contex_memory_dict:
|
||||
continue
|
||||
self._contex_memory_dict[node.memory_id] = node
|
||||
self.set_context(key, [n.memory_id for n in nodes])
|
||||
|
||||
@property
|
||||
def monitor(self) -> BaseMonitor:
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@ import re
|
|||
from typing import Dict
|
||||
|
||||
from memory_scope.constants.common_constants import DATATIME_KEY_MAP, QUERY_WITH_TS, EXTRACT_TIME_DICT
|
||||
from memory_scope.constants.language_constants import DATATIME_WORD_LIST
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
|
|
@ -14,23 +14,21 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
query, query_timestamp = self.get_context(QUERY_WITH_TS)
|
||||
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
contain_datetime = DatetimeHandler.has_time_word(query)
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# prepare prompt
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format)
|
||||
extract_time_prompt: str = self.prompt_handler.extract_time_prompt.format(query=query,
|
||||
query_time_str=query_time_str)
|
||||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
system_prompt = self.prompt_handler.extract_time_system
|
||||
few_shot = self.prompt_handler.extract_time_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str)
|
||||
extract_time_message = prompt_to_msg(system_prompt=system_prompt, few_shot=few_shot, user_query=user_query)
|
||||
self.logger.info(f"extract_time_message={extract_time_message}")
|
||||
|
||||
# call sft model
|
||||
response = self.generation_model.call(prompt=extract_time_prompt, top_k=self.generation_model_top_k)
|
||||
# call llm
|
||||
response = self.generation_model.call(messages=extract_time_message, top_k=self.generation_model_top_k)
|
||||
|
||||
# if empty, return
|
||||
if not response.status or not response.message.content:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
extract_time_system:
|
||||
cn: |
|
||||
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。回答的格式严格遵照示例中的已有格式规范。
|
||||
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。回答的格式严格遵照示例中的已有格式规范。若语句不涉及时间则回答无。
|
||||
|
||||
en: |
|
||||
Instructions: From the sentences and the time when they occurred, infer and extract the time periods indicated in the content of the sentences. Answer with the most complete time periods possible. The format of the answers must strictly adhere to the specifications in the examples provided. If the sentence does not involve time, respond with "none."
|
||||
|
||||
|
||||
extract_time_few_shot:
|
||||
|
|
@ -47,6 +50,60 @@ extract_time_few_shot:
|
|||
回答:
|
||||
- 年:1995 - 月:10 - 日:24
|
||||
|
||||
示例8:
|
||||
句子:我的朋友非常喜欢运动,他认为运动有助于增强身体素质。
|
||||
时间:2015年1月23日,2015年第4周,周四,7时38分0秒。
|
||||
回答:
|
||||
无
|
||||
|
||||
en: |
|
||||
Example 1:
|
||||
Sentence: I remember you went to the UAE in April the year before last. Which places in the UAE are fun? Which do you prefer, Dubai or Abu Dhabi? Are the desert views spectacular?
|
||||
Time: August 20, 1992, 34th week of 1992, Monday, 18:46:25.
|
||||
Answer:
|
||||
- Year: 1990 - Month: 4
|
||||
|
||||
Example 2:
|
||||
Sentence: Remember to attend the meeting at 3 PM the day after tomorrow. I carefully marked this important day on my calendar to remind myself not to miss it. The meeting will be held in the company conference room, and it's an important opportunity to discuss future development directions.
|
||||
Time: June 19, 2024, 25th week of 2024, Tuesday, 13:30:0.
|
||||
Answer:
|
||||
- Year: 2024 - Month: 6 - Day: 21 - Hour: 15
|
||||
|
||||
Example 3:
|
||||
Sentence: Next month on the first Saturday, let's go to Hangzhou.
|
||||
Time: July 15, 2005, 28th week of 2005, Saturday, 0:0:0.
|
||||
Answer:
|
||||
- Year: 2005 - Month: 8 - Week: 31 - Day of Week: 6
|
||||
|
||||
Example 4:
|
||||
Sentence: The small town we visited last weekend was truly beautiful.
|
||||
Time: December 2, 1999, 48th week of 1999, Tuesday, 8:40:10.
|
||||
Answer:
|
||||
- Year: 1999 - Week: 47 - Day of Week: 6, 7
|
||||
|
||||
Example 5:
|
||||
Sentence: The presentation will start in half an hour, remember to prepare the materials.
|
||||
Time: June 22, 2020, 25th week of 2020, Monday, 9:30:0
|
||||
Answer:
|
||||
- Year: 2020 - Month: 6 - Day: 22 - Hour: 10 - Minute: 0 - Second: 0
|
||||
|
||||
Example 6:
|
||||
Sentence: The start time for the 10,000-meter race was 3 minutes and 47 seconds ago.
|
||||
Time: February 17, 1987, 7th week of 1987, Wednesday, 19:54:43.
|
||||
Answer:
|
||||
- Year: 1987 - Month: 2 - Day: 17 - Hour: 19 - Minute: 50 - Second: 56
|
||||
|
||||
Example 7:
|
||||
Sentence: At this time last month, we were still preparing for the concert. Every day was busy and fulfilling, and we worked hard for the successful holding of the concert. Rehearsals, set design, and program arrangements - every detail needed careful planning and preparation.
|
||||
Time: November 24, 1995, 48th week of 1995, Tuesday, 17:45:0.
|
||||
Answer:
|
||||
- Year: 1995 - Month: 10 - Day: 24
|
||||
|
||||
Example 8:
|
||||
Sentence: My friend loves sports very much and believes that exercise helps improve physical fitness.
|
||||
Time: January 23, 2015, 4th week of 2015, Thursday, 7:38:0.
|
||||
Answer:
|
||||
None
|
||||
|
||||
|
||||
extract_time_user_query:
|
||||
|
|
@ -55,6 +112,10 @@ extract_time_user_query:
|
|||
时间:{query_time_str}
|
||||
回答:
|
||||
|
||||
en: |
|
||||
Sentence: {query}
|
||||
Time: {query_time_str}
|
||||
Answer:
|
||||
|
||||
time_string_format:
|
||||
cn: |
|
||||
|
|
|
|||
72
memory_scope/memory/worker/read/print_memory_worker.py
Normal file
72
memory_scope/memory/worker/read/print_memory_worker.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
from memory_scope.utils.timer import timer
|
||||
|
||||
|
||||
class PrintMemoryWorker(MemoryBaseWorker):
|
||||
|
||||
@timer
|
||||
def retrieve_expired_memory(self, query: str):
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.EXPIRED.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
return self.memory_store.retrieve_memories(query=query,
|
||||
top_k=self.retrieve_expired_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
expired_memories: List[MemoryNode] = self.retrieve_expired_memory(query="_")
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
|
||||
|
||||
obs_content_list: List[str] = []
|
||||
insight_content_list: List[str] = []
|
||||
expired_content_list: List[str] = []
|
||||
i = 0
|
||||
j = 0
|
||||
for node in memory_node_list:
|
||||
|
||||
if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
|
||||
i += 1
|
||||
dt_handler = DatetimeHandler(node.timestamp)
|
||||
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")
|
||||
line = f" {i} {dt} {node.content}"
|
||||
obs_content_list.append(line)
|
||||
|
||||
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.INSIGHT, ]:
|
||||
j += 1
|
||||
line = f" {j} {node.content}"
|
||||
insight_content_list.append(line)
|
||||
|
||||
for i, node in enumerate(expired_memories):
|
||||
line = f" {j} {node.content}"
|
||||
expired_content_list.append(line)
|
||||
|
||||
obs_content = "\n".join(obs_content_list)
|
||||
insight_content = "\n".join(insight_content_list)
|
||||
expired_content = "\n".join(expired_content_list)
|
||||
result: str = f"""
|
||||
The memories of {self.user_name} about {self.target_name}.
|
||||
|
||||
----- observation -----
|
||||
{obs_content}
|
||||
----- observation -----
|
||||
|
||||
----- insight -----
|
||||
{insight_content}
|
||||
----- insight -----
|
||||
|
||||
----- expired -----
|
||||
{expired_content}
|
||||
----- expired -----
|
||||
""".strip()
|
||||
self.set_context(RESULT, result)
|
||||
|
|
@ -16,26 +16,28 @@ class RetrieveStoreWorker(MemoryBaseWorker):
|
|||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
return await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_obs_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
|
||||
filter_dict = {
|
||||
"user_name": self.user_name,
|
||||
"target_name": self.target_name,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.INSIGHT.value, MemoryTypeEnum.PROFILE.value],
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
return await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
filter_dict=filter_dict)
|
||||
return await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_ins_pf_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
self.submit_async_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
|
||||
|
||||
memory_node_list: List[MemoryNode] = []
|
||||
fn_list = [self.retrieve_from_observation, self.retrieve_from_insight_and_profile]
|
||||
for result in self.async_run(fn_list=fn_list, query=query):
|
||||
for result in self.gather_async_result():
|
||||
if result:
|
||||
memory_node_list.extend(result)
|
||||
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
|
||||
|
|
@ -43,4 +45,4 @@ class RetrieveStoreWorker(MemoryBaseWorker):
|
|||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}")
|
||||
self.set_context(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from memory_scope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
|
@ -11,39 +10,27 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
# query
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
memory_node_list: List[MemoryNode] = self.get_context(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
self.logger.warning(f"retrieve memory nodes is empty!")
|
||||
return
|
||||
|
||||
# solve content repeat, insight & profile has higher priority
|
||||
memory_node_dict: Dict[str, MemoryNode] = {}
|
||||
memory_type_selected = [MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value]
|
||||
for node in memory_node_list:
|
||||
if node.memory_type in memory_type_selected:
|
||||
return
|
||||
memory_node_dict[node.content] = node
|
||||
for node in memory_node_list:
|
||||
if node.memory_type not in memory_type_selected:
|
||||
return
|
||||
memory_node_dict[node.content] = node
|
||||
# drop repeated
|
||||
memory_node_dict: Dict[str, MemoryNode] = {n.content: n for n in memory_node_list}
|
||||
memory_node_list = list(memory_node_dict.values())
|
||||
|
||||
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
|
||||
|
||||
if not response.status or not response.rank_scores:
|
||||
return
|
||||
|
||||
# set score
|
||||
rank_memory_nodes: List[MemoryNode] = []
|
||||
rank_scores = sorted(response.rank_scores.items(), key=lambda x: x[1], reverse=True)
|
||||
for idx, score in rank_scores:
|
||||
for idx, score in response.rank_scores.items():
|
||||
if idx >= len(memory_node_list):
|
||||
self.logger.warning(f"idx={idx} exceeds the maximum length of the array")
|
||||
self.logger.warning(f"idx={idx} exceeds the maximum length of rank_scores!")
|
||||
continue
|
||||
node = memory_node_list[idx]
|
||||
node.score_rank = score
|
||||
memory_node_list[idx].score_rank = score
|
||||
|
||||
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
|
||||
for node in memory_node_list:
|
||||
self.logger.info(f"rank_stage: content={node.content} score={node.score_rank}")
|
||||
self.get_context(RANKED_MEMORY_NODES, rank_memory_nodes)
|
||||
self.set_memories(RANKED_MEMORY_NODES, memory_node_list)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
get_reflect_system:
|
||||
get_reflection_subject_system:
|
||||
cn: |
|
||||
任务:从下面的信息中提取出最重要的最多{num_questions}条{user_name}属性,要求不与已有的{user_name}属性语义重复。
|
||||
要求1:{user_name}属性可以是一般的{user_name}偏好,也可以是运动偏好,旅游偏好,饮食偏好等等,也可以是重要事件性质,比如最近重要的事情,也可以是一些高度概括的人生理想,价值观,人生观,性格, 也可以是和朋友的人际关系等等。
|
||||
|
|
@ -6,7 +6,7 @@ get_reflect_system:
|
|||
输出格式:每一行输出一个{user_name}属性,每个{user_name}属性推荐4个字,如果没有信息请回答无,最多输出{num_questions}条。
|
||||
|
||||
|
||||
get_reflect_few_shot:
|
||||
get_reflection_subject_few_shot:
|
||||
cn: |
|
||||
示例1
|
||||
信息:
|
||||
|
|
@ -85,7 +85,7 @@ get_reflect_few_shot:
|
|||
新增{user_name}属性:
|
||||
无
|
||||
|
||||
get_reflect_user_query:
|
||||
get_reflection_subject_user_query:
|
||||
cn: |
|
||||
信息:
|
||||
{user_query}
|
||||
|
|
@ -11,7 +11,7 @@ from memory_scope.utils.response_text_parser import ResponseTextParser
|
|||
from memory_scope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
class GetReflectionSubjectWorker(MemoryBaseWorker):
|
||||
|
||||
def new_insight_node(self, insight_key: str) -> MemoryNode:
|
||||
dt_handler = DatetimeHandler()
|
||||
|
|
@ -22,7 +22,7 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
meta_data=meta_data,
|
||||
key=insight_key,
|
||||
memory_type=MemoryTypeEnum.INSIGHT.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value)
|
||||
status=MemoryNodeStatus.NEW.value)
|
||||
|
||||
def _run(self):
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
|
|
@ -40,10 +40,11 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
|
||||
# gen reflect prompt
|
||||
user_query_list = [n.content for n in not_reflected_nodes]
|
||||
system_prompt = self.prompt_handler.get_reflect_system.format(user_name=self.target_name,
|
||||
num_questions=self.reflect_num_questions)
|
||||
few_shot = self.prompt_handler.get_reflect_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.get_reflect_user_query.format(
|
||||
system_prompt = self.prompt_handler.get_reflection_subject_system.format(
|
||||
user_name=self.target_name,
|
||||
num_questions=self.reflect_num_questions)
|
||||
few_shot = self.prompt_handler.get_reflection_subject_few_shot.format(user_name=self.target_name)
|
||||
user_query = self.prompt_handler.get_reflection_subject_user_query.format(
|
||||
user_name=self.target_name,
|
||||
exist_keys=self.get_language_value(COMMA_WORD).join(exist_keys),
|
||||
user_query="\n".join(user_query_list))
|
||||
|
|
@ -59,7 +60,7 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# parse text & save
|
||||
new_insight_keys = ResponseTextParser(response.message.content).parse_v2("get_reflection")
|
||||
new_insight_keys = ResponseTextParser(response.message.content).parse_v2(self.__class__.__name__)
|
||||
if new_insight_keys:
|
||||
for insight_key in new_insight_keys:
|
||||
insight_nodes.append(self.new_insight_node(insight_key))
|
||||
|
|
@ -19,9 +19,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_reflected": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
|
|
@ -33,9 +33,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"obs_updated": False,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
|
|
@ -46,9 +46,9 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
}
|
||||
nodes: List[MemoryNode] = await self.vector_store.async_retrieve(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
|
||||
top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.set_context(INSIGHT_NODES, nodes)
|
||||
|
||||
async def _run(self):
|
||||
|
|
|
|||
|
|
@ -1,11 +1,18 @@
|
|||
long_contra_repeat_system:
|
||||
cn: |
|
||||
对下面的{num_obs}句句子,逐一判断是否与“前面序号”的任意句子存在信息的矛盾,或者句子的主要信息被“前面序号”的任意句子中的信息包含。只判断与“前面序号”的句子的关系。
|
||||
如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
|
||||
对每个句子都做一个判断,最后一共输出{num_obs}条判断。如果句子与前面序号的句子存在矛盾,则以前面序号的句子中的信息为准,修改句子中矛盾的部分。
|
||||
请一步步思考,并按如下格式输出:
|
||||
思考:思考的依据和过程,30字以内。
|
||||
判断:<句子序号> <矛盾,被包含,无> <修改后的内容>,一定加<>
|
||||
|
||||
en: |
|
||||
For the following {num_obs} sentences, determine one by one if there is any contradiction with the information in any "previously numbered" sentences, or if the main information of the sentence is contained within the information of any "previously numbered" sentences. Only evaluate the relation to "previously numbered" sentences.
|
||||
Make an evaluation for each sentence, resulting in a total of {num_obs} evaluations. If the sentence conflicts with a "previously numbered" sentence, the information in the earlier sentence takes precedence, and the conflicting part of the current sentence should be modified accordingly.
|
||||
Please think step by step and output in the following format:
|
||||
Thought: The basis and process of the thought, within 30 words.
|
||||
Evaluation: <sentence number> <Contradiction, Contained or None>, <Revised content>, enclosed in <>.
|
||||
|
||||
|
||||
long_contra_repeat_few_shot:
|
||||
cn: |
|
||||
|
|
@ -50,6 +57,47 @@ long_contra_repeat_few_shot:
|
|||
思考:第6句中所有信息都被前面序号中第5句的信息完全包含。
|
||||
判断:<2> <被包含> <>
|
||||
|
||||
en: |
|
||||
Example 1
|
||||
Sentences:
|
||||
1 {user_name} suffers from insomnia frequently and is interested in the effects of sleeping pills, suggesting a possible consideration of their use.
|
||||
2 {user_name} suffers from insomnia frequently and seeks remedies.
|
||||
3 Charles is {user_name}'s supervisor.
|
||||
4 Charles is {user_name}'s supervisor.
|
||||
5 Charles is {user_name}'s supervisor and the branch manager of a bank.
|
||||
Thought: The first sentence does not have any contradictions or complete repetitions with the previously numbered sentences.
|
||||
Evaluation: <1> <None> <>
|
||||
Thought: All information in the second sentence is completely contained within the information of the first sentence.
|
||||
Evaluation: <2> <Contained> <>
|
||||
Thought: The information in the third sentence does not appear in the previously numbered sentences.
|
||||
Evaluation: <3> <None> <>
|
||||
Thought: The fourth sentence is completely repetitive of the information in the third sentence, i.e., it is completely contained.
|
||||
Evaluation: <4> <Contained> <>
|
||||
Thought: The information that Charles is {user_name}'s supervisor in the fifth sentence is contained within the information of the third sentence, but the new information that Charles is the branch manager of a bank is not, so it is not contained.
|
||||
Evaluation: <5> <None> <>
|
||||
|
||||
Example 2
|
||||
Sentences:
|
||||
1 {user_name}'s child does not perform well academically.
|
||||
2 {user_name}'s child often skips school.
|
||||
3 {user_name}'s father's birthday is on June 2, 2024, and {user_name} plans to prepare a gift.
|
||||
4 {user_name}'s father's birthday is on May 1, 2024.
|
||||
5 {user_name} loves playing basketball with classmates.
|
||||
6 {user_name} likes playing basketball.
|
||||
|
||||
Thought: The first sentence does not have any contradictions or complete repetitions with the previously numbered sentences.
|
||||
Evaluation: <1> <None> <>
|
||||
Thought: The second sentence neither contradicts nor repeats any of the previously numbered sentences.
|
||||
Evaluation: <2> <None> <>
|
||||
Thought: The third sentence neither contradicts nor repeats any of the previously numbered sentences.
|
||||
Evaluation: <3> <None> <>
|
||||
Thought: The date of {user_name}'s father's birthday in the fourth sentence contradicts the information in the third sentence.
|
||||
Evaluation: <4> <Contradiction> <{user_name}'s father's birthday is on June 2, 2024.>
|
||||
Thought: The fifth sentence neither contradicts nor repeats any of the previously numbered sentences.
|
||||
Evaluation: <5> <None> <>
|
||||
Thought: All information in the sixth sentence is completely contained within the information of the fifth sentence.
|
||||
Evaluation: <6> <Contained> <>
|
||||
|
||||
long_contra_repeat_user_query:
|
||||
cn: |
|
||||
句子:
|
||||
|
|
|
|||
|
|
@ -21,9 +21,9 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value]
|
||||
}
|
||||
retrieve_nodes = await self.vector_store.async_retrieve(query=node.content,
|
||||
top_k=self.long_contra_repeat_top_k,
|
||||
filter_dict=filter_dict)
|
||||
retrieve_nodes = await self.memory_store.a_retrieve_memories(query=node.content,
|
||||
top_k=self.long_contra_repeat_top_k,
|
||||
filter_dict=filter_dict)
|
||||
return node, [n for n in retrieve_nodes if n.score_similar >= self.long_contra_repeat_threshold]
|
||||
|
||||
def _run(self):
|
||||
|
|
|
|||
|
|
@ -1,46 +1,40 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from memory_scope.constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
NEW,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
)
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \
|
||||
NOT_UPDATED_NODES
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
update_memories: Dict[str, MemoryNode] = {}
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update(
|
||||
{n.id: n for n in insight_nodes if n.obs_updated}
|
||||
)
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.obs_updated = "0"
|
||||
all_node_dict.update({n.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
update_memories.update({n.memory_id: n for n in insight_nodes})
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
if not_reflected_nodes:
|
||||
update_memories.update({n.memory_id: n for n in not_reflected_nodes})
|
||||
|
||||
not_updated_nodes: List[MemoryNode] = self.get_context(NOT_UPDATED_NODES)
|
||||
if not_updated_nodes:
|
||||
for node in not_updated_nodes:
|
||||
if node.memory_id in update_memories:
|
||||
|
||||
|
||||
|
||||
|
||||
keys = [
|
||||
INSIGHT_NODES,
|
||||
MERGE_OBS_NODES,
|
||||
NOT_UPDATED_NODES,
|
||||
NOT_REFLECTED_NODES,
|
||||
]
|
||||
|
||||
memory_nodes: List[MemoryNode] = []
|
||||
for key in keys:
|
||||
memory_nodes.extend(self.get_context(key))
|
||||
|
||||
self.memory_store.update_memories(update_memories)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ from typing import List
|
|||
|
||||
from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES
|
||||
from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
|
|
@ -48,6 +49,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
|
||||
insight_node.timestamp = dt_handler.timestamp
|
||||
insight_node.dt = dt_handler.datetime_format()
|
||||
if insight_node.status == MemoryNodeStatus.ACTIVE.value:
|
||||
insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
|
||||
self.logger.info(f"after_update_{insight_node.key} value={insight_value}")
|
||||
return insight_node
|
||||
|
||||
|
|
@ -89,6 +92,10 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"update_{insight_node.key} insight_value={insight_value} is invalid.")
|
||||
return insight_node
|
||||
|
||||
if insight_node.value == insight_value:
|
||||
self.logger.info(f"value={insight_value} is same!")
|
||||
return insight_node
|
||||
|
||||
self.update_insight_node(insight_node=insight_node, insight_value=insight_value)
|
||||
return insight_node
|
||||
|
||||
|
|
@ -102,7 +109,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
for node in insight_nodes:
|
||||
if node.content:
|
||||
if node.status == MemoryNodeStatus.ACTIVE.value:
|
||||
self.submit_async_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_updated_nodes)
|
||||
|
|
@ -126,3 +133,6 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# get result
|
||||
self.gather_async_result()
|
||||
|
||||
for node in not_updated_nodes:
|
||||
node.obs_updated = True
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
"dt": dt_handler.datetime_format(),
|
||||
}
|
||||
return self.vector_store.retrieve(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
|
||||
return self.memory_store.retrieve_memories(query=message.content, top_k=self.today_obs_top_k, filter_dict=filter_dict)
|
||||
|
||||
def _run(self):
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ contra_repeat_system:
|
|||
Make an evaluation for each sentence, resulting in a total of {num_obs} evaluations.
|
||||
Please think step by step and output in the following format:
|
||||
Thought: The basis and process of the thought, within 30 words.
|
||||
Evaluation: <Contradiction, Contained, None>, enclosed in <>.
|
||||
Evaluation: <sentence number> <Contradiction, Contained or None>, enclosed in <>.
|
||||
|
||||
|
||||
contra_repeat_few_shot:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import NEW_OBS_WITH_TIME_NODES
|
||||
from memory_scope.constants.language_constants import DATATIME_WORD_LIST, COLON_WORD
|
||||
from memory_scope.constants.language_constants import COLON_WORD
|
||||
from memory_scope.memory.worker.write.get_observation_worker import GetObservationWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.scheme.message import Message
|
||||
|
|
@ -16,12 +16,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.chat_messages:
|
||||
match = False
|
||||
for time_keyword in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
if DatetimeHandler.has_time_word(query=msg.content):
|
||||
dt_handler = DatetimeHandler(dt=msg.time_created)
|
||||
dt = dt_handler.string_format(self.prompt_handler.time_string_format)
|
||||
user_query_list.append(f"{i} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ time_string_format:
|
|||
|
||||
get_observation_with_time_system:
|
||||
cn: |
|
||||
任务:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。
|
||||
任务指令:从下面的{num_obs}句{user_name}句子中依次提取出关于{user_name}的重要信息,相应的关键词与时间信息。
|
||||
每一句{user_name}句子的格式是:<序号> <对话时间> {user_name}:<句子>
|
||||
对每一句句子,只提取非常明确的信息和进行非常确定的推断,不要进行任何猜测。不要提取重复的信息,如果句子中的所有信息与已经提取出的信息重复了则回答“重复“,如果没有重要信息则回答“无”。
|
||||
如果{user_name}信息涉及时间,则结合对话时间推断{user_name}信息的时间信息,没有则不输出。注意区分,对于句子中包含{user_name}假设的信息或者{user_name}虚构的内容比如{user_name}创作的小说或剧本,不要提取信息。
|
||||
|
|
@ -13,6 +13,16 @@ get_observation_with_time_system:
|
|||
思考:思考的依据和过程,50字以内。
|
||||
信息:<句子序号> <时间信息或“无”> <明确的重要信息或“重复”或”无“> <关键词>
|
||||
|
||||
en: |
|
||||
Instruction: Extract important information about {user_name}, corresponding keywords, and time information from the following {num_obs} sentences by {user_name}, one by one.
|
||||
Each sentence from {user_name} is formatted as: <sentence number> <conversation time> {user_name}: <sentence>
|
||||
For each sentence, only extract very clear information and make very certain inferences without any speculation. Do not extract repeated information. If all information in the sentence repeats what has already been extracted, respond with "repeat." If there is no important information, respond with "none."
|
||||
If the information about {user_name} involves time, combine it with the conversation time to infer the time information of {user_name}'s information; if not, do not output time information. Be sure to distinguish information; for example, do not extract hypothetical or fictional content from {user_name} such as {user_name}'s novels or scripts.
|
||||
Perform information extraction for each sentence, resulting in a total of {num_obs} pieces of information.
|
||||
Please think step by step, and be sure to output in the following format, with the final output enclosed in <>:
|
||||
Thought: The basis and process of the thought, within 50 words.
|
||||
Information: <sentence number> <Time information or "none"> <Clear important information or "repeat" or "none"> <keywords>
|
||||
|
||||
|
||||
get_observation_with_time_few_shot:
|
||||
cn: |
|
||||
|
|
@ -62,7 +72,6 @@ get_observation_with_time_few_shot:
|
|||
4 2023年5月21日周六14点 {user_name}:有人说兴趣是最好的老师,也建议兴趣和职业联系起来,但我发现喜欢打篮球的人很多,但靠打篮球成职业的稀少,赚钱的更少,此外,怎么分辨兴趣和喜欢
|
||||
5 2018年3月6日周四19点 {user_name}:李增杰:这个是星座蛙设,但是我是处女座的,我妈感觉因为我的不正常,我妈不让我看了\n雌猴摸了摸李增杰的头,这样啊\n雌猴打开了哔哩哔哩看了看\n雌猴:要不换个设吧,我听你未来的你说,有一个叫难忘的朱古力232这个人,他弄的设是Windows设\n这是剧本1,剧本2未完待续
|
||||
|
||||
|
||||
思考:从第1句可以得知{user_name}和家人上个月去杭州旅游了,这是关于{user_name}的经历的重要信息。其余信息重要性不足。{user_name}信息涉及时间,结合对话时间为2023年6月推断{user_name}和家人2023年5月去杭州旅游了。
|
||||
信息:<1> <2023年5月> <{user_name}和家人2023年5月去杭州旅游了。> <家人, 杭州, 旅游>
|
||||
思考:从第2句可以得知{user_name}的生日是昨天,这是关于{user_name}重要纪念日的信息。其余信息重要性不足。{user_name}信息涉及时间,结合对话时间为2023年7月2日,
|
||||
|
|
@ -75,9 +84,70 @@ get_observation_with_time_few_shot:
|
|||
思考:第5句是{user_name}创作的剧本内容,无法提取{user_name}个人信息。
|
||||
信息:<5> <> <无> <>
|
||||
|
||||
en: |
|
||||
Example 1:
|
||||
{user_name} sentences:
|
||||
1 May 1, 2022, Tuesday, 3 PM {user_name}: Please help me write a birthday greeting for my colleague Jason's daughter who is turning three.
|
||||
2 May 2, 2022, Tuesday, 5 PM {user_name}: Chronology of major events in Chinese history from 1400 to 1550 AD.
|
||||
3 May 3, 2022, Tuesday, 6 PM {user_name}: Can you compile a list of tips on how to use large models for me, and try to keep the content concise?
|
||||
4 July 3, 2022, Thursday, 12 PM {user_name}: I got a swimming pass two months ago.
|
||||
|
||||
Thought: From the first sentence, it can be inferred that Jason is {user_name}'s colleague, which is important information about {user_name}'s interpersonal relationships. The remaining information is of insufficient importance. {user_name}'s information does not involve time.
|
||||
Information: <1> <> <Zhang San is {user_name}'s colleague> <Zhang San, colleague>
|
||||
Thought: The second sentence is a request made by {user_name}, with no clear mention of {user_name}'s personal information.
|
||||
Information: <2> <> <none> <>
|
||||
Thought: The third sentence is a request made by {user_name}, with no clear mention of {user_name}'s personal information.
|
||||
Information: <3> <> <none> <>
|
||||
Thought: From the fourth sentence, it can be inferred that {user_name} got a swimming pass two months ago. {user_name}'s information involves time. Combining it with the conversation time of July 2022, it can be inferred that {user_name} got the swimming pass in May 2022.
|
||||
Information: <4> <May 2022> <{user_name} got a swimming pass in May 2022> <swimming pass>
|
||||
|
||||
|
||||
Example 2:
|
||||
{user_name} sentences:
|
||||
1 January 4, 2020, Sunday, 10 AM {user_name}: I spent $5000 to buy 100 shares of General Motors.
|
||||
2 April 27, 2023, Friday, 8 AM {user_name}: Tomorrow is my wedding anniversary with my wife. Could you recommend a restaurant?
|
||||
3 January 4, 2020, Sunday, 10 AM {user_name}: I spent $5000 to buy 100 shares of General Motors.
|
||||
4 June 2, 2021, Thursday, 11 PM {user_name}: Thanks. I'm having lunch near the company at noon; can you recommend a restaurant near Alibaba Xuhui Riverside Campus for me?
|
||||
5 July 9, 2021, Saturday, 11 AM {user_name}: Two pieces of bad news: I broke my badminton racket while playing... Then I went to my friend's house to pet the cat and ended up having an allergic reaction to the cat fur, sneezing like crazy today...
|
||||
|
||||
Thought: From the first sentence, it can be inferred that {user_name} bought 100 shares of General Motors stock for $5000. This is important information about {user_name}'s investment decision. {user_name}'s information does not involve time.
|
||||
Information: <1> <> <{user_name} bought 100 shares of General Motors stock for $5000> <General Motors, stock>
|
||||
Thought: From the second sentence, it can be inferred that {user_name}'s wedding anniversary with his wife is tomorrow, which is important information about {user_name}'s significant dates. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation date of April 27, 2023, and knowing that the anniversary is a recurring date, it can be inferred that {user_name}'s wedding anniversary is on April 28th each year.
|
||||
Information: <2> <April 28 each year> <{user_name}'s wedding anniversary with his wife is on April 28 each year> <wife, wedding anniversary>
|
||||
Thought: The information in the third sentence is a repeat of the first sentence.
|
||||
Information: <3> <> <repeat> <>
|
||||
Thought: From the fourth sentence, it can be inferred that {user_name} works at Alibaba Xuhui Riverside Campus, which is important information about {user_name}'s job. The remaining information is of insufficient importance. {user_name}'s information does not involve time.
|
||||
Information: <4> <> <{user_name} works at Alibaba Xuhui Riverside Campus> <Alibaba, Xuhui Riverside Campus, job>
|
||||
Thought: From the fifth sentence, it can be inferred that {user_name} broke their badminton racket the other day while playing, but this is not important information. It can also be inferred that {user_name} is allergic to cat fur, which is important information about {user_name}'s health. {user_name}'s information does not involve time.
|
||||
Information: <5> <> <{user_name} is allergic to cat fur> <cat fur, allergy>
|
||||
|
||||
|
||||
Example 3:
|
||||
{user_name} sentences:
|
||||
1 June 30, 2023, Friday, 3 PM {user_name}: Last month, my family and I went to San Jose for a trip. The scenery was very nice.
|
||||
2 July 2, 2023, Tuesday, 10 AM {user_name}: Yesterday was my birthday. I spent it alone.
|
||||
3 July 3, 2020, Thursday, 11 AM {user_name}: Remind me to go for a medical check-up next Monday.
|
||||
4 May 21, 2023, Saturday, 2 PM {user_name}: Someone said that passion is the best teacher and suggested linking passion with a career, but I found that many people like playing basketball, but few make a career out of it, and even fewer make money from it. Also, how do you distinguish passion from liking?
|
||||
5 March 6, 2018, Thursday, 7 PM {user_name}: Zack: This is a constellation frog setting, but I am a Virgo. My mom feels I am abnormal and doesn't let me watch it. \n The female monkey patted Zack's head, "Is that so?" \n The female monkey opened Bilibili and took a look. \n Female monkey: "Why don't you switch the setting? I heard from your future self that there's someone called 'Unforgettable Chocolate 232' who created a Windows setting." \n This is script 1; script 2 is to be continued.
|
||||
|
||||
Thought: From the first sentence, it can be inferred that {user_name} and their family went to San Jose for a trip last month. This is important information about {user_name}'s experience. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation time of June 2023, it can be inferred that {user_name} and their family went to San Jose for a trip in May 2023.
|
||||
Information: <1> <May 2023> <{user_name} and their family went to San Jose for a trip in May 2023> <family, San Jose, trip>
|
||||
Thought: From the second sentence, it can be inferred that {user_name}'s birthday was yesterday. This is important information about {user_name}'s significant dates. The remaining information is of insufficient importance. {user_name}'s information involves time. Combining it with the conversation time of July 2, 2023, and knowing that the birthday is a recurring date, it can be inferred that {user_name}'s birthday is on July 2 each year.
|
||||
Information: <2> <July 2 each year> <{user_name}'s birthday is on July 2 each year> <birthday>
|
||||
Thought: From the third sentence, it can be inferred that {user_name} will go for a medical check-up next Monday, which is an important reminder for {user_name}. {user_name}'s information involves time. Combining it with the conversation time of July 3, 2020, Thursday, it can be inferred that {user_name} will go for a check-up on July 6, 2020, Monday.
|
||||
Information: <3> <July 6, 2020, Monday> <{user_name} will go for a medical check-up on July 6, 2020, Monday> <medical check-up>
|
||||
Thought: The fourth sentence is a discussion and query about other people's opinions by {user_name}, with no clear mention of {user_name}'s personal information.
|
||||
Information: <4> <> <none> <>
|
||||
Thought: The fifth sentence is content from a script written by {user_name}, with no extractable personal information about {user_name}.
|
||||
Information: <5> <> <none> <>
|
||||
|
||||
|
||||
get_observation_with_time_user_query:
|
||||
cn: |
|
||||
{user_name}句子:
|
||||
{user_query}
|
||||
|
||||
en: |
|
||||
{user_name} sentences:
|
||||
{user_query}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
|
||||
from memory_scope.constants.language_constants import DATATIME_WORD_LIST, REPEATED_WORD, NONE_WORD, COLON_WORD
|
||||
from memory_scope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
|
@ -33,7 +33,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
meta_data=meta_data,
|
||||
content=obs_content,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
status=MemoryNodeStatus.NEW.value,
|
||||
timestamp=message.time_created,
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
|
|
@ -43,12 +43,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.chat_messages:
|
||||
match = False
|
||||
for time_keyword in self.get_language_value(DATATIME_WORD_LIST):
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
if not DatetimeHandler.has_time_word(query=msg.content):
|
||||
user_query_list.append(f"{i} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
|
||||
i += 1
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ get_observation_system:
|
|||
Perform information extraction for each sentence, resulting in a total of {num_obs} pieces of information.
|
||||
Please output the results in the following format, with the final output enclosed in <>:
|
||||
Thought: The basis and process of the thought, within 50 words.
|
||||
Information: <> <Clear important information or “Repeat” or “None”> <keywords>
|
||||
Information: <sentence number> <> <Clear important information or “Repeat” or “None”> <keywords>
|
||||
|
||||
|
||||
get_observation_few_shot:
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
if self.has_content(store_key):
|
||||
memory_nodes: List[MemoryNode] = self.get_context(store_key)
|
||||
self.vector_store.update_batch(memory_nodes)
|
||||
self.memory_store.update_memories(memory_nodes)
|
||||
|
||||
elif store_key in self.chat_kwargs:
|
||||
query = self.chat_kwargs[store_key]
|
||||
|
|
@ -27,8 +27,8 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
status=MemoryNodeStatus.NEW.value,
|
||||
timestamp=dt_handler.timestamp,
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
self.vector_store.update(node)
|
||||
self.memory_store.update_memories(node)
|
||||
|
|
|
|||
|
|
@ -1,13 +1,12 @@
|
|||
import datetime
|
||||
from typing import Dict, List
|
||||
from uuid import uuid4
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
from memory_scope.utils.tool_functions import md5_hash
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
memory_id: str = Field("", description="unique id for memory")
|
||||
memory_id: str = Field(uuid4(), description="unique id for memory")
|
||||
|
||||
user_name: str = Field("", description="the user who owns the memory")
|
||||
|
||||
|
|
@ -43,7 +42,6 @@ class MemoryNode(BaseModel):
|
|||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_id = f"{self.user_name}_{self.target_name}_{self.timestamp}_{md5_hash(self.content)[:8]}"
|
||||
self.dt = datetime.datetime.fromtimestamp(self.timestamp).strftime("%Y%m%d")
|
||||
|
||||
@property
|
||||
|
|
@ -52,4 +50,3 @@ class MemoryNode(BaseModel):
|
|||
|
||||
def __getitem__(self, key: str):
|
||||
return self.model_dump().get(key)
|
||||
|
||||
|
|
|
|||
33
memory_scope/storage/base_memory_store.py
Normal file
33
memory_scope/storage/base_memory_store.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class BaseMemoryStore(metaclass=ABCMeta):
|
||||
|
||||
@abstractmethod
|
||||
def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
|
||||
"""
|
||||
status:
|
||||
1. new: emb & insert
|
||||
2. modified: update
|
||||
3. content_modified: emb & update
|
||||
4. active: do nothing
|
||||
5. expired: update
|
||||
"""
|
||||
|
||||
def flush(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def close(self):
|
||||
pass
|
||||
|
|
@ -1,37 +0,0 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class BaseVectorStore(metaclass=ABCMeta):
|
||||
|
||||
@abstractmethod
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def insert(self, node: MemoryNode):
|
||||
pass
|
||||
|
||||
def insert_batch(self, nodes: List[MemoryNode]):
|
||||
pass
|
||||
|
||||
def delete(self, node: MemoryNode):
|
||||
pass
|
||||
|
||||
def update(self, node: MemoryNode):
|
||||
pass
|
||||
|
||||
def update_batch(self, nodes: List[MemoryNode]):
|
||||
pass
|
||||
|
||||
def flush(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
24
memory_scope/storage/dummy_memory_store.py
Normal file
24
memory_scope/storage/dummy_memory_store.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
|
||||
|
||||
class DummyMemoryStore(BaseMemoryStore):
|
||||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
async def a_retrieve_memories(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
|
@ -1,21 +0,0 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class DummyVectorStore(BaseVectorStore):
|
||||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
def insert(self, node: MemoryNode):
|
||||
pass
|
||||
|
|
@ -1,143 +0,0 @@
|
|||
from typing import Dict, List, Any
|
||||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class _ElasticsearchStore(ElasticsearchStore):
|
||||
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
|
||||
"""
|
||||
Async delete node from Elasticsearch index.
|
||||
|
||||
Args:
|
||||
ref_doc_id: ID of the node to delete.
|
||||
delete_kwargs: Optional. Additional arguments to
|
||||
pass to AsyncElasticsearch delete_by_query.
|
||||
|
||||
Raises:
|
||||
Exception: If AsyncElasticsearch delete_by_query fails.
|
||||
"""
|
||||
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
|
||||
|
||||
|
||||
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert standard filters to Elasticsearch filter.
|
||||
|
||||
Args:
|
||||
standard_filters: Standard Llama-index filters.
|
||||
|
||||
Returns:
|
||||
Elasticsearch filter.
|
||||
"""
|
||||
|
||||
result = {
|
||||
"bool": {}
|
||||
}
|
||||
for key, value in standard_filters.items():
|
||||
if isinstance(value, list):
|
||||
operands = []
|
||||
for v in value:
|
||||
operands.append(
|
||||
{
|
||||
"term":
|
||||
{
|
||||
f"metadata.{key}.keyword": {"value": v}
|
||||
}
|
||||
}
|
||||
)
|
||||
result['bool'].update({"should": operands})
|
||||
result['bool'].update({"minimum_should_match": 1})
|
||||
else:
|
||||
operand = [{
|
||||
"term": {
|
||||
f"metadata.{key}.keyword": {
|
||||
"value": value,
|
||||
}
|
||||
}
|
||||
}]
|
||||
if "must" in result['bool']:
|
||||
result['bool']['must'].extend(operand)
|
||||
else:
|
||||
result['bool'].update({"must": operand})
|
||||
return result
|
||||
|
||||
|
||||
class LlamaIndexElasticSearchStore(BaseVectorStore):
|
||||
def __init__(self,
|
||||
embedding_model: BaseModel,
|
||||
index_name: str,
|
||||
es_url: str,
|
||||
use_hybrid: bool = True,
|
||||
**kwargs):
|
||||
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = _ElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
**kwargs)
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k)
|
||||
text_nodes = retriever.retrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
async def async_retrieve(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}")
|
||||
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(
|
||||
vector_store_kwargs={"es_filter": es_filter},
|
||||
similarity_top_k=top_k)
|
||||
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
def insert(self, node: MemoryNode):
|
||||
self.index.insert_nodes([self._memory_node_2_text_node(node)])
|
||||
|
||||
def delete(self, node: MemoryNode):
|
||||
memory_id = node.memory_id
|
||||
return self.es_store.delete(memory_id)
|
||||
|
||||
def update(self, node: MemoryNode):
|
||||
self.delete(node)
|
||||
self.insert(node)
|
||||
|
||||
def update_batch(self, nodes: List[MemoryNode]):
|
||||
for node in nodes:
|
||||
self.update(node)
|
||||
|
||||
def close(self):
|
||||
self.es_store.close()
|
||||
|
||||
@staticmethod
|
||||
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
|
||||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
metadata=memory_node.model_dump(exclude={"content"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
return MemoryNode(content=text_node.text, **text_node.metadata)
|
||||
252
memory_scope/storage/llama_index_es_memory_store.py
Normal file
252
memory_scope/storage/llama_index_es_memory_store.py
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
from typing import Dict, List, Any, Optional, cast
|
||||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
|
||||
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
|
||||
def _hybrid(
|
||||
self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int,
|
||||
) -> Dict[str, Any]:
|
||||
# Add a query to the knn query.
|
||||
# RRF is used to even the score from the knn query and text query
|
||||
# RRF has two optional parameters: {'rank_constant':int, 'window_size':int}
|
||||
# https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html
|
||||
query_body = {
|
||||
"knn": knn,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": [
|
||||
{
|
||||
"match": {
|
||||
self.text_field: {
|
||||
"query": query,
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"filter": filter,
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
if isinstance(self.rrf, Dict):
|
||||
query_body["rank"] = {"rrf": self.rrf}
|
||||
elif isinstance(self.rrf, bool) and self.rrf is True:
|
||||
query_body["rank"] = {"rrf": {"window_size": top_k}}
|
||||
return query_body
|
||||
|
||||
def es_query(
|
||||
self,
|
||||
*,
|
||||
query: Optional[str],
|
||||
query_vector: Optional[List[float]],
|
||||
text_field: str,
|
||||
vector_field: str,
|
||||
k: int,
|
||||
num_candidates: int,
|
||||
filter: List[Dict[str, Any]] = [],
|
||||
) -> Dict[str, Any]:
|
||||
knn = {
|
||||
"filter": filter,
|
||||
"field": vector_field,
|
||||
"k": k,
|
||||
"num_candidates": num_candidates,
|
||||
}
|
||||
|
||||
if query_vector is not None:
|
||||
knn["query_vector"] = query_vector
|
||||
else:
|
||||
# Inference in Elasticsearch. When initializing we make sure to always have
|
||||
# a model_id if don't have an embedding_service.
|
||||
knn["query_vector_builder"] = {
|
||||
"text_embedding": {
|
||||
"model_id": self.model_id,
|
||||
"model_text": query,
|
||||
}
|
||||
}
|
||||
|
||||
if self.hybrid:
|
||||
return self._hybrid(query=cast(str, query), knn=knn, filter=filter, top_k=k)
|
||||
|
||||
return {"knn": knn}
|
||||
|
||||
class _ElasticsearchStore(ElasticsearchStore):
|
||||
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
|
||||
"""
|
||||
Async delete node from Elasticsearch index.
|
||||
|
||||
Args:
|
||||
ref_doc_id: ID of the node to delete.
|
||||
delete_kwargs: Optional. Additional arguments to
|
||||
pass to AsyncElasticsearch delete_by_query.
|
||||
|
||||
Raises:
|
||||
Exception: If AsyncElasticsearch delete_by_query fails.
|
||||
"""
|
||||
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
|
||||
|
||||
|
||||
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert standard filters to Elasticsearch filter.
|
||||
|
||||
Args:
|
||||
standard_filters: Standard Llama-index filters.
|
||||
|
||||
Returns:
|
||||
Elasticsearch filter.
|
||||
"""
|
||||
|
||||
result = {
|
||||
"bool": {}
|
||||
}
|
||||
for key, value in standard_filters.items():
|
||||
if isinstance(value, list):
|
||||
operands = []
|
||||
for v in value:
|
||||
operands.append(
|
||||
{
|
||||
"term":
|
||||
{
|
||||
f"metadata.{key}.keyword": {"value": v}
|
||||
}
|
||||
}
|
||||
)
|
||||
result['bool'].update({"should": operands})
|
||||
result['bool'].update({"minimum_should_match": 1})
|
||||
else:
|
||||
operand = [{
|
||||
"term": {
|
||||
f"metadata.{key}.keyword": {
|
||||
"value": value,
|
||||
}
|
||||
}
|
||||
}]
|
||||
if "must" in result['bool']:
|
||||
result['bool']['must'].extend(operand)
|
||||
else:
|
||||
result['bool'].update({"must": operand})
|
||||
return result
|
||||
|
||||
|
||||
class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
||||
def __init__(self,
|
||||
embedding_model: BaseModel,
|
||||
index_name: str,
|
||||
es_url: str,
|
||||
use_hybrid: bool = True,
|
||||
**kwargs):
|
||||
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.es_store = _ElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
**kwargs)
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
embed_model=self.embedding_model.model)
|
||||
self.index.build_index_from_nodes([TextNode(text="text")])
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter}, similarity_top_k=top_k, sparse_top_k=top_k)
|
||||
text_nodes = retriever.retrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
async def a_retrieve_memories(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
self.logger.info(f"query={query} top_k={top_k} filter_dict={filter_dict}")
|
||||
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(
|
||||
vector_store_kwargs={"es_filter": es_filter},
|
||||
similarity_top_k=top_k)
|
||||
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
def insert(self, node: MemoryNode):
|
||||
self.index.insert_nodes([self._memory_node_2_text_node(node)])
|
||||
|
||||
def delete(self, node: MemoryNode):
|
||||
memory_id = node.memory_id
|
||||
return self.es_store.delete(memory_id)
|
||||
|
||||
def update(self, node: MemoryNode):
|
||||
self.delete(node)
|
||||
self.insert(node)
|
||||
|
||||
def update_batch(self, nodes: List[MemoryNode]):
|
||||
for node in nodes:
|
||||
self.update(node)
|
||||
|
||||
def close(self):
|
||||
self.es_store.close()
|
||||
|
||||
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
|
||||
if not nodes:
|
||||
self.logger.warning("empty nodes!")
|
||||
return
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
# emb & insert new memories
|
||||
# TODO batch insert
|
||||
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW]
|
||||
if new_memories:
|
||||
for n in new_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
self.insert(n)
|
||||
|
||||
# emb & update new memories
|
||||
# TODO insert overwrite
|
||||
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED]
|
||||
if c_modified_memories:
|
||||
for n in c_modified_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
self.delete(n)
|
||||
self.insert(n)
|
||||
|
||||
# update new memories
|
||||
# TODO no emb
|
||||
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED]
|
||||
if modified_memories:
|
||||
for n in modified_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
self.delete(n)
|
||||
self.insert(n)
|
||||
|
||||
# set memories expired
|
||||
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED]
|
||||
if expired_memories:
|
||||
for n in expired_memories:
|
||||
n.status = MemoryNodeStatus.ACTIVE.value
|
||||
self.delete(n)
|
||||
self.insert(n)
|
||||
|
||||
@staticmethod
|
||||
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
|
||||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
metadata=memory_node.model_dump(exclude={"content"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
return MemoryNode(content=text_node.text, **text_node.metadata)
|
||||
|
|
@ -2,7 +2,7 @@ import datetime
|
|||
import re
|
||||
from typing import Dict
|
||||
|
||||
from memory_scope.constants.language_constants import WEEKDAYS
|
||||
from memory_scope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
|
@ -65,6 +65,82 @@ class DatetimeHandler(object):
|
|||
extracted_data[key] = int(match.group(1))
|
||||
return extracted_data
|
||||
|
||||
@classmethod
|
||||
def extract_date_parts_en(cls, input_string: str) -> dict:
|
||||
date_info = {
|
||||
"year": -1,
|
||||
"month": -1,
|
||||
"day": -1,
|
||||
"hour": -1,
|
||||
"minute": -1,
|
||||
"second": -1,
|
||||
"weekday": -1
|
||||
}
|
||||
|
||||
# Patterns to extract the parts of the date/time
|
||||
patterns = {
|
||||
"year": r"\b(\d{4})\b",
|
||||
"month": r"\b(January|February|March|April|May|June|July|August|September|October|November|December)\b",
|
||||
"day_month_year": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October"
|
||||
r"|November|December) (?P<day>\d{1,2}),? (?P<year>\d{4})\b",
|
||||
"day_month": r"\b(?P<month>January|February|March|April|May|June|July|August|September|October|November"
|
||||
r"|December) (?P<day>\d{1,2})\b",
|
||||
"hour_12": r"\b(\d{1,2})\s*(AM|PM|am|pm)\b",
|
||||
"hour_24": r"\b(\d{1,2}):(\d{2}):(\d{2})\b"
|
||||
}
|
||||
|
||||
month_mapping = {
|
||||
"January": 1, "February": 2, "March": 3, "April": 4, "May": 5, "June": 6, "July": 7, "August": 8,
|
||||
"September": 9, "October": 10, "November": 11, "December": 12
|
||||
}
|
||||
|
||||
weekday_mapping = {
|
||||
"Monday": 1, "Tuesday": 2, "Wednesday": 3, "Thursday": 4, "Friday": 5, "Saturday": 6, "Sunday": 7
|
||||
}
|
||||
|
||||
day_month_year_match = re.search(patterns["day_month_year"], input_string)
|
||||
if day_month_year_match:
|
||||
date_info["year"] = int(day_month_year_match.group("year"))
|
||||
date_info["month"] = month_mapping[day_month_year_match.group("month")]
|
||||
date_info["day"] = int(day_month_year_match.group("day"))
|
||||
|
||||
# Extract month and day without year
|
||||
elif date_info["year"] == -1:
|
||||
day_month_match = re.search(patterns["day_month"], input_string)
|
||||
if day_month_match:
|
||||
date_info["month"] = month_mapping[day_month_match.group("month")]
|
||||
date_info["day"] = int(day_month_match.group("day"))
|
||||
|
||||
# Extract year
|
||||
if date_info["year"] == -1:
|
||||
year_match = re.search(patterns["year"], input_string)
|
||||
if year_match:
|
||||
date_info["year"] = int(year_match.group(0))
|
||||
|
||||
# Extract month
|
||||
if date_info["month"] == -1:
|
||||
month_match = re.search(patterns["month"], input_string)
|
||||
if month_match:
|
||||
date_info["month"] = month_mapping[month_match.group(0)]
|
||||
|
||||
# Extract 12-hour format time
|
||||
hour_12_match = re.search(patterns["hour_12"], input_string)
|
||||
if hour_12_match:
|
||||
hour, period = int(hour_12_match.group(1)), hour_12_match.group(2).lower()
|
||||
if period == 'pm' and hour != 12:
|
||||
hour += 12
|
||||
elif period == 'am' and hour == 12:
|
||||
hour = 0
|
||||
date_info["hour"] = hour
|
||||
|
||||
# Extract weekday
|
||||
for week_day, value in weekday_mapping.items():
|
||||
if week_day in input_string:
|
||||
date_info["weekday"] = value
|
||||
break
|
||||
|
||||
return date_info
|
||||
|
||||
@classmethod
|
||||
def extract_date_parts(cls, input_string: str) -> dict:
|
||||
func_name = f"extract_date_parts_{G_CONTEXT.language}"
|
||||
|
|
@ -96,6 +172,16 @@ class DatetimeHandler(object):
|
|||
return ""
|
||||
return getattr(cls, func_name)(extract_time_dict, meta_data)
|
||||
|
||||
@classmethod
|
||||
def has_time_word(cls, query: str) -> bool:
|
||||
contain_datetime = False
|
||||
# TODO use re
|
||||
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
return contain_datetime
|
||||
|
||||
def datetime_format(self, dt_format: str = "%Y%m%d"):
|
||||
return self._dt.strftime(dt_format)
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from memory_scope.enumeration.language_enum import LanguageEnum
|
|||
from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
|
||||
|
||||
class GlobalContext(object):
|
||||
|
|
@ -18,7 +18,7 @@ class GlobalContext(object):
|
|||
self.model_dict: Dict[str, BaseModel] = {}
|
||||
self.memory_chat_dict: Dict[str, BaseMemoryChat] = {}
|
||||
|
||||
self.vector_store: BaseVectorStore | None = None
|
||||
self.memory_store: BaseMemoryStore | None = None
|
||||
self.monitor: BaseMonitor | None = None
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.language: LanguageEnum = LanguageEnum.EN
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ def init_instance_by_config(config: dict,
|
|||
**kwargs: Additional keyword arguments to pass to the class constructor.
|
||||
|
||||
Returns:
|
||||
object: An instance of the class initialized with the provided config and kwargs.
|
||||
instance: An instance of the class initialized with the provided config and kwargs.
|
||||
"""
|
||||
|
||||
config_copy = deepcopy(config)
|
||||
|
|
|
|||
|
|
@ -1,124 +0,0 @@
|
|||
import json
|
||||
import time
|
||||
from http import HTTPStatus
|
||||
|
||||
import requests
|
||||
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from enumeration.env_type import EnvType
|
||||
|
||||
|
||||
class DashClient(object):
|
||||
|
||||
def __init__(self,
|
||||
request_id: str,
|
||||
dash_scope_uid: str,
|
||||
authorization: str,
|
||||
workspace: str,
|
||||
model_name: str,
|
||||
env_type: EnvType | str = EnvType.DAILY,
|
||||
timeout: int = None,
|
||||
max_retry_count: int = 2,
|
||||
retry_sleep_time: float = 1.0,
|
||||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
self.env_type: EnvType = EnvType(env_type)
|
||||
self.timeout: int = timeout
|
||||
self.max_retry_count: int = max_retry_count
|
||||
self.retry_sleep_time: float = retry_sleep_time
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
# 20240506 update by 泉雨
|
||||
# if authorization:
|
||||
# workspace = ""
|
||||
# dash_scope_uid = ""
|
||||
|
||||
self.headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': authorization,
|
||||
'X-Request-Id': request_id,
|
||||
'X-DashScope-Uid': dash_scope_uid,
|
||||
'X-DashScope-WorkSpace': workspace,
|
||||
}
|
||||
|
||||
self.url: str = ""
|
||||
self.data = {}
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
pass
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
pass
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"url={self.url} header={self.headers} data={self.data} timeout={self.timeout}")
|
||||
response = requests.post(url=self.url,
|
||||
headers=self.headers,
|
||||
data=json.dumps(self.data),
|
||||
timeout=self.timeout)
|
||||
|
||||
if response.status_code == HTTPStatus.OK:
|
||||
response_obj = json.loads(response.text)
|
||||
self.logger.info(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
|
||||
f"call model={model_name} success! retry_cnt={retry_cnt}",
|
||||
stacklevel=3)
|
||||
return self.after_call(response_obj, **kwargs), True
|
||||
|
||||
else:
|
||||
self.logger.warning(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
|
||||
f"call model={model_name} failed! retry_cnt={retry_cnt} details={response.text}",
|
||||
stacklevel=3)
|
||||
return None, False
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class LLIClient(object):
|
||||
|
||||
def __init__(self,
|
||||
model_name: str,
|
||||
timeout: int = None,
|
||||
max_retry_count: int = 2,
|
||||
retry_sleep_time: float = 1.0,
|
||||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
self.timeout: int = timeout
|
||||
self.max_retry_count: int = max_retry_count
|
||||
self.retry_sleep_time: float = retry_sleep_time
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.data = {}
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
pass
|
||||
|
||||
def after_call(self, **kwargs):
|
||||
pass
|
||||
|
||||
def call_once(self, **kwargs):
|
||||
pass
|
||||
|
||||
def call(self, **kwargs):
|
||||
pass
|
||||
|
||||
|
|
@ -1,103 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
import dashscope
|
||||
import time
|
||||
|
||||
from models import EMB
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
|
||||
from typing import List, Dict
|
||||
from utils.registry import build_from_cfg
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
|
||||
class DashEmbeddingClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2782232.html?spm=a2c4g.2782227.0.0.76195b1d9UeBAk#a6a39590fegqx
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.TextEmbedding.Models.text_embedding_v2, **kwargs):
|
||||
super(DashEmbeddingClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.EMBEDDING)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
# text_type: query or document
|
||||
text_type: str = kwargs.pop("text_type", "query")
|
||||
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
|
||||
self.kwargs["text_type"] = text_type
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": {
|
||||
"texts": text,
|
||||
},
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs) -> Dict[int, List[float]] | List[float]:
|
||||
embedding_results = {}
|
||||
for emb in response_obj["output"]["embeddings"]:
|
||||
embedding_results[emb["text_index"]] = emb["embedding"]
|
||||
|
||||
if len(embedding_results) == 1:
|
||||
embedding_results = list(embedding_results.values())[0]
|
||||
return embedding_results
|
||||
|
||||
class LLIEmbedding(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name, **kwargs):
|
||||
super(LLIEmbedding, self).__init__(model_name, **kwargs)
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.embedder = build_from_cfg(self.config, EMB)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
self.data = dict(texts=text)
|
||||
|
||||
def after_call(self, emb: Dict[int, List[float]], **kwargs) -> Dict[int, List[float]] | List[float]:
|
||||
embedding_results = {}
|
||||
for idx, e in enumerate(emb):
|
||||
embedding_results[idx] = e
|
||||
|
||||
if len(embedding_results) == 1:
|
||||
embedding_results = list(embedding_results.values())[0]
|
||||
return embedding_results
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
try:
|
||||
results = self.embedder.get_text_embedding_batch(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
except Exception as e:
|
||||
self.logger.debug(f"Get Error in Embedding: {e}")
|
||||
return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
|
@ -1,129 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
import dashscope
|
||||
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
import time
|
||||
from typing import List, Dict
|
||||
from utils.timer import Timer
|
||||
from models import LLM
|
||||
from utils.registry import build_from_cfg
|
||||
from llama_index.core.base.llms.types import ChatMessage
|
||||
from llama_index.core.base.llms.types import (
|
||||
ChatResponse,
|
||||
CompletionResponse,
|
||||
)
|
||||
|
||||
|
||||
class DashGenerateClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2712576.html
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.Generation.Models.qwen_max, **kwargs):
|
||||
super(DashGenerateClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.GENERATION)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
|
||||
|
||||
input_text = {}
|
||||
if prompt:
|
||||
input_text["prompt"] = prompt
|
||||
elif messages:
|
||||
input_text["messages"] = messages
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": input_text,
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
self.logger.debug(f"response_obj={response_obj}")
|
||||
output = response_obj["output"]
|
||||
if "text" in output:
|
||||
return output["text"]
|
||||
elif "choices" in output:
|
||||
return output["choices"][0]["message"]["content"]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class LLILLM(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name: str, **kwargs):
|
||||
super(LLILLM, self).__init__(model_name, **kwargs)
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.llm = build_from_cfg(self.config, LLM)
|
||||
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
|
||||
|
||||
if prompt:
|
||||
input_text = prompt
|
||||
input_type = 'prompt'
|
||||
llama_input = input_text
|
||||
elif messages:
|
||||
input_text = messages
|
||||
input_type = 'messages'
|
||||
llama_input = [ChatMessage(
|
||||
role=x['role'], content=x['content']
|
||||
) for x in input_text]
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
self.data = {
|
||||
input_type: llama_input,
|
||||
}
|
||||
|
||||
def after_call(self, response_obj: ChatResponse | CompletionResponse, **kwargs) -> str:
|
||||
self.logger.debug(f"response_obj={response_obj}")
|
||||
if isinstance(response_obj, CompletionResponse):
|
||||
return response_obj.text
|
||||
elif isinstance(response_obj, ChatResponse):
|
||||
return response_obj.message.content
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
if True:
|
||||
# try:
|
||||
if 'prompt' in self.data:
|
||||
results = self.llm.complete(**self.data)
|
||||
else:
|
||||
results = self.llm.chat(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
# except:
|
||||
# return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
print("dashscope llm results:",result)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
|
@ -1,119 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
import dashscope
|
||||
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
|
||||
import time
|
||||
from typing import List
|
||||
from models import RERANKER
|
||||
from utils.timer import Timer
|
||||
|
||||
from utils.registry import build_from_cfg
|
||||
from llama_index.core.data_structs import Node
|
||||
from llama_index.core.schema import NodeWithScore # type: ignore
|
||||
|
||||
|
||||
class DashReRankClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2780059.html
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.TextReRank.Models.gte_rerank, **kwargs):
|
||||
super(DashReRankClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.RERANK)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
top_n: int | None = kwargs.pop("top_n", None)
|
||||
return_documents: bool = kwargs.pop("return_documents", False)
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
self.kwargs.update({
|
||||
"top_n": top_n,
|
||||
"return_documents": return_documents,
|
||||
})
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": {
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
},
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
return response_obj["output"]["results"]
|
||||
|
||||
class LLIReRank(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name, **kwargs):
|
||||
super(LLIReRank, self).__init__(model_name, **kwargs)
|
||||
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.reranker = build_from_cfg(self.config, RERANKER)
|
||||
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
top_n: int | None = kwargs.pop("top_n", None)
|
||||
return_documents: bool = kwargs.pop("return_documents", False)
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
nodes = [NodeWithScore(node=Node(text=text), score=1.0) for text in documents]
|
||||
|
||||
self.data = {
|
||||
"nodes": nodes,
|
||||
"query_str": query,
|
||||
}
|
||||
|
||||
|
||||
def after_call(self, nodes, **kwargs):
|
||||
results = []
|
||||
for idx, node in enumerate(nodes):
|
||||
results.append(dict(index=idx,
|
||||
relevance_score=node.score,
|
||||
document=node.node.text))
|
||||
return results
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
try:
|
||||
results = self.reranker.postprocess_nodes(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
except Exception as e:
|
||||
self.logger.debug(f"Rerank falls, data={self.data}")
|
||||
# return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
||||
|
|
@ -1,419 +0,0 @@
|
|||
from elasticsearch import Elasticsearch
|
||||
from elasticsearch.helpers import bulk
|
||||
|
||||
|
||||
from models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
|
||||
from common.dash_embedding_client import DashEmbeddingClient
|
||||
from common.logger import Logger
|
||||
|
||||
from constants.common_constants import ES_ENV_URL_DICT
|
||||
from enumeration.env_type import EnvType
|
||||
from utils.logger import Logger
|
||||
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore
|
||||
from llama_index.core.schema import TextNode
|
||||
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
|
||||
|
||||
class ElasticSearchClient(object):
|
||||
def __init__(self,
|
||||
es_user_name: str,
|
||||
es_password: str,
|
||||
es_index_name: str,
|
||||
embedding_client: DashEmbeddingClient | None = None,
|
||||
env_type: EnvType | str = EnvType.DAILY,
|
||||
content_key: str = "content",
|
||||
vector_key: str = "vector",
|
||||
**kwargs):
|
||||
|
||||
self.es_index_name: str = es_index_name
|
||||
self.embedding_client: DashEmbeddingClient = embedding_client
|
||||
self.content_key: str = content_key
|
||||
self.vector_key: str = vector_key
|
||||
|
||||
self.es_client = Elasticsearch(
|
||||
hosts=[ES_ENV_URL_DICT.get(EnvType(env_type))],
|
||||
basic_auth=(es_user_name, es_password),
|
||||
**kwargs)
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
self.logger.debug(f"connect es_client info={self.es_client.info()}")
|
||||
|
||||
def log_index_info(self):
|
||||
index_info = self.es_client.indices.get(index=self.es_index_name)
|
||||
self.logger.info(f"index={self.es_index_name} exists. index_info={index_info}")
|
||||
|
||||
def insert(self, _id: str, body: dict):
|
||||
assert body and self.content_key in body, f"body={body} is illegal!"
|
||||
|
||||
# text_type: document
|
||||
content = body[self.content_key]
|
||||
vector = self.embedding_client.call(text=content, text_type="document")
|
||||
if not vector:
|
||||
self.logger.warning(f"embedding_client call failed, stop es insert!")
|
||||
return
|
||||
|
||||
body[self.vector_key] = vector
|
||||
response = self.es_client.index(id=_id, index=self.es_index_name, body=body)
|
||||
self.logger.info(f"insert response={response}")
|
||||
|
||||
def insert_batch(self, doc_list: list):
|
||||
"""
|
||||
doc_list = [
|
||||
{
|
||||
"_id": 2,
|
||||
"_source": {
|
||||
"author": "john",
|
||||
"text": "Elasticsearch: cool.",
|
||||
"timestamp": "2023-03-23T10:00:00"
|
||||
}
|
||||
},
|
||||
{
|
||||
"_id": 3,
|
||||
"_source": {
|
||||
"author": "jane",
|
||||
"text": "Elasticsearch: very cool.",
|
||||
"timestamp": "2023-03-23T11:00:00"
|
||||
}
|
||||
}
|
||||
]
|
||||
"""
|
||||
text_list = []
|
||||
for doc in doc_list:
|
||||
assert "_id" in doc and "_source" in doc
|
||||
content = doc["_source"][self.content_key]
|
||||
text_list.append(content)
|
||||
|
||||
vector_dict = self.embedding_client.call(text=text_list, text_type="document")
|
||||
if not vector_dict:
|
||||
self.logger.warning(f"embedding_client call failed, stop es insert!")
|
||||
return
|
||||
|
||||
# add _index
|
||||
for i, doc in enumerate(doc_list):
|
||||
doc["_index"] = self.es_index_name
|
||||
vector = vector_dict[i]
|
||||
doc["_source"][self.vector_key] = vector
|
||||
|
||||
# 执行批量插入
|
||||
responses = bulk(self.es_client, doc_list)
|
||||
|
||||
# 输出批量插入的响应
|
||||
for response in responses[1]:
|
||||
self.logger.info(f"insert_batch response={response}")
|
||||
|
||||
def print_hits(self, hits: list):
|
||||
for hit in hits:
|
||||
print_kwargs = {
|
||||
"_id": hit['_id'],
|
||||
"_score": hit['_score'],
|
||||
}
|
||||
for k, v in hit['_source'].items():
|
||||
# 不打印vector
|
||||
if k == self.vector_key:
|
||||
v = len(v)
|
||||
print_kwargs[k] = v
|
||||
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
|
||||
|
||||
def exact_search(self,
|
||||
size: int,
|
||||
exact_filters: dict = None,
|
||||
wildcard_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
"""
|
||||
{
|
||||
"match": {
|
||||
"category": "electronics" # 一级字段过滤
|
||||
}
|
||||
},
|
||||
{
|
||||
"match": {
|
||||
"product.name": "laptop" # 二级字段过滤
|
||||
}
|
||||
}
|
||||
{
|
||||
"terms": {
|
||||
"product.keyA": ["a", "b", "c"] # 二级字段keyA的精确值必须为a、b、c中的一
|
||||
}
|
||||
}
|
||||
"""
|
||||
must_list = []
|
||||
for key, value in exact_filters.items():
|
||||
if not key:
|
||||
continue
|
||||
if isinstance(value, str):
|
||||
must_list.append({"match": {key: value}})
|
||||
elif isinstance(value, list):
|
||||
must_list.append({"terms": {key: value}})
|
||||
|
||||
query = {
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": must_list
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if wildcard_filters:
|
||||
should_list = []
|
||||
for key, value in wildcard_filters.items():
|
||||
if not key:
|
||||
continue
|
||||
if isinstance(value, str):
|
||||
should_list.append({"wildcard": {key: f"*{value}*"}})
|
||||
elif isinstance(value, list):
|
||||
for v in value:
|
||||
should_list.append({"wildcard": {key: f"*{v}*"}})
|
||||
|
||||
query["query"]["bool"].update({
|
||||
"should": should_list,
|
||||
"minimum_should_match": 1,
|
||||
})
|
||||
self.logger.info(f"query={query}")
|
||||
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"exact_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']} "
|
||||
f"exact_filters={exact_filters}", stacklevel=2)
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
def exact_search_v2(self,
|
||||
size: int,
|
||||
term_filters: dict = None,
|
||||
match_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
|
||||
"""
|
||||
"bool": {
|
||||
"must": [
|
||||
{"term": {"field1": "固定值"}}, # 一级目录关键字过滤(等于某个值)
|
||||
{"terms": {"field2": ["a", "b", "c"]}} # 二级目录关键字过滤(等于三个中的任意一个)
|
||||
],
|
||||
"should": [ # 至少匹配其中之一
|
||||
{"match": {"key": "ccc"}}, # key包含"ccc"
|
||||
{"match": {"key": "bbb"}} # 或者key包含"bbb"
|
||||
],
|
||||
"minimum_should_match": 1 # 至少有一个`should`条件匹配
|
||||
}
|
||||
"""
|
||||
|
||||
query = {
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if term_filters:
|
||||
must_list = []
|
||||
for k, v in term_filters.items():
|
||||
if isinstance(v, list):
|
||||
must_list.append({"terms": {k: v}})
|
||||
elif isinstance(v, str):
|
||||
must_list.append({"term": {k: v}})
|
||||
else:
|
||||
raise NotImplemented
|
||||
query["query"]["bool"]["must"] = must_list
|
||||
|
||||
if match_filters:
|
||||
match_list = []
|
||||
for k, v in match_filters.items():
|
||||
if isinstance(v, list):
|
||||
for v_sub in v:
|
||||
match_list.append({"match": {k: v_sub}})
|
||||
elif isinstance(v, str):
|
||||
match_list.append({"match": {k: v}})
|
||||
else:
|
||||
raise NotImplemented
|
||||
query["query"]["bool"]["should"] = match_list
|
||||
query["query"]["bool"]["minimum_should_match"] = 1
|
||||
|
||||
self.logger.info(query)
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"exact_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']}", stacklevel=2)
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
def similar_search(self,
|
||||
text: str,
|
||||
size: int,
|
||||
exact_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
|
||||
if exact_filters is None:
|
||||
exact_filters = {}
|
||||
|
||||
# 过滤or
|
||||
or_filters = {}
|
||||
for k in list(exact_filters.keys()):
|
||||
v = exact_filters[k]
|
||||
if isinstance(v, list):
|
||||
exact_filters.pop(k)
|
||||
or_filters[k] = v
|
||||
|
||||
vector = self.embedding_client.call(text=text)
|
||||
if not vector:
|
||||
self.logger.warning(f"embedding_client call failed, stop select from es!")
|
||||
return
|
||||
|
||||
query = {
|
||||
# 返回最相似的top_k个文档
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": {
|
||||
"script_score": {
|
||||
# 对所有文档执行
|
||||
"query": {
|
||||
"match_all": {}
|
||||
},
|
||||
"script": {
|
||||
# 使用余弦相似度+1,es不能返回负数
|
||||
"source": f"cosineSimilarity(params.query_vector, '{self.vector_key}') + 1.0",
|
||||
"params": {"query_vector": vector}
|
||||
}
|
||||
}
|
||||
},
|
||||
"filter": [
|
||||
{"term": {k: v}} for k, v in exact_filters.items()
|
||||
],
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if or_filters:
|
||||
k_v_pair = []
|
||||
for k, v_list in or_filters.items():
|
||||
for v in v_list:
|
||||
k_v_pair.append((k, v))
|
||||
query["query"]["bool"]["should"] = [{"term": {k: v}} for k, v in k_v_pair]
|
||||
query["query"]["bool"]["minimum_should_match"] = 1
|
||||
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"similar_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']} "
|
||||
f"text={text} "
|
||||
f"exact_filters={exact_filters}", stacklevel=2)
|
||||
|
||||
# 还原打分
|
||||
for hit in hits:
|
||||
hit['_score'] -= 1
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
|
||||
class LLIElasticSearch(object):
|
||||
def __init__(self,
|
||||
es_index_name: str,
|
||||
embedding_client: LLIEmbedding | None = None,
|
||||
retrieve_topk: int = 3,
|
||||
content_key: str = "text",
|
||||
):
|
||||
self.es_index_name = es_index_name
|
||||
self.content_key = content_key
|
||||
self.embedding_client: LLIEmbedding = embedding_client
|
||||
# using local es for debug convenient
|
||||
self.es_client = ElasticsearchStore(index_name="my_index",
|
||||
es_url="http://localhost:9200",
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True))
|
||||
|
||||
self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_client, llm=None)
|
||||
self.storage_context = StorageContext.from_defaults(vector_store=self.es_client)
|
||||
self.index = VectorStoreIndex(storage_context=self.storage_context,
|
||||
service_context=self.service_context)
|
||||
|
||||
self.retriever = self.index.as_retriever(similarity_top_k=retrieve_topk)
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def log_index_info(self, ):
|
||||
pass
|
||||
|
||||
def print_hits(self, hits: list):
|
||||
for hit in hits:
|
||||
print_kwargs = {
|
||||
"_id": hit['_id'],
|
||||
"_score": hit['_score'],
|
||||
}
|
||||
for k, v in hit['_source'].items():
|
||||
# 不打印vector
|
||||
if k == self.vector_key:
|
||||
v = len(v)
|
||||
print_kwargs[k] = v
|
||||
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
|
||||
|
||||
def similar_search(self,
|
||||
text: str,
|
||||
size: int, ):
|
||||
|
||||
ret_nodes = self.retriever.retrieve(text)
|
||||
return ret_nodes
|
||||
|
||||
def insert_batch(self, doc_list:list[str]):
|
||||
node_list = []
|
||||
for doc in doc_list:
|
||||
assert "_id" in doc and "_source" in doc
|
||||
content = doc["_source"]["text"]
|
||||
doc["_source"].pop("text")
|
||||
meta = doc["_source"]
|
||||
node = TextNode(text=content, metadata=meta)
|
||||
node.node_id(doc['_id'])
|
||||
node_list.append(node)
|
||||
self.index.insert_nodes(node_list)
|
||||
|
||||
|
||||
def insert(self, _id: str, body: dict):
|
||||
assert body and self.content_key in body, f"body={body} is illegal!"
|
||||
content = body[self.content_key]
|
||||
body.pop(self.content_key)
|
||||
meta = body
|
||||
node = TextNode(text=content, metadata=meta)
|
||||
self.index.insert_nodes([node])
|
||||
|
||||
|
||||
|
||||
|
|
@ -1,71 +0,0 @@
|
|||
import re
|
||||
from typing import Dict, List
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
"""
|
||||
除了 content_modified,其他均和数据库字段保持统一
|
||||
根据code判断,如果code是空,则为新增的memoryNode,如果有值,则为更新
|
||||
if content_modified is true,则需要调用embedding服务
|
||||
"""
|
||||
id: str = Field("", description="唯一主键 uuid64")
|
||||
|
||||
code: str = Field("", description="和id保持一致,为空则是新增")
|
||||
|
||||
# 0520新增
|
||||
timeCreated: str = Field("", description="Memory创建时间(算法不关注)")
|
||||
|
||||
# 0520新增
|
||||
timeModified: str = Field("", description="Memory更新时间(算法不关注)")
|
||||
|
||||
content: str = Field("", description="记忆内容")
|
||||
|
||||
memoryId: str = Field("", description="记忆 id,检索区分字段")
|
||||
|
||||
# 0520新增
|
||||
scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT")
|
||||
|
||||
# 0520新增
|
||||
# NOTE 百炼服务端只召回observation, insight, profile, obs_customized, profile_customized
|
||||
memoryType: str = Field("", description="conversation, observation, insight, "
|
||||
"profile, obs_customized, profile_customized")
|
||||
|
||||
# 0520新增,但不是数据库字段
|
||||
content_modified: bool = Field(False, description="content是否被更新,if true,则需要调用embedding服务")
|
||||
|
||||
# reflected: 1 is reflected before, 0 has not reflected, 如果是用户自定义,写入空值"".
|
||||
metaData: Dict[str, str] = Field({}, description="元信息: infoScore, algoVersion, datetime, reflected")
|
||||
|
||||
status: str = Field("active", description="active or expired")
|
||||
|
||||
tenantId: str = Field("", description="request id")
|
||||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
|
||||
def get_time_info(self, time_format: str):
|
||||
pattern = re.compile(r'\{([^}]*)}')
|
||||
keys = pattern.findall(time_format)
|
||||
|
||||
match_flag = True
|
||||
kv_dict = {}
|
||||
for k in keys:
|
||||
if k not in self.metaData:
|
||||
match_flag = False
|
||||
break
|
||||
v = self.metaData[k]
|
||||
if not v:
|
||||
match_flag = False
|
||||
break
|
||||
|
||||
kv_dict[k] = v
|
||||
|
||||
if match_flag:
|
||||
return time_format.format(**kv_dict)
|
||||
return ""
|
||||
|
||||
def to_dict(self):
|
||||
res = {"content": self.content, "memoryId": self.memoryId, "memoryType": self.memoryType,
|
||||
"status": self.status, "metaData": self.metaData}
|
||||
return res
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from pydantic import Field, BaseModel
|
||||
|
||||
from scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
id: str = Field("", description="uuid64")
|
||||
|
||||
score_similar: float = Field(0, description="相似度打分")
|
||||
|
||||
score_rank: float = Field(0, description="排序打分")
|
||||
|
||||
score_rerank: float = Field(0, description="重排打分")
|
||||
|
||||
memory_node: MemoryNode = Field(None, description="memory node 核心,返回给上游的结构")
|
||||
|
||||
@classmethod
|
||||
def init_from_es(cls, hit: dict):
|
||||
memory_node = MemoryNode(**hit['_source'])
|
||||
return cls(id=hit['_id'], score_similar=hit['_score'], memory_node=memory_node)
|
||||
|
||||
@classmethod
|
||||
def init_from_attrs(cls, **kwargs):
|
||||
_id: str = kwargs.get("_id", "")
|
||||
score_similar: float = kwargs.pop("score_similar", 0)
|
||||
score_rank: float = kwargs.pop("score_rank", 0)
|
||||
score_rerank: float = kwargs.pop("score_rerank", 0)
|
||||
memory_node = MemoryNode(**kwargs)
|
||||
return cls(id=_id,
|
||||
score_similar=score_similar,
|
||||
score_rank=score_rank,
|
||||
score_rerank=score_rerank,
|
||||
memory_node=memory_node)
|
||||
|
|
@ -1,140 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERGE_NODES, NEW_INSIGHT_KEYS, INSIGHT_KEY, \
|
||||
INSIGHT_VALUE, REFLECTED
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
def __init__(self, insight_obs_max_cnt, es_insight_similar_top_k, get_insight_model, get_insight_max_token, get_insight_temperature, get_insight_top_k, **kwargs):
|
||||
super(GetInsightWorker,self).__init__(*args,**kwargs)
|
||||
self.insight_obs_max_cnt = insight_obs_max_cnt
|
||||
self.get_insight_model = get_insight_model
|
||||
self.get_insight_max_token = get_insight_max_token
|
||||
self.get_insight_temperature = get_insight_temperature
|
||||
self.get_insight_top_k = get_insight_top_k
|
||||
self.es_insight_similar_top_k = es_insight_similar_top_k
|
||||
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
DT: dt,
|
||||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
}
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryNode.init_from_attrs(content=content,
|
||||
memoryId=self.memory_id,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.INSIGHT.value,
|
||||
content_modified=True, # 新增的insight需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def reflect_new_insight_key(self,
|
||||
insight_key: str,
|
||||
not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None:
|
||||
|
||||
# 检索历史memory
|
||||
hits = self.es_client.similar_search(text=insight_key,
|
||||
size=self.es_insight_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
# 转化成 MemoryNodeWrap 合并新增nodes
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
related_nodes.extend(not_reflected_merge_nodes)
|
||||
|
||||
# content去重
|
||||
related_node_dict = {n.memory_node.content: n for n in related_nodes}
|
||||
related_nodes = sorted(list(related_node_dict.values()), key=lambda x: x.memory_node.id)
|
||||
documents = [n.memory_node.content for n in related_nodes]
|
||||
|
||||
# 重排所有记忆
|
||||
result = self.rerank_client.call(query=insight_key, documents=documents)
|
||||
if not result:
|
||||
self.add_run_info(f"reflect insight_key={insight_key} call rerank client failed!")
|
||||
return
|
||||
|
||||
# 根据打分过滤
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
related_nodes[index].score_rank = score
|
||||
related_nodes_sorted = sorted(related_nodes, key=lambda x: x.score_rank, reverse=True)[
|
||||
:self.insight_obs_max_cnt]
|
||||
|
||||
# 生成prompt
|
||||
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
|
||||
get_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_insight_system,
|
||||
few_shot=self.prompt_config.get_insight_few_shot,
|
||||
user_query=self.prompt_config.get_insight_user_query.format(
|
||||
insight_key=insight_key, user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"get_insight_message={get_insight_message}")
|
||||
|
||||
# call LLM, 提取insight
|
||||
response_text = self.gene_client.call(messages=get_insight_message,
|
||||
model_name=self.get_insight_model,
|
||||
max_token=self.get_insight_max_token,
|
||||
temperature=self.get_insight_temperature,
|
||||
top_k=self.get_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_upon_user_attr call llm failed!")
|
||||
return
|
||||
response_text = response_text.strip()
|
||||
if response_text in ["无"]:
|
||||
return
|
||||
return self.new_insight_node(insight_key=insight_key, insight_value=response_text)
|
||||
|
||||
def _run(self):
|
||||
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
|
||||
if not new_insight_keys:
|
||||
self.add_run_info("new_insight_keys is empty! stop insight.")
|
||||
return
|
||||
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
if not not_reflected_merge_nodes:
|
||||
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
|
||||
return
|
||||
|
||||
# submit insight task
|
||||
for insight_key in new_insight_keys:
|
||||
self.submit_thread(self.reflect_new_insight_key,
|
||||
sleep_time=1,
|
||||
insight_key=insight_key,
|
||||
not_reflected_merge_nodes=not_reflected_merge_nodes)
|
||||
|
||||
# save output
|
||||
new_insight_nodes: List[MemoryNode] = []
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
new_insight_nodes.append(result)
|
||||
assert isinstance(result, MemoryNode)
|
||||
insight_key = result.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = result.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"after_get_insight insight_key={insight_key} insight_value={insight_value}")
|
||||
|
||||
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
|
||||
|
||||
# set REFLECTED
|
||||
for node in not_reflected_merge_nodes:
|
||||
scheme.memory_node.metaData[REFLECTED] = "1"
|
||||
|
|
@ -1,80 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NOT_REFLECTED_OBS_NODES, REFLECTED, INSIGHT_NODES, INSIGHT_KEY, \
|
||||
NEW_INSIGHT_KEYS, NOT_REFLECTED_MERGE_NODES
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
def __init__(self, reflect_obs_cnt_threshold, reflect_num_questions, reflect_obs_model, reflect_obs_max_token, reflect_obs_temperature, reflect_obs_top_k, *args, **kwargs):
|
||||
super(GetReflectionWorker,self).__init__(*args, **kwargs)
|
||||
self.reflect_obs_cnt_threshold = reflect_obs_cnt_threshold
|
||||
self.reflect_num_questions = reflect_num_questions
|
||||
self.reflect_obs_model = reflect_obs_model
|
||||
self.reflect_obs_max_token = reflect_obs_max_token
|
||||
self.reflect_obs_temperature = reflect_obs_temperature
|
||||
self.reflect_obs_top_k = reflect_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
not_reflected_merge_nodes.extend(new_obs_nodes)
|
||||
if not_reflected_nodes:
|
||||
not_reflected_merge_nodes.extend(not_reflected_nodes)
|
||||
not_reflected_merge_nodes = [node for node in not_reflected_merge_nodes
|
||||
if scheme.memory_node.metaData.get(REFLECTED, "") == "0"]
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop reflect.")
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = []
|
||||
profile_keys: List[str] = list(self.user_profile_dict.keys())
|
||||
exist_keys.extend(profile_keys)
|
||||
self.logger.info(f"profile_keys={profile_keys}")
|
||||
|
||||
# get insight_keys
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if insight_nodes:
|
||||
insight_keys = [n.memory_node.metaData.get(INSIGHT_KEY) for n in insight_nodes]
|
||||
insight_keys = [x.strip() for x in insight_keys if x]
|
||||
exist_keys.extend(insight_keys)
|
||||
self.logger.info(f"insight_keys={insight_keys}")
|
||||
|
||||
# gen reflect prompt
|
||||
user_query_list = [n.memory_node.content for n in not_reflected_merge_nodes]
|
||||
reflect_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_reflect_system.format(
|
||||
num_questions=self.reflect_num_questions),
|
||||
few_shot=self.prompt_config.get_reflect_few_shot,
|
||||
user_query=self.prompt_config.get_reflect_user_query.format(exist_keys=",".join(exist_keys),
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# # call LLM
|
||||
response_text = self.gene_client.call(messages=reflect_message,
|
||||
model_name=self.reflect_obs_model,
|
||||
max_token=self.reflect_obs_max_token,
|
||||
temperature=self.reflect_obs_temperature,
|
||||
top_k=self.reflect_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_obs_questions call llm failed!")
|
||||
return
|
||||
|
||||
# parse text & save
|
||||
new_insight_keys = ResponseTextParser(response_text).parse_v2("get_insight_keys")
|
||||
if new_insight_keys:
|
||||
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)
|
||||
|
|
@ -1,118 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(es_contra_repeat_similar_top_k, long_contra_repeat_threshold, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(LongContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.es_contra_repeat_similar_top_k = es_contra_repeat_similar_top_k
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
# new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
# oday_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_scheme.memory_node.content
|
||||
hits = self.es_client.similar_search(text=text,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
all_obs_nodes.append(related_node)
|
||||
if has_match:
|
||||
all_obs_nodes.append(new_obs_node)
|
||||
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.memory_node.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.long_contra_repeat_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.long_contra_repeat_few_shot,
|
||||
user_query=self.prompt_config.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants.common_constants import NEW_INSIGHT_NODES, MODIFIED_MEMORIES, INSIGHT_NODES, NEW_OBS_NODES, \
|
||||
NOT_REFLECTED_OBS_NODES, NEW, NOT_REFLECTED_MERGE_NODES
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update({n.id: n for n in insight_nodes if n.memory_node.content_modified})
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.memory_node.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.memory_node.metaData[NEW] = "0"
|
||||
all_node_dict.update({n.memory_node.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
|
|
@ -1,151 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
def __init__(update_insight_threshold, update_insight_max_thread, update_insight_model, update_insight_max_token, update_insight_temperature, update_insight_top_k,*args, **kwargs):
|
||||
super(UpdateInsightWorker, self).__init__(*args, **kwargs)
|
||||
self.update_insight_threshold = update_insight_threshold
|
||||
self.update_insight_max_thread = update_insight_max_thread
|
||||
self.update_insight_model = update_insight_model
|
||||
self.update_insight_max_token = update_insight_max_token
|
||||
self.update_insight_temperature = update_insight_temperature
|
||||
self.update_insight_top_k = update_insight_top_k
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
new_obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
if not insight_key or not insight_value:
|
||||
self.logger.warning(f"insight_key={insight_key} insight_value={insight_value} is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
result = self.rerank_client.call(query=insight_key,
|
||||
documents=[x.memory_node.content for x in new_obs_nodes])
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_insight_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(f"insight_key={insight_key} insight_value={insight_value} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
|
||||
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(self,
|
||||
insight_node: MemoryNode,
|
||||
filtered_nodes: List[MemoryNode]) -> MemoryNode:
|
||||
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"update_insight insight_key={insight_key} insight_value={insight_value} "
|
||||
f"doc.size={len(filtered_nodes)}")
|
||||
|
||||
# gen prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{scheme.memory_node.content}")
|
||||
update_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_insight_system,
|
||||
few_shot=self.prompt_config.update_insight_few_shot,
|
||||
user_query=self.prompt_config.update_insight_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key + ":" + insight_value))
|
||||
self.logger.info(f"update_insight_message={update_insight_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_insight_message,
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} call llm failed!")
|
||||
return insight_node
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(f"update_profile {insight_key}")
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 1!")
|
||||
return insight_node
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 2")
|
||||
return insight_node
|
||||
insight_value = profile_list[0]
|
||||
|
||||
if not insight_value or insight_value in ["无", "重复"]:
|
||||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
insight_scheme.memory_node.metaData[INSIGHT_VALUE] = insight_value
|
||||
insight_scheme.memory_node.content_modified = True
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop update sights!")
|
||||
return
|
||||
if not insight_nodes:
|
||||
self.logger.info("insight_nodes is empty, stop update sights!")
|
||||
return
|
||||
|
||||
# 提交打分任务
|
||||
for node in insight_nodes:
|
||||
self.submit_thread(self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
insight_node=node,
|
||||
new_obs_nodes=new_obs_nodes)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
insight_node, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[:update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(self.update_insight,
|
||||
sleep_time=1,
|
||||
insight_node=insight_node,
|
||||
filtered_nodes=filtered_nodes)
|
||||
|
||||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryNode = result
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"after_update_insight insight_key={insight_key} insight_value={insight_value}")
|
||||
|
|
@ -1,210 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
def __init__(update_profile_max_thread, update_profile_threshold, extra_user_attrs, update_profile_model, update_profile_max_token, update_profile_temperature, update_profile_top_k, *args, **kwargs):
|
||||
super(UpdateProfileWorker,self).__init__(*args, **kwargs)
|
||||
self.update_profile_max_thread = update_profile_max_thread
|
||||
self.extra_user_attrs = extra_user_attrs
|
||||
self.update_profile_threshold = update_profile_threshold
|
||||
self.update_profile_model = update_profile_model
|
||||
self.update_profile_max_token = update_profile_max_token
|
||||
self.update_profile_temperature = update_profile_temperature
|
||||
self.update_profile_top_k = update_profile_top_k
|
||||
|
||||
# @property
|
||||
# def extra_user_attrs(self):
|
||||
# return self.request.extra_user_attrs
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
user_attr: UserAttribute,
|
||||
new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
result = self.rerank_client.call(query=user_attr.description,
|
||||
documents=[x.memory_node.content for x in new_obs_nodes])
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_user_attr={user_attr.memory_key} call rerank failed!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_profile_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(f"key={user_attr.memory_key} desc={user_attr.description} "
|
||||
f"content={scheme.memory_node.content} score={score} keep_flag={keep_flag}")
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
def update_user_attr(self, user_attr: UserAttribute, filtered_nodes: List[MemoryNode]) -> UserAttribute:
|
||||
self.logger.info(f"update_user_attr memory_key={user_attr.memory_key} desc={user_attr.description} "
|
||||
f"value={user_attr.value} doc.size={len(filtered_nodes)}")
|
||||
|
||||
# 根据不同的参数类型是否多值,分别给出prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{scheme.memory_node.content}")
|
||||
update_profile = f"{user_attr.memory_key}({user_attr.description})"
|
||||
update_profile_value = update_profile + ":" + ",".join(user_attr.value)
|
||||
|
||||
if user_attr.is_unique == 1:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_unique_profile_system,
|
||||
few_shot=self.prompt_config.update_unique_profile_few_shot,
|
||||
user_query=self.prompt_config.update_unique_profile_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value))
|
||||
else:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_plural_profile_system,
|
||||
few_shot=self.prompt_config.update_plural_profile_few_shot,
|
||||
user_query=self.prompt_config.update_plural_profile_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value))
|
||||
self.logger.info(f"update_profile_message={update_profile_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_profile_message,
|
||||
model_name=self.update_profile_model,
|
||||
max_token=self.update_profile_max_token,
|
||||
temperature=self.update_profile_temperature,
|
||||
top_k=self.update_profile_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} call llm failed!")
|
||||
return user_attr
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(f"update_attr {user_attr.memory_key}")
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 1!")
|
||||
return user_attr
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 2")
|
||||
return user_attr
|
||||
profile = profile_list[0]
|
||||
|
||||
if not profile or profile in ["无", "重复"]:
|
||||
self.logger.info(f"profile={profile}, skip.")
|
||||
return user_attr
|
||||
|
||||
# check 英文中午逗号
|
||||
if user_attr.is_unique == 1:
|
||||
user_attr.value = [profile.strip()]
|
||||
else:
|
||||
attr_value_list = profile.replace(",", ",").split(",")
|
||||
user_attr.value = [x.strip() for x in sorted(list(set(user_attr.value + attr_value_list)))]
|
||||
return user_attr
|
||||
|
||||
def add_extra_user_attrs(self):
|
||||
# 解析为空返回
|
||||
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
|
||||
if not extra_user_attr_list:
|
||||
return
|
||||
|
||||
for user_attr_info in extra_user_attr_list:
|
||||
user_attr_split = user_attr_info.split(":")
|
||||
|
||||
# 格式不对返回
|
||||
if len(user_attr_split) < 1:
|
||||
continue
|
||||
user_attr_key = user_attr_split[0]
|
||||
|
||||
user_attr_desc = ""
|
||||
if len(user_attr_split) >= 2:
|
||||
user_attr_desc = user_attr_split[1]
|
||||
|
||||
user_attr_unique = 0
|
||||
if len(user_attr_split) >= 3:
|
||||
user_attr_unique = int(user_attr_split[2])
|
||||
|
||||
# 已经包含返回
|
||||
if user_attr_key in self.user_profile_dict:
|
||||
user_attr = self.user_profile_dict[user_attr_key]
|
||||
# description为空,补充description
|
||||
if not user_attr.description:
|
||||
user_attr.description = user_attr_desc
|
||||
continue
|
||||
|
||||
# 增加新属性
|
||||
new_attr = UserAttribute(memory_id=self.config.memory_id,
|
||||
scene=self.scene,
|
||||
memory_key=user_attr_key,
|
||||
is_unique=int(user_attr_unique),
|
||||
is_mutable=1,
|
||||
memory_type=MemoryTypeEnum.PROFILE,
|
||||
description=user_attr_desc,
|
||||
status=1)
|
||||
self.user_profile_dict[user_attr_key] = new_attr
|
||||
|
||||
def _run(self):
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop user profile!")
|
||||
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
|
||||
return
|
||||
|
||||
# 增加环境变量配置的属性
|
||||
if self.extra_user_attrs:
|
||||
self.add_extra_user_attrs()
|
||||
|
||||
new_user_profile: List[UserAttribute] = []
|
||||
self.set_context(NEW_USER_PROFILE, new_user_profile)
|
||||
|
||||
for user_attr_key, user_attr in self.user_profile_dict.items():
|
||||
# 不可修改直接跳过
|
||||
if user_attr.is_mutable != 1:
|
||||
new_user_profile.append(user_attr)
|
||||
self.logger.info(f"{user_attr_key} is not mutable! continue")
|
||||
continue
|
||||
|
||||
self.submit_thread(self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
user_attr=user_attr,
|
||||
new_obs_nodes=new_obs_nodes)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
user_attr, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_profile_max_thread:
|
||||
result_sorted = result_sorted[:self.update_profile_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for user_attr, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(self.update_user_attr,
|
||||
sleep_time=1,
|
||||
user_attr=user_attr,
|
||||
filtered_nodes=filtered_nodes)
|
||||
|
||||
# collect result & save
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
user_attribute: UserAttribute = result
|
||||
self.logger.info(f"after_update_profile memory_key={user_attribute.memory_key} "
|
||||
f"desc={user_attribute.description} value={user_attribute.value}")
|
||||
new_user_profile.append(user_attribute)
|
||||
|
|
@ -1,98 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(self, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(ContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
all_obs_nodes.extend(new_obs_nodes)
|
||||
if new_obs_with_time_nodes:
|
||||
all_obs_nodes.extend(new_obs_with_time_nodes)
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.memory_node.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.contra_repeat_few_shot,
|
||||
user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,134 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict, extract_date_parts
|
||||
from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME, KEY_WORD, DATATIME_WORD_LIST, \
|
||||
NEW_OBS_WITH_TIME_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWithTimeWorker, self).__init__(*args, **kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: time_infer, # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
|
||||
# 事件时间
|
||||
meta_data.update({f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()})
|
||||
# 对话时间
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
dt = time_to_formatted_str(time=msg.time_created,
|
||||
date_format="",
|
||||
string_format="{year}年{month}月{day}日{weekday}{hour}点")
|
||||
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs with time user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.get_observation_with_time_few_shot,
|
||||
user_query=self.prompt_config.get_observation_with_time_user_query.format(
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if time_infer == "无":
|
||||
time_infer = ""
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
time_infer=time_infer,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
|
||||
|
|
@ -1,122 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER, NEW, MSG_TIME, KEY_WORD, \
|
||||
DATATIME_WORD_LIST
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWorker, self).__init__(*args,**kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: "", # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
user_query_list.append(f"{i} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_observation_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.get_observation_few_shot,
|
||||
user_query=self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,64 +0,0 @@
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
def __init__(self, info_filter_msg_max_size, info_filter_model, info_filter_max_token, info_filter_temperature, info_filter_top_k, *args, **kwargs):
|
||||
super(InfoFilterWorker,self).__init__(*args,**kwargs)
|
||||
self.info_filter_msg_max_size
|
||||
self.info_filter_model = info_filter_model
|
||||
self.info_filter_max_token = info_filter_max_token
|
||||
self.info_filter_temperature = info_filter_temperature
|
||||
self.info_filter_top_k = info_filter_top_k
|
||||
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
info_messages = []
|
||||
for msg in self.messages:
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
continue
|
||||
info_messages.append(msg)
|
||||
|
||||
# gene prompt
|
||||
user_query = "\n".join([f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)])
|
||||
info_filter_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.info_filter_system.format(batch_size=len(info_messages)),
|
||||
few_shot=self.prompt_config.info_filter_few_shot,
|
||||
user_query=self.prompt_config.info_filter_user_query.format(user_query=user_query))
|
||||
self.logger.info(f"info_filter_message={info_filter_message}")
|
||||
|
||||
# call llm
|
||||
response_text = self.gene_client.call(messages=info_filter_message,
|
||||
model_name=self.info_filter_model,
|
||||
max_token=self.info_filter_max_token,
|
||||
temperature=self.info_filter_temperature,
|
||||
top_k=self.info_filter_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("info score call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.add_run_info(f"info_score_size != info_messages_size, "
|
||||
f"{len(info_score_list)} vs {len(info_messages)}", continue_run=False)
|
||||
return
|
||||
|
||||
# 过滤value=0的messages
|
||||
filtered_messages = []
|
||||
for msg, info_score in zip(info_messages, info_score_list):
|
||||
if not info_score:
|
||||
continue
|
||||
score = info_score[0]
|
||||
# if score in ("1", "2",):
|
||||
if score in ("2",):
|
||||
msg.info_score = score
|
||||
filtered_messages.append(msg)
|
||||
|
||||
# 后续不会关注为0的msg,直接丢弃
|
||||
self.messages = filtered_messages
|
||||
|
|
@ -1,198 +0,0 @@
|
|||
import re
|
||||
from datetime import datetime
|
||||
from importlib import import_module
|
||||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import WEEKDAYS
|
||||
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
|
||||
def under_line_to_hump(underline_str):
|
||||
sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str)
|
||||
return sub[0:1].upper() + sub[1:]
|
||||
|
||||
|
||||
def parse_response_text_v1(response_text: str) -> dict:
|
||||
"""
|
||||
parse text like:
|
||||
<1> <AAA>
|
||||
<2> <BBB> ddd
|
||||
<4> <CCC> dddd<555>
|
||||
|
||||
result = {1: "AAA", 2: "BBB", 4: "CCC"}
|
||||
"""
|
||||
result_dict: Dict[int, str] = {}
|
||||
|
||||
# 确保第一个数字,后面是string
|
||||
matches = re.findall(r'<(\d+)>\s*<([^>]+)>', response_text.strip())
|
||||
|
||||
# matches 为空返回
|
||||
for key, value in matches:
|
||||
result_dict[int(key)] = value
|
||||
|
||||
return result_dict
|
||||
|
||||
|
||||
def parse_response_text_v2(response_text: str) -> Dict[int, List[str]]:
|
||||
"""
|
||||
parse text like:
|
||||
XXX
|
||||
<1> <AAA> <222>
|
||||
<2> <BBB>
|
||||
<4,5> <CCC>
|
||||
|
||||
result = {1: ["AAA", "222"], 2: "BBB", 4: "CCC"}
|
||||
"""
|
||||
result_dict: Dict[int, List[str]] = {}
|
||||
for line in response_text.strip().split("\n"):
|
||||
if "> <" not in line:
|
||||
continue
|
||||
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
|
||||
idx: str = ll[0]
|
||||
values: List[str] = ll[1:]
|
||||
if idx.isdigit():
|
||||
idx_int = int(idx)
|
||||
else:
|
||||
idx_split = idx.split(",")
|
||||
if len(idx_split) == 0:
|
||||
continue
|
||||
idx = idx_split[0]
|
||||
if idx.isdigit():
|
||||
idx_int = int(idx)
|
||||
else:
|
||||
continue
|
||||
if values:
|
||||
result_dict[idx_int] = values
|
||||
|
||||
return result_dict
|
||||
|
||||
|
||||
def parse_response_text_v3(response_text: str) -> List[List[str]]:
|
||||
"""
|
||||
parse text like:
|
||||
XXX
|
||||
<1> <AAA>
|
||||
<2c> <BBB>
|
||||
<41> <CCC> <BBB>
|
||||
|
||||
result = [["1", "AAA"], ["2c", "BBB"], ["41", "CCC", "BBB"]]
|
||||
"""
|
||||
result_list: List[List[str]] = []
|
||||
for line in response_text.strip().split("\n"):
|
||||
if "> <" not in line:
|
||||
continue
|
||||
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
|
||||
result_list.append(ll)
|
||||
return result_list
|
||||
|
||||
|
||||
def get_datetime_info_dict(parse_dt: datetime):
|
||||
return {
|
||||
"year": parse_dt.year,
|
||||
"month": parse_dt.month,
|
||||
"day": parse_dt.day,
|
||||
"hour": parse_dt.hour,
|
||||
"minute": parse_dt.minute,
|
||||
"second": parse_dt.second,
|
||||
"week": parse_dt.isocalendar().week,
|
||||
"weekday": WEEKDAYS[parse_dt.isocalendar().weekday - 1],
|
||||
}
|
||||
|
||||
|
||||
def extract_date_parts(input_string: str):
|
||||
# Extending our pattern to handle "每" (every) as a possible value.
|
||||
patterns = {
|
||||
'year': r'(\d+|每)年',
|
||||
'month': r'(\d+|每)月',
|
||||
'day': r'(\d+|每)日',
|
||||
'weekday': r'周([一二三四五六日])?',
|
||||
'hour': r'(\d+)点'
|
||||
}
|
||||
weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7}
|
||||
extracted_data = {}
|
||||
|
||||
# Search for patterns in the input string and populate the dictionary
|
||||
for key, pattern in patterns.items():
|
||||
match = re.search(pattern, input_string)
|
||||
if match: # If there is a match, include it in the output dictionary
|
||||
if match.group(1) == "每":
|
||||
extracted_data[key] = -1
|
||||
elif match.group(1) in weekday_dict.keys():
|
||||
extracted_data[key] = weekday_dict[match.group(1)]
|
||||
else:
|
||||
extracted_data[key] = int(match.group(1))
|
||||
return extracted_data
|
||||
|
||||
|
||||
def time_to_formatted_str(time: datetime | str | int | float = None,
|
||||
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
|
||||
string_format: str = "") -> str:
|
||||
if isinstance(time, str | int | float):
|
||||
if isinstance(time, str):
|
||||
time = float(time)
|
||||
current_dt = datetime.fromtimestamp(time)
|
||||
elif isinstance(time, datetime):
|
||||
current_dt = time
|
||||
else:
|
||||
current_dt = datetime.now()
|
||||
|
||||
return_str = ""
|
||||
if date_format:
|
||||
return_str = current_dt.strftime(date_format)
|
||||
elif string_format:
|
||||
return_str = string_format.format(**get_datetime_info_dict(current_dt))
|
||||
|
||||
return return_str
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict|object, default_module_path: str = None, try_kwargs: dict = {}, accept_types: type = None):
|
||||
if isinstance(config, accept_types):
|
||||
return config
|
||||
|
||||
import_module(config.pop("path", default_module_path))
|
||||
clazz = getattr(module, config.pop("name"))
|
||||
try:
|
||||
return clazz(**config, **try_kwargs)
|
||||
except:
|
||||
return clazz(**config)
|
||||
|
||||
|
||||
def init_instance_by_config_v2(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs):
|
||||
clazz_path = config.pop("clazz")
|
||||
if not clazz_path:
|
||||
raise RuntimeError("empty clazz_path!")
|
||||
clazz_name_split = clazz_path.split(".")
|
||||
clazz_name: str = clazz_name_split[-1]
|
||||
if suffix_name and not clazz_name.endswith(suffix_name):
|
||||
clazz_name = f"{clazz_name}_{suffix_name}"
|
||||
|
||||
# 构造path
|
||||
clazz_paths = []
|
||||
if default_clazz_path:
|
||||
clazz_paths.append(default_clazz_path)
|
||||
clazz_paths.extend(clazz_name_split[:-1])
|
||||
clazz_paths.append(clazz_name)
|
||||
module = import_module(".".join(clazz_paths))
|
||||
|
||||
cls_name = under_line_to_hump(clazz_name)
|
||||
return getattr(module, cls_name)(**config, **kwargs)
|
||||
|
||||
|
||||
def complete_config_name(config_name: str, suffix: str = ".json"):
|
||||
if not config_name.endswith(suffix):
|
||||
config_name += suffix
|
||||
return config_name
|
||||
|
||||
|
||||
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
|
||||
return [
|
||||
{
|
||||
"role": MessageRoleEnum.SYSTEM.value,
|
||||
"content": system_prompt.strip(),
|
||||
},
|
||||
{
|
||||
"role": MessageRoleEnum.USER.value,
|
||||
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
|
||||
},
|
||||
]
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
||||
class UserAttribute(BaseModel):
|
||||
"""
|
||||
用户画像的一条属性,和数据库保持一致,只会选择status为1的属性透传过来。
|
||||
status会透传过来。
|
||||
如果code为空,则为新增,否则是更新。
|
||||
确保请求是10条,返回是原始10条+加上新增的条数(如果可以新增)。只会对正确的请求操作数据库。
|
||||
"""
|
||||
id: str = Field("", description="唯一主键")
|
||||
|
||||
memory_id: str = Field("", description="memory id")
|
||||
|
||||
# 从key改成memory_key
|
||||
memory_key: str = Field("", description="memory key")
|
||||
|
||||
value: List[str] = Field([], description="value")
|
||||
|
||||
is_unique: int = Field(1, description="属性是否唯一,if 1 value只有一个,if 0, value 可以很多个")
|
||||
|
||||
is_mutable: int = Field(1, description="是否可变,if 1,value可变,if 1,不可变(用户定义)")
|
||||
|
||||
memory_type: str = Field("", description="profile, profile_customized")
|
||||
|
||||
description: str = Field("", description="memory id")
|
||||
|
||||
status: int = Field(1,
|
||||
description="0为删除,1为active,状态,算法不感知,只为了保存用户删除的画像,给算法传status为valid的用户画像")
|
||||
|
||||
ext_info: Dict[str, str] = Field({}, description="占位符字典")
|
||||
|
|
@ -1,102 +0,0 @@
|
|||
import json
|
||||
from typing import List, Dict
|
||||
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
|
||||
|
||||
class UserProfileHandler(object):
|
||||
@classmethod
|
||||
def format_content(cls, key: str, description: str, value: str | List[str] = None):
|
||||
if not key.startswith("用户"):
|
||||
key = f"用户的{key}"
|
||||
|
||||
if not description.startswith("用户"):
|
||||
description = f"用户{description}"
|
||||
|
||||
content = f"{key}({description})"
|
||||
|
||||
if value:
|
||||
if isinstance(value, list):
|
||||
value = ",".join(value)
|
||||
content = f"{content}:{value}"
|
||||
|
||||
return content
|
||||
|
||||
"""
|
||||
提供UserAttribute 和 MemoryNode 的相互转化
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def to_nodes(cls,
|
||||
user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None,
|
||||
split_value: bool = False) -> List[MemoryNode]:
|
||||
|
||||
user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
if user_profile:
|
||||
if isinstance(user_profile, list):
|
||||
for user_attr in user_profile:
|
||||
user_profile_dict[user_attr.memory_key] = user_attr
|
||||
elif isinstance(user_profile, dict):
|
||||
user_profile_dict.update(user_profile)
|
||||
|
||||
user_profile_nodes: List[MemoryNode] = []
|
||||
for _, user_attr in user_profile_dict.items():
|
||||
# 获取id
|
||||
_id = user_attr.code
|
||||
if not _id:
|
||||
_id = f"{user_attr.memory_id}_{user_attr.scene}_profile_{user_attr.memory_key}"
|
||||
|
||||
attr_node = MemoryNode.init_from_attrs(id=_id,
|
||||
code=_id,
|
||||
content="",
|
||||
memoryId=user_attr.memory_id,
|
||||
scene=user_attr.scene,
|
||||
memoryType=user_attr.memory_type,
|
||||
content_modified=True,
|
||||
metaData={
|
||||
"memory_key": user_attr.memory_key,
|
||||
"value": json.dumps(user_attr.value, ensure_ascii=False),
|
||||
"is_unique": str(user_attr.is_unique),
|
||||
"is_mutable": str(user_attr.is_mutable),
|
||||
"description": user_attr.description,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"ext_info": json.dumps(user_attr.ext_info,
|
||||
ensure_ascii=False),
|
||||
},
|
||||
status=MemoryNodeStatus.ACTIVE.value)
|
||||
|
||||
if split_value:
|
||||
for value in user_attr.value:
|
||||
content = cls.format_content(user_attr.memory_key, user_attr.description, value)
|
||||
attr_node_copy = attr_node.copy(deep=True)
|
||||
attr_node_copy.memory_node.content = content
|
||||
user_profile_nodes.append(attr_node_copy)
|
||||
else:
|
||||
content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value)
|
||||
attr_scheme.memory_node.content = content
|
||||
user_profile_nodes.append(attr_node)
|
||||
|
||||
return user_profile_nodes
|
||||
|
||||
@classmethod
|
||||
def to_user_attr(cls, user_profile_nodes: List[MemoryNode]) -> Dict[str, UserAttribute]:
|
||||
user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
|
||||
for node in user_profile_nodes:
|
||||
user_attr = UserAttribute(
|
||||
code=node.id,
|
||||
memory_id=scheme.memory_node.memoryId,
|
||||
scene=scheme.memory_node.scene,
|
||||
memory_key=scheme.memory_node.metaData["memory_key"],
|
||||
value=json.loads(scheme.memory_node.metaData["value"]),
|
||||
is_unique=int(scheme.memory_node.metaData["is_unique"]),
|
||||
is_mutable=int(scheme.memory_node.metaData["is_mutable"]),
|
||||
memory_type=scheme.memory_node.memoryType,
|
||||
description=scheme.memory_node.metaData["description"],
|
||||
status=1 if scheme.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
|
||||
ext_info=json.loads(scheme.memory_node.metaData["ext_info"]),
|
||||
)
|
||||
user_profile_dict[user_attr.memory_key] = user_attr
|
||||
return user_profile_dict
|
||||
|
|
@ -1,73 +0,0 @@
|
|||
from typing import Any, Dict
|
||||
|
||||
from ..utils.logger import Logger
|
||||
from ..utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
||||
def __init__(self, raise_exception: bool = True, **kwargs):
|
||||
super(BaseWorker, self).__init__(**kwargs)
|
||||
# 异常是否继续执行
|
||||
self.raise_exception: bool = raise_exception
|
||||
|
||||
# True 为正常运行,False会结束整个pipeline
|
||||
self.continue_run: bool = True
|
||||
|
||||
# 短name
|
||||
self._name_simple: str = ""
|
||||
|
||||
# 是否多线程环境
|
||||
self.is_multi_thread: bool = False
|
||||
|
||||
# pipeline 上下文
|
||||
self.context_dict: Dict[str, Any] | None = None
|
||||
self.context_lock = None
|
||||
|
||||
# 日志
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
# worker 参数
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
def _run(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
self.logger.info(f"----- Begin {self.name_simple} -----")
|
||||
with Timer(self.name_simple, log_time=False) as t:
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
else:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as e:
|
||||
self.logger.exception(f"run {self.name_simple} failed! args={e.args}")
|
||||
|
||||
self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----")
|
||||
|
||||
def set_context_dict(self, context_dict: dict, context_lock=None):
|
||||
self.context_dict = context_dict
|
||||
if context_lock is not None:
|
||||
self.context_lock = context_lock
|
||||
self.is_multi_thread = True
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context_dict.get(key, default)
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
# add lock to multi thread
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
else:
|
||||
self.context_dict[key] = value
|
||||
|
||||
def __getattr__(self, key):
|
||||
return self.kwargs[key]
|
||||
|
||||
@property
|
||||
def name_simple(self) -> str:
|
||||
if not self._name_simple:
|
||||
self._name_simple = self.__class__.__name__.replace("Worker", "")
|
||||
return self._name_simple
|
||||
|
|
@ -1,6 +0,0 @@
|
|||
from memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
pass
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import INSIGHT_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
from cli import GLOBAL_CONTEXT
|
||||
|
||||
|
||||
class EsInsightWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
insight_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_insight_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
},
|
||||
)
|
||||
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
|
||||
self.set_context(INSIGHT_NODES, insight_nodes)
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import NEW, NEW_OBS_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNewObsWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
new_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_data.{NEW}": "1",
|
||||
},
|
||||
)
|
||||
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,29 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNotReflectedWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
|
||||
not_reflected_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
f"meta_data.{REFLECTED}": "0",
|
||||
},
|
||||
)
|
||||
self.logger.info(
|
||||
f"retrieve_not_reflected_obs.size={len(not_reflected_obs_nodes)}"
|
||||
)
|
||||
self.set_context(NOT_REFLECTED_OBS_NODES, not_reflected_obs_nodes)
|
||||
|
|
@ -1,37 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsSimilarWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_similar_top_k, *args, **kwargs):
|
||||
super(EsSimilarWorker, self).__init__(*args, **kwargs)
|
||||
self.es_similar_top_k = es_similar_top_k
|
||||
|
||||
def _run(self):
|
||||
query = self.messages[-1].content
|
||||
similar_obs_nodes = self.vector_store.retrieve(
|
||||
text=query,
|
||||
size=self.es_similar_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
for node in similar_obs_nodes:
|
||||
node.meta_data[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
|
||||
for node in similar_obs_nodes:
|
||||
self.logger.info(f"node={node.content} score_similar={node.score_similar}")
|
||||
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)
|
||||
|
|
@ -1,32 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from utils.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsTodayObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_today_obs_top_k, *args, **kwargs):
|
||||
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
|
||||
self.es_today_obs_top_k = es_today_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
if not self.messages:
|
||||
self.logger.warning("messages is empty!")
|
||||
return
|
||||
msg_time_created = self.messages[-1].time_created
|
||||
today_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.es_today_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_Data.{DT}": time_to_formatted_str(msg_time_created),
|
||||
},
|
||||
)
|
||||
|
||||
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
|
||||
self.set_context(TODAY_OBS_NODES, today_obs_nodes)
|
||||
|
|
@ -1,25 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants import common_constants
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LoadProfileWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
user_profile_node = self.vector_store(
|
||||
size=10000,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
self.set_context(common_constants.USER_PROFILE, user_profile_node)
|
||||
self.logger.info(f"retrieve_user_profile.size={len(user_profile_node)}")
|
||||
|
|
@ -1,74 +0,0 @@
|
|||
import re
|
||||
|
||||
from utils.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import (
|
||||
DATATIME_WORD_LIST,
|
||||
DATATIME_KEY_MAP,
|
||||
EXTRACT_TIME_DICT,
|
||||
)
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
# TODO add en version
|
||||
@staticmethod
|
||||
def get_parse_time_prompt(query: str, query_time_str: str):
|
||||
return f"""
|
||||
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。
|
||||
语句:{query}
|
||||
时间:{query_time_str}
|
||||
回答:
|
||||
""".strip()
|
||||
|
||||
def _run(self):
|
||||
# save to context
|
||||
extract_time_dict = {}
|
||||
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)
|
||||
|
||||
# get query & time_created_dt
|
||||
query = self.messages[-1].content
|
||||
time_created = self.messages[-1].time_created
|
||||
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in DATATIME_WORD_LIST:
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# prepare prompt
|
||||
# TODO add en version
|
||||
time_format = "{year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。"
|
||||
query_time_str = time_to_formatted_str(
|
||||
time=time_created, date_format="", string_format=time_format
|
||||
)
|
||||
extract_time_prompt = self.get_parse_time_prompt(
|
||||
query=query, query_time_str=query_time_str
|
||||
)
|
||||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
|
||||
# call sft model
|
||||
|
||||
response_text = self.generation_model.call(
|
||||
prompt=extract_time_prompt,
|
||||
model_name=self.parse_time_model,
|
||||
max_token=self.parse_time_max_token,
|
||||
temperature=self.parse_time_temperature,
|
||||
top_k=self.parse_time_top_k,
|
||||
)
|
||||
|
||||
# if empty, return
|
||||
if not response_text:
|
||||
return
|
||||
|
||||
# re-match time info to dict
|
||||
pattern = r"-\s*(\S+):(\d+)"
|
||||
matches = re.findall(pattern, response_text)
|
||||
for key, value in matches:
|
||||
if key in DATATIME_KEY_MAP.keys():
|
||||
extract_time_dict[DATATIME_KEY_MAP[key]] = value
|
||||
|
||||
self.logger.info(f"response_text={response_text} filters={extract_time_dict}")
|
||||
|
|
@ -1,120 +0,0 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import (
|
||||
RELATED_MEMORIES,
|
||||
EXTRACT_TIME_DICT,
|
||||
ALL_ONLINE_NODES,
|
||||
TIME_MATCHED,
|
||||
)
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
@staticmethod
|
||||
def format_time_infer(
|
||||
time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]
|
||||
):
|
||||
if time_infer:
|
||||
return time_infer
|
||||
|
||||
time_infer = ""
|
||||
if "year" in extract_time_dict:
|
||||
value = meta_data.get("msg_year")
|
||||
if value:
|
||||
time_infer += f"{value}年"
|
||||
elif value == "-1":
|
||||
time_infer += "每年"
|
||||
|
||||
if "month" in extract_time_dict:
|
||||
value = meta_data.get("msg_month")
|
||||
if value:
|
||||
time_infer += f"{value}月"
|
||||
elif value == "-1":
|
||||
time_infer += "每月"
|
||||
|
||||
if "day" in extract_time_dict:
|
||||
value = meta_data.get("msg_day")
|
||||
if value:
|
||||
time_infer += f"{value}日"
|
||||
elif value == "-1":
|
||||
time_infer += "每日"
|
||||
|
||||
if "weekday" in extract_time_dict:
|
||||
value = meta_data.get("msg_weekday")
|
||||
if value:
|
||||
time_infer += value
|
||||
|
||||
return time_infer
|
||||
|
||||
def _run(self):
|
||||
# 解析时间meta信息
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
all_online_nodes: List[MemoryNode] = self.get_context(ALL_ONLINE_NODES)
|
||||
|
||||
if not all_online_nodes:
|
||||
self.add_run_info("all_online_nodes is empty, stop")
|
||||
return
|
||||
|
||||
filtered_nodes = []
|
||||
for node in all_online_nodes:
|
||||
if node.score_rank < self.fuse_score_threshold:
|
||||
continue
|
||||
|
||||
# 根据类型给ratio
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
|
||||
|
||||
# 时间系数,完全匹配才行
|
||||
fuse_time_ratio: float = 1.0
|
||||
match_event_flag = False
|
||||
match_msg_flag = False
|
||||
if extract_time_dict:
|
||||
match_event_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
event_value = node.meta_data.get(f"event_{k}", "")
|
||||
if event_value in ["-1", v]:
|
||||
continue
|
||||
else:
|
||||
match_event_flag = False
|
||||
break
|
||||
|
||||
match_msg_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
msg_value = node.meta_data.get(f"msg_{k}", "")
|
||||
if msg_value == v:
|
||||
continue
|
||||
else:
|
||||
match_msg_flag = False
|
||||
break
|
||||
|
||||
if match_event_flag or match_msg_flag:
|
||||
fuse_time_ratio = self.fuse_time_ratio
|
||||
node.meta_data[TIME_MATCHED] = "1"
|
||||
|
||||
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
|
||||
self.logger.info(
|
||||
f"content={node.content} f_event={int(match_event_flag)} "
|
||||
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}"
|
||||
)
|
||||
filtered_nodes.append(node)
|
||||
|
||||
# get output & save context
|
||||
filtered_nodes = sorted(
|
||||
filtered_nodes, key=lambda x: x.score_rerank, reverse=True
|
||||
)
|
||||
filtered_nodes = filtered_nodes[: self.output_max_count]
|
||||
related_memories: List[str] = []
|
||||
for node in filtered_nodes:
|
||||
content = node.content
|
||||
|
||||
# 如果命中时间逻辑
|
||||
if node.meta_data.get(TIME_MATCHED, "") == "1":
|
||||
time_infer = self.format_time_infer(
|
||||
time_infer="",
|
||||
extract_time_dict=extract_time_dict,
|
||||
meta_data=node.meta_data,
|
||||
)
|
||||
content = f"{time_infer}: {content}"
|
||||
related_memories.append(content)
|
||||
|
||||
self.set_context(RELATED_MEMORIES, related_memories)
|
||||
|
|
@ -1,35 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class MemoryStoreWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(
|
||||
MODIFIED_MEMORIES
|
||||
)
|
||||
if modified_memories:
|
||||
if isinstance(modified_memories[0], MemoryNode):
|
||||
modified_memories = [n.memory_node for n in modified_memories]
|
||||
|
||||
for n in modified_memories:
|
||||
if not n.id:
|
||||
n.id = f"{n.memory_id}_content_{n.content}"
|
||||
n.code = n.id
|
||||
# TODO add batch insert
|
||||
n.meta_data.pop(CONTENT_MODIFIED)
|
||||
self.vector_store.insert(n)
|
||||
else:
|
||||
self.logger.warning("modified_memories is empty!")
|
||||
|
||||
new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE)
|
||||
if new_user_profile:
|
||||
for n in new_user_profile:
|
||||
n.meta_data.pop(CONTENT_MODIFIED)
|
||||
self.vector_store.insert(n)
|
||||
else:
|
||||
self.logger.warning("new_user_profile is empty!")
|
||||
|
|
@ -1,62 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
|
||||
QUERY_KEYWORDS
|
||||
from enumeration.memory_recall_enum import MemoryRecallType
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
||||
def user_profile_to_nodes(self) -> List[MemoryNode]:
|
||||
user_profile_nodes: List[MemoryNode] = self.user_profile_dict
|
||||
for node in user_profile_nodes:
|
||||
# 从画像侧召回
|
||||
node.meta_data[RECALL_TYPE] = MemoryRecallType.PROFILE
|
||||
self.logger.info(f"user profile node={node.content}")
|
||||
return user_profile_nodes
|
||||
|
||||
def _run(self):
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
# 优先级: similar_obs_nodes < profile_nodes
|
||||
similar_obs_nodes: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES)
|
||||
if similar_obs_nodes:
|
||||
for node in similar_obs_nodes:
|
||||
all_node_dict[node.content] = node
|
||||
|
||||
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
|
||||
if profile_nodes:
|
||||
for node in profile_nodes:
|
||||
all_node_dict[node.content] = node
|
||||
|
||||
if not all_node_dict:
|
||||
self.add_run_info("all_node_dict is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# call recall model
|
||||
query_keywords = self.get_context(QUERY_KEYWORDS)
|
||||
# TODO 根据效果更改
|
||||
# query: str = "用户:" + self.messages[-1].content
|
||||
query: str = self.messages[-1].content
|
||||
if query_keywords:
|
||||
query_keyword_join = ",".join(query_keywords)
|
||||
query = f"{query} 用户的{query_keyword_join}。"
|
||||
documents = list(all_node_dict.keys())
|
||||
result = self.rank_model.call(query=query, documents=documents)
|
||||
|
||||
if not result:
|
||||
self.add_run_info("semantic call recall model failed!")
|
||||
return
|
||||
|
||||
# set score
|
||||
for index, score in result.rank_scores.items():
|
||||
content = documents[index]
|
||||
node = all_node_dict[content]
|
||||
node.score_rank = score
|
||||
self.logger.info(f"query={query} content={node.content} score_rank={node.score_rank}")
|
||||
|
||||
# save context
|
||||
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())
|
||||
self.set_context(ALL_ONLINE_NODES, all_online_nodes)
|
||||
|
|
@ -1,166 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from ...constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
DT,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
NEW_INSIGHT_KEYS,
|
||||
INSIGHT_KEY,
|
||||
INSIGHT_VALUE,
|
||||
REFLECTED,
|
||||
CONTENT_MODIFIED
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_insight_prompt import (
|
||||
GET_INSIGHT_FEW_SHOT_PROMPT,
|
||||
GET_INSIGHT_SYSTEM_PROMPT,
|
||||
GET_INSIGHT_USER_QUERY_PROMPT
|
||||
)
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
DT: dt,
|
||||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
CONTENT_MODIFIED: True, # 新增的insight需要置为true
|
||||
}
|
||||
meta_data.update(
|
||||
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryNode(
|
||||
content=content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.INSIGHT.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def reflect_new_insight_key(
|
||||
self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode]
|
||||
) -> MemoryNode | None:
|
||||
|
||||
# 检索历史memory
|
||||
hits = self.vector_store.similar_search(
|
||||
text=insight_key,
|
||||
size=self.es_insight_similar_top_k,
|
||||
exact_filters={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
# 转化成 MemoryNodeWrap 合并新增nodes
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
related_nodes.extend(not_reflected_merge_nodes)
|
||||
|
||||
# content去重
|
||||
related_node_dict = {n.memory_node.content: n for n in related_nodes}
|
||||
related_nodes = sorted(
|
||||
list(related_node_dict.values()), key=lambda x: x.memory_node.id
|
||||
)
|
||||
documents = [n.memory_node.content for n in related_nodes]
|
||||
|
||||
# 重排所有记忆
|
||||
result = self.rank_model.call(query=insight_key, documents=documents)
|
||||
if not result:
|
||||
self.add_run_info(
|
||||
f"reflect insight_key={insight_key} call rerank client failed!"
|
||||
)
|
||||
return
|
||||
|
||||
# 根据打分过滤
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
related_nodes[index].score_rank = score
|
||||
related_nodes_sorted = sorted(
|
||||
related_nodes, key=lambda x: x.score_rank, reverse=True
|
||||
)[: self.insight_obs_max_cnt]
|
||||
|
||||
# 生成prompt
|
||||
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
|
||||
get_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_INSIGHT_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(GET_INSIGHT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_INSIGHT_USER_QUERY_PROMPT).format(
|
||||
insight_key=insight_key, user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"get_insight_message={get_insight_message}")
|
||||
|
||||
# call LLM, 提取insight
|
||||
response_text = self.generation_model.call(
|
||||
messages=get_insight_message,
|
||||
model_name=self.get_insight_model,
|
||||
max_token=self.get_insight_max_token,
|
||||
temperature=self.get_insight_temperature,
|
||||
top_k=self.get_insight_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_upon_user_attr call llm failed!")
|
||||
return
|
||||
response_text = response_text.strip()
|
||||
if response_text in ["无"]:
|
||||
return
|
||||
return self.new_insight_node(
|
||||
insight_key=insight_key, insight_value=response_text
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
|
||||
if not new_insight_keys:
|
||||
self.add_run_info("new_insight_keys is empty! stop insight.")
|
||||
return
|
||||
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
if not not_reflected_merge_nodes:
|
||||
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
|
||||
return
|
||||
|
||||
# submit insight task
|
||||
for insight_key in new_insight_keys:
|
||||
self.submit_thread(
|
||||
self.reflect_new_insight_key,
|
||||
sleep_time=1,
|
||||
insight_key=insight_key,
|
||||
not_reflected_merge_nodes=not_reflected_merge_nodes,
|
||||
)
|
||||
|
||||
# save output
|
||||
new_insight_nodes: List[MemoryNode] = []
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
new_insight_nodes.append(result)
|
||||
assert isinstance(result, MemoryNode)
|
||||
insight_key = result.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = result.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"after_get_insight insight_key={insight_key} insight_value={insight_value}"
|
||||
)
|
||||
|
||||
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
|
||||
|
||||
# set REFLECTED
|
||||
for node in not_reflected_merge_nodes:
|
||||
node.meta_data[REFLECTED] = "1"
|
||||
|
|
@ -1,99 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utilsresponse_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
REFLECTED,
|
||||
INSIGHT_NODES,
|
||||
INSIGHT_KEY,
|
||||
NEW_INSIGHT_KEYS,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_reflection_prompt import (
|
||||
GET_REFLECTION_FEW_SHOT_PROMPT,
|
||||
GET_REFLECTION_SYSTEM_PROMPT,
|
||||
GET_REFLECTION_USER_QUERY_PROMPT
|
||||
)
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
not_reflected_merge_nodes.extend(new_obs_nodes)
|
||||
if not_reflected_nodes:
|
||||
not_reflected_merge_nodes.extend(not_reflected_nodes)
|
||||
not_reflected_merge_nodes = [
|
||||
node
|
||||
for node in not_reflected_merge_nodes
|
||||
if node.meta_data.get(REFLECTED, "") == "0"
|
||||
]
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(
|
||||
f"not_reflected_count={not_reflected_count} is not enough, stop reflect."
|
||||
)
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = []
|
||||
profile_keys: List[str] = list(self.user_profile_dict.keys())
|
||||
exist_keys.extend(profile_keys)
|
||||
self.logger.info(f"profile_keys={profile_keys}")
|
||||
|
||||
# get insight_keys
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if insight_nodes:
|
||||
insight_keys = [
|
||||
n.meta_data.get(INSIGHT_KEY) for n in insight_nodes
|
||||
]
|
||||
insight_keys = [x.strip() for x in insight_keys if x]
|
||||
exist_keys.extend(insight_keys)
|
||||
self.logger.info(f"insight_keys={insight_keys}")
|
||||
|
||||
# gen reflect prompt
|
||||
user_query_list = [n.content for n in not_reflected_merge_nodes]
|
||||
reflect_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_REFLECTION_SYSTEM_PROMPT).format(
|
||||
num_questions=self.reflect_num_questions
|
||||
),
|
||||
few_shot=self.get_prompt(GET_REFLECTION_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_REFLECTION_USER_QUERY_PROMPT).format(
|
||||
exist_keys=",".join(exist_keys), user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# # call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=reflect_message,
|
||||
model_name=self.reflect_obs_model,
|
||||
max_token=self.reflect_obs_max_token,
|
||||
temperature=self.reflect_obs_temperature,
|
||||
top_k=self.reflect_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_obs_questions call llm failed!")
|
||||
return
|
||||
|
||||
# parse text & save
|
||||
new_insight_keys = ResponseTextParser(response_text).parse_v2(
|
||||
"get_insight_keys"
|
||||
)
|
||||
if new_insight_keys:
|
||||
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)
|
||||
|
|
@ -1,129 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
MSG_TIME,
|
||||
MODIFIED_MEMORIES,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.long_contra_repeat_prompt import (
|
||||
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT,
|
||||
LONG_CONTRA_REPEAT_SYSTEM_PROMPT,
|
||||
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_node.content
|
||||
related_nodes = self.vector_store.similar_search(
|
||||
text=text,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
all_obs_nodes.append(related_node)
|
||||
if has_match:
|
||||
all_obs_nodes.append(new_obs_node)
|
||||
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(
|
||||
all_obs_nodes,
|
||||
key=lambda x: x.meta_data.get(MSG_TIME, ""),
|
||||
reverse=True,
|
||||
)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(LONG_CONTRA_REPEAT_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(LONG_CONTRA_REPEAT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {node.content} {node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,47 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from ...constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
NEW,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update(
|
||||
{n.id: n for n in insight_nodes if n.meta_data.get(CONTENT_MODIFIED, False)}
|
||||
)
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.meta_data[NEW] = "0"
|
||||
all_node_dict.update({n.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
|
|
@ -1,177 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
INSIGHT_KEY,
|
||||
INSIGHT_VALUE,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.update_insight_prompt import (
|
||||
UPDATE_INSIGHT_FEW_SHOT_PROMPT,
|
||||
UPDATE_INSIGHT_SYSTEM_PROMPT,
|
||||
UPDATE_INSIGHT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
||||
def filter_obs_nodes(
|
||||
self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode]
|
||||
) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
if not insight_key or not insight_value:
|
||||
self.logger.warning(
|
||||
f"insight_key={insight_key} insight_value={insight_value} is empty!"
|
||||
)
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
result = self.rank_model.call(
|
||||
query=insight_key, documents=[x.content for x in new_obs_nodes]
|
||||
)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
|
||||
for index, score in result.rank_scores.items():
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_insight_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(
|
||||
f"insight_key={insight_key} insight_value={insight_value} "
|
||||
f"score={score} keep_flag={keep_flag}"
|
||||
)
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
|
||||
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(
|
||||
self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]
|
||||
) -> MemoryNode:
|
||||
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"update_insight insight_key={insight_key} insight_value={insight_value} "
|
||||
f"doc.size={len(filtered_nodes)}"
|
||||
)
|
||||
|
||||
# gen prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{node.content}")
|
||||
update_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_INSIGHT_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_INSIGHT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_INSIGHT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key + ":" + insight_value,
|
||||
),
|
||||
)
|
||||
self.logger.info(f"update_insight_message={update_insight_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=update_insight_message,
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} call llm failed!"
|
||||
)
|
||||
return insight_node
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(
|
||||
f"update_profile {insight_key}"
|
||||
)
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 1!"
|
||||
)
|
||||
return insight_node
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 2"
|
||||
)
|
||||
return insight_node
|
||||
insight_value = profile_list[0]
|
||||
|
||||
if not insight_value or insight_value in ["无", "重复"]:
|
||||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
insight_node.meta_data[INSIGHT_VALUE] = insight_value
|
||||
insight_node.meta_data[CONTENT_MODIFIED] = True
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop update sights!")
|
||||
return
|
||||
if not insight_nodes:
|
||||
self.logger.info("insight_nodes is empty, stop update sights!")
|
||||
return
|
||||
|
||||
# 提交打分任务
|
||||
for node in insight_nodes:
|
||||
self.submit_thread(
|
||||
self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
insight_node=node,
|
||||
new_obs_nodes=new_obs_nodes,
|
||||
)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
insight_node, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[: self.update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(
|
||||
self.update_insight,
|
||||
sleep_time=1,
|
||||
insight_node=insight_node,
|
||||
filtered_nodes=filtered_nodes,
|
||||
)
|
||||
|
||||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryNode = result
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"after_update_insight insight_key={insight_key} insight_value={insight_value}"
|
||||
)
|
||||
|
|
@ -1,241 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ....memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.update_profile_prompt import (
|
||||
UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT,
|
||||
UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT,
|
||||
UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT
|
||||
)
|
||||
from ...chat.global_context import GlobalContext
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
@property
|
||||
def extra_user_attrs(self):
|
||||
return GlobalContext.global_configs.get("extra_user_attrs", [])
|
||||
|
||||
def filter_obs_nodes(
|
||||
self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode]
|
||||
) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
result = self.rank_model.call(
|
||||
query=user_attr.meta_data.get("description", ""),
|
||||
documents=[x.content for x in new_obs_nodes],
|
||||
)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(
|
||||
f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!"
|
||||
)
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
for index, score in result.rank_scores.items():
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_profile_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(
|
||||
f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
|
||||
f"content={node.content} score={score} keep_flag={keep_flag}"
|
||||
)
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
def update_user_attr(
|
||||
self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode]
|
||||
) -> MemoryNode:
|
||||
self.logger.info(
|
||||
f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
|
||||
f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}"
|
||||
)
|
||||
|
||||
# 根据不同的参数类型是否多值,分别给出prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{node.content}")
|
||||
update_profile = f"{user_attr.meta_data.get("memory_key", "")}({user_attr.meta_data.get("description", "")})"
|
||||
update_profile_value = update_profile + ":" + ",".join(user_attr.meta_data.get("value", ""))
|
||||
|
||||
if user_attr.meta_data.get("is_unique", 0) == 1:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value,
|
||||
),
|
||||
)
|
||||
else:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value,
|
||||
),
|
||||
)
|
||||
self.logger.info(f"update_profile_message={update_profile_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=update_profile_message,
|
||||
model_name=self.update_profile_model,
|
||||
max_token=self.update_profile_max_token,
|
||||
temperature=self.update_profile_temperature,
|
||||
top_k=self.update_profile_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!"
|
||||
)
|
||||
return user_attr
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(
|
||||
f"update_attr {user_attr.meta_data.get("memory_key", "")}"
|
||||
)
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!"
|
||||
)
|
||||
return user_attr
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2"
|
||||
)
|
||||
return user_attr
|
||||
profile = profile_list[0]
|
||||
|
||||
if not profile or profile in ["无", "重复"]:
|
||||
self.logger.info(f"profile={profile}, skip.")
|
||||
return user_attr
|
||||
|
||||
# check 英文中午逗号
|
||||
if user_attr.meta_data.get("is_unique", 0) == 1:
|
||||
user_attr.meta_data["value"] = [profile.strip()]
|
||||
else:
|
||||
attr_value_list = profile.replace(",", ",").split(",")
|
||||
user_attr.meta_data["value"] = [
|
||||
x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list)))
|
||||
]
|
||||
return user_attr
|
||||
|
||||
def add_extra_user_attrs(self):
|
||||
# 解析为空返回
|
||||
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
|
||||
if not extra_user_attr_list:
|
||||
return
|
||||
|
||||
for user_attr_info in extra_user_attr_list:
|
||||
user_attr_split = user_attr_info.split(":")
|
||||
|
||||
# 格式不对返回
|
||||
if len(user_attr_split) < 1:
|
||||
continue
|
||||
user_attr_key = user_attr_split[0]
|
||||
|
||||
user_attr_desc = ""
|
||||
if len(user_attr_split) >= 2:
|
||||
user_attr_desc = user_attr_split[1]
|
||||
|
||||
user_attr_unique = 0
|
||||
if len(user_attr_split) >= 3:
|
||||
user_attr_unique = int(user_attr_split[2])
|
||||
|
||||
# 已经包含返回
|
||||
if user_attr_key in self.user_profile_dict:
|
||||
user_attr = self.user_profile_dict[user_attr_key]
|
||||
# description为空,补充description
|
||||
if not user_attr.meta_data.get("description", ""):
|
||||
user_attr.meta_data["description"] = user_attr_desc
|
||||
continue
|
||||
|
||||
# 增加新属性
|
||||
new_attr = MemoryNode(
|
||||
memory_id=self.memory_id,
|
||||
meta_data={
|
||||
"memory_key": user_attr_key,
|
||||
"is_unique": int(user_attr_unique),
|
||||
"is_mutable": 1,
|
||||
"description": user_attr_desc
|
||||
},
|
||||
memory_type=MemoryTypeEnum.PROFILE,
|
||||
status=1,
|
||||
)
|
||||
self.user_profile_dict[user_attr_key] = new_attr
|
||||
|
||||
def _run(self):
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop user profile!")
|
||||
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
|
||||
return
|
||||
|
||||
# 增加环境变量配置的属性
|
||||
if self.extra_user_attrs:
|
||||
self.add_extra_user_attrs()
|
||||
|
||||
new_user_profile: List[MemoryNode] = []
|
||||
self.set_context(NEW_USER_PROFILE, new_user_profile)
|
||||
|
||||
for user_attr_key, user_attr in self.user_profile_dict.items():
|
||||
# 不可修改直接跳过
|
||||
if user_attr.meta_data.get("is_mutable", 0) != 1:
|
||||
new_user_profile.append(user_attr)
|
||||
self.logger.info(f"{user_attr_key} is not mutable! continue")
|
||||
continue
|
||||
|
||||
self.submit_thread(
|
||||
self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
user_attr=user_attr,
|
||||
new_obs_nodes=new_obs_nodes,
|
||||
)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
user_attr, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_profile_max_thread:
|
||||
result_sorted = result_sorted[: self.update_profile_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for user_attr, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(
|
||||
self.update_user_attr,
|
||||
sleep_time=1,
|
||||
user_attr=user_attr,
|
||||
filtered_nodes=filtered_nodes,
|
||||
)
|
||||
|
||||
# collect result & save
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
user_attribute: MemoryNode = result
|
||||
self.logger.info(
|
||||
f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} "
|
||||
f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}"
|
||||
)
|
||||
new_user_profile.append(user_attribute)
|
||||
|
|
@ -1,117 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
TODAY_OBS_NODES,
|
||||
MSG_TIME,
|
||||
NEW_OBS_WITH_TIME_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.contra_repeat_prompt import (
|
||||
CONTRA_REPEAT_FEW_SHOT_PROMPT,
|
||||
CONTRA_REPEAT_SYSTEM_PROMPT,
|
||||
CONTRA_REPEAT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(
|
||||
NEW_OBS_WITH_TIME_NODES
|
||||
)
|
||||
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
all_obs_nodes.extend(new_obs_nodes)
|
||||
if new_obs_with_time_nodes:
|
||||
all_obs_nodes.extend(new_obs_with_time_nodes)
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(
|
||||
all_obs_nodes,
|
||||
key=lambda x: x.meta_data.get(MSG_TIME, ""),
|
||||
reverse=True,
|
||||
)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(CONTRA_REPEAT_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(CONTRA_REPEAT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(CONTRA_REPEAT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(
|
||||
f"after contra repeat: {node.content} {node.status}"
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,167 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...utils.tool_functions import (
|
||||
time_to_formatted_str,
|
||||
get_datetime_info_dict,
|
||||
extract_date_parts,
|
||||
)
|
||||
from ...constants.common_constants import (
|
||||
REFLECTED,
|
||||
DT,
|
||||
TIME_INFER,
|
||||
NEW,
|
||||
MSG_TIME,
|
||||
KEY_WORD,
|
||||
DATATIME_WORD_LIST,
|
||||
NEW_OBS_WITH_TIME_NODES,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ...scheme.message import Message
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_observation_with_time_prompt import (
|
||||
GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT,
|
||||
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT,
|
||||
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
|
||||
def add_observation(
|
||||
self, message: Message, obs_content: str, time_infer: str, keywords: str
|
||||
):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: time_infer, # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
CONTENT_MODIFIED: True, # 新增的obs需要置为true
|
||||
}
|
||||
|
||||
# 事件时间
|
||||
meta_data.update(
|
||||
{f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()}
|
||||
)
|
||||
# 对话时间
|
||||
meta_data.update(
|
||||
{f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
return MemoryNode.init_from_attrs(
|
||||
content=obs_content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
dt = time_to_formatted_str(
|
||||
time=msg.time_created,
|
||||
date_format="",
|
||||
string_format="{year}年{month}月{day}日{weekday}{hour}点",
|
||||
)
|
||||
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(
|
||||
f"get obs with time user_query_list={user_query_list} is empty"
|
||||
)
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(
|
||||
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT
|
||||
).format(num_obs=len(user_query_list)),
|
||||
few_shot=self.get_prompt(GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(
|
||||
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT
|
||||
).format(user_query="\n".join(user_query_list)),
|
||||
)
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if time_infer == "无":
|
||||
time_infer = ""
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(
|
||||
f"idx={idx} is invalid! messages.size={len(self.messages)}"
|
||||
)
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(
|
||||
self.add_observation(
|
||||
message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
time_infer=time_infer,
|
||||
keywords=keywords,
|
||||
)
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
|
||||
|
|
@ -1,144 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from ...constants.common_constants import (
|
||||
REFLECTED,
|
||||
DT,
|
||||
NEW_OBS_NODES,
|
||||
TIME_INFER,
|
||||
NEW,
|
||||
MSG_TIME,
|
||||
KEY_WORD,
|
||||
DATATIME_WORD_LIST,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ...scheme.message import Message
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_observation_prompt import (
|
||||
GET_OBSERVATION_FEW_SHOT_PROMPT,
|
||||
GET_OBSERVATION_SYSTEM_PROMPT,
|
||||
GET_OBSERVATION_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: "", # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
CONTENT_MODIFIED: True, # 新增的obs需要置为true
|
||||
}
|
||||
meta_data.update(
|
||||
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
return MemoryNode(
|
||||
content=obs_content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
user_query_list.append(f"{i} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_OBSERVATION_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(GET_OBSERVATION_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_OBSERVATION_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(
|
||||
f"idx={idx} is invalid! messages.size={len(self.messages)}"
|
||||
)
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(
|
||||
self.add_observation(
|
||||
message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
keywords=keywords,
|
||||
)
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,70 +0,0 @@
|
|||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
from ...chat.global_context import GlobalContext
|
||||
from ...prompts.info_filter_prompt import INFO_FILTER_FEW_SHOT_PROMPT, INFO_FILTER_SYSTEM_PROMPT, INFO_FILTER_USER_QUERY_PROMPT
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
info_messages = []
|
||||
for msg in self.messages:
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
continue
|
||||
info_messages.append(msg)
|
||||
|
||||
# gene prompt
|
||||
user_query = "\n".join(
|
||||
[f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)]
|
||||
)
|
||||
info_filter_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(INFO_FILTER_SYSTEM_PROMPT).format(
|
||||
batch_size=len(info_messages)
|
||||
),
|
||||
few_shot=self.get_prompt(INFO_FILTER_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(INFO_FILTER_USER_QUERY_PROMPT).format(
|
||||
user_query=user_query
|
||||
),
|
||||
)
|
||||
self.logger.info(f"info_filter_message={info_filter_message}")
|
||||
|
||||
# call llm
|
||||
response_text = self.generation_model.call(
|
||||
messages=info_filter_message,
|
||||
model_name=self.info_filter_model,
|
||||
max_token=self.info_filter_max_token,
|
||||
temperature=self.info_filter_temperature,
|
||||
top_k=self.info_filter_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("info score call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.add_run_info(
|
||||
f"info_score_size != info_messages_size, "
|
||||
f"{len(info_score_list)} vs {len(info_messages)}",
|
||||
continue_run=False,
|
||||
)
|
||||
return
|
||||
|
||||
# 过滤value=0的messages
|
||||
filtered_messages = []
|
||||
for msg, info_score in zip(info_messages, info_score_list):
|
||||
if not info_score:
|
||||
continue
|
||||
score = info_score[0]
|
||||
# if score in ("1", "2",):
|
||||
if score in ("2",):
|
||||
msg.info_score = score
|
||||
filtered_messages.append(msg)
|
||||
|
||||
# 后续不会关注为0的msg,直接丢弃
|
||||
self.messages = filtered_messages
|
||||
|
|
@ -2,7 +2,7 @@ import unittest
|
|||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore
|
||||
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
|
||||
|
||||
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
||||
|
|
@ -12,16 +12,18 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
config = {
|
||||
"module_name": "dashscope_embedding",
|
||||
"model_name": "text-embedding-v2",
|
||||
"clazz": "models.llama_index_embedding_model"
|
||||
"clazz": "models.llama_index_embedding_model",
|
||||
}
|
||||
emb = LlamaIndexEmbeddingModel(**config)
|
||||
|
||||
config = {
|
||||
"index_name": "0626_1",
|
||||
"index_name": "0708_2",
|
||||
"es_url": "http://localhost:9200",
|
||||
"embedding_model": emb,
|
||||
"use_hybrid": True
|
||||
|
||||
}
|
||||
self.es_store = LlamaIndexElasticSearchStore(**config)
|
||||
self.es_store = LlamaIndexEsMemoryStore(**config)
|
||||
self.data = [
|
||||
MemoryNode(
|
||||
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
|
||||
|
|
@ -100,15 +102,53 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="hhh234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="iii234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="jjj234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="kkk234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
]
|
||||
|
||||
def test_retrieve(self):
|
||||
filter_dict = {
|
||||
"user_id": "6",
|
||||
}
|
||||
# filter_dict = {
|
||||
# "user_id": "6",
|
||||
# }
|
||||
filter_dict = {}
|
||||
|
||||
for node in self.data:
|
||||
self.es_store.insert(node)
|
||||
|
||||
self.es_store.insert(MemoryNode(
|
||||
content="xxxxxx",
|
||||
memory_type="profile",
|
||||
|
|
@ -117,9 +157,10 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
memory_id="ggg567",
|
||||
meta_data={"5": "5"}
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
||||
self.es_store.update(MemoryNode(
|
||||
content="test update",
|
||||
|
|
@ -128,7 +169,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
status="invalid",
|
||||
memory_id="ggg567"
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
|
@ -140,7 +181,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
memory_id="ggg567"
|
||||
))
|
||||
import asyncio
|
||||
res = asyncio.run(self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10))
|
||||
res = asyncio.run(self.es_store.a_retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15))
|
||||
#res = self.es_store.async_retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
|
|
|||
31
tt.py
Normal file
31
tt.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
import asyncio
|
||||
|
||||
|
||||
class TT(object):
|
||||
def __init__(self):
|
||||
self.task_list = []
|
||||
|
||||
async def async_func(self, i: int):
|
||||
await asyncio.sleep(i) # 模拟异步操作
|
||||
print(f"函数{i}的结果")
|
||||
|
||||
def submit_async_task(self, fn, *args, **kwargs):
|
||||
self.task_list.append((fn, args, kwargs))
|
||||
|
||||
def gather_async_result(self):
|
||||
async def async_gather():
|
||||
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.task_list])
|
||||
|
||||
results = asyncio.run(async_gather())
|
||||
self.task_list.clear()
|
||||
return results
|
||||
|
||||
def run(self):
|
||||
self.submit_async_task(self.async_func, i=1)
|
||||
self.submit_async_task(self.async_func, i=2)
|
||||
self.submit_async_task(self.async_func, i=3)
|
||||
|
||||
self.gather_async_result()
|
||||
|
||||
|
||||
TT().run()
|
||||
Loading…
Add table
Reference in a new issue