mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-12 23:01:15 +00:00
[dev] add memoryscope to class path
This commit is contained in:
parent
1b088fcca2
commit
6d30e4b18a
49 changed files with 1911 additions and 540 deletions
|
|
@ -5,23 +5,23 @@ global_config:
|
|||
open_ai_apikey:
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat_v2.cli_memory_chat
|
||||
class: memory_scope.chat_v2.cli_memory_chat
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.service.chat_memory_service
|
||||
class: memory_scope.memory.service.chat_memory_service
|
||||
history_msg_count: 32
|
||||
contextual_msg_count: 6
|
||||
read_memory_key: read_memory
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: memory.operation.read_memory
|
||||
class: memory_scope.memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
description: "read session messages of the user"
|
||||
contextual_msg_count: 0
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
class: memory_scope.memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
|
|
@ -34,7 +34,7 @@ memory_service:
|
|||
description: "write observation memories of the user"
|
||||
interval_time: 60
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
class: memory_scope.memory.operation.summary_memory
|
||||
workflow: dummy_worker
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 300
|
||||
|
|
@ -61,4 +61,5 @@ workers:
|
|||
clazz: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
rank_model: dashscope_rank
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,3 @@
|
|||
""" Version of MemoryScope."""
|
||||
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ class BaseMemoryChat(metaclass=ABCMeta):
|
|||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,4 +1,54 @@
|
|||
class BaseMemoryService(object):
|
||||
def __init__(self, **kwargs):
|
||||
import threading
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List, Dict
|
||||
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
def __init__(self,
|
||||
memory_operations: Dict[str, dict],
|
||||
read_memory_key: str = "read_memory",
|
||||
**kwargs):
|
||||
self.memory_operations: Dict[str, dict] = memory_operations
|
||||
self.read_memory_key: str = read_memory_key
|
||||
|
||||
self._operation_dict: Dict[str, BaseOperation] = {}
|
||||
self._op_description_dict: Dict[str, str] = {}
|
||||
self.chat_messages: List[Message] = []
|
||||
self.message_lock = threading.Lock
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
self.kwargs = kwargs
|
||||
|
||||
self._init_operation(memory_operations)
|
||||
|
||||
@abstractmethod
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
||||
def prepare_service(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def do_operation(self, op_name: str):
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def op_description_dict(self) -> Dict[str, str]:
|
||||
if not self._op_description_dict:
|
||||
self._op_description_dict = {k: v.description for k, v in self._operation_dict.items()}
|
||||
return self._op_description_dict
|
||||
|
||||
def read_memory(self):
|
||||
assert self.read_memory_key in self._operation_dict, f"op={self.read_memory_key} is not inited!"
|
||||
return self.operate(self.read_memory_key)
|
||||
|
||||
# def __getattr__(self, key):
|
||||
# return self.kwargs[key]
|
||||
|
|
|
|||
50
memory_scope/chat/chat_memory_service.py
Normal file
50
memory_scope/chat/chat_memory_service.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class ChatMemoryService(BaseMemoryService):
|
||||
def __init__(self,
|
||||
history_msg_count: int = 32,
|
||||
contextual_msg_count: int = 6,
|
||||
**kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.history_msg_count: int = history_msg_count
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
assert self.history_msg_count >= self.contextual_msg_count
|
||||
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
for name, operation_config in memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
self._operation_dict[name] = init_instance_by_config(config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count)
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
if isinstance(messages, Message):
|
||||
messages = [messages]
|
||||
|
||||
messages = sorted(messages, key=lambda x: x.time_created)
|
||||
self.chat_messages.extend(messages)
|
||||
if len(self.chat_messages) > self.history_msg_count:
|
||||
gap_size = len(self.chat_messages) - self.history_msg_count
|
||||
for _ in range(gap_size):
|
||||
self.chat_messages.pop(0)
|
||||
|
||||
def prepare_service(self):
|
||||
for _, operation in self._operation_dict.items():
|
||||
operation.init_workflow()
|
||||
if operation.operation_type == "backend":
|
||||
operation.run_operation_backend()
|
||||
|
||||
def do_operation(self, op_name: str):
|
||||
if op_name not in self._operation_dict:
|
||||
self.logger.warning(f"op_name={op_name} is not inited!")
|
||||
return
|
||||
return self._operation_dict[op_name].run_operation()
|
||||
|
|
@ -1,83 +1,139 @@
|
|||
import datetime
|
||||
import time
|
||||
from typing import Dict, List
|
||||
|
||||
import questionary
|
||||
from rich.console import Console
|
||||
|
||||
from .memory_chat import MemoryChat
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from scheme.message import Message
|
||||
from memory_scope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.chat.global_context import GlobalContext
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT
|
||||
from memory_scope.scheme.message import Message
|
||||
from ..models.model_response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
||||
class CliMemoryChat(MemoryChat):
|
||||
|
||||
class CliMemoryChat(BaseMemoryChat):
|
||||
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
|
||||
"exit": "exit the CLI",
|
||||
"help": "get cli commands help",
|
||||
"stream": "get stream response"
|
||||
}
|
||||
|
||||
def chat_with_memory(self, query): # for testing
|
||||
def __init__(self, memory_service: str, generation_model: str, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.stream: bool = True
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
if isinstance(self._memory_service, str):
|
||||
self._memory_service = GlobalContext.memory_service_dict[self._memory_service]
|
||||
self._memory_service.prepare_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = GlobalContext.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
@staticmethod
|
||||
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[GlobalContext.language]
|
||||
if related_memories:
|
||||
memory_prompt = MEMORY_PROMPT[GlobalContext.language]
|
||||
all_prompt_list = [system_prompt, memory_prompt]
|
||||
all_prompt_list.extend(related_memories)
|
||||
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
|
||||
|
||||
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
|
||||
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)
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created)
|
||||
self.submit_messages(new_message)
|
||||
related_memories: List[str] = self.memory_service.read_memory()
|
||||
system_message: Message = self.get_system_prompt(related_memories, time_created)
|
||||
if self.stream:
|
||||
for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream):
|
||||
yield result
|
||||
|
||||
def retrieve_all(self): # for testing
|
||||
return "memory 1. 2. 3."
|
||||
self.submit_messages(result.text)
|
||||
|
||||
def run(self):
|
||||
console = Console()
|
||||
op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict()
|
||||
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
|
||||
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Enter your message or command:",
|
||||
"Please enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query = query.rstrip()
|
||||
query: str = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Try again!")
|
||||
print("Empty input received. Please try again!")
|
||||
continue
|
||||
|
||||
# Handle CLI commands
|
||||
# handle cli / commands with memory ops
|
||||
if query.startswith("/"):
|
||||
if query.lower() == "/exit":
|
||||
query_split = query.lstrip("/").lower().split(" ")
|
||||
query = query_split[0]
|
||||
args = query_split[1:]
|
||||
if query == "exit":
|
||||
break
|
||||
elif query.lower() == "/memory":
|
||||
console.print(self.memory_service.retrieve_all())
|
||||
elif query.lower() == "/help":
|
||||
elif query == "help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
print(f" {desc}")
|
||||
elif query == "stream":
|
||||
questionary.print(f"stream: {self.stream}")
|
||||
self.stream = ~self.stream
|
||||
elif query in op_description_dict:
|
||||
if not args:
|
||||
result = self.memory_service.do_operation(op_name=query)
|
||||
print(result)
|
||||
|
||||
elif args[0].isdigit():
|
||||
refresh_time = int(args[0])
|
||||
try:
|
||||
while True:
|
||||
time.sleep(refresh_time)
|
||||
result = self.memory_service.do_operation(op_name=query)
|
||||
print(result, flush=True)
|
||||
except KeyboardInterrupt:
|
||||
print("stop refresh!")
|
||||
else:
|
||||
print("unknown command received. Please try again!")
|
||||
else:
|
||||
print("unknown command received. Please try again!")
|
||||
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()
|
||||
if self.stream:
|
||||
for msg in self.chat_with_memory(query=query):
|
||||
print(msg.text, flush=True)
|
||||
print()
|
||||
else:
|
||||
msg = self.chat_with_memory(query=query)
|
||||
print(msg.text)
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
console.print("User interrupt occurred.")
|
||||
questionary.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}"
|
||||
)
|
||||
questionary.print(f"An exception occurred when running chat_with_memory(): {e}")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
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
|
||||
from .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):
|
||||
def __init__(self):
|
||||
self.global_configs: Dict[str, Any] = {}
|
||||
|
||||
self.worker_dict: Dict[str, Dict[str, BaseWorker]] = {}
|
||||
self.worker_config: Dict[str, Dict[str, BaseWorker]] = {}
|
||||
|
||||
self.model_dict: Dict[str, BaseModel] = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ class MemoryChat(BaseMemoryChat):
|
|||
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,70 +0,0 @@
|
|||
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(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,
|
||||
)
|
||||
|
||||
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_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)
|
||||
self.summary_short_pipeline.submit_message(message)
|
||||
self.summary_long_pipeline.submit_message(message)
|
||||
return self.retrieve_pipeline.run(RELATED_MEMORIES)
|
||||
|
||||
def retrieve_all(self):
|
||||
return self.retrieve_all_pipeline.run(RELATED_MEMORIES)
|
||||
|
||||
def start_memory_backend(self):
|
||||
self.summary_short_pipeline.start_loop_run()
|
||||
self.summary_long_pipeline.start_loop_run()
|
||||
|
||||
def get_worker_list(self) -> list:
|
||||
worker_set = set()
|
||||
worker_set.update(self.retrieve_pipeline.worker_set)
|
||||
worker_set.update(self.retrieve_all_pipeline.worker_set)
|
||||
worker_set.update(self.summary_short_pipeline.worker_set)
|
||||
worker_set.update(self.summary_long_pipeline.worker_set)
|
||||
return sorted(worker_set)
|
||||
|
|
@ -3,8 +3,6 @@ import time
|
|||
from typing import Dict, List
|
||||
|
||||
import questionary
|
||||
from rich.console import Console
|
||||
|
||||
from memory_scope.chat_v2.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.chat_v2.global_context import G_CONTEXT
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
|
|
@ -12,18 +10,21 @@ from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
|||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.prompts.memory_chat_prompt import MEMORY_PROMPT, SYSTEM_PROMPT
|
||||
from memory_scope.scheme.message import Message
|
||||
from ..models.model_response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
||||
class CliMemoryChat(BaseMemoryChat):
|
||||
USER_COMMANDS = {
|
||||
"exit": "exit the CLI",
|
||||
"help": "get cli commands help",
|
||||
"stream": "get stream response"
|
||||
}
|
||||
|
||||
def __init__(self, memory_service: str, generation_model: str, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.stream: bool = True
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
|
|
@ -48,79 +49,91 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
system_prompt = "\n".join([x.strip() for x in all_prompt_list])
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt, time_created=time_created)
|
||||
|
||||
def chat_with_memory(self, query: str):
|
||||
def chat_with_memory(self, query: str) -> ModelResponse | ModelResponseGen:
|
||||
query = query.strip()
|
||||
if not query:
|
||||
return
|
||||
|
||||
time_created = int(datetime.datetime.now().timestamp())
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER, content=query, time_created=time_created)
|
||||
self.submit_messages(new_message)
|
||||
related_memories: List[str] = self.memory_service.read_memory()
|
||||
system_message: Message = self.get_system_prompt(related_memories, time_created)
|
||||
return self.generation_model.call(messages=[system_message, new_message], stream=True)
|
||||
if self.stream:
|
||||
for result in self.generation_model.call(messages=[system_message, new_message], stream=self.stream):
|
||||
yield result
|
||||
|
||||
self.submit_messages(result.text)
|
||||
|
||||
def run(self):
|
||||
op_description_dict: Dict[str, str] = self.memory_service.op_description_dict
|
||||
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
|
||||
|
||||
console = Console()
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Please enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query: str = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Please try again!")
|
||||
continue
|
||||
|
||||
# handle cli / commands with memory ops
|
||||
if query.startswith("/"):
|
||||
query_split = query.lstrip("/").lower().split(" ")
|
||||
query = query_split[0]
|
||||
args = query_split[1:]
|
||||
if query == "exit":
|
||||
break
|
||||
elif query == "help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
elif query in op_description_dict:
|
||||
if not args:
|
||||
result = self.memory_service.operate(op_name=query)
|
||||
questionary.print(result)
|
||||
|
||||
elif args[0].isdigit():
|
||||
refresh_time = int(args[0])
|
||||
while True:
|
||||
time.sleep(refresh_time)
|
||||
result = self.memory_service.operate(op_name=query)
|
||||
questionary.print(result)
|
||||
else:
|
||||
console.print("unknown command received. Please try again!")
|
||||
else:
|
||||
console.print("unknown command received. Please try again!")
|
||||
continue
|
||||
def run(self):
|
||||
op_description_dict: Dict[str, str] = self.memory_service.get_op_description_dict()
|
||||
self.USER_COMMANDS.update({f"/{k}": v for k, v in op_description_dict.items()})
|
||||
|
||||
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:
|
||||
query = questionary.text(
|
||||
"Please enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
|
||||
query: str = query.rstrip()
|
||||
|
||||
if query == "":
|
||||
print("Empty input received. Please try again!")
|
||||
continue
|
||||
|
||||
# handle cli / commands with memory ops
|
||||
if query.startswith("/"):
|
||||
query_split = query.lstrip("/").lower().split(" ")
|
||||
query = query_split[0]
|
||||
args = query_split[1:]
|
||||
if query == "exit":
|
||||
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:
|
||||
elif query == "help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
print(f" {desc}")
|
||||
elif query == "stream":
|
||||
questionary.print(f"stream: {self.stream}")
|
||||
self.stream = ~self.stream
|
||||
elif query in op_description_dict:
|
||||
if not args:
|
||||
result = self.memory_service.do_operation(op_name=query)
|
||||
print(result)
|
||||
|
||||
elif args[0].isdigit():
|
||||
refresh_time = int(args[0])
|
||||
try:
|
||||
while True:
|
||||
time.sleep(refresh_time)
|
||||
result = self.memory_service.do_operation(op_name=query)
|
||||
print(result, flush=True)
|
||||
except KeyboardInterrupt:
|
||||
print("stop refresh!")
|
||||
else:
|
||||
print("unknown command received. Please try again!")
|
||||
else:
|
||||
print("unknown command received. Please try again!")
|
||||
continue
|
||||
|
||||
while True:
|
||||
try:
|
||||
if self.stream:
|
||||
for msg in self.chat_with_memory(query=query):
|
||||
print(msg.text, flush=True)
|
||||
print()
|
||||
else:
|
||||
msg = self.chat_with_memory(query=query)
|
||||
print(msg.text)
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
questionary.print("User interrupt occurred.")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
except Exception as e:
|
||||
questionary.print(f"An exception occurred when running chat_with_memory(): {e}")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
|
|
|
|||
|
|
@ -24,5 +24,3 @@ class GlobalContext(pydantic.BaseModel):
|
|||
thread_pool: ThreadPoolExecutor | None = pydantic.Field(None, description="global thread_pool")
|
||||
language: LanguageEnum = pydantic.Field(LanguageEnum.CN, description="language: cn / en")
|
||||
|
||||
|
||||
G_CONTEXT = GlobalContext()
|
||||
|
|
|
|||
|
|
@ -1,135 +1,76 @@
|
|||
import json
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any, List
|
||||
import sys
|
||||
import time
|
||||
import fire
|
||||
from datetime import datetime
|
||||
from typing import Dict, Any
|
||||
|
||||
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 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
|
||||
import yaml
|
||||
import fire
|
||||
|
||||
from .chat_v2.global_context import G_CONTEXT
|
||||
from .enumeration.language_enum import LanguageEnum
|
||||
from .utils.logger import Logger
|
||||
from .utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class CliJob(object):
|
||||
|
||||
def __init__(self, config_path: str):
|
||||
def __init__(self, config_path: str, config_suffix: str = ".yaml"):
|
||||
self.config_path: str = config_path
|
||||
self.config_base_dir: str = os.path.dirname(config_path)
|
||||
self.config_suffix: str = config_suffix
|
||||
self.config: Dict[str, Any] = {}
|
||||
|
||||
self.worker_chat_dict: Dict[str, List[str]] = {}
|
||||
self.logger: Logger = Logger.get_logger("memory_chat")
|
||||
|
||||
def init_memory_chat(self):
|
||||
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
|
||||
)
|
||||
GLOBAL_CONTEXT.memory_chat_dict[chat_name] = memory_chat
|
||||
|
||||
for worker_name in memory_chat.memory_service.get_worker_list():
|
||||
if worker_name not in self.worker_chat_dict:
|
||||
self.worker_chat_dict[worker_name] = []
|
||||
self.worker_chat_dict[worker_name].append(chat_name)
|
||||
|
||||
generation_model = memory_chat_config[ModelEnum.GENERATION_MODEL.value]
|
||||
self.init_model(generation_model)
|
||||
|
||||
def init_model(self, model_name: str):
|
||||
if not model_name or model_name in GLOBAL_CONTEXT.model_dict:
|
||||
return
|
||||
|
||||
with open(
|
||||
os.path.join(
|
||||
self.config_base_dir, "model", complete_config_name(model_name)
|
||||
)
|
||||
) as f:
|
||||
model_config = json.load(f)
|
||||
GLOBAL_CONTEXT.model_dict[model_name] = init_instance_by_config(model_config)
|
||||
|
||||
def init_workers(self):
|
||||
"""load worker config & init workers"""
|
||||
worker_config_name: str = self.config["workers"]
|
||||
with open(
|
||||
os.path.join(self.config_base_dir, complete_config_name(worker_config_name))
|
||||
) as f:
|
||||
worker_config_dict = json.load(f)
|
||||
|
||||
for worker_name, worker_config in worker_config_dict.items():
|
||||
if worker_name not in self.worker_chat_dict:
|
||||
continue
|
||||
|
||||
chat_name_list = self.worker_chat_dict[worker_name]
|
||||
for chat_name in chat_name_list:
|
||||
if chat_name not in GLOBAL_CONTEXT.worker_dict:
|
||||
GLOBAL_CONTEXT.worker_dict[chat_name] = {}
|
||||
GLOBAL_CONTEXT.worker_dict[chat_name][worker_name] = (
|
||||
init_instance_by_config(
|
||||
worker_config,
|
||||
suffix_name="worker",
|
||||
**GLOBAL_CONTEXT.global_configs,
|
||||
)
|
||||
)
|
||||
|
||||
self.init_model(worker_config.get(ModelEnum.EMBEDDING_MODEL.value))
|
||||
self.init_model(worker_config.get(ModelEnum.GENERATION_MODEL.value))
|
||||
self.init_model(worker_config.get(ModelEnum.RANK_MODEL.value))
|
||||
self.logger: Logger = Logger.get_logger("cli_job")
|
||||
|
||||
@staticmethod
|
||||
def set_global_config():
|
||||
"""TODO set global_configs & set apikey into env"""
|
||||
GLOBAL_CONTEXT.language = LanguageEnum(
|
||||
GLOBAL_CONTEXT.global_configs["language"]
|
||||
)
|
||||
GLOBAL_CONTEXT.thread_pool = ThreadPoolExecutor(
|
||||
max_workers=int(GLOBAL_CONTEXT.global_configs["max_workers"])
|
||||
def set_global_config(global_config: Dict[str, Any]):
|
||||
"""set global_configs & set apikey into env
|
||||
:return:
|
||||
TODO at sen
|
||||
"""
|
||||
G_CONTEXT.global_config = global_config
|
||||
G_CONTEXT.language = LanguageEnum(global_config["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(
|
||||
max_workers=int(global_config["max_workers"])
|
||||
)
|
||||
|
||||
def init_global_content_by_config(self):
|
||||
with open(complete_config_name(self.config_path)) as f:
|
||||
self.config = json.load(f)
|
||||
# load config
|
||||
config_path = self.config_path
|
||||
if not self.config_path.endswith(self.config_suffix):
|
||||
config_path += self.config_suffix
|
||||
with open(config_path) as f:
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
|
||||
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
|
||||
self.set_global_config()
|
||||
# set global_config
|
||||
self.set_global_config(self.config["global_config"])
|
||||
|
||||
self.init_memory_chat()
|
||||
# init memory_chat
|
||||
for name, conf in self.config["memory_chat"].items():
|
||||
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(
|
||||
conf, name=name
|
||||
)
|
||||
|
||||
self.init_workers()
|
||||
# set memory_service
|
||||
for name, conf in self.config["memory_service"].items():
|
||||
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(
|
||||
conf, name=name
|
||||
)
|
||||
|
||||
## 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"])
|
||||
# init models
|
||||
for name, conf in self.config["models"].items():
|
||||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
G_CONTEXT.vector_store = init_instance_by_config(
|
||||
self.config["vector_store"]
|
||||
)
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["workers"]
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
with GLOBAL_CONTEXT.thread_pool:
|
||||
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
|
||||
with G_CONTEXT.thread_pool:
|
||||
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
memory_chat.run()
|
||||
|
||||
|
||||
def main(config_path: str):
|
||||
job = CliJob(config_path=config_path)
|
||||
job.init_global_content_by_config()
|
||||
job.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
fire.Fire(main)
|
||||
|
|
@ -1,67 +0,0 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
import yaml
|
||||
|
||||
from chat_v2.global_context import G_CONTEXT
|
||||
from enumeration.language_enum import LanguageEnum
|
||||
from utils.logger import Logger
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class CliJob(object):
|
||||
|
||||
def __init__(self, config_path: str, config_suffix: str = ".yaml"):
|
||||
self.config_path: str = config_path
|
||||
self.config_suffix: str = config_suffix
|
||||
self.config: Dict[str, Any] = {}
|
||||
|
||||
self.logger: Logger = Logger.get_logger("cli_job")
|
||||
|
||||
@staticmethod
|
||||
def set_global_config(global_config: Dict[str, Any]):
|
||||
""" set global_configs & set apikey into env
|
||||
:return:
|
||||
TODO at sen
|
||||
"""
|
||||
G_CONTEXT.global_config = global_config
|
||||
G_CONTEXT.language = LanguageEnum(global_config["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
|
||||
|
||||
def init_global_content_by_config(self):
|
||||
# load config
|
||||
config_path = self.config_path
|
||||
if not self.config_path.endswith(self.config_suffix):
|
||||
config_path += self.config_suffix
|
||||
with open(config_path) as f:
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
|
||||
# set global_config
|
||||
self.set_global_config(self.config["global_config"])
|
||||
|
||||
# init memory_chat
|
||||
for name, conf in self.config["memory_chat"].items():
|
||||
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# set memory_service
|
||||
for name, conf in self.config["memory_service"].items():
|
||||
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init models
|
||||
for name, conf in self.config["models"].items():
|
||||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
G_CONTEXT.vector_store = init_instance_by_config(self.config["vector_store"])
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["workers"]
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
with G_CONTEXT.thread_pool:
|
||||
memory_chat = list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
memory_chat.run()
|
||||
|
|
@ -22,10 +22,6 @@ 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"
|
||||
|
|
@ -116,3 +112,5 @@ DATATIME_KEY_MAP = {
|
|||
"周": "week",
|
||||
"星期几": "weekday",
|
||||
}
|
||||
|
||||
CONTENT_MODIFIED = "content_modified"
|
||||
|
|
@ -23,8 +23,6 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
self.logger = Logger.get_logger()
|
||||
self.kwargs = kwargs
|
||||
|
||||
self._init_operation(memory_operations)
|
||||
|
||||
@abstractmethod
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
|
|
@ -6,26 +6,28 @@ from memory_scope.utils.tool_functions import init_instance_by_config
|
|||
|
||||
|
||||
class ChatMemoryService(BaseMemoryService):
|
||||
|
||||
def __init__(self,
|
||||
history_msg_count: int = 32,
|
||||
contextual_msg_count: int = 6,
|
||||
**kwargs):
|
||||
def __init__(
|
||||
self, history_msg_count: int = 32, contextual_msg_count: int = 6, **kwargs
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.history_msg_count: int = history_msg_count
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
assert self.history_msg_count >= self.contextual_msg_count
|
||||
|
||||
self._init_operation(self.memory_operations)
|
||||
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
for name, operation_config in memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
self._operation_dict[name] = init_instance_by_config(config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count)
|
||||
self._operation_dict[name] = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_count=self.contextual_msg_count,
|
||||
)
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
if isinstance(messages, Message):
|
||||
|
|
|
|||
|
|
@ -19,7 +19,9 @@ class LlamaIndexRerankModel(BaseModel):
|
|||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
assert (
|
||||
query and documents
|
||||
), f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
|
||||
# using -1.0 as dummy scores
|
||||
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
|
||||
|
|
@ -41,7 +43,9 @@ class LlamaIndexRerankModel(BaseModel):
|
|||
return model_response
|
||||
|
||||
def _call(self, **kwargs) -> ModelResponse:
|
||||
return ModelResponse(m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data))
|
||||
return ModelResponse(
|
||||
m_type=self.m_type, raw=self.model.postprocess_nodes(**self.data)
|
||||
)
|
||||
|
||||
async def _async_call(self, **kwargs) -> ModelResponse:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
8
memory_scope/prompts/get_insight_prompt.py
Normal file
8
memory_scope/prompts/get_insight_prompt.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
GET_INSIGHT_SYSTEM_PROMPT = {}
|
||||
|
||||
GET_INSIGHT_FEW_SHOT_PROMPT = {}
|
||||
|
||||
GET_INSIGHT_USER_QUERY_PROMPT = {}
|
||||
8
memory_scope/prompts/get_reflection_prompt.py
Normal file
8
memory_scope/prompts/get_reflection_prompt.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
GET_REFLECTION_SYSTEM_PROMPT = {}
|
||||
|
||||
GET_REFLECTION_FEW_SHOT_PROMPT = {}
|
||||
|
||||
GET_REFLECTION_USER_QUERY_PROMPT = {}
|
||||
8
memory_scope/prompts/long_contra_repeat_prompt.py
Normal file
8
memory_scope/prompts/long_contra_repeat_prompt.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
LONG_CONTRA_REPEAT_SYSTEM_PROMPT = {}
|
||||
|
||||
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT = {}
|
||||
|
||||
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT = {}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
from enumeration.language_enum import LanguageEnum
|
||||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
SYSTEM_PROMPT = {
|
||||
LanguageEnum.CN: """
|
||||
|
|
|
|||
7
memory_scope/prompts/update_insight_prompt.py
Normal file
7
memory_scope/prompts/update_insight_prompt.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
from ..enumeration.language_enum import LanguageEnum
|
||||
|
||||
UPDATE_INSIGHT_SYSTEM_PROMPT = {}
|
||||
|
||||
UPDATE_INSIGHT_FEW_SHOT_PROMPT = {}
|
||||
|
||||
UPDATE_INSIGHT_USER_QUERY_PROMPT = {}
|
||||
|
|
@ -4,7 +4,7 @@ from utils.logger import Logger
|
|||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
pattern_v1 = re.compile(r'<(.*?)>')
|
||||
pattern_v1 = re.compile(r"<(.*?)>")
|
||||
|
||||
def __init__(self, response_text: str):
|
||||
self.response_text: str = response_text.strip()
|
||||
|
|
@ -12,22 +12,26 @@ class ResponseTextParser(object):
|
|||
|
||||
def parse_v1(self, prefix: str = ""):
|
||||
result = []
|
||||
for line in self.response_text.split('\n'):
|
||||
for line in self.response_text.split("\n"):
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
matches = [match.group(1) for match in self.pattern_v1.finditer(line)]
|
||||
if matches:
|
||||
result.append(matches)
|
||||
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
self.logger.info(
|
||||
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
|
||||
)
|
||||
return result
|
||||
|
||||
def parse_v2(self, prefix: str = ""):
|
||||
result = []
|
||||
for line in self.response_text.split('\n'):
|
||||
for line in self.response_text.split("\n"):
|
||||
line = line.strip()
|
||||
if not line or line == "无":
|
||||
continue
|
||||
result.append(line)
|
||||
self.logger.info(f"{prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
self.logger.info(
|
||||
f"{prefix} response_text={self.response_text} result={result}", stacklevel=2
|
||||
)
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from typing import Any, Dict
|
||||
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from ..utils.logger import Logger
|
||||
from ..utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
|
|
|||
|
|
@ -13,9 +13,9 @@ class EsInsightWorker(MemoryBaseWorker):
|
|||
insight_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_insight_top_k,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.INSIGHT.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
},
|
||||
)
|
||||
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
|
||||
|
|
|
|||
|
|
@ -12,10 +12,10 @@ class EsNewObsWorker(MemoryBaseWorker):
|
|||
new_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{NEW}": "1",
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_data.{NEW}": "1",
|
||||
},
|
||||
)
|
||||
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
|
||||
|
|
|
|||
|
|
@ -14,13 +14,13 @@ class EsNotReflectedWorker(MemoryBaseWorker):
|
|||
not_reflected_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
f"metaData.{REFLECTED}": "0",
|
||||
f"meta_data.{REFLECTED}": "0",
|
||||
},
|
||||
)
|
||||
self.logger.info(
|
||||
|
|
|
|||
|
|
@ -19,9 +19,9 @@ class EsSimilarWorker(MemoryBaseWorker):
|
|||
text=query,
|
||||
size=self.es_similar_top_k,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
|
|
@ -30,7 +30,7 @@ class EsSimilarWorker(MemoryBaseWorker):
|
|||
)
|
||||
|
||||
for node in similar_obs_nodes:
|
||||
node.metaData[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
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}")
|
||||
|
|
|
|||
|
|
@ -21,10 +21,10 @@ class EsTodayObsWorker(MemoryBaseWorker):
|
|||
today_obs_nodes = self.vector_store.retrieve(
|
||||
size=self.es_today_obs_top_k,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"metaData.{DT}": time_to_formatted_str(msg_time_created),
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_Data.{DT}": time_to_formatted_str(msg_time_created),
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -13,9 +13,9 @@ class LoadProfileWorker(MemoryBaseWorker):
|
|||
user_profile_node = self.vector_store(
|
||||
size=10000,
|
||||
filter_dict={
|
||||
"memoryId": self.memory_id,
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memoryType": [
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
|
||||
],
|
||||
|
|
|
|||
|
|
@ -1,20 +1,20 @@
|
|||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
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
|
||||
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
|
||||
from ..scheme.memory_node import MemoryNode
|
||||
from ..constants import common_constants
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
def __init__(self,
|
||||
embedding_model: str,
|
||||
generation_model: str,
|
||||
rank_model: str,
|
||||
**kwargs):
|
||||
def __init__(
|
||||
self, embedding_model: str, generation_model: str, rank_model: str, **kwargs
|
||||
):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
self.embedding_model_name: str = embedding_model
|
||||
self.generation_model_name: str = generation_model
|
||||
|
|
@ -40,25 +40,29 @@ class MemoryBaseWorker(BaseWorker):
|
|||
return self.get_context(CHAT_NAME)
|
||||
|
||||
@property
|
||||
def embedding_model(self):
|
||||
def embedding_model(self) -> BaseModel:
|
||||
if self._embedding_model is None:
|
||||
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
|
||||
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(
|
||||
self.embedding_model_name
|
||||
)
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self):
|
||||
def generation_model(self) -> BaseModel:
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict.get(
|
||||
self.generation_model_name
|
||||
)
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self):
|
||||
def rank_model(self) -> BaseModel:
|
||||
if self._rank_model is None:
|
||||
self._rank_model = GLOBAL_CONTEXT.model_dict.get(self.rank_model_name)
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self):
|
||||
def vector_store(self) -> BaseVectorStore:
|
||||
if self._vector_store is None:
|
||||
self._vector_store = GLOBAL_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
|
@ -68,3 +72,22 @@ class MemoryBaseWorker(BaseWorker):
|
|||
if self._monitor is None:
|
||||
self._monitor = GLOBAL_CONTEXT.monitor
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
def user_profile_dict(self) -> Dict[str, MemoryNode]:
|
||||
if not self._user_profile_dict:
|
||||
self._user_profile_dict = {
|
||||
user_attr.meta_data.get("memory_key", ""): user_attr
|
||||
for user_attr in self.get_context(common_constants.USER_PROFILE)
|
||||
}
|
||||
return self._user_profile_dict
|
||||
|
||||
@property
|
||||
def memory_id(self) -> str:
|
||||
pass
|
||||
|
||||
def __getattr__(self, key):
|
||||
return self.kwargs[key]
|
||||
|
||||
def get_prompt(self, x):
|
||||
return x[GLOBAL_CONTEXT.global_configs["language"]]
|
||||
|
|
@ -1,11 +1,16 @@
|
|||
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 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"""
|
||||
|
|
@ -35,27 +40,32 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
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)
|
||||
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
|
||||
|
||||
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)
|
||||
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+)'
|
||||
pattern = r"-\s*(\S+):(\d+)"
|
||||
matches = re.findall(pattern, response_text)
|
||||
for key, value in matches:
|
||||
if key in DATATIME_KEY_MAP.keys():
|
||||
|
|
|
|||
|
|
@ -1,25 +1,20 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
|
||||
TIME_MATCHED
|
||||
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):
|
||||
def __init__(self, fuse_time_ratio, fuse_score_threshold, fuse_ratio_dict, *args, **kwargs):
|
||||
super(FuseRerankWorker, self).__init__(*args, **kwargs)
|
||||
self.fuse_score_threshold = fuse_score_threshold
|
||||
self.fuse_ratio_dict = fuse_ratio_dict
|
||||
# self.default_system_prompt = default_system_prompt
|
||||
self.fuse_time_ratio = fuse_time_ratio
|
||||
|
||||
@property
|
||||
def output_max_count(self):
|
||||
return self.request.user.output_max_count
|
||||
|
||||
@staticmethod
|
||||
def format_time_infer(time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]):
|
||||
def format_time_infer(
|
||||
time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]
|
||||
):
|
||||
if time_infer:
|
||||
return time_infer
|
||||
|
||||
|
|
@ -29,21 +24,21 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
if value:
|
||||
time_infer += f"{value}年"
|
||||
elif value == "-1":
|
||||
time_infer += f"每年"
|
||||
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 += f"每月"
|
||||
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 += f"每日"
|
||||
time_infer += "每日"
|
||||
|
||||
if "weekday" in extract_time_dict:
|
||||
value = meta_data.get("msg_weekday")
|
||||
|
|
@ -67,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
continue
|
||||
|
||||
# 根据类型给ratio
|
||||
type_ratio: float = self.fuse_ratio_dict.get(scheme.memory_node.memoryType, 0.1)
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
|
||||
|
||||
# 时间系数,完全匹配才行
|
||||
fuse_time_ratio: float = 1.0
|
||||
|
|
@ -76,7 +71,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
if extract_time_dict:
|
||||
match_event_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
event_value = scheme.memory_node.metaData.get(f"event_{k}", "")
|
||||
event_value = node.meta_data.get(f"event_{k}", "")
|
||||
if event_value in ["-1", v]:
|
||||
continue
|
||||
else:
|
||||
|
|
@ -85,7 +80,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
match_msg_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
msg_value = scheme.memory_node.metaData.get(f"msg_{k}", "")
|
||||
msg_value = node.meta_data.get(f"msg_{k}", "")
|
||||
if msg_value == v:
|
||||
continue
|
||||
else:
|
||||
|
|
@ -94,32 +89,32 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
|
||||
if match_event_flag or match_msg_flag:
|
||||
fuse_time_ratio = self.fuse_time_ratio
|
||||
scheme.memory_node.metaData[TIME_MATCHED] = "1"
|
||||
node.meta_data[TIME_MATCHED] = "1"
|
||||
|
||||
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
|
||||
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}")
|
||||
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 = 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 = scheme.memory_node.content
|
||||
content = node.content
|
||||
|
||||
# 如果命中时间逻辑
|
||||
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=scheme.memory_node.metaData)
|
||||
time_infer = self.format_time_infer(time_infer="",
|
||||
extract_time_dict=extract_time_dict,
|
||||
meta_data=scheme.memory_node.metaData)
|
||||
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)
|
||||
# self.set_context(DEFAULT_SYSTEM_PROMPT, self.default_system_prompt)
|
||||
|
|
|
|||
|
|
@ -1,33 +1,35 @@
|
|||
from typing import List
|
||||
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED
|
||||
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[MemoryNode] | 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], MemoryNode):
|
||||
modified_memories = [n.memory_node for n in modified_memories]
|
||||
|
||||
for n in modified_memories:
|
||||
if not n.id:
|
||||
n.id = f"{n.memoryId}_{n.scene}_content_{n.content}"
|
||||
n.id = f"{n.memory_id}_content_{n.content}"
|
||||
n.code = n.id
|
||||
# TODO add batch insert
|
||||
self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", )))
|
||||
n.meta_data.pop(CONTENT_MODIFIED)
|
||||
self.vector_store.insert(n)
|
||||
else:
|
||||
self.logger.warning("modified_memories is empty!")
|
||||
|
||||
new_user_profile: List[UserAttribute] = self.get_context(NEW_USER_PROFILE)
|
||||
new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE)
|
||||
if new_user_profile:
|
||||
new_user_nodes: List[MemoryNode] = [n.memory_node for n in UserProfileHandler.to_nodes(new_user_profile)]
|
||||
for n in new_user_nodes:
|
||||
self.es_client.insert(n.id, body=n.model_dump(exclude=set("content_modified", )))
|
||||
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,36 +0,0 @@
|
|||
import json
|
||||
|
||||
from config.bailian_memory_config import BailianMemoryConfig
|
||||
from constants.common_constants import REQUEST, CONFIG
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class ParseParamsWorker(BaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 参数合并
|
||||
memory_config = {}
|
||||
|
||||
# 更新环境变量
|
||||
memory_config.update(self.context_handler.env_configs)
|
||||
|
||||
# 更新请求参数
|
||||
request: MemoryServiceRequestModel = self.context_handler.get_context(REQUEST)
|
||||
memory_config.update(request.model_dump(exclude=set("ext_info", )))
|
||||
|
||||
# 更新ext_info
|
||||
if request.ext_info:
|
||||
memory_config.update(request.ext_info)
|
||||
|
||||
# 存入上下文
|
||||
memory_config_model: BailianMemoryConfig = BailianMemoryConfig(**memory_config)
|
||||
self.context_handler.set_context(CONFIG, memory_config_model)
|
||||
|
||||
# 打印
|
||||
self.logger.info(f"memory_config_model={json.dumps(memory_config_model.model_dump(), ensure_ascii=False)}")
|
||||
|
||||
# 上游可能没有传这个参数,可能隐藏在memory_id做区分
|
||||
if request.user_profile:
|
||||
for user_attr in request.user_profile:
|
||||
user_attr.scene = request.scene
|
||||
|
|
@ -1,9 +1,8 @@
|
|||
from typing import List, Dict
|
||||
|
||||
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 enumeration.memory_recall_enum import MemoryRecallType
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
|
@ -11,29 +10,29 @@ from worker.memory_base_worker import MemoryBaseWorker
|
|||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
||||
def user_profile_to_nodes(self) -> List[MemoryNode]:
|
||||
user_profile_nodes: List[MemoryNode] = UserProfileHandler.to_nodes(self.user_profile_dict, split_value=True)
|
||||
user_profile_nodes: List[MemoryNode] = self.user_profile_dict
|
||||
for node in user_profile_nodes:
|
||||
# 从画像侧召回
|
||||
scheme.memory_node.metaData[RECALL_TYPE] = MemoryRecallType.PROFILE
|
||||
self.logger.info(f"user profile node={scheme.memory_node.content}")
|
||||
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 < 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[scheme.memory_node.content] = node
|
||||
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[scheme.memory_node.content] = node
|
||||
all_node_dict[node.content] = node
|
||||
|
||||
if not all_node_dict:
|
||||
self.add_run_info(f"all_node_dict is empty!", continue_run=False)
|
||||
self.add_run_info("all_node_dict is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# call recall model
|
||||
|
|
@ -45,18 +44,18 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
query_keyword_join = ",".join(query_keywords)
|
||||
query = f"{query} 用户的{query_keyword_join}。"
|
||||
documents = list(all_node_dict.keys())
|
||||
result = self.rerank_client.call(query=query, documents=documents)
|
||||
result = self.rank_model.call(query=query, documents=documents)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"semantic call recall model failed!")
|
||||
self.add_run_info("semantic call recall model failed!")
|
||||
return
|
||||
|
||||
# set score
|
||||
for rank_node in result:
|
||||
content = documents[rank_node["index"]]
|
||||
for index, score in result.rank_scores.items():
|
||||
content = documents[index]
|
||||
node = all_node_dict[content]
|
||||
node.score_rank = rank_node["relevance_score"]
|
||||
self.logger.info(f"query={query} content={scheme.memory_node.content} score_rank={node.score_rank}")
|
||||
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())
|
||||
|
|
|
|||
0
memory_scope/worker/summary_long/__init__.py
Normal file
0
memory_scope/worker/summary_long/__init__.py
Normal file
166
memory_scope/worker/summary_long/get_insight_worker.py
Normal file
166
memory_scope/worker/summary_long/get_insight_worker.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
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"
|
||||
99
memory_scope/worker/summary_long/get_reflection_worker.py
Normal file
99
memory_scope/worker/summary_long/get_reflection_worker.py
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
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)
|
||||
129
memory_scope/worker/summary_long/long_contra_repeat_worker.py
Normal file
129
memory_scope/worker/summary_long/long_contra_repeat_worker.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
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)
|
||||
47
memory_scope/worker/summary_long/summary_collect_worker.py
Normal file
47
memory_scope/worker/summary_long/summary_collect_worker.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
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()))
|
||||
177
memory_scope/worker/summary_long/update_insight_worker.py
Normal file
177
memory_scope/worker/summary_long/update_insight_worker.py
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
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}"
|
||||
)
|
||||
241
memory_scope/worker/summary_long/update_profile_worker.py
Normal file
241
memory_scope/worker/summary_long/update_profile_worker.py
Normal file
|
|
@ -0,0 +1,241 @@
|
|||
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)
|
||||
0
memory_scope/worker/summary_short/__init__.py
Normal file
0
memory_scope/worker/summary_short/__init__.py
Normal file
117
memory_scope/worker/summary_short/contra_repeat_worker.py
Normal file
117
memory_scope/worker/summary_short/contra_repeat_worker.py
Normal file
|
|
@ -0,0 +1,117 @@
|
|||
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)
|
||||
|
|
@ -0,0 +1,167 @@
|
|||
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)
|
||||
144
memory_scope/worker/summary_short/get_observation_worker.py
Normal file
144
memory_scope/worker/summary_short/get_observation_worker.py
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
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)
|
||||
70
memory_scope/worker/summary_short/info_filter_worker.py
Normal file
70
memory_scope/worker/summary_short/info_filter_worker.py
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
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
|
||||
13
test.py
Normal file
13
test.py
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
from memory_scope.cli import CliJob
|
||||
import fire
|
||||
|
||||
|
||||
def main(config_path: str):
|
||||
job = CliJob(config_path=config_path)
|
||||
job.init_global_content_by_config()
|
||||
job.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# fire.Fire(main)
|
||||
main("config/config.yaml")
|
||||
Loading…
Add table
Reference in a new issue