mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
memory stream & some workers
This commit is contained in:
parent
55ca839e95
commit
e84733f810
62 changed files with 730 additions and 617 deletions
16
.vscode/launch.json
vendored
Normal file
16
.vscode/launch.json
vendored
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
{
|
||||
// 使用 IntelliSense 了解相关属性。
|
||||
// 悬停以查看现有属性的描述。
|
||||
// 欲了解更多信息,请访问: https://go.microsoft.com/fwlink/?linkid=830387
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "Python 调试程序: 当前文件",
|
||||
"type": "debugpy",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
{
|
||||
"clazz": "models.base_generation_model",
|
||||
"clazz": "models.llama_index_generation_model",
|
||||
"model_name": "qwen-max",
|
||||
"method": "DashScope"
|
||||
"method_type": "DashScope"
|
||||
}
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
{
|
||||
"clazz": "models.base_rank_model",
|
||||
"model_name": "gte-rerank",
|
||||
"method": "DashScopeRerank"
|
||||
"method_type": "DashScopeRerank"
|
||||
}
|
||||
|
|
@ -1,3 +1,3 @@
|
|||
""" Version of MemoryScope."""
|
||||
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
|
|
@ -1,12 +1,10 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
from memory_scope.chat.memory_service import MemoryService
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
def __init__(self, chat_name: str, **kwargs):
|
||||
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
|
|
|
|||
4
memory_scope/chat/base_memory_service.py
Normal file
4
memory_scope/chat/base_memory_service.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
class BaseMemoryService(object):
|
||||
def __init__(self, **kwargs):
|
||||
|
||||
self.kwargs = kwargs
|
||||
83
memory_scope/chat/cli_memory_chat.py
Normal file
83
memory_scope/chat/cli_memory_chat.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
import datetime
|
||||
|
||||
import questionary
|
||||
from rich.console import Console
|
||||
|
||||
from .memory_chat import MemoryChat
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from scheme.message import Message
|
||||
|
||||
|
||||
class CliMemoryChat(MemoryChat):
|
||||
|
||||
USER_COMMANDS = {
|
||||
"/exit": "exit the CLI",
|
||||
"/memory": "print the current contents of agent memory",
|
||||
"/retrieve": "retrieve related memory",
|
||||
"/log": "log chat progress",
|
||||
# TODO add more commands
|
||||
}
|
||||
|
||||
def chat_with_memory(self, query): # for testing
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
message = Message(
|
||||
role=MessageRoleEnum.USER, content=query, time_created=time_created
|
||||
)
|
||||
messages = [message]
|
||||
return self.generation_model.call(messages=messages, stream=True)
|
||||
|
||||
def retrieve_all(self): # for testing
|
||||
return "memory 1. 2. 3."
|
||||
|
||||
def run(self):
|
||||
console = Console()
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Try again!")
|
||||
continue
|
||||
|
||||
# Handle CLI commands
|
||||
if query.startswith("/"):
|
||||
if query.lower() == "/exit":
|
||||
break
|
||||
elif query.lower() == "/memory":
|
||||
console.print(self.memory_service.retrieve_all())
|
||||
elif query.lower() == "/help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
|
||||
continue
|
||||
|
||||
while True:
|
||||
try:
|
||||
# with console.status("[bold cyan]Thinking..."):
|
||||
for msg in self.chat_with_memory(query=query):
|
||||
console.print(msg.delta, end="")
|
||||
console.print()
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
console.print("User interrupt occurred.")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
except Exception as e:
|
||||
console.print(
|
||||
f"An exception occurred when running chat_with_memory(): {e}"
|
||||
)
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
|
|
@ -1,12 +1,12 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
from memory_scope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
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.worker.base_worker import BaseWorker
|
||||
from chat.base_memory_chat import BaseMemoryChat
|
||||
from enumeration.language_enum import LanguageEnum
|
||||
from models.base_model import BaseModel
|
||||
from storage.base_monitor import BaseMonitor
|
||||
from storage.base_vector_store import BaseVectorStore
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class GlobalContext(object):
|
||||
|
|
|
|||
|
|
@ -1,21 +1,20 @@
|
|||
import datetime
|
||||
from typing import List
|
||||
|
||||
from memory_scope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.chat.global_context import GLOBAL_CONTEXT
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
|
||||
from memory_scope.scheme.message import Message
|
||||
from .base_memory_chat import BaseMemoryChat
|
||||
from .global_context import GLOBAL_CONTEXT
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from models.base_model import BaseModel
|
||||
from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
|
||||
from scheme.message import Message
|
||||
from .memory_service import MemoryService
|
||||
|
||||
|
||||
class MemoryChat(BaseMemoryChat):
|
||||
|
||||
def __init__(self,
|
||||
generation_model: str,
|
||||
history_msg_count: int,
|
||||
**kwargs):
|
||||
def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
|
||||
self.generation_model_name: str = generation_model
|
||||
self.history_msg_count: int = history_msg_count
|
||||
|
||||
|
|
@ -25,7 +24,9 @@ class MemoryChat(BaseMemoryChat):
|
|||
@property
|
||||
def generation_model(self):
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict[self.generation_model_name]
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict[
|
||||
self.generation_model_name
|
||||
]
|
||||
return self._generation_model
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -34,7 +35,11 @@ class MemoryChat(BaseMemoryChat):
|
|||
if related_memories:
|
||||
memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language]
|
||||
system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories)
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt.strip(), time_created=time_created)
|
||||
return Message(
|
||||
role=MessageRoleEnum.SYSTEM,
|
||||
content=system_prompt.strip(),
|
||||
time_created=time_created,
|
||||
)
|
||||
|
||||
def chat_with_memory(self, query: str):
|
||||
query = query.strip()
|
||||
|
|
@ -42,11 +47,13 @@ class MemoryChat(BaseMemoryChat):
|
|||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created)
|
||||
new_message: Message = Message(
|
||||
role=MessageRoleEnum.USER, content=query, time_created=time_created
|
||||
)
|
||||
related_memories: List[str] = self.memory_service.retrieve(message=new_message)
|
||||
system_message = self.get_system_prompt(related_memories, time_created)
|
||||
self.history_message_list.append(new_message)
|
||||
self.history_message_list = self.history_message_list[-self.history_msg_count:]
|
||||
self.history_message_list = self.history_message_list[-self.history_msg_count :]
|
||||
all_messages = [system_message] + self.history_message_list
|
||||
# TODO at xian zhe
|
||||
return self.generation_model.call(messages=all_messages, stream=True)
|
||||
|
|
|
|||
|
|
@ -1,43 +1,52 @@
|
|||
from memory_scope.constants.common_constants import RELATED_MEMORIES
|
||||
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.pipeline import Pipeline
|
||||
from constants.common_constants import RELATED_MEMORIES
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from scheme.message import Message
|
||||
from utils.pipeline import Pipeline
|
||||
from .base_memory_service import BaseMemoryService
|
||||
|
||||
|
||||
class MemoryService(object):
|
||||
class MemoryService(BaseMemoryService):
|
||||
def __init__(
|
||||
self,
|
||||
chat_name: str,
|
||||
retrieve_pipeline: str,
|
||||
retrieve_all_pipeline: str,
|
||||
summary_short_pipeline: str,
|
||||
summary_long_pipeline: str,
|
||||
summary_short_interval_time: int = 60,
|
||||
summary_short_minimum_count: int = 5,
|
||||
summary_long_interval_time: int = 60 * 5,
|
||||
summary_long_minimum_count: int = 5 * 5,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.retrieve_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE,
|
||||
pipeline_str=retrieve_pipeline,
|
||||
)
|
||||
|
||||
def __init__(self,
|
||||
chat_name: str,
|
||||
retrieve_pipeline: str,
|
||||
retrieve_all_pipeline: str,
|
||||
summary_short_pipeline: str,
|
||||
summary_long_pipeline: str,
|
||||
summary_short_interval_time: int = 60,
|
||||
summary_short_minimum_count: int = 5,
|
||||
summary_long_interval_time: int = 60 * 5,
|
||||
summary_long_minimum_count: int = 5 * 5,
|
||||
**kwargs):
|
||||
self.retrieve_pipeline = Pipeline(chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE,
|
||||
pipeline_str=retrieve_pipeline)
|
||||
self.retrieve_all_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
|
||||
pipeline_str=retrieve_all_pipeline,
|
||||
)
|
||||
|
||||
self.retrieve_all_pipeline = Pipeline(chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.RETRIEVE_ALL,
|
||||
pipeline_str=retrieve_all_pipeline)
|
||||
self.summary_short_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
|
||||
pipeline_str=summary_short_pipeline,
|
||||
loop_interval_time=summary_short_interval_time,
|
||||
loop_minimum_count=summary_short_minimum_count,
|
||||
)
|
||||
|
||||
self.summary_short_pipeline = Pipeline(chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_SHORT,
|
||||
pipeline_str=summary_short_pipeline,
|
||||
loop_interval_time=summary_short_interval_time,
|
||||
loop_minimum_count=summary_short_minimum_count)
|
||||
|
||||
self.summary_long_pipeline = Pipeline(chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
|
||||
pipeline_str=summary_long_pipeline,
|
||||
loop_interval_time=summary_long_interval_time,
|
||||
loop_minimum_count=summary_long_minimum_count)
|
||||
|
||||
self.kwargs = kwargs
|
||||
self.summary_long_pipeline = Pipeline(
|
||||
chat_name=chat_name,
|
||||
memory_method_type=MemoryMethodEnum.SUMMARY_LONG,
|
||||
pipeline_str=summary_long_pipeline,
|
||||
loop_interval_time=summary_long_interval_time,
|
||||
loop_minimum_count=summary_long_minimum_count,
|
||||
)
|
||||
|
||||
def retrieve(self, message: Message):
|
||||
self.retrieve_pipeline.submit_message(message, with_lock=False)
|
||||
|
|
|
|||
|
|
@ -2,88 +2,27 @@ import json
|
|||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any, List
|
||||
import questionary
|
||||
from rich.console import Console
|
||||
import sys
|
||||
import time
|
||||
import fire
|
||||
from datetime import datetime
|
||||
|
||||
from memory_scope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.chat.global_context import GLOBAL_CONTEXT
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.tool_functions import (
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from enumeration.language_enum import LanguageEnum
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from utils.logger import Logger
|
||||
from utils.tool_functions import (
|
||||
complete_config_name,
|
||||
init_instance_by_config,
|
||||
under_line_to_hump,
|
||||
)
|
||||
from memory_scope.chat.memory_chat import MemoryChat
|
||||
|
||||
|
||||
class CliMemoryChat(object): # object -> MemoryChat
|
||||
|
||||
USER_COMMANDS = {
|
||||
"/exit": "exit the CLI",
|
||||
"/memory": "print the current contents of agent memory",
|
||||
"/retrieve": "retrieve related memory",
|
||||
"/log": "log chat progress"
|
||||
# TODO add more commands
|
||||
}
|
||||
|
||||
def chat_with_memory(self, query): # for testing
|
||||
return query
|
||||
|
||||
def retrieve_all(self): # for testing
|
||||
return "memory 1. 2. 3."
|
||||
|
||||
def run(self):
|
||||
console = Console()
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Try again!")
|
||||
continue
|
||||
|
||||
# Handle CLI commands
|
||||
if query.startswith("/"):
|
||||
if query.lower() == "/exit":
|
||||
break
|
||||
elif query.lower() == "/memory":
|
||||
console.print(self.memory_service.retrieve_all())
|
||||
elif query.lower() == "/help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
continue
|
||||
|
||||
continue
|
||||
|
||||
while True:
|
||||
try:
|
||||
with console.status("[bold cyan]Thinking..."):
|
||||
messages = self.chat_with_memory(query=query)
|
||||
console.print(messages)
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
console.print("User interrupt occurred.")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
except Exception as e:
|
||||
console.print(
|
||||
f"An exception occurred when running chat_with_memory(): {e}"
|
||||
)
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
from chat.memory_chat import MemoryChat
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from scheme.message import Message
|
||||
from chat.base_memory_chat import BaseMemoryChat
|
||||
from models.llama_index_generation_model import LlamaIndexGenerationModel
|
||||
from models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from models.llama_index_rerank_model import LlamaIndexRerankModel
|
||||
|
||||
|
||||
class CliJob(object):
|
||||
|
|
@ -97,7 +36,7 @@ class CliJob(object):
|
|||
self.logger: Logger = Logger.get_logger("memory_chat")
|
||||
|
||||
def init_memory_chat(self):
|
||||
for chat_name in self.config["chat_list"]:
|
||||
for chat_name in GLOBAL_CONTEXT.global_configs["chat_list"]:
|
||||
memory_chat_config = self.config[chat_name]
|
||||
memory_chat: BaseMemoryChat = init_instance_by_config(
|
||||
memory_chat_config, chat_name=chat_name
|
||||
|
|
@ -172,10 +111,12 @@ class CliJob(object):
|
|||
self.init_memory_chat()
|
||||
|
||||
self.init_workers()
|
||||
GLOBAL_CONTEXT.vector_store = init_instance_by_config(
|
||||
self.config["vector_store"]
|
||||
)
|
||||
GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
## TODO no db and monitor now
|
||||
# GLOBAL_CONTEXT.vector_store = init_instance_by_config(
|
||||
# self.config["vector_store"]
|
||||
# )
|
||||
# GLOBAL_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
|
|
@ -191,4 +132,4 @@ def main(config_path: str):
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
fire.Fire(main)
|
||||
fire.Fire(main)
|
||||
|
|
@ -3,3 +3,110 @@ RELATED_MEMORIES = "related_memories"
|
|||
MESSAGES = "messages"
|
||||
|
||||
CHAT_NAME = "chat_name"
|
||||
|
||||
PIPELINE = "pipeline"
|
||||
|
||||
WORKER = "worker"
|
||||
|
||||
MEMORY = "memory"
|
||||
|
||||
DEFAULT_SYSTEM_PROMPT = "default_system_prompt"
|
||||
|
||||
RELATED_MEMORIES = "related_memories"
|
||||
|
||||
MODIFIED_MEMORIES = "modified_memories"
|
||||
|
||||
RESPONSE_EXT_INFO = "response_ext_info"
|
||||
|
||||
PROMPT_CONFIG = "prompt_config"
|
||||
|
||||
MESSAGES = "messages"
|
||||
|
||||
EXTRACT_TIME_DICT = "extract_time_dict"
|
||||
|
||||
NEW_OBS_NODES = "new_obs_nodes"
|
||||
|
||||
NEW_OBS_WITH_TIME_NODES = "new_obs_with_time_nodes"
|
||||
|
||||
INSIGHT_NODES = "insight_nodes"
|
||||
|
||||
MERGE_OBS_NODES = "merge_obs_nodes"
|
||||
|
||||
NEW_INSIGHT_NODES = "new_insight_nodes"
|
||||
|
||||
TODAY_OBS_NODES = "today_obs_nodes"
|
||||
|
||||
SIMILAR_OBS_NODES = "similar_obs_nodes"
|
||||
|
||||
KEYWORD_OBS_NODES = "keyword_obs_nodes"
|
||||
|
||||
NOT_REFLECTED_OBS_NODES = "not_reflected_obs_nodes"
|
||||
|
||||
NOT_REFLECTED_MERGE_NODES = "not_reflected_merge_nodes"
|
||||
|
||||
NEW_INSIGHT_KEYS = "new_insight_keys"
|
||||
|
||||
INSIGHT_KEY = "insight_key"
|
||||
|
||||
INSIGHT_VALUE = "insight_value"
|
||||
|
||||
DT = "dt"
|
||||
|
||||
MSG_TIME = "msg_time"
|
||||
|
||||
NEW = "new"
|
||||
|
||||
TIME_INFER = "time_infer"
|
||||
|
||||
KEY_WORD = "key_word"
|
||||
|
||||
REFLECTED = "reflected"
|
||||
|
||||
NEW_USER_PROFILE = "new_user_profile"
|
||||
|
||||
RECALL_TYPE = "recall_type"
|
||||
|
||||
ALL_ONLINE_NODES = "all_online_nodes"
|
||||
|
||||
MAX_WORKERS = "max_workers"
|
||||
|
||||
TIME_MATCHED = "time_matched"
|
||||
|
||||
QUERY_KEYWORDS = "query_keywords"
|
||||
|
||||
|
||||
WEEKDAYS = ["周一", "周二", "周三", "周四", "周五", "周六", "周日"]
|
||||
|
||||
DATATIME_WORD_LIST = [
|
||||
"天",
|
||||
"周",
|
||||
"月",
|
||||
"年",
|
||||
"星期",
|
||||
"点",
|
||||
"分钟",
|
||||
"小时",
|
||||
"秒",
|
||||
"上午",
|
||||
"下午",
|
||||
"早上",
|
||||
"早晨",
|
||||
"晚上",
|
||||
"中午",
|
||||
"日",
|
||||
"夜",
|
||||
"清晨",
|
||||
"傍晚",
|
||||
"凌晨",
|
||||
"岁",
|
||||
]
|
||||
|
||||
TIME_FORMAT_V1 = "{year}年{month}月{day}日{weekday}{hour}点"
|
||||
|
||||
DATATIME_KEY_MAP = {
|
||||
"年": "year",
|
||||
"月": "month",
|
||||
"日": "day",
|
||||
"周": "week",
|
||||
"星期几": "weekday",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from memory_scope.utils.registry import Registry
|
||||
from utils.registry import Registry
|
||||
|
||||
# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"]
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
|
|
|||
|
|
@ -2,11 +2,11 @@ import inspect
|
|||
import time
|
||||
from abc import abstractmethod, ABCMeta
|
||||
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models import MODEL_REGISTRY
|
||||
from memory_scope.models.response import ModelResponse, ModelResponseGen
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from . import MODEL_REGISTRY
|
||||
from .response import ModelResponse, ModelResponseGen
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
class BaseModel(metaclass=ABCMeta):
|
||||
|
|
|
|||
|
|
@ -2,18 +2,20 @@ from typing import List
|
|||
|
||||
from llama_index.embeddings.dashscope import DashScopeEmbedding
|
||||
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models import MODEL_REGISTRY
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.models.response import ModelResponse
|
||||
from models import MODEL_REGISTRY
|
||||
from models.base_model import BaseModel
|
||||
from models.response import ModelResponse, ModelResponseGen
|
||||
from enumeration.model_enum import ModelEnum
|
||||
|
||||
|
||||
class LlamaIndexEmbeddingModel(BaseModel):
|
||||
m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
|
||||
|
||||
MODEL_REGISTRY.batch_register([
|
||||
DashScopeEmbedding,
|
||||
])
|
||||
MODEL_REGISTRY.batch_register(
|
||||
[
|
||||
DashScopeEmbedding,
|
||||
]
|
||||
)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
|
|
@ -40,11 +42,16 @@ class LlamaIndexEmbeddingModel(BaseModel):
|
|||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data))
|
||||
|
||||
return ModelResponse(
|
||||
m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data)
|
||||
)
|
||||
|
||||
async def _async_call(self, **kwargs) -> ModelResponse:
|
||||
"""
|
||||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data))
|
||||
return ModelResponse(
|
||||
m_type=self.m_type,
|
||||
raw=await self.model.aget_text_embedding_batch(**self.data),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,16 +1,15 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from llama_index.core.base.llms.types import ChatMessage
|
||||
from llama_index.core.base.llms.types import (
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
CompletionResponse,
|
||||
)
|
||||
from llama_index.llms.dashscope import DashScope
|
||||
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models import MODEL_REGISTRY
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.models.response import ModelResponse, ModelResponseGen
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from . import MODEL_REGISTRY
|
||||
from .base_model import BaseModel
|
||||
from .response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
||||
class LlamaIndexGenerationModel(BaseModel):
|
||||
|
|
@ -32,7 +31,7 @@ class LlamaIndexGenerationModel(BaseModel):
|
|||
elif messages:
|
||||
input_text = messages
|
||||
input_type = 'messages'
|
||||
llama_input = [ChatMessage(role=x['role'], content=x['content']) for x in input_text]
|
||||
llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text]
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
|
|
|
|||
|
|
@ -4,10 +4,11 @@ from llama_index.core.data_structs import Node
|
|||
from llama_index.core.schema import NodeWithScore
|
||||
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
|
||||
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models import MODEL_REGISTRY
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.models.response import ModelResponse
|
||||
from models import MODEL_REGISTRY
|
||||
from models.base_model import BaseModel
|
||||
from models.response import ModelResponse, ModelResponseGen
|
||||
from enumeration.model_enum import ModelEnum
|
||||
|
||||
|
||||
|
||||
class LlamaIndexRerankModel(BaseModel):
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from typing import Generator, List, Dict, Any
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from enumeration.model_enum import ModelEnum
|
||||
|
||||
|
||||
class ModelResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from enumeration.language_enum import LanguageEnum
|
||||
|
||||
SYSTEM_PROMPT = {
|
||||
LanguageEnum.CN: """
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from models.base_model import BaseModel
|
||||
|
||||
|
||||
class BaseVectorStore(metaclass=ABCMeta):
|
||||
|
|
@ -21,6 +20,7 @@ class BaseVectorStore(metaclass=ABCMeta):
|
|||
:param filter_dict:
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
|
||||
|
|
@ -30,27 +30,32 @@ class BaseVectorStore(metaclass=ABCMeta):
|
|||
:param filter_dict:
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def insert(self, node: MemoryNode):
|
||||
""" TODO 是否overwrite
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def insert_batch(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def flush(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -5,13 +5,13 @@ from concurrent.futures import as_completed
|
|||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from memory_scope.chat.global_context import GLOBAL_CONTEXT
|
||||
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
from memory_scope.worker.base_worker import BaseWorker
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from scheme.message import Message
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class Pipeline(object):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import re
|
||||
|
||||
from memory_scope.utils.logger import Logger
|
||||
from utils.logger import Logger
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import time
|
||||
|
||||
from memory_scope.utils.logger import Logger
|
||||
from .logger import Logger
|
||||
|
||||
|
||||
class Timer(object):
|
||||
|
|
|
|||
|
|
@ -1,15 +1,18 @@
|
|||
import re
|
||||
from importlib import import_module
|
||||
from datetime import datetime
|
||||
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
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)
|
||||
sub = re.sub(r"(_\w)", lambda x: x.group(1)[1].upper(), underline_str)
|
||||
return sub[0:1].upper() + sub[1:]
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs):
|
||||
def init_instance_by_config(
|
||||
config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs
|
||||
):
|
||||
clazz_path = config.pop("clazz")
|
||||
if not clazz_path:
|
||||
raise RuntimeError("empty clazz_path!")
|
||||
|
|
@ -44,6 +47,44 @@ def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
|
|||
},
|
||||
{
|
||||
"role": MessageRoleEnum.USER.value,
|
||||
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
|
||||
"content": "\n".join(
|
||||
[x.strip() for x in [few_shot, system_prompt, user_query]]
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
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 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
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from typing import Any, Dict
|
||||
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
|
@ -63,6 +63,9 @@ class BaseWorker(object):
|
|||
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:
|
||||
|
|
|
|||
6
memory_scope/worker/dummy_worker.py
Normal file
6
memory_scope/worker/dummy_worker.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
pass
|
||||
22
memory_scope/worker/es/es_insight_worker.py
Normal file
22
memory_scope/worker/es/es_insight_worker.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.INSIGHT.value,
|
||||
},
|
||||
)
|
||||
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
|
||||
self.set_context(INSIGHT_NODES, insight_nodes)
|
||||
22
memory_scope/worker/es/es_new_obs_worker.py
Normal file
22
memory_scope/worker/es/es_new_obs_worker.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{NEW}": "1",
|
||||
},
|
||||
)
|
||||
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
29
memory_scope/worker/es/es_not_reflected_worker.py
Normal file
29
memory_scope/worker/es/es_not_reflected_worker.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
f"metaData.{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)
|
||||
37
memory_scope/worker/es/es_similar_worker.py
Normal file
37
memory_scope/worker/es/es_similar_worker.py
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
for node in similar_obs_nodes:
|
||||
node.metaData[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)
|
||||
32
memory_scope/worker/es/es_today_obs_worker.py
Normal file
32
memory_scope/worker/es/es_today_obs_worker.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{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)
|
||||
25
memory_scope/worker/es/load_profile_worker.py
Normal file
25
memory_scope/worker/es/load_profile_worker.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
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={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
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,12 +1,12 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.chat.global_context import GLOBAL_CONTEXT
|
||||
from memory_scope.constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.worker.base_worker import BaseWorker
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from models.base_model import BaseModel
|
||||
from scheme.message import Message
|
||||
from storage.base_monitor import BaseMonitor
|
||||
from storage.base_vector_store import BaseVectorStore
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
|
|
|
|||
|
|
@ -1,19 +1,11 @@
|
|||
import re
|
||||
|
||||
from common.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
|
||||
from constants.common_constants import EXTRACT_TIME_DICT
|
||||
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):
|
||||
def __init__(self, parse_time_model, parse_time_max_token, parse_time_temperature, parse_time_top_k, *args, **kwargs):
|
||||
super(ExtractTimeWorker, self).__init__(*args, **kwargs)
|
||||
self.parse_time_model = parse_time_model
|
||||
self.parse_time_max_token = parse_time_max_token
|
||||
self.parse_time_temperature = parse_time_temperature
|
||||
self.parse_time_top_k = parse_time_top_k
|
||||
|
||||
@staticmethod
|
||||
def get_parse_time_prompt(query: str, query_time_str: str):
|
||||
return f"""
|
||||
|
|
@ -51,7 +43,8 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
|
||||
# call sft model
|
||||
response_text = self.gene_client.call(prompt=extract_time_prompt,
|
||||
|
||||
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,
|
||||
|
|
@ -2,7 +2,7 @@ from typing import Dict, List
|
|||
|
||||
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
|
||||
TIME_MATCHED
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -55,7 +55,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
def _run(self):
|
||||
# 解析时间meta信息
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
all_online_nodes: List[MemoryWrapNode] = self.get_context(ALL_ONLINE_NODES)
|
||||
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")
|
||||
|
|
@ -67,7 +67,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
continue
|
||||
|
||||
# 根据类型给ratio
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_node.memoryType, 0.1)
|
||||
type_ratio: float = self.fuse_ratio_dict.get(scheme.memory_node.memoryType, 0.1)
|
||||
|
||||
# 时间系数,完全匹配才行
|
||||
fuse_time_ratio: float = 1.0
|
||||
|
|
@ -76,7 +76,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
if extract_time_dict:
|
||||
match_event_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
event_value = node.memory_node.metaData.get(f"event_{k}", "")
|
||||
event_value = scheme.memory_node.metaData.get(f"event_{k}", "")
|
||||
if event_value in ["-1", v]:
|
||||
continue
|
||||
else:
|
||||
|
|
@ -85,7 +85,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
match_msg_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
msg_value = node.memory_node.metaData.get(f"msg_{k}", "")
|
||||
msg_value = scheme.memory_node.metaData.get(f"msg_{k}", "")
|
||||
if msg_value == v:
|
||||
continue
|
||||
else:
|
||||
|
|
@ -94,10 +94,10 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
if match_event_flag or match_msg_flag:
|
||||
fuse_time_ratio = self.fuse_time_ratio
|
||||
node.memory_node.metaData[TIME_MATCHED] = "1"
|
||||
scheme.memory_node.metaData[TIME_MATCHED] = "1"
|
||||
|
||||
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
|
||||
self.logger.info(f"content={node.memory_node.content} f_event={int(match_event_flag)} "
|
||||
self.logger.info(f"content={scheme.memory_node.content} f_event={int(match_event_flag)} "
|
||||
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}")
|
||||
filtered_nodes.append(node)
|
||||
|
||||
|
|
@ -106,18 +106,18 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
filtered_nodes = filtered_nodes[: self.output_max_count]
|
||||
related_memories: List[str] = []
|
||||
for node in filtered_nodes:
|
||||
content = node.memory_node.content
|
||||
content = scheme.memory_node.content
|
||||
|
||||
# 如果命中时间逻辑
|
||||
if node.memory_node.metaData.get(TIME_MATCHED, "") == "1":
|
||||
# time_infer = node.memory_node.metaData.get(TIME_INFER)
|
||||
if scheme.memory_node.metaData.get(TIME_MATCHED, "") == "1":
|
||||
# time_infer = scheme.memory_node.metaData.get(TIME_INFER)
|
||||
# if not time_infer:
|
||||
# time_infer = self.format_time_infer(time_infer=time_infer,
|
||||
# extract_time_dict=extract_time_dict,
|
||||
# meta_data=node.memory_node.metaData)
|
||||
# meta_data=scheme.memory_node.metaData)
|
||||
time_infer = self.format_time_infer(time_infer="",
|
||||
extract_time_dict=extract_time_dict,
|
||||
meta_data=node.memory_node.metaData)
|
||||
meta_data=scheme.memory_node.metaData)
|
||||
content = f"{time_infer}: {content}"
|
||||
related_memories.append(content)
|
||||
|
||||
|
|
@ -2,17 +2,17 @@ from typing import List
|
|||
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
|
||||
from node.memory_node import MemoryNode
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
class MemoryStoreWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
modified_memories: List[MemoryWrapNode] | List[MemoryNode] = self.get_context(MODIFIED_MEMORIES)
|
||||
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(MODIFIED_MEMORIES)
|
||||
if modified_memories:
|
||||
if isinstance(modified_memories[0], MemoryWrapNode):
|
||||
if isinstance(modified_memories[0], MemoryNode):
|
||||
modified_memories = [n.memory_node for n in modified_memories]
|
||||
|
||||
for n in modified_memories:
|
||||
|
|
@ -4,38 +4,33 @@ from utils.user_profile_handler import UserProfileHandler
|
|||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
|
||||
QUERY_KEYWORDS
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
||||
def user_profile_to_nodes(self) -> List[MemoryWrapNode]:
|
||||
user_profile_nodes: List[MemoryWrapNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
|
||||
def user_profile_to_nodes(self) -> List[MemoryNode]:
|
||||
user_profile_nodes: List[MemoryNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
|
||||
for node in user_profile_nodes:
|
||||
# 从画像侧召回
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
|
||||
self.logger.info(f"user profile node={node.memory_node.content}")
|
||||
scheme.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
|
||||
self.logger.info(f"user profile node={scheme.memory_node.content}")
|
||||
return user_profile_nodes
|
||||
|
||||
def _run(self):
|
||||
all_node_dict: Dict[str, MemoryWrapNode] = {}
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
# 优先级: similar_obs_nodes < keyword_obs_nodes < profile_nodes
|
||||
similar_obs_nodes: List[MemoryWrapNode] = self.get_context(SIMILAR_OBS_NODES)
|
||||
# 优先级: 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.memory_node.content] = node
|
||||
all_node_dict[scheme.memory_node.content] = node
|
||||
|
||||
keyword_obs_nodes: List[MemoryWrapNode] = self.get_context(KEYWORD_OBS_NODES)
|
||||
if keyword_obs_nodes:
|
||||
for node in keyword_obs_nodes:
|
||||
all_node_dict[node.memory_node.content] = node
|
||||
|
||||
profile_nodes: List[MemoryWrapNode] = self.user_profile_to_nodes()
|
||||
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
|
||||
if profile_nodes:
|
||||
for node in profile_nodes:
|
||||
all_node_dict[node.memory_node.content] = node
|
||||
all_node_dict[scheme.memory_node.content] = node
|
||||
|
||||
if not all_node_dict:
|
||||
self.add_run_info(f"all_node_dict is empty!", continue_run=False)
|
||||
|
|
@ -61,8 +56,8 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
content = documents[rank_node["index"]]
|
||||
node = all_node_dict[content]
|
||||
node.score_rank = rank_node["relevance_score"]
|
||||
self.logger.info(f"query={query} content={node.memory_node.content} score_rank={node.score_rank}")
|
||||
self.logger.info(f"query={query} content={scheme.memory_node.content} score_rank={node.score_rank}")
|
||||
|
||||
# save context
|
||||
all_online_nodes: List[MemoryWrapNode] = list(all_node_dict.values())
|
||||
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())
|
||||
self.set_context(ALL_ONLINE_NODES, all_online_nodes)
|
||||
|
|
@ -2,13 +2,13 @@ from elasticsearch import Elasticsearch
|
|||
from elasticsearch.helpers import bulk
|
||||
|
||||
|
||||
from memory_scope.models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
|
||||
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 memory_scope.utils.logger import Logger
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1,26 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import INSIGHT_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsInsightWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_insight_top_k, *args, **kwargs):
|
||||
super(EsInsightWorker, self).__init__(*args, **kwargs)
|
||||
self.es_insight_top_k = es_insight_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.es_insight_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.INSIGHT.value,
|
||||
})
|
||||
|
||||
insight_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
|
||||
self.set_context(INSIGHT_NODES, insight_nodes)
|
||||
|
|
@ -1,53 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import KEY_WORD, KEYWORD_OBS_NODES, RECALL_TYPE, QUERY_KEYWORDS
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsKeywordWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
query = self.messages[-1].content
|
||||
# keywords = jieba.analyse.extract_tags(query, topK=3, withWeight=False, allowPOS=()) # 分解关键词
|
||||
# keywords = jieba.cut(query, cut_all=False) # 使用精确模式分词
|
||||
|
||||
# 查询相关关键词
|
||||
keywords = set()
|
||||
query_keywords = set()
|
||||
for key, values in self.config.key_word_relate_dict.items():
|
||||
if key in query:
|
||||
keywords.add(key)
|
||||
keywords.update(values)
|
||||
query_keywords.add(values[0])
|
||||
keywords = list(keywords)
|
||||
|
||||
self.set_context(QUERY_KEYWORDS, query_keywords)
|
||||
|
||||
# 任意一个匹配都算
|
||||
hits = self.es_client.exact_search_v2(size=self.config.es_keyword_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.config.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
},
|
||||
match_filters={
|
||||
f"metaData.{KEY_WORD}": keywords,
|
||||
})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
keyword_obs_nodes: List[MemoryWrapNode] = []
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.KEYWORD.value
|
||||
keyword_obs_nodes.append(node)
|
||||
self.logger.info(f"keyword_obs_nodes size={len(keyword_obs_nodes)}")
|
||||
for node in keyword_obs_nodes:
|
||||
self.logger.info(f"node={node.memory_node.content} score_similar={node.score_similar}")
|
||||
self.set_context(KEYWORD_OBS_NODES, keyword_obs_nodes)
|
||||
|
|
@ -1,27 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import NEW, NEW_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNewObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_new_obs_top_k, *args, **kwargs):
|
||||
super(EsNewObsWorker, self).__init__(*args, **kwargs)
|
||||
self.es_new_obs_top_k = es_new_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.es_new_obs_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{NEW}": "1",
|
||||
})
|
||||
|
||||
new_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,28 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNotReflectedWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_not_reflected_top_k, *args, **kwargs):
|
||||
super(EsNotReflectedWorker, self).__init__(*args, **kwargs)
|
||||
self.es_new_obs_top_k = es_new_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=self.es_not_reflected_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
f"metaData.{REFLECTED}": "0",
|
||||
})
|
||||
|
||||
not_reflected_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
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,30 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import ALL_NODES, ALL_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsRetrieveAllWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# msg_time_created = self.messages[-1].time_created
|
||||
hits = self.es_client.exact_search_v2(size=1000,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
# "memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
# f"metaData.{DT}": time_to_formatted_str(msg_time_created),
|
||||
})
|
||||
|
||||
all_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
self.logger.info(f"retrieve_all.size={len(all_nodes)}")
|
||||
self.set_context(ALL_NODES, all_nodes)
|
||||
|
||||
all_memories = []
|
||||
if all_nodes:
|
||||
for node in all_nodes:
|
||||
all_memories.append(node.memory_node.to_dict())
|
||||
self.set_context(ALL_MEMORIES, all_memories)
|
||||
|
|
@ -1,38 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
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
|
||||
hits = self.es_client.similar_search(text=query,
|
||||
size=self.es_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
# 初始化成MemoryWrapNode,并加入召回源的参数
|
||||
similar_obs_nodes: List[MemoryWrapNode] = []
|
||||
for hit in hits:
|
||||
node = MemoryWrapNode.init_from_es(hit)
|
||||
node.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
similar_obs_nodes.append(node)
|
||||
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
|
||||
for node in similar_obs_nodes:
|
||||
self.logger.info(f"node={node.memory_node.content} score_similar={node.score_similar}")
|
||||
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)
|
||||
|
|
@ -1,31 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
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
|
||||
hits = self.es_client.exact_search_v2(size=self.es_today_obs_top_k,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
|
||||
})
|
||||
|
||||
today_obs_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
|
||||
self.set_context(TODAY_OBS_NODES, today_obs_nodes)
|
||||
|
|
@ -1,34 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants import common_constants
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LoadProfileWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
hits = self.es_client.exact_search_v2(size=10000,
|
||||
term_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
user_profile_node: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(hit) for hit in hits]
|
||||
user_profile_dict: Dict[str, UserAttribute] = UserProfileHandler.to_user_attr(user_profile_node)
|
||||
|
||||
request: MemoryServiceRequestModel = self.get_context(common_constants.REQUEST)
|
||||
for user_attr in request.user_profile:
|
||||
user_profile_dict[user_attr.memory_key] = user_attr
|
||||
request.user_profile = list(user_profile_dict.values())
|
||||
self.logger.info(f"retrieve_user_profile.size={len(user_profile_dict)}")
|
||||
for key, user_attr in user_profile_dict.items():
|
||||
self.logger.info(f"{key}: {user_attr.description}: {user_attr.value}")
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
from pydantic import Field, BaseModel
|
||||
|
||||
from node.memory_node import MemoryNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class MemoryWrapNode(BaseModel):
|
||||
class MemoryNode(BaseModel):
|
||||
id: str = Field("", description="uuid64")
|
||||
|
||||
score_similar: float = Field(0, description="相似度打分")
|
||||
|
|
|
|||
|
|
@ -4,9 +4,9 @@ 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_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -20,7 +20,7 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
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) -> MemoryWrapNode:
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
|
|
@ -33,7 +33,7 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryWrapNode.init_from_attrs(content=content,
|
||||
return MemoryNode.init_from_attrs(content=content,
|
||||
memoryId=self.memory_id,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.INSIGHT.value,
|
||||
|
|
@ -44,7 +44,7 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
|
||||
def reflect_new_insight_key(self,
|
||||
insight_key: str,
|
||||
not_reflected_merge_nodes: List[MemoryWrapNode]) -> MemoryWrapNode | None:
|
||||
not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None:
|
||||
|
||||
# 检索历史memory
|
||||
hits = self.es_client.similar_search(text=insight_key,
|
||||
|
|
@ -58,7 +58,7 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
})
|
||||
|
||||
# 转化成 MemoryNodeWrap 合并新增nodes
|
||||
related_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(x) for x in hits]
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
related_nodes.extend(not_reflected_merge_nodes)
|
||||
|
||||
# content去重
|
||||
|
|
@ -106,12 +106,12 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
return self.new_insight_node(insight_key=insight_key, insight_value=response_text)
|
||||
|
||||
def _run(self):
|
||||
new_insight_keys: List[MemoryWrapNode] = self.get_context(NEW_INSIGHT_KEYS)
|
||||
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[MemoryWrapNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
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
|
||||
|
|
@ -124,11 +124,11 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
not_reflected_merge_nodes=not_reflected_merge_nodes)
|
||||
|
||||
# save output
|
||||
new_insight_nodes: List[MemoryWrapNode] = []
|
||||
new_insight_nodes: List[MemoryNode] = []
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
new_insight_nodes.append(result)
|
||||
assert isinstance(result, MemoryWrapNode)
|
||||
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}")
|
||||
|
|
@ -137,4 +137,4 @@ class GetInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# set REFLECTED
|
||||
for node in not_reflected_merge_nodes:
|
||||
node.memory_node.metaData[REFLECTED] = "1"
|
||||
scheme.memory_node.metaData[REFLECTED] = "1"
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ 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 node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -19,15 +19,15 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryWrapNode] = []
|
||||
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.memory_node.metaData.get(REFLECTED, "") == "0"]
|
||||
if scheme.memory_node.metaData.get(REFLECTED, "") == "0"]
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
|
|
@ -45,7 +45,7 @@ class GetReflectionWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"profile_keys={profile_keys}")
|
||||
|
||||
# get insight_keys
|
||||
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
|
||||
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]
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@ 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_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -20,12 +20,12 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
# new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
# oday_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryWrapNode] = []
|
||||
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_node.memory_node.content
|
||||
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={
|
||||
|
|
@ -36,7 +36,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
related_nodes: List[MemoryWrapNode] = [MemoryWrapNode.init_from_es(x) for x in hits]
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
|
|
@ -82,7 +82,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryWrapNode] = []
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
|
@ -108,11 +108,11 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryWrapNode = all_obs_nodes[idx]
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}")
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -2,21 +2,21 @@ 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 node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryWrapNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryWrapNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
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, MemoryWrapNode] = {}
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ 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 node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -17,13 +17,13 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.update_insight_top_k = update_insight_top_k
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryWrapNode,
|
||||
new_obs_nodes: List[MemoryWrapNode]) -> (MemoryWrapNode, List[MemoryWrapNode], float):
|
||||
insight_node: MemoryNode,
|
||||
new_obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryWrapNode] = []
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
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
|
||||
|
|
@ -55,18 +55,18 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(self,
|
||||
insight_node: MemoryWrapNode,
|
||||
filtered_nodes: List[MemoryWrapNode]) -> MemoryWrapNode:
|
||||
insight_node: MemoryNode,
|
||||
filtered_nodes: List[MemoryNode]) -> MemoryNode:
|
||||
|
||||
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
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"句子:{node.memory_node.content}")
|
||||
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,
|
||||
|
|
@ -102,14 +102,14 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
insight_node.memory_node.metaData[INSIGHT_VALUE] = insight_value
|
||||
insight_node.memory_node.content_modified = True
|
||||
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[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryWrapNode] = self.get_context(INSIGHT_NODES)
|
||||
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
|
||||
|
|
@ -145,7 +145,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryWrapNode = result
|
||||
insight_key = insight_node.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
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}")
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ 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 node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
|
@ -25,9 +25,9 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
|
||||
def filter_obs_nodes(self,
|
||||
user_attr: UserAttribute,
|
||||
new_obs_nodes: List[MemoryWrapNode]) -> (UserAttribute, List[MemoryWrapNode], float):
|
||||
new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryWrapNode] = []
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
result = self.rerank_client.call(query=user_attr.description,
|
||||
documents=[x.memory_node.content for x in new_obs_nodes])
|
||||
|
||||
|
|
@ -36,7 +36,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
filtered_nodes: List[MemoryWrapNode] = []
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
|
|
@ -47,20 +47,20 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(f"key={user_attr.memory_key} desc={user_attr.description} "
|
||||
f"content={node.memory_node.content} score={score} keep_flag={keep_flag}")
|
||||
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[MemoryWrapNode]) -> UserAttribute:
|
||||
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"句子:{node.memory_node.content}")
|
||||
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)
|
||||
|
||||
|
|
@ -158,7 +158,7 @@ class UpdateProfileWorker(MemoryBaseWorker):
|
|||
self.user_profile_dict[user_attr_key] = new_attr
|
||||
|
||||
def _run(self):
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
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()))
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ 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_node_status import MemoryNodeStatus
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
|
|
@ -18,10 +18,10 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryWrapNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
today_obs_nodes: List[MemoryWrapNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryWrapNode] = []
|
||||
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:
|
||||
|
|
@ -62,7 +62,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryWrapNode] = []
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
|
@ -88,11 +88,11 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryWrapNode = all_obs_nodes[idx]
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {node.memory_node.content} {node.memory_node.status}")
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -5,9 +5,9 @@ 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_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
|
@ -40,7 +40,7 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
|
|||
# 对话时间
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryWrapNode.init_from_attrs(content=obs_content,
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
|
|
@ -97,7 +97,7 @@ class GetObservationWithTimeWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryWrapNode] = []
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -5,9 +5,9 @@ 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_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
|
@ -36,7 +36,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
}
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryWrapNode.init_from_attrs(content=obs_content,
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
|
|
@ -89,7 +89,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryWrapNode] = []
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from typing import Dict, List
|
|||
|
||||
from constants.common_constants import WEEKDAYS
|
||||
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
|
||||
def under_line_to_hump(underline_str):
|
||||
|
|
@ -150,7 +150,7 @@ def init_instance_by_config(config: dict|object, default_module_path: str = None
|
|||
if isinstance(config, accept_types):
|
||||
return config
|
||||
|
||||
import_module(config.pop("path", default_module_path))
|
||||
import_module(config.pop("path", default_module_path))
|
||||
clazz = getattr(module, config.pop("name"))
|
||||
try:
|
||||
return clazz(**config, **try_kwargs)
|
||||
|
|
|
|||
|
|
@ -10,13 +10,10 @@ class UserAttribute(BaseModel):
|
|||
如果code为空,则为新增,否则是更新。
|
||||
确保请求是10条,返回是原始10条+加上新增的条数(如果可以新增)。只会对正确的请求操作数据库。
|
||||
"""
|
||||
code: str = Field("", description="唯一主键 code")
|
||||
id: str = Field("", description="唯一主键")
|
||||
|
||||
memory_id: str = Field("", description="memory id")
|
||||
|
||||
# 上游可能没有传这个参数,可能隐藏在memory_id做区分
|
||||
scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT")
|
||||
|
||||
# 从key改成memory_key
|
||||
memory_key: str = Field("", description="memory key")
|
||||
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
import json
|
||||
from typing import List, Dict
|
||||
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
|
||||
|
||||
|
|
@ -25,13 +25,13 @@ class UserProfileHandler(object):
|
|||
return content
|
||||
|
||||
"""
|
||||
提供UserAttribute 和 MemoryWrapNode 的相互转化
|
||||
提供UserAttribute 和 MemoryNode 的相互转化
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def to_nodes(cls,
|
||||
user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None,
|
||||
split_value: bool = False) -> List[MemoryWrapNode]:
|
||||
split_value: bool = False) -> List[MemoryNode]:
|
||||
|
||||
user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
if user_profile:
|
||||
|
|
@ -41,14 +41,14 @@ class UserProfileHandler(object):
|
|||
elif isinstance(user_profile, dict):
|
||||
user_profile_dict.update(user_profile)
|
||||
|
||||
user_profile_nodes: List[MemoryWrapNode] = []
|
||||
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 = MemoryWrapNode.init_from_attrs(id=_id,
|
||||
attr_node = MemoryNode.init_from_attrs(id=_id,
|
||||
code=_id,
|
||||
content="",
|
||||
memoryId=user_attr.memory_id,
|
||||
|
|
@ -75,28 +75,28 @@ class UserProfileHandler(object):
|
|||
user_profile_nodes.append(attr_node_copy)
|
||||
else:
|
||||
content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value)
|
||||
attr_node.memory_node.content = content
|
||||
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[MemoryWrapNode]) -> Dict[str, UserAttribute]:
|
||||
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=node.memory_node.memoryId,
|
||||
scene=node.memory_node.scene,
|
||||
memory_key=node.memory_node.metaData["memory_key"],
|
||||
value=json.loads(node.memory_node.metaData["value"]),
|
||||
is_unique=int(node.memory_node.metaData["is_unique"]),
|
||||
is_mutable=int(node.memory_node.metaData["is_mutable"]),
|
||||
memory_type=node.memory_node.memoryType,
|
||||
description=node.memory_node.metaData["description"],
|
||||
status=1 if node.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
|
||||
ext_info=json.loads(node.memory_node.metaData["ext_info"]),
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue