mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
feat: Resolve conflict, auto committed by CodeFlow
This commit is contained in:
commit
509f11be84
83 changed files with 2787 additions and 1140 deletions
|
|
@ -1,27 +0,0 @@
|
|||
{
|
||||
"global_configs": {
|
||||
"thread_pool_max_count": 5,
|
||||
"dash_scope_apikey": "",
|
||||
"open_ai_apikey": "",
|
||||
"language": "en",
|
||||
"chat_list": [
|
||||
"memory_chat"
|
||||
]
|
||||
},
|
||||
"memory_chat": {
|
||||
"clazz": "chat.memory_chat",
|
||||
"retrieve": "parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank",
|
||||
"generation_model": "dashscope_generation",
|
||||
"history_msg_count": 3
|
||||
},
|
||||
"vector_store": {
|
||||
"clazz": "storage.base_vector_store",
|
||||
"index_name": "memory_test",
|
||||
"password": ""
|
||||
},
|
||||
"monitor": {
|
||||
"clazz": "storage.base_monitor",
|
||||
"index_name": "memory_test"
|
||||
},
|
||||
"workers": "workers"
|
||||
}
|
||||
|
|
@ -1,57 +1,68 @@
|
|||
global_config:
|
||||
thread_pool_max_count: 5
|
||||
language: en
|
||||
max_workers: 5
|
||||
dash_scope_apikey:
|
||||
open_ai_apikey:
|
||||
language: en
|
||||
chat_list:
|
||||
- memory_chat
|
||||
memory_chat:
|
||||
memory_service: memory_chat_service
|
||||
memory_chat_service:
|
||||
class: memory.base_memory_service
|
||||
memory_operations:
|
||||
- name: read_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: parse_params,load_profile,extract_time,es_similar,es_keyword,semantic_rank,fuse_rerank
|
||||
work_type: frontend
|
||||
- name: list_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: frontend
|
||||
- name: extract_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 60
|
||||
min_count: 5
|
||||
- name: reflect_memory
|
||||
class: memory.workflow.base_workflow
|
||||
workflow: dummy
|
||||
work_type: backend
|
||||
interval_time: 300
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
human_name: human
|
||||
assistant_name: assistant
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: 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
|
||||
workflow: dummy_worker
|
||||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_worker
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: dummy_worker
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 60
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: dummy_worker
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 300
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
dashscope_embedding:
|
||||
class: models.llama_index_embedding_model
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
class: models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
vector_store:
|
||||
clazz: storage.base_vector_store
|
||||
index_name: memory_test
|
||||
password: ''
|
||||
class: storage.llama_index_elastic_search_store
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
monitor:
|
||||
clazz: storage.base_monitor
|
||||
index_name: memory_test
|
||||
workers:
|
||||
- name: update_insight
|
||||
clazz: worker.summary_long.update_insight
|
||||
class: storage.dummy_monitor
|
||||
worker:
|
||||
dummy_worker:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
models:
|
||||
- name: dashscope_generation
|
||||
clazz: models.llama_index_generation_model
|
||||
module_name: DashScope
|
||||
model_name: qwen-max
|
||||
- name: dashscope_embedding
|
||||
clazz: models.base_embedding_model
|
||||
module_name: DashScopeEmbedding
|
||||
model_name: text-embedding-v2
|
||||
- name: dashscope_rank
|
||||
clazz: models.base_rank_model
|
||||
module_name: DashScopeRerank
|
||||
model_name: gte-rerank
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +0,0 @@
|
|||
{
|
||||
"clazz": "models.base_embedding_model",
|
||||
"model_name": "text-embedding-v2",
|
||||
"method_type": "DashScopeEmbedding"
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
{
|
||||
"clazz": "models.llama_index_generation_model",
|
||||
"model_name": "qwen-max",
|
||||
"method_type": "DashScope"
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
{
|
||||
"clazz": "models.base_rank_model",
|
||||
"model_name": "gte-rerank",
|
||||
"method_type": "DashScopeRerank"
|
||||
}
|
||||
|
|
@ -1,8 +0,0 @@
|
|||
{
|
||||
"update_insight": {
|
||||
"clazz": "worker.summary_long.update_insight",
|
||||
"generation_model": "dashscope_generation",
|
||||
"embedding_model": "dashscope_embedding",
|
||||
"rank_model": "dashscope_rank"
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,3 @@
|
|||
""" Version of MemoryScope."""
|
||||
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
|
|
|
|||
|
|
@ -2,9 +2,6 @@ from abc import ABCMeta, abstractmethod
|
|||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
|
|
|
|||
|
|
@ -1,4 +0,0 @@
|
|||
class BaseMemoryService(object):
|
||||
def __init__(self, **kwargs):
|
||||
|
||||
self.kwargs = kwargs
|
||||
|
|
@ -1,83 +1,177 @@
|
|||
import datetime
|
||||
import os
|
||||
import time
|
||||
from typing import 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 G_CONTEXT
|
||||
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 memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.tool_functions import char_logo
|
||||
|
||||
|
||||
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,
|
||||
stream: bool = True,
|
||||
human_name: str = "human",
|
||||
assistant_name: str = "assistant",
|
||||
**kwargs):
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.stream: bool = stream
|
||||
self.human_name: str = human_name
|
||||
self.assistant_name: str = assistant_name
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self._logo = char_logo("MemoryScope")
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def print_logo(self):
|
||||
for line in self._logo:
|
||||
print(line)
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
if isinstance(self._memory_service, str):
|
||||
self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service]
|
||||
self._memory_service.start_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def get_system_prompt(self) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[G_CONTEXT.language].strip()
|
||||
|
||||
memories: str = self.memory_service.read_memory()
|
||||
if memories:
|
||||
memory_prompt = MEMORY_PROMPT[G_CONTEXT.language]
|
||||
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
|
||||
|
||||
return Message(role=MessageRoleEnum.SYSTEM, content=system_prompt)
|
||||
|
||||
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.value,
|
||||
role_name=self.human_name,
|
||||
content=query)
|
||||
|
||||
def retrieve_all(self): # for testing
|
||||
return "memory 1. 2. 3."
|
||||
self.memory_service.add_messages(new_message)
|
||||
system_message: Message = self.get_system_prompt()
|
||||
|
||||
model_response = self.generation_model.call(messages=[system_message, new_message], stream=self.stream)
|
||||
if self.stream:
|
||||
for _ in model_response:
|
||||
_.message.role_name = self.assistant_name
|
||||
yield _
|
||||
else:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
return model_response
|
||||
|
||||
def process_commands(self, query: str) -> bool:
|
||||
continue_run = True
|
||||
query_split = query.lstrip("/").lower().split(" ")
|
||||
query = query_split[0]
|
||||
args = query_split[1:]
|
||||
if query == "exit":
|
||||
self.memory_service.stop_service()
|
||||
continue_run = False
|
||||
|
||||
elif query == "help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(text=f" /{cmd}:", style="bold")
|
||||
questionary.print(text=f" {desc}")
|
||||
|
||||
elif query == "stream":
|
||||
self.stream = bool(args[0])
|
||||
questionary.print(f"stream: {self.stream}")
|
||||
|
||||
elif query in self.memory_service.op_description_dict:
|
||||
if not args:
|
||||
result = self.memory_service.do_operation(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.do_operation(op_name=query)
|
||||
os.system('clear')
|
||||
self.print_logo()
|
||||
questionary.print(result)
|
||||
|
||||
else:
|
||||
questionary.print("unknown command received. Please try again!")
|
||||
|
||||
else:
|
||||
questionary.print("unknown command received. Please try again!")
|
||||
|
||||
return continue_run
|
||||
|
||||
def run(self):
|
||||
console = Console()
|
||||
self.print_logo()
|
||||
self.USER_COMMANDS.update(self.memory_service.op_description_dict)
|
||||
|
||||
while True:
|
||||
query = questionary.text(
|
||||
"Enter your message or command:",
|
||||
multiline=False,
|
||||
qmark=">",
|
||||
).ask()
|
||||
try:
|
||||
query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask()
|
||||
if not query:
|
||||
continue
|
||||
|
||||
query = query.rstrip()
|
||||
query: str = query.strip()
|
||||
|
||||
if query == "":
|
||||
console.print("Empty input received. Try again!")
|
||||
continue
|
||||
# handle cli / commands with memory ops
|
||||
if query.startswith("/"):
|
||||
if self.process_commands(query=query):
|
||||
continue
|
||||
else:
|
||||
break
|
||||
|
||||
# Handle CLI commands
|
||||
if query.startswith("/"):
|
||||
if query.lower() == "/exit":
|
||||
break
|
||||
elif query.lower() == "/memory":
|
||||
console.print(self.memory_service.retrieve_all())
|
||||
elif query.lower() == "/help":
|
||||
questionary.print("CLI commands", "bold")
|
||||
for cmd, desc in self.USER_COMMANDS.items():
|
||||
questionary.print(cmd, "bold")
|
||||
questionary.print(f" {desc}")
|
||||
|
||||
continue
|
||||
|
||||
while True:
|
||||
try:
|
||||
# with console.status("[bold cyan]Thinking..."):
|
||||
msg = None
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
if self.stream:
|
||||
for msg in self.chat_with_memory(query=query):
|
||||
console.print(msg.delta, end="")
|
||||
console.print()
|
||||
questionary.print(msg.delta, end="")
|
||||
questionary.print("")
|
||||
else:
|
||||
msg = self.chat_with_memory(query=query)
|
||||
questionary.print(msg.message.content)
|
||||
self.memory_service.add_messages(msg.message)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
questionary.print("User interrupt occurred.")
|
||||
is_exit = questionary.confirm("continue exit").unsafe_ask()
|
||||
if is_exit:
|
||||
self.memory_service.stop_service()
|
||||
break
|
||||
except KeyboardInterrupt:
|
||||
console.print("User interrupt occurred.")
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
except Exception as e:
|
||||
console.print(
|
||||
f"An exception occurred when running chat_with_memory(): {e}"
|
||||
)
|
||||
retry = questionary.confirm("Retry chat_with_memory()?").ask()
|
||||
if not retry:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
line = f"An exception occurred when running cli memory chat. args={e.args}"
|
||||
questionary.print(line)
|
||||
self.logger.exception(line)
|
||||
continue
|
||||
|
||||
questionary.print(f"A memory writing thread is still running, please be patient and wait!")
|
||||
|
|
|
|||
|
|
@ -1,31 +1,27 @@
|
|||
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 memory_scope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from memory_scope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class GlobalContext(object):
|
||||
def __init__(self):
|
||||
self.global_configs: Dict[str, Any] = {}
|
||||
|
||||
self.worker_dict: Dict[str, Dict[str, BaseWorker]] = {}
|
||||
self.global_config: Dict[str, Any] = {}
|
||||
self.worker_config: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
self.memory_service_dict: Dict[str, BaseMemoryService] = {}
|
||||
self.model_dict: Dict[str, BaseModel] = {}
|
||||
|
||||
self.memory_chat_dict: Dict[str, BaseMemoryChat] = {}
|
||||
|
||||
self.vector_store: BaseVectorStore | None = None
|
||||
|
||||
self.monitor: BaseMonitor | None = None
|
||||
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
|
||||
self.language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
|
||||
GLOBAL_CONTEXT = GlobalContext()
|
||||
G_CONTEXT = GlobalContext()
|
||||
|
|
|
|||
|
|
@ -1,67 +0,0 @@
|
|||
import datetime
|
||||
from typing import List
|
||||
|
||||
from .base_memory_chat import BaseMemoryChat
|
||||
from .global_context import GLOBAL_CONTEXT
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from models.base_model import BaseModel
|
||||
from prompts.memory_chat_prompt import SYSTEM_PROMPT, MEMORY_PROMPT
|
||||
from scheme.message import Message
|
||||
from .memory_service import MemoryService
|
||||
|
||||
|
||||
class MemoryChat(BaseMemoryChat):
|
||||
|
||||
def __init__(self, generation_model: str, history_msg_count: int, chat_name: str, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_service = MemoryService(chat_name=chat_name, **kwargs)
|
||||
self.generation_model_name: str = generation_model
|
||||
self.history_msg_count: int = history_msg_count
|
||||
|
||||
self._generation_model: BaseModel | None = None
|
||||
self.history_message_list: List[Message] = []
|
||||
|
||||
@property
|
||||
def generation_model(self):
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict[
|
||||
self.generation_model_name
|
||||
]
|
||||
return self._generation_model
|
||||
|
||||
@staticmethod
|
||||
def get_system_prompt(related_memories: List[str], time_created: int) -> Message:
|
||||
system_prompt = SYSTEM_PROMPT[GLOBAL_CONTEXT.language]
|
||||
if related_memories:
|
||||
memory_prompt = MEMORY_PROMPT[GLOBAL_CONTEXT.language]
|
||||
system_prompt = "\n".join([system_prompt, memory_prompt] + related_memories)
|
||||
return Message(
|
||||
role=MessageRoleEnum.SYSTEM,
|
||||
content=system_prompt.strip(),
|
||||
time_created=time_created,
|
||||
)
|
||||
|
||||
def chat_with_memory(self, query: str):
|
||||
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
|
||||
)
|
||||
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 :]
|
||||
all_messages = [system_message] + self.history_message_list
|
||||
# TODO at xian zhe
|
||||
return self.generation_model.call(messages=all_messages, stream=True)
|
||||
|
||||
def run(self):
|
||||
self.memory_service.start_memory_backend()
|
||||
while True:
|
||||
query = input("wait for input:")
|
||||
if query in ["stop", "停止"]:
|
||||
break
|
||||
self.chat_with_memory(query=query)
|
||||
|
|
@ -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)
|
||||
|
|
@ -1,135 +1,82 @@
|
|||
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 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
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
import fire
|
||||
import yaml
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.enumeration.language_enum import LanguageEnum
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
from memory_scope.utils.timer import timer
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
|
||||
|
||||
class CliJob(object):
|
||||
|
||||
def __init__(self, config_path: str):
|
||||
self.config_path: str = config_path
|
||||
self.config_base_dir: str = os.path.dirname(config_path)
|
||||
def __init__(self):
|
||||
self.config: Dict[str, Any] = {}
|
||||
self.logger: Logger = Logger.get_logger("cli_job", to_stream=False)
|
||||
|
||||
self.worker_chat_dict: Dict[str, List[str]] = {}
|
||||
self.logger: Logger = Logger.get_logger("memory_chat")
|
||||
def load_config(self, path: str):
|
||||
with open(path) as f:
|
||||
if path.endswith("yaml"):
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
elif path.endswith("json"):
|
||||
self.config = json.load(f)
|
||||
else:
|
||||
raise RuntimeError("not supported config file type!")
|
||||
|
||||
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))
|
||||
|
||||
@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(self):
|
||||
G_CONTEXT.global_config = global_config = self.config["global_config"]
|
||||
G_CONTEXT.language = LanguageEnum(global_config["language"])
|
||||
G_CONTEXT.thread_pool = ThreadPoolExecutor(max_workers=int(global_config["max_workers"]))
|
||||
|
||||
@timer
|
||||
def init_global_content_by_config(self):
|
||||
with open(complete_config_name(self.config_path)) as f:
|
||||
self.config = json.load(f)
|
||||
|
||||
GLOBAL_CONTEXT.global_configs = self.config["global_configs"]
|
||||
# set global config
|
||||
self.set_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)
|
||||
|
||||
@staticmethod
|
||||
def run():
|
||||
with GLOBAL_CONTEXT.thread_pool:
|
||||
memory_chat = list(GLOBAL_CONTEXT.memory_chat_dict.values())[0]
|
||||
# init vector_store
|
||||
vector_store_config = self.config["vector_store"]
|
||||
embedding_model = G_CONTEXT.model_dict[vector_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.vector_store = init_instance_by_config(vector_store_config, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["worker"]
|
||||
|
||||
def run(self, config: str):
|
||||
self.load_config(config)
|
||||
self.init_global_content_by_config()
|
||||
|
||||
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()
|
||||
G_CONTEXT.vector_store.close()
|
||||
G_CONTEXT.monitor.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
fire.Fire(main)
|
||||
cli_job = CliJob()
|
||||
fire.Fire(cli_job.run)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,9 @@
|
|||
WORKFLOW_NAME = "workflow_name"
|
||||
|
||||
RESULT = "result"
|
||||
|
||||
CHAT_MESSAGES = "chat_messages"
|
||||
|
||||
RELATED_MEMORIES = "related_memories"
|
||||
|
||||
MESSAGES = "messages"
|
||||
|
|
@ -16,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"
|
||||
|
|
@ -110,3 +112,5 @@ DATATIME_KEY_MAP = {
|
|||
"周": "week",
|
||||
"星期几": "weekday",
|
||||
}
|
||||
|
||||
CONTENT_MODIFIED = "content_modified"
|
||||
25
memory_scope/memory/operation/base_operation.py
Normal file
25
memory_scope/memory/operation/base_operation.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Literal
|
||||
|
||||
OPERATION_TYPE = Literal["frontend", "backend"]
|
||||
|
||||
|
||||
class BaseOperation(metaclass=ABCMeta):
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
||||
def __init__(self, name: str, description: str = "", **kwargs):
|
||||
self.name: str = name
|
||||
self.description: str = description
|
||||
|
||||
def init_workflow(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def run_operation(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run_operation_backend(self):
|
||||
pass
|
||||
|
||||
def stop_operation_backend(self):
|
||||
pass
|
||||
124
memory_scope/memory/operation/base_workflow.py
Normal file
124
memory_scope/memory/operation/base_workflow.py
Normal file
|
|
@ -0,0 +1,124 @@
|
|||
import re
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.constants.common_constants import WORKFLOW_NAME
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
from memory_scope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class BaseWorkflow(object):
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
workflow: str,
|
||||
thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool,
|
||||
**kwargs):
|
||||
|
||||
self.name: str = name
|
||||
self.workflow: str = workflow
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.kwargs = kwargs
|
||||
|
||||
self.workflow_worker_list: List[List[List[str]]] = []
|
||||
self.worker_dict: Dict[str, BaseWorker | bool] = {}
|
||||
self.context: Dict[str, Any] = {}
|
||||
self.context_lock = threading.Lock()
|
||||
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
if self.workflow:
|
||||
self._parse_workflow()
|
||||
self._print_workflow()
|
||||
|
||||
def _parse_workflow(self):
|
||||
# re-match e.g., [a|b],c,[d,e,f|g,h],j
|
||||
pattern = r'(\[[^\]]*\]|[^,]+)'
|
||||
workflow_split = re.findall(pattern, self.workflow)
|
||||
for workflow_part in workflow_split:
|
||||
# e.g., [d,e,f|g,h]
|
||||
workflow_part = workflow_part.strip()
|
||||
if '[' in workflow_part or ']' in workflow_part:
|
||||
workflow_part = workflow_part.replace('[', '').replace(']', '')
|
||||
|
||||
# e.g., ["d,e,f", "g,h"]
|
||||
line_split = [x.strip() for x in workflow_part.split("|") if x]
|
||||
if len(line_split) <= 0:
|
||||
continue
|
||||
|
||||
# is under multi thread cond
|
||||
is_multi_thread: bool = len(line_split) > 1
|
||||
|
||||
# e.g., ["d","e","f"]
|
||||
line_split_split: List[List[str]] = []
|
||||
for sub_line_split in line_split:
|
||||
sub_split = [x.strip() for x in sub_line_split.split(",")]
|
||||
line_split_split.append(sub_split)
|
||||
# add workers
|
||||
for sub_item in sub_split:
|
||||
self.worker_dict[sub_item] = is_multi_thread
|
||||
self.workflow_worker_list.append(line_split_split)
|
||||
|
||||
def _print_workflow(self):
|
||||
self.logger.info(f"----- print_workflow_{self.name}_begin -----")
|
||||
i: int = 0
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
for w in workflow_part[0]:
|
||||
self.logger.info(f"stage{i}: {w}")
|
||||
i += 1
|
||||
else:
|
||||
for w_zip in zip_longest(*workflow_part, fillvalue="-"):
|
||||
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
|
||||
i += 1
|
||||
for w in w_zip:
|
||||
if w == "-":
|
||||
continue
|
||||
self.logger.info(f"----- print_workflow_{self.name}_end -----")
|
||||
|
||||
def init_workers(self):
|
||||
for name in list(self.worker_dict.keys()):
|
||||
if name not in G_CONTEXT.worker_config:
|
||||
raise RuntimeError(f"worker={name} is not exists in worker_config!")
|
||||
|
||||
self.worker_dict[name] = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
name=name,
|
||||
is_multi_thread=self.worker_dict[name],
|
||||
context=self.context,
|
||||
context_lock=self.context_lock)
|
||||
|
||||
def _run_sub_workflow(self, worker_list: List[str]) -> bool:
|
||||
for name in worker_list:
|
||||
worker = self.worker_dict[name]
|
||||
worker.run()
|
||||
if not worker.continue_run:
|
||||
return False
|
||||
return True
|
||||
|
||||
def run_workflow(self):
|
||||
with Timer(f"run_workflow_{self.name}"):
|
||||
self.context[WORKFLOW_NAME] = self.name
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
if len(workflow_part) == 1:
|
||||
if not self._run_sub_workflow(workflow_part[0]):
|
||||
break
|
||||
else:
|
||||
t_list = []
|
||||
for sub_workflow in workflow_part:
|
||||
t_list.append(G_CONTEXT.thread_pool.submit(
|
||||
self._run_sub_workflow, sub_workflow))
|
||||
|
||||
flag = True
|
||||
for future in as_completed(t_list):
|
||||
if not future.result():
|
||||
flag = False
|
||||
break
|
||||
if not flag:
|
||||
break
|
||||
34
memory_scope/memory/operation/read_memory.py
Normal file
34
memory_scope/memory/operation/read_memory.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import RESULT, CHAT_MESSAGES
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memory_scope.memory.operation.base_workflow import BaseWorkflow
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
|
||||
class ReadMemory(BaseWorkflow, BaseOperation):
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
description: str,
|
||||
chat_messages: List[Message],
|
||||
his_msg_count: int = 0, # supplement to the current query
|
||||
contextual_msg_count: int = 0, # for the current context dialogue
|
||||
**kwargs):
|
||||
super().__init__(name=name, **kwargs)
|
||||
BaseOperation.__init__(self, name=name, description=description)
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
self.his_msg_count: int = his_msg_count
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
|
||||
def init_workflow(self):
|
||||
self.init_workers()
|
||||
|
||||
def run_operation(self):
|
||||
max_count = 1 + max(self.his_msg_count, self.contextual_msg_count)
|
||||
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
return result
|
||||
56
memory_scope/memory/operation/summary_memory.py
Normal file
56
memory_scope/memory/operation/summary_memory.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
import time
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.constants.common_constants import RESULT
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memory_scope.memory.operation.base_workflow import BaseWorkflow
|
||||
|
||||
|
||||
class SummaryMemory(BaseWorkflow, BaseOperation):
|
||||
operation_type: OPERATION_TYPE = "backend"
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
description: str,
|
||||
interval_time: int = 300,
|
||||
**kwargs):
|
||||
super().__init__(name=name, **kwargs)
|
||||
BaseOperation.__init__(self, name=name, description=description)
|
||||
|
||||
self.interval_time: int = interval_time
|
||||
|
||||
self._operation_status_run: bool = False
|
||||
self._loop_switch: bool = False
|
||||
self._run_thread = None
|
||||
|
||||
def init_workflow(self):
|
||||
self.init_workers()
|
||||
|
||||
def run_operation(self):
|
||||
if self._operation_status_run:
|
||||
return
|
||||
|
||||
self._operation_status_run = True
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
self._operation_status_run = False
|
||||
return result
|
||||
|
||||
def _loop_operation(self):
|
||||
while self._loop_switch:
|
||||
for _ in range(self.interval_time):
|
||||
if self._loop_switch:
|
||||
time.sleep(1)
|
||||
else:
|
||||
break
|
||||
if self._loop_switch:
|
||||
self.run_operation()
|
||||
|
||||
def run_operation_backend(self):
|
||||
if not self._loop_switch:
|
||||
self._loop_switch = True
|
||||
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
|
||||
def stop_operation_backend(self):
|
||||
self._loop_switch = False
|
||||
83
memory_scope/memory/operation/write_memory.py
Normal file
83
memory_scope/memory/operation/write_memory.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
import time
|
||||
from typing import List
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES, RESULT
|
||||
from memory_scope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memory_scope.memory.operation.base_workflow import BaseWorkflow
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
|
||||
class WriteMemory(BaseWorkflow, BaseOperation):
|
||||
operation_type: OPERATION_TYPE = "backend"
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
description: str,
|
||||
chat_messages: List[Message],
|
||||
his_msg_count: int = 0,
|
||||
message_lock=None,
|
||||
interval_time: int = 60,
|
||||
contextual_msg_count: int = 6,
|
||||
**kwargs):
|
||||
|
||||
super().__init__(name=name, **kwargs)
|
||||
BaseOperation.__init__(self, name=name, description=description)
|
||||
|
||||
self.chat_messages: List[Message] = chat_messages
|
||||
self.his_msg_count: int = his_msg_count
|
||||
self.message_lock = message_lock
|
||||
self.interval_time: int = interval_time
|
||||
self.contextual_msg_count: int = contextual_msg_count
|
||||
|
||||
self._operation_status_run: bool = False
|
||||
self._loop_switch: bool = False
|
||||
|
||||
@property
|
||||
def not_memorized_size(self):
|
||||
return sum([not x.memorized for x in self.chat_messages])
|
||||
|
||||
def set_memorized(self):
|
||||
if self.message_lock:
|
||||
with self.message_lock:
|
||||
for msg in self.chat_messages:
|
||||
msg.memorized = True
|
||||
|
||||
def init_workflow(self):
|
||||
self.init_workers()
|
||||
|
||||
def run_operation(self):
|
||||
if self._operation_status_run:
|
||||
return
|
||||
|
||||
self._operation_status_run = True
|
||||
not_memorized_size = self.not_memorized_size
|
||||
if not_memorized_size < self.contextual_msg_count:
|
||||
return
|
||||
|
||||
max_count = not_memorized_size + self.his_msg_count
|
||||
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
|
||||
self.run_workflow()
|
||||
result = self.context.get(RESULT)
|
||||
self.context.clear()
|
||||
self.set_memorized()
|
||||
self._operation_status_run = False
|
||||
return result
|
||||
|
||||
def _loop_operation(self):
|
||||
while self._loop_switch:
|
||||
for _ in range(self.interval_time):
|
||||
if self._loop_switch:
|
||||
time.sleep(1)
|
||||
else:
|
||||
break
|
||||
if self._loop_switch:
|
||||
self.run_operation()
|
||||
|
||||
def run_operation_backend(self):
|
||||
if not self._loop_switch:
|
||||
self._loop_switch = True
|
||||
return G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
|
||||
def stop_operation_backend(self):
|
||||
self._loop_switch = False
|
||||
52
memory_scope/memory/service/base_memory_service.py
Normal file
52
memory_scope/memory/service/base_memory_service.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
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
|
||||
|
||||
@abstractmethod
|
||||
def _init_operation(self, memory_operations: Dict[str, dict]):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_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.do_operation(self.read_memory_key)
|
||||
|
||||
def stop_service(self):
|
||||
pass
|
||||
55
memory_scope/memory/service/chat_memory_service.py
Normal file
55
memory_scope/memory/service/chat_memory_service.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
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
|
||||
|
||||
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)
|
||||
|
||||
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 start_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()
|
||||
|
||||
def stop_service(self):
|
||||
for _, operation in self._operation_dict.items():
|
||||
if operation.operation_type == "backend":
|
||||
operation.stop_operation_backend()
|
||||
56
memory_scope/memory/worker/base_worker.py
Normal file
56
memory_scope/memory/worker/base_worker.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Any, Dict
|
||||
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(metaclass=ABCMeta):
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
context: Dict[str, Any],
|
||||
context_lock=None,
|
||||
raise_exception: bool = True,
|
||||
is_multi_thread: bool = False,
|
||||
**kwargs):
|
||||
|
||||
self.name: str = name
|
||||
self.context: Dict[str, Any] = context
|
||||
self.context_lock = context_lock
|
||||
self.raise_exception: bool = raise_exception
|
||||
self.is_multi_thread: bool = is_multi_thread
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.continue_run: bool = True
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
@abstractmethod
|
||||
def _run(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
self.logger.info(f"----- worker_{self.name}_begin -----")
|
||||
with Timer(self.name, log_time=False) as t:
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
else:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as e:
|
||||
self.logger.exception(f"run {self.name} failed! args={e.args}")
|
||||
|
||||
self.logger.info(f"----- worker_{self.name}_end cost={t.cost_str}-----")
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context.get(key, default)
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
else:
|
||||
self.context[key] = value
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
12
memory_scope/memory/worker/dummy_worker.py
Normal file
12
memory_scope/memory/worker/dummy_worker.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
import datetime
|
||||
|
||||
from memory_scope.constants.common_constants import RESULT, WORKFLOW_NAME
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class DummyWorker(BaseWorker):
|
||||
def _run(self):
|
||||
workflow_name = self.get_context(WORKFLOW_NAME)
|
||||
self.logger.info(f"enter workflow={workflow_name}.dummy_worker!")
|
||||
ts = int(datetime.datetime.now().timestamp())
|
||||
self.set_context(RESULT, f"test {workflow_name} \nts={ts}")
|
||||
82
memory_scope/memory/worker/memory_base_worker.py
Normal file
82
memory_scope/memory/worker/memory_base_worker.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
from abc import ABCMeta
|
||||
from typing import List
|
||||
|
||||
from memory_scope.chat.global_context import G_CONTEXT
|
||||
from memory_scope.constants.common_constants import CHAT_MESSAGES
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
||||
|
||||
def __init__(self,
|
||||
embedding_model: str = "",
|
||||
generation_model: str = "",
|
||||
rank_model: str = "",
|
||||
**kwargs):
|
||||
super(MemoryBaseWorker, self).__init__(**kwargs)
|
||||
|
||||
self._embedding_model: BaseModel | str = embedding_model
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self._rank_model: BaseModel | str = rank_model
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
self._user_id: str | None = None
|
||||
|
||||
@property
|
||||
def messages(self) -> List[Message]:
|
||||
return self.get_context(CHAT_MESSAGES)
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value):
|
||||
self.set_context(CHAT_MESSAGES, value)
|
||||
|
||||
@property
|
||||
def embedding_model(self) -> BaseModel:
|
||||
if isinstance(self._embedding_model, str):
|
||||
self._embedding_model = G_CONTEXT.model_dict[self._embedding_model]
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self) -> BaseModel:
|
||||
if isinstance(self._rank_model, str):
|
||||
self._rank_model = G_CONTEXT.model_dict[self._rank_model]
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
def vector_store(self) -> BaseVectorStore:
|
||||
if self._vector_store is None:
|
||||
self._vector_store = G_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
||||
@property
|
||||
def monitor(self):
|
||||
if self._monitor is None:
|
||||
self._monitor = G_CONTEXT.monitor
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
def user_id(self) -> str:
|
||||
if self._user_id is None:
|
||||
message = [x for x in self.messages if x.role == MessageRoleEnum.USER.value][-1]
|
||||
self._user_id = message.role_name
|
||||
return self._user_id
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
|
||||
@staticmethod
|
||||
def get_prompt(prompt: dict) -> str:
|
||||
return prompt[G_CONTEXT.global_configs["language"]]
|
||||
|
|
@ -1,4 +1 @@
|
|||
from utils.registry import Registry
|
||||
|
||||
# __all__ = ["LlamaIndexEmbeddingModel", "LlamaIndexGenerationModel", "LlamaIndexRerankModel"]
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
|
|
|||
|
|
@ -1,12 +1,15 @@
|
|||
import inspect
|
||||
import time
|
||||
from abc import abstractmethod, ABCMeta
|
||||
from typing import Any
|
||||
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from . import MODEL_REGISTRY
|
||||
from .response import ModelResponse, ModelResponseGen
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.utils.registry import Registry
|
||||
from memory_scope.utils.timer import Timer
|
||||
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
||||
|
||||
class BaseModel(metaclass=ABCMeta):
|
||||
|
|
@ -14,7 +17,7 @@ class BaseModel(metaclass=ABCMeta):
|
|||
|
||||
def __init__(self,
|
||||
model_name: str,
|
||||
method_type: str,
|
||||
module_name: str,
|
||||
timeout: int = None,
|
||||
max_retries: int = 3,
|
||||
retry_interval: float = 1.0,
|
||||
|
|
@ -22,24 +25,32 @@ class BaseModel(metaclass=ABCMeta):
|
|||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
self.method_type: str = method_type
|
||||
self.module_name: str = module_name
|
||||
self.timeout: int = timeout
|
||||
self.max_retries: int = max_retries
|
||||
self.retry_interval: float = retry_interval
|
||||
self.kwargs_filter: bool = kwargs_filter
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.data = {}
|
||||
self._model: Any = None
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
obj_cls = MODEL_REGISTRY.get(self.method_type)
|
||||
if not obj_cls:
|
||||
raise RuntimeError(f"method_type={self.method_type} is not supported!")
|
||||
@property
|
||||
def model(self):
|
||||
if self._model is None:
|
||||
if self.module_name not in MODEL_REGISTRY.module_dict:
|
||||
raise RuntimeError(f"method_type={self.module_name} is not supported!")
|
||||
obj_cls = MODEL_REGISTRY[self.module_name]
|
||||
|
||||
if kwargs_filter:
|
||||
allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys())
|
||||
kwargs = {key: value for key, value in kwargs.items() if key in allowed_kwargs}
|
||||
|
||||
self.model = obj_cls(**kwargs)
|
||||
if self.kwargs_filter:
|
||||
allowed_kwargs = list(inspect.signature(obj_cls.__init__).parameters.keys())
|
||||
kwargs = {key: value for key, value in self.kwargs.items() if key in allowed_kwargs}
|
||||
else:
|
||||
kwargs = self.kwargs
|
||||
self._model = obj_cls(**kwargs)
|
||||
return self._model
|
||||
|
||||
@abstractmethod
|
||||
def before_call(self, **kwargs) -> None:
|
||||
|
|
@ -70,8 +81,8 @@ class BaseModel(metaclass=ABCMeta):
|
|||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
self.before_call(stream=stream, **kwargs)
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.before_call(stream=stream, **kwargs)
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
model_response = self._call(stream=stream, **kwargs)
|
||||
|
|
@ -97,8 +108,8 @@ class BaseModel(metaclass=ABCMeta):
|
|||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
self.before_call(**kwargs)
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.before_call(**kwargs)
|
||||
for i in range(self.max_retries):
|
||||
try:
|
||||
model_response = await self._async_call(**kwargs)
|
||||
|
|
|
|||
|
|
@ -2,20 +2,15 @@ from typing import List
|
|||
|
||||
from llama_index.embeddings.dashscope import DashScopeEmbedding
|
||||
|
||||
from models import MODEL_REGISTRY
|
||||
from models.base_model import BaseModel
|
||||
from models.response import ModelResponse, ModelResponseGen
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memory_scope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
class LlamaIndexEmbeddingModel(BaseModel):
|
||||
m_type: ModelEnum = ModelEnum.EMBEDDING_MODEL
|
||||
|
||||
MODEL_REGISTRY.batch_register(
|
||||
[
|
||||
DashScopeEmbedding,
|
||||
]
|
||||
)
|
||||
MODEL_REGISTRY.register("dashscope_embedding", DashScopeEmbedding)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
|
|
@ -42,16 +37,11 @@ class LlamaIndexEmbeddingModel(BaseModel):
|
|||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
return ModelResponse(
|
||||
m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data)
|
||||
)
|
||||
return ModelResponse(m_type=self.m_type, raw=self.model.get_text_embedding_batch(**self.data))
|
||||
|
||||
async def _async_call(self, **kwargs) -> ModelResponse:
|
||||
"""
|
||||
:param kwargs:
|
||||
:return:
|
||||
"""
|
||||
return ModelResponse(
|
||||
m_type=self.m_type,
|
||||
raw=await self.model.aget_text_embedding_batch(**self.data),
|
||||
)
|
||||
return ModelResponse(m_type=self.m_type, raw=await self.model.aget_text_embedding_batch(**self.data))
|
||||
|
|
|
|||
|
|
@ -1,73 +1,62 @@
|
|||
from typing import List, Dict
|
||||
from llama_index.core.base.llms.types import (
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
CompletionResponse,
|
||||
)
|
||||
import datetime
|
||||
from typing import List
|
||||
|
||||
from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse
|
||||
from llama_index.llms.dashscope import DashScope
|
||||
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from . import MODEL_REGISTRY
|
||||
from .base_model import BaseModel
|
||||
from .response import ModelResponse, ModelResponseGen
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memory_scope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
|
||||
class LlamaIndexGenerationModel(BaseModel):
|
||||
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
|
||||
|
||||
# TODO rename module name at xianzhe
|
||||
MODEL_REGISTRY.batch_register([
|
||||
DashScope,
|
||||
])
|
||||
MODEL_REGISTRY.register("dashscope_generation", DashScope)
|
||||
|
||||
def before_call(self, **kwargs) -> None:
|
||||
def before_call(self, **kwargs):
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
|
||||
messages: List[Message] = kwargs.pop("messages", [])
|
||||
|
||||
if prompt:
|
||||
input_text = prompt
|
||||
input_type = 'prompt'
|
||||
llama_input = input_text
|
||||
self.data = {"prompt": prompt}
|
||||
elif messages:
|
||||
input_text = messages
|
||||
input_type = 'messages'
|
||||
llama_input = [ChatMessage(role=x.role, content=x.content) for x in input_text]
|
||||
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
self.data = {input_type: llama_input}
|
||||
|
||||
def after_call(self,
|
||||
model_response: ModelResponse,
|
||||
stream: bool = False,
|
||||
**kwargs) -> ModelResponse | ModelResponseGen:
|
||||
|
||||
model_response.message = Message(role=MessageRoleEnum.ASSISTANT, content="")
|
||||
|
||||
call_result = model_response.raw
|
||||
if stream:
|
||||
def gen() -> ModelResponseGen:
|
||||
text = ""
|
||||
for response in call_result:
|
||||
delta = response.delta
|
||||
text += delta
|
||||
model_response.text = text
|
||||
model_response.delta = delta
|
||||
model_response.message.content += response.delta
|
||||
model_response.delta = response.delta
|
||||
yield model_response
|
||||
return gen()
|
||||
else:
|
||||
if isinstance(call_result, CompletionResponse):
|
||||
content = call_result.text
|
||||
model_response.message.content = call_result.text
|
||||
elif isinstance(call_result, ChatResponse):
|
||||
content = call_result.message.content
|
||||
model_response.message.content = call_result.message.content
|
||||
else:
|
||||
raise NotImplementedError
|
||||
model_response.text = content
|
||||
|
||||
return model_response
|
||||
|
||||
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
|
||||
|
||||
assert "prompt" in self.data or "messages" in self.data
|
||||
results = ModelResponse(m_type=self.m_type)
|
||||
|
||||
if 'prompt' in self.data:
|
||||
if "prompt" in self.data:
|
||||
if stream:
|
||||
response = self.model.stream_complete(**self.data)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -4,24 +4,21 @@ from llama_index.core.data_structs import Node
|
|||
from llama_index.core.schema import NodeWithScore
|
||||
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
|
||||
|
||||
from models import MODEL_REGISTRY
|
||||
from models.base_model import BaseModel
|
||||
from models.response import ModelResponse, ModelResponseGen
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memory_scope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
|
||||
class LlamaIndexRerankModel(BaseModel):
|
||||
class LlamaIndexRankModel(BaseModel):
|
||||
m_type: ModelEnum = ModelEnum.RANK_MODEL
|
||||
|
||||
MODEL_REGISTRY.batch_register([
|
||||
DashScopeRerank
|
||||
])
|
||||
MODEL_REGISTRY.register("dashscope_rank", DashScopeRerank)
|
||||
|
||||
def before_call(self, **kwargs) -> None:
|
||||
assert "query" in kwargs or "documents" in kwargs
|
||||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
if isinstance(documents, str):
|
||||
documents = [documents]
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
|
||||
|
|
@ -29,10 +26,7 @@ class LlamaIndexRerankModel(BaseModel):
|
|||
nodes = [NodeWithScore(node=Node(text=doc), score=-1.0) for doc in documents]
|
||||
self._get_documents_mapping(documents)
|
||||
|
||||
self.data = {
|
||||
"nodes": nodes,
|
||||
"query_str": query,
|
||||
}
|
||||
self.data = {"nodes": nodes, "query_str": query}
|
||||
|
||||
def after_call(self, model_response: ModelResponse, **kwargs) -> ModelResponse:
|
||||
if not model_response.rank_scores:
|
||||
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: """
|
||||
|
|
|
|||
|
|
@ -1,13 +1,17 @@
|
|||
import datetime
|
||||
from typing import Dict, List
|
||||
|
||||
import
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
from memory_scope.utils.tool_functions import md5_hash
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
user_id: str = Field("", description="unique memory id for user")
|
||||
|
||||
memory_id: str = Field("", description="unique id for memory item")
|
||||
|
||||
user_id: str = Field("", description="unique memory id for user")
|
||||
|
||||
content: str = Field("", description="memory content")
|
||||
|
||||
score_similar: float = Field(0, description="es similar score")
|
||||
|
|
@ -24,3 +28,14 @@ class MemoryNode(BaseModel):
|
|||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
|
||||
timestamp: int = Field(int(datetime.datetime.now().timestamp()), description="timestamp of the memory node")
|
||||
|
||||
@property
|
||||
def node_keys(self):
|
||||
return list(self.model_json_schema()["properties"].keys())
|
||||
|
||||
def __getitem__(self, key: str):
|
||||
return self.model_dump().get(key)
|
||||
|
||||
def gen_memory_id(self):
|
||||
self.memory_id = f"{self.user_id}_{self.timestamp}_{md5_hash(self.content)[:8]}"
|
||||
|
|
|
|||
|
|
@ -1,9 +1,19 @@
|
|||
import datetime
|
||||
from typing import Dict
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
||||
class Message(BaseModel):
|
||||
role: str = Field(..., description="The role of the message sender (user, assistant, system)")
|
||||
|
||||
role_name: str = Field("", description="role name")
|
||||
|
||||
content: str = Field(..., description="The body of the message")
|
||||
|
||||
time_created: int = Field("", description="Timestamp when the message was created")
|
||||
time_created: int = Field(int(datetime.datetime.now().timestamp()),
|
||||
description="Timestamp when the message was created")
|
||||
|
||||
memorized: bool = Field(False, description="indicate whether message is memorized")
|
||||
|
||||
meta_data: Dict[str, str] = Field({}, description="meta data for msg")
|
||||
|
|
|
|||
|
|
@ -3,11 +3,12 @@ from typing import Generator, List, Dict, Any
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from enumeration.model_enum import ModelEnum
|
||||
from memory_scope.enumeration.model_enum import ModelEnum
|
||||
from memory_scope.scheme.message import Message
|
||||
|
||||
|
||||
class ModelResponse(BaseModel):
|
||||
text: str = Field("", description="generation model result")
|
||||
message: Message | None = Field(None, description="generation model result")
|
||||
|
||||
delta: str = Field("", description="New text that just streamed in (only used when streaming)")
|
||||
|
||||
|
|
@ -27,12 +28,7 @@ class ModelResponse(BaseModel):
|
|||
|
||||
def __str__(self, max_size=100, **kwargs):
|
||||
result = {}
|
||||
try:
|
||||
all_dict = self.model_dump()
|
||||
except Exception:
|
||||
all_dict = self.dict()
|
||||
|
||||
for key, value in all_dict.items():
|
||||
for key, value in self.model_dump().items():
|
||||
if key == "raw" or not value:
|
||||
continue
|
||||
|
||||
|
|
@ -18,8 +18,8 @@ class BaseMonitor(metaclass=ABCMeta):
|
|||
:return:
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def flush(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -1,61 +1,34 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from models.base_model import BaseModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class BaseVectorStore(metaclass=ABCMeta):
|
||||
|
||||
def __init__(self, index_name: str, embedding_model: BaseModel, content_key: str = "text", **kwargs):
|
||||
self.index_name: str = index_name
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.content_key: str = content_key
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
@abstractmethod
|
||||
def retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
|
||||
"""
|
||||
:param text:
|
||||
:param limit_size:
|
||||
:param filter_dict:
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def async_retrieve(self, query: str, top_k: int, filter_dict: Dict[str, List[str]]):
|
||||
"""
|
||||
:param text:
|
||||
:param limit_size:
|
||||
:param filter_dict:
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def insert(self, node: MemoryNode):
|
||||
""" TODO 是否overwrite
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def insert_batch(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
def insert_batch(self, nodes: List[MemoryNode]):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
def delete(self, node: MemoryNode):
|
||||
pass
|
||||
|
||||
def update(self, node: MemoryNode):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def flush(self):
|
||||
"""
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
|
|
|||
12
memory_scope/storage/dummy_monitor.py
Normal file
12
memory_scope/storage/dummy_monitor.py
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
|
||||
|
||||
class DummyMonitor(BaseMonitor):
|
||||
def add(self):
|
||||
pass
|
||||
|
||||
def add_token(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
|
@ -1,15 +1,28 @@
|
|||
from typing import Dict, List, Any
|
||||
|
||||
from llama_index.core.schema import TextNode
|
||||
from llama_index.core.vector_stores import VectorStoreQuery
|
||||
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
|
||||
from llama_index.core.vector_stores.types import MetadataFilters, ExactMatchFilter, VectorStoreQueryMode
|
||||
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.base_vector_store import BaseVectorStore
|
||||
|
||||
|
||||
class _ElasticsearchStore(ElasticsearchStore):
|
||||
async def adelete(self, ref_doc_id: str, **delete_kwargs: Any) -> None:
|
||||
"""
|
||||
Async delete node from Elasticsearch index.
|
||||
|
||||
Args:
|
||||
ref_doc_id: ID of the node to delete.
|
||||
delete_kwargs: Optional. Additional arguments to
|
||||
pass to AsyncElasticsearch delete_by_query.
|
||||
|
||||
Raises:
|
||||
Exception: If AsyncElasticsearch delete_by_query fails.
|
||||
"""
|
||||
return await self._store.delete(query={"term": {"_id": ref_doc_id}}, **delete_kwargs)
|
||||
|
||||
|
||||
def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str, Any]:
|
||||
|
|
@ -24,7 +37,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
|
|||
"""
|
||||
|
||||
result = {
|
||||
"bool" : {}
|
||||
"bool": {}
|
||||
}
|
||||
for key, value in standard_filters.items():
|
||||
if isinstance(value, list):
|
||||
|
|
@ -32,10 +45,10 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
|
|||
for v in value:
|
||||
operands.append(
|
||||
{
|
||||
"term":
|
||||
{
|
||||
f"metadata.{key}.keyword": {"value": v}
|
||||
}
|
||||
"term":
|
||||
{
|
||||
f"metadata.{key}.keyword": {"value": v}
|
||||
}
|
||||
}
|
||||
)
|
||||
result['bool'].update({"should": operands})
|
||||
|
|
@ -56,67 +69,69 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
|
|||
|
||||
|
||||
class LlamaIndexElasticSearchStore(BaseVectorStore):
|
||||
def __init__(self,
|
||||
index_name: str,
|
||||
embedding_model: BaseModel,
|
||||
content_key: str = "text",
|
||||
def __init__(self,
|
||||
embedding_model: BaseModel,
|
||||
index_name: str,
|
||||
es_url: str,
|
||||
use_hybrid: bool = True,
|
||||
**kwargs):
|
||||
|
||||
self.index_name: str = index_name
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
|
||||
self.es_store = ElasticsearchStore(index_name=self.index_name,
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True),
|
||||
**kwargs)
|
||||
|
||||
self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_model, llm=None)
|
||||
self.es_store = _ElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=use_hybrid),
|
||||
**kwargs)
|
||||
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
|
||||
service_context=self.service_context)
|
||||
|
||||
|
||||
def retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode:
|
||||
|
||||
filter = _to_elasticsearch_filter(filter_dict)
|
||||
embed_model=self.embedding_model.model)
|
||||
|
||||
def retrieve(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(
|
||||
vector_store_kwargs={
|
||||
"es_filter": filter
|
||||
},
|
||||
similarity_top_k=top_k
|
||||
)
|
||||
textnodes = retriever.retrieve(query)
|
||||
results = self._textnodes2memorynodes(textnodes)
|
||||
vector_store_kwargs={"es_filter": es_filter},
|
||||
similarity_top_k=top_k)
|
||||
text_nodes = retriever.retrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
async def async_retrieve(self,
|
||||
query: str,
|
||||
top_k: int,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
if filter_dict is None:
|
||||
filter_dict = {}
|
||||
|
||||
es_filter = _to_elasticsearch_filter(filter_dict)
|
||||
retriever = self.index.as_retriever(
|
||||
vector_store_kwargs={"es_filter": es_filter},
|
||||
similarity_top_k=top_k)
|
||||
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
|
||||
return [self._text_node_2_memory_node(n) for n in text_nodes]
|
||||
|
||||
return results
|
||||
|
||||
async def async_retrieve(self, query: str, top_k: int = 3, filter_dict: Dict[str, List[str]] = {}) -> MemoryNode:
|
||||
raise NotImplementedError
|
||||
## return await super().async_retrieve(text, limit_size, filter_dict)
|
||||
|
||||
def insert(self, node: MemoryNode):
|
||||
node = self._memorynode2textnode(node)
|
||||
self.index.insert_nodes([node])
|
||||
self.index.insert_nodes([self._memory_node_2_text_node(node)])
|
||||
|
||||
def insert_batch(self, node: MemoryNode) -> None:
|
||||
raise NotImplementedError
|
||||
def delete(self, node: MemoryNode):
|
||||
memory_id = node.memory_id
|
||||
return self.es_store.delete(memory_id)
|
||||
|
||||
def delete(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def flush(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def _memorynode2textnode(self, memory_node: MemoryNode) -> TextNode:
|
||||
content = memory_node.content
|
||||
meta = memory_node.model_dump(exclude={"content"})
|
||||
return TextNode(text=content, metadata=meta)
|
||||
def update(self, node: MemoryNode):
|
||||
self.delete(node)
|
||||
self.insert(node)
|
||||
|
||||
def _textnode2memorynode(self, text_node: TextNode) -> MemoryNode:
|
||||
content = text_node.text
|
||||
meta = text_node.metadata
|
||||
mem_node = MemoryNode(content=content, **meta)
|
||||
return mem_node
|
||||
def close(self):
|
||||
self.es_store.close()
|
||||
|
||||
def _textnodes2memorynodes(self, text_nodes: TextNode) -> MemoryNode:
|
||||
mem_nodes = [self._textnode2memorynode(node) for node in text_nodes]
|
||||
return mem_nodes
|
||||
|
||||
@staticmethod
|
||||
def _memory_node_2_text_node(memory_node: MemoryNode) -> TextNode:
|
||||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
metadata=memory_node.model_dump(exclude={"content"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
return MemoryNode(content=text_node.text, **text_node.metadata)
|
||||
|
|
|
|||
|
|
@ -1,185 +0,0 @@
|
|||
import re
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import as_completed
|
||||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from scheme.message import Message
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class Pipeline(object):
|
||||
def __init__(self,
|
||||
chat_name: str,
|
||||
memory_method_type: MemoryMethodEnum,
|
||||
pipeline_str: str,
|
||||
history_msg_count: int = 3,
|
||||
loop_interval_time: int = 300,
|
||||
loop_minimum_count: int = 20):
|
||||
|
||||
self.chat_name: str = chat_name
|
||||
self.memory_method_type: MemoryMethodEnum = memory_method_type
|
||||
self.pipeline_str: str = pipeline_str
|
||||
self.history_msg_count: int = history_msg_count
|
||||
self.loop_interval_time: int = loop_interval_time
|
||||
self.loop_minimum_count: int = loop_minimum_count
|
||||
|
||||
# pipeline上下文和锁
|
||||
self.context: Dict[str, Any] = {}
|
||||
self.context_lock = threading.Lock()
|
||||
|
||||
# pipeline run config
|
||||
self.loop_switch: bool = False
|
||||
self.pipeline_list: list[list] = []
|
||||
self.worker_set: set[str] = set()
|
||||
self.worker_dict: Dict[str, BaseWorker] = {}
|
||||
self.injected: bool = False
|
||||
|
||||
# message list
|
||||
self.history_message_list: List[Message] = []
|
||||
self.current_message_list: List[Message] = []
|
||||
self.message_lock = threading.Lock()
|
||||
|
||||
# 日志
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
self._parse_pipeline()
|
||||
|
||||
def _parse_pipeline(self):
|
||||
if not self.pipeline_str:
|
||||
return
|
||||
|
||||
# re-match e.g., [a|b],c,[d,e,f|g,h],j
|
||||
pattern = r'(\[[^\]]*\]|[^,]+)'
|
||||
pipeline_split = re.findall(pattern, self.pipeline_str)
|
||||
|
||||
self.pipeline_list = []
|
||||
for pipeline_part in pipeline_split:
|
||||
# e.g., [d,e,f|g,h]
|
||||
pipeline_part = pipeline_part.strip()
|
||||
if '[' in pipeline_part or ']' in pipeline_part:
|
||||
pipeline_part = pipeline_part.replace('[', '').replace(']', '')
|
||||
|
||||
# e.g., ["d,e,f", "g,h"]
|
||||
line_split = [x.strip() for x in pipeline_part.split("|") if x]
|
||||
if len(line_split) <= 0:
|
||||
continue
|
||||
|
||||
# e.g., ["d","e","f"]
|
||||
line_split_split = []
|
||||
for sub_line_split in line_split:
|
||||
sub_split = [x.strip() for x in sub_line_split.split(",")]
|
||||
line_split_split.append(sub_split)
|
||||
# add to workers
|
||||
self.worker_set.update(sub_split)
|
||||
self.pipeline_list.append(line_split_split)
|
||||
|
||||
def _visit_and_inject_workers(self):
|
||||
if self.injected:
|
||||
return
|
||||
|
||||
self.worker_dict = GLOBAL_CONTEXT.worker_dict[self.chat_name]
|
||||
|
||||
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline Begin -----")
|
||||
i: int = 0
|
||||
for pipeline_part in self.pipeline_list:
|
||||
if len(pipeline_part) == 1:
|
||||
for w in pipeline_part[0]:
|
||||
self.logger.info(f"stage{i}: {w}")
|
||||
i += 1
|
||||
if w not in self.worker_dict:
|
||||
raise RuntimeError(f"worker={w} is not inited.")
|
||||
# 注入context
|
||||
self.worker_dict[w].set_context_dict(self.context)
|
||||
else:
|
||||
for w_zip in zip_longest(*pipeline_part, fillvalue="-"):
|
||||
self.logger.info(f"stage{i}: {' | '.join(w_zip)}")
|
||||
i += 1
|
||||
for w in w_zip:
|
||||
if w == "-":
|
||||
continue
|
||||
if w not in self.worker_dict:
|
||||
raise RuntimeError(f"worker={w} is not inited.")
|
||||
|
||||
# 注入context & lock
|
||||
self.worker_dict[w].set_context_dict(self.context, self.context_lock)
|
||||
|
||||
self.logger.info(f"----- {self.chat_name} {self.memory_method_type.value} Pipeline End -----")
|
||||
self.injected = True
|
||||
|
||||
def _worker_run(self, worker_list: list[str]) -> bool:
|
||||
for worker_name in worker_list:
|
||||
worker = self.worker_dict[worker_name]
|
||||
worker.run()
|
||||
if not worker.continue_run:
|
||||
return False
|
||||
return True
|
||||
|
||||
def _run(self):
|
||||
self._visit_and_inject_workers()
|
||||
|
||||
with Timer(f"pipeline_{self.chat_name}_{self.memory_method_type.value}"):
|
||||
self.context[MESSAGES] = self.history_message_list + self.current_message_list
|
||||
self.context[CHAT_NAME] = self.chat_name
|
||||
|
||||
for pipeline_part in self.pipeline_list:
|
||||
if len(pipeline_part) == 1:
|
||||
if not self._worker_run(pipeline_part[0]):
|
||||
break
|
||||
else:
|
||||
t_list = []
|
||||
for worker_list in pipeline_part:
|
||||
t_list.append(GLOBAL_CONTEXT.thread_pool.submit(self._worker_run, worker_list))
|
||||
|
||||
flag = True
|
||||
for future in as_completed(t_list):
|
||||
if not future.result():
|
||||
flag = False
|
||||
break
|
||||
if not flag:
|
||||
break
|
||||
|
||||
def _thread_loop(self):
|
||||
while self.loop_switch:
|
||||
time.sleep(self.loop_interval_time)
|
||||
if len(self.current_message_list) < self.loop_minimum_count:
|
||||
continue
|
||||
self._run()
|
||||
self.context.clear()
|
||||
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
|
||||
-self.history_msg_count:]
|
||||
with self.message_lock:
|
||||
self.current_message_list.clear()
|
||||
|
||||
def start_loop_run(self):
|
||||
if not self.loop_switch:
|
||||
self.loop_switch = True
|
||||
return GLOBAL_CONTEXT.thread_pool.submit(self._thread_loop)
|
||||
|
||||
def run(self, result_key: str = None):
|
||||
self._run()
|
||||
|
||||
# 获取result
|
||||
result = None
|
||||
if result_key:
|
||||
result = self.context.get(result_key)
|
||||
self.context.clear()
|
||||
|
||||
# 清理 msg
|
||||
self.history_message_list = self.history_message_list.extend(self.current_message_list)[
|
||||
-self.history_msg_count:]
|
||||
self.current_message_list.clear()
|
||||
return result
|
||||
|
||||
def submit_message(self, message: Message, with_lock=True):
|
||||
if with_lock:
|
||||
with self.message_lock:
|
||||
self.current_message_list.append(message)
|
||||
else:
|
||||
self.current_message_list.append(message)
|
||||
|
|
@ -10,7 +10,8 @@ class Registry(object):
|
|||
self.name: str = name
|
||||
self.module_dict: Dict[str, Any] = {}
|
||||
|
||||
def register(self, module: Any, module_name: str = None):
|
||||
def register(self, module_name: str = None, module: Any = None):
|
||||
assert module is not None
|
||||
if module_name is None:
|
||||
module_name = module.__name__
|
||||
|
||||
|
|
@ -27,6 +28,6 @@ class Registry(object):
|
|||
raise NotImplementedError
|
||||
self.module_dict.update(module_name_dict)
|
||||
|
||||
def get(self, module_name: str):
|
||||
assert module_name in self.module_dict, f'{module_name} not found in {self.name}'
|
||||
def __getitem__(self, module_name: str):
|
||||
assert module_name in self.module_dict, f"{module_name} not found in {self.name}"
|
||||
return self.module_dict[module_name]
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
import re
|
||||
|
||||
from utils.logger import Logger
|
||||
from memory_scope.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,6 +1,6 @@
|
|||
import time
|
||||
|
||||
from .logger import Logger
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class Timer(object):
|
||||
|
|
|
|||
|
|
@ -1,8 +1,16 @@
|
|||
import hashlib
|
||||
import random
|
||||
import re
|
||||
from importlib import import_module
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from datetime import datetime
|
||||
from importlib import import_module
|
||||
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
import pyfiglet
|
||||
from termcolor import colored, COLORS
|
||||
|
||||
from memory_scope.constants.common_constants import WEEKDAYS
|
||||
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
|
||||
def under_line_to_hump(underline_str):
|
||||
|
|
@ -10,27 +18,27 @@ def under_line_to_hump(underline_str):
|
|||
return sub[0:1].upper() + sub[1:]
|
||||
|
||||
|
||||
def init_instance_by_config(
|
||||
config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs
|
||||
):
|
||||
clazz_path = config.pop("clazz")
|
||||
if not clazz_path:
|
||||
raise RuntimeError("empty clazz_path!")
|
||||
clazz_name_split = clazz_path.split(".")
|
||||
clazz_name: str = clazz_name_split[-1]
|
||||
if suffix_name and not clazz_name.endswith(suffix_name):
|
||||
clazz_name = f"{clazz_name}_{suffix_name}"
|
||||
def init_instance_by_config(config: dict, default_class_path: str = "memory_scope", suffix_name: str = "", **kwargs):
|
||||
config_copy = deepcopy(config)
|
||||
origin_class_path: str = config_copy.pop("class")
|
||||
if not origin_class_path:
|
||||
raise RuntimeError("empty class path!")
|
||||
|
||||
# 构造path
|
||||
clazz_paths = []
|
||||
if default_clazz_path:
|
||||
clazz_paths.append(default_clazz_path)
|
||||
clazz_paths.extend(clazz_name_split[:-1])
|
||||
clazz_paths.append(clazz_name)
|
||||
module = import_module(".".join(clazz_paths))
|
||||
class_name_split = origin_class_path.split(".")
|
||||
class_name: str = class_name_split[-1]
|
||||
if suffix_name and not class_name.lower().endswith(suffix_name.lower()):
|
||||
class_name = f"{class_name}_{suffix_name}"
|
||||
class_name_split[-1] = class_name
|
||||
|
||||
cls_name = under_line_to_hump(clazz_name)
|
||||
return getattr(module, cls_name)(**config, **kwargs)
|
||||
class_paths = []
|
||||
if default_class_path and not origin_class_path.startswith(default_class_path):
|
||||
class_paths.append(default_class_path)
|
||||
class_paths.extend(class_name_split)
|
||||
module = import_module(".".join(class_paths))
|
||||
|
||||
cls_name = under_line_to_hump(class_name)
|
||||
config_copy.update(kwargs)
|
||||
return getattr(module, cls_name)(**config_copy)
|
||||
|
||||
|
||||
def complete_config_name(config_name: str, suffix: str = ".json"):
|
||||
|
|
@ -67,17 +75,16 @@ def get_datetime_info_dict(parse_dt: datetime):
|
|||
}
|
||||
|
||||
|
||||
def time_to_formatted_str(
|
||||
time: datetime | str | int | float = None,
|
||||
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
|
||||
string_format: str = "",
|
||||
) -> str:
|
||||
if isinstance(time, str | int | float):
|
||||
if isinstance(time, str):
|
||||
time = float(time)
|
||||
current_dt = datetime.fromtimestamp(time)
|
||||
elif isinstance(time, datetime):
|
||||
current_dt = time
|
||||
def time_to_formatted_str(dt: datetime | str | int | float = None,
|
||||
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
|
||||
string_format: str = "") -> str:
|
||||
|
||||
if isinstance(dt, str | int | float):
|
||||
if isinstance(dt, str):
|
||||
dt = float(dt)
|
||||
current_dt = datetime.fromtimestamp(dt)
|
||||
elif isinstance(dt, datetime):
|
||||
current_dt = dt
|
||||
else:
|
||||
current_dt = datetime.now()
|
||||
|
||||
|
|
@ -88,3 +95,28 @@ def time_to_formatted_str(
|
|||
return_str = string_format.format(**get_datetime_info_dict(current_dt))
|
||||
|
||||
return return_str
|
||||
|
||||
|
||||
def char_logo(words: str, seed: int = time.time_ns(), color=None):
|
||||
font = pyfiglet.Figlet()
|
||||
rendered_text = font.renderText(words)
|
||||
colored_lines = []
|
||||
all_colors = list(COLORS.keys())
|
||||
random.seed = seed
|
||||
for line in rendered_text.splitlines():
|
||||
line_color = color
|
||||
if line_color is None:
|
||||
random.shuffle(all_colors)
|
||||
line_color = all_colors[0]
|
||||
colored_line = ""
|
||||
for char in line:
|
||||
colored_char = colored(char, line_color, attrs=['bold'])
|
||||
colored_line += colored_char
|
||||
colored_lines.append(colored_line)
|
||||
return colored_lines
|
||||
|
||||
|
||||
def md5_hash(input_string: str):
|
||||
m = hashlib.md5()
|
||||
m.update(input_string.encode('utf-8'))
|
||||
return m.hexdigest()
|
||||
|
|
|
|||
|
|
@ -1,70 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from chat.global_context import GLOBAL_CONTEXT
|
||||
from constants.common_constants import MESSAGES, CHAT_NAME
|
||||
from models.base_model import BaseModel
|
||||
from scheme.message import Message
|
||||
from storage.base_monitor import BaseMonitor
|
||||
from storage.base_vector_store import BaseVectorStore
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
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
|
||||
self.rank_model_name: str = rank_model
|
||||
|
||||
self._embedding_model: BaseModel | None = None
|
||||
self._generation_model: BaseModel | None = None
|
||||
self._rank_model: BaseModel | None = None
|
||||
|
||||
self._vector_store: BaseVectorStore | None = None
|
||||
self._monitor: BaseMonitor | None = None
|
||||
|
||||
@property
|
||||
def messages(self) -> List[Message]:
|
||||
return self.get_context(MESSAGES)
|
||||
|
||||
@messages.setter
|
||||
def messages(self, value):
|
||||
self.set_context(MESSAGES, value)
|
||||
|
||||
@property
|
||||
def chat_name(self):
|
||||
return self.get_context(CHAT_NAME)
|
||||
|
||||
@property
|
||||
def embedding_model(self):
|
||||
if self._embedding_model is None:
|
||||
self._embedding_model = GLOBAL_CONTEXT.model_dict.get(self.embedding_model_name)
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
def generation_model(self):
|
||||
if self._generation_model is None:
|
||||
self._generation_model = GLOBAL_CONTEXT.model_dict.get(self.generation_model_name)
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
def rank_model(self):
|
||||
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):
|
||||
if self._vector_store is None:
|
||||
self._vector_store = GLOBAL_CONTEXT.vector_store
|
||||
return self._vector_store
|
||||
|
||||
@property
|
||||
def monitor(self):
|
||||
if self._monitor is None:
|
||||
self._monitor = GLOBAL_CONTEXT.monitor
|
||||
return self._monitor
|
||||
|
|
@ -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
|
||||
0
old/worker/__init__.py
Normal file
0
old/worker/__init__.py
Normal file
|
|
@ -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):
|
||||
0
old/worker/es/__init__.py
Normal file
0
old/worker/es/__init__.py
Normal file
|
|
@ -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,
|
||||
],
|
||||
0
old/worker/retrieve/__init__.py
Normal file
0
old/worker/retrieve/__init__.py
Normal file
|
|
@ -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,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
old/worker/summary_long/__init__.py
Normal file
0
old/worker/summary_long/__init__.py
Normal file
166
old/worker/summary_long/get_insight_worker.py
Normal file
166
old/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
old/worker/summary_long/get_reflection_worker.py
Normal file
99
old/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
old/worker/summary_long/long_contra_repeat_worker.py
Normal file
129
old/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
old/worker/summary_long/summary_collect_worker.py
Normal file
47
old/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
old/worker/summary_long/update_insight_worker.py
Normal file
177
old/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
old/worker/summary_long/update_profile_worker.py
Normal file
241
old/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
old/worker/summary_short/__init__.py
Normal file
0
old/worker/summary_short/__init__.py
Normal file
117
old/worker/summary_short/contra_repeat_worker.py
Normal file
117
old/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)
|
||||
167
old/worker/summary_short/get_observation_with_time_worker.py
Normal file
167
old/worker/summary_short/get_observation_with_time_worker.py
Normal file
|
|
@ -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
old/worker/summary_short/get_observation_worker.py
Normal file
144
old/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
old/worker/summary_short/info_filter_worker.py
Normal file
70
old/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")
|
||||
|
|
@ -1,7 +1,12 @@
|
|||
import sys
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import asyncio
|
||||
import unittest
|
||||
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class TestLLIEmbedding(unittest.TestCase):
|
||||
|
|
@ -9,26 +14,31 @@ class TestLLIEmbedding(unittest.TestCase):
|
|||
|
||||
def setUp(self):
|
||||
config = {
|
||||
"method_type": "DashScopeEmbedding",
|
||||
"module_name": "dashscope_embedding",
|
||||
"model_name": "text-embedding-v2",
|
||||
"clazz": "models.base_embedding_model"
|
||||
}
|
||||
self.emb = LlamaIndexEmbeddingModel(**config)
|
||||
print()
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def test_single_embedding(self):
|
||||
text = "您吃了吗?"
|
||||
result = self.emb.call(text=text)
|
||||
print(result)
|
||||
self.logger.info(result.m_type)
|
||||
self.logger.info(len(result.embedding_results))
|
||||
|
||||
def test_batch_embedding(self):
|
||||
texts = ["您吃了吗?",
|
||||
"吃了吗您?"]
|
||||
result = self.emb.call(text=texts)
|
||||
print(result)
|
||||
print()
|
||||
self.logger.info(result)
|
||||
|
||||
def test_async_embedding(self):
|
||||
texts = ["您吃了吗?",
|
||||
"吃了吗您?"]
|
||||
# 调用异步函数并等待其结果
|
||||
result = asyncio.run(self.emb.async_call(text=texts))
|
||||
print(result)
|
||||
print()
|
||||
self.logger.info(result)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
import unittest
|
||||
import sys
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import unittest
|
||||
import time
|
||||
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.models.llama_index_generation_model import LlamaIndexGenerationModel
|
||||
from memory_scope.utils.logger import Logger
|
||||
|
||||
|
||||
class TestLLILLM(unittest.TestCase):
|
||||
|
|
@ -8,52 +15,41 @@ class TestLLILLM(unittest.TestCase):
|
|||
|
||||
def setUp(self):
|
||||
config = {
|
||||
"method_type": "DashScope",
|
||||
"module_name": "dashscope_generation",
|
||||
"model_name": "qwen-max",
|
||||
"clazz": "models.llama_index_generation_model"
|
||||
}
|
||||
self.llm = LlamaIndexGenerationModel(**config)
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def test_llm_prompt(self):
|
||||
prompt = "你是谁?"
|
||||
ans = self.llm.call(
|
||||
stream=False,
|
||||
prompt=prompt
|
||||
)
|
||||
print(ans.text)
|
||||
@unittest.skip("tmp")
|
||||
ans = self.llm.call(stream=False, prompt=prompt)
|
||||
self.logger.info(ans.message.content)
|
||||
|
||||
def test_llm_messages(self):
|
||||
messages = [{"role": "system", "content": "you are a helpful assistant."},
|
||||
{"role": "user", "content": "你是谁?"}]
|
||||
ans = self.llm.call(
|
||||
stream=False,
|
||||
messages=messages
|
||||
)
|
||||
print(ans.text)
|
||||
@unittest.skip("tmp")
|
||||
messages = [Message(role="system", content="you are a helpful assistant."),
|
||||
Message(role="user", content="你是谁?")]
|
||||
ans = self.llm.call(stream=False, messages=messages)
|
||||
self.logger.info(ans.message.content)
|
||||
|
||||
def test_llm_prompt_stream(self):
|
||||
prompt = "你如何看待黄金上涨?"
|
||||
ans = self.llm.call(
|
||||
stream=True,
|
||||
prompt=prompt
|
||||
)
|
||||
import sys
|
||||
import time
|
||||
ans = self.llm.call(stream=True, prompt=prompt)
|
||||
self.logger.info("-----start-----")
|
||||
for a in ans:
|
||||
sys.stdout.write(a.delta)
|
||||
sys.stdout.flush()
|
||||
time.sleep(0.1)
|
||||
@unittest.skip("tmp")
|
||||
def test_llm_messages(self):
|
||||
messages = [{"role": "system", "content": "you are a helpful assistant."},
|
||||
{"role": "user", "content": "你如何看待黄金上涨?"}]
|
||||
ans = self.llm.call(
|
||||
stream=True,
|
||||
messages=messages
|
||||
)
|
||||
import sys
|
||||
import time
|
||||
self.logger.info("-----end-----")
|
||||
|
||||
def test_llm_messages_stream(self):
|
||||
messages = [Message(role="system", content="you are a helpful assistant."),
|
||||
Message(role="user", content="你如何看待黄金上涨?")]
|
||||
ans = self.llm.call(stream=True, messages=messages)
|
||||
self.logger.info("-----start-----")
|
||||
for a in ans:
|
||||
sys.stdout.write(a.delta)
|
||||
sys.stdout.flush()
|
||||
time.sleep(0.1)
|
||||
time.sleep(0.1)
|
||||
self.logger.info("-----end-----")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
import json
|
||||
import unittest
|
||||
|
||||
from memory_scope.models.llama_index_rerank_model import LlamaIndexRerankModel
|
||||
from memory_scope.models.llama_index_rank_model import LlamaIndexRankModel
|
||||
|
||||
|
||||
class TestLLIReRank(unittest.TestCase):
|
||||
|
|
@ -9,11 +8,11 @@ class TestLLIReRank(unittest.TestCase):
|
|||
|
||||
def setUp(self):
|
||||
config = {
|
||||
"method_type": "DashScopeRerank",
|
||||
"module_name": "dashscope_rank",
|
||||
"model_name": "gte-rerank",
|
||||
"clazz": "models.llama_index_rerank_model"
|
||||
}
|
||||
self.reranker = LlamaIndexRerankModel(**config)
|
||||
self.reranker = LlamaIndexRankModel(**config)
|
||||
|
||||
def test_rerank(self):
|
||||
query = "吃啥?"
|
||||
|
|
@ -1,90 +1,147 @@
|
|||
import unittest
|
||||
|
||||
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters, FilterCondition, FilterOperator
|
||||
from llama_index.core.schema import TextNode
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.storage.llama_index_elastic_search_store import LlamaIndexElasticSearchStore
|
||||
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
|
||||
|
||||
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
||||
"""Tests for LLIEmbedding"""
|
||||
|
||||
def setUp(self):
|
||||
config = {
|
||||
"method_type": "DashScopeEmbedding",
|
||||
"module_name": "dashscope_embedding",
|
||||
"model_name": "text-embedding-v2",
|
||||
"clazz": "models.llama_index_embedding_model"
|
||||
}
|
||||
emb = LlamaIndexEmbeddingModel(**config).model
|
||||
emb = LlamaIndexEmbeddingModel(**config)
|
||||
|
||||
config = {
|
||||
"index_name" : "0625_3",
|
||||
"es_url" : "http://localhost:9200",
|
||||
"embedding_model" : emb,
|
||||
|
||||
"index_name": "0626_1",
|
||||
"es_url": "http://localhost:9200",
|
||||
"embedding_model": emb,
|
||||
}
|
||||
self.es_store = LlamaIndexElasticSearchStore(**config)
|
||||
self.data = [
|
||||
MemoryNode(
|
||||
content="The lives of two mob hitmen, a boxer, a gangster and his wife, and a pair of diner bandits intertwine in four tales of violence and redemption.",
|
||||
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "
|
||||
"and a pair of diner bandits intertwine in four tales of violence and redemption.",
|
||||
memory_type="observation",
|
||||
id="0"
|
||||
user_id="0",
|
||||
status="valid",
|
||||
memory_id="aaa123",
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, Batman must accept one of the greatest psychological and physical tests of his ability to fight injustice.",
|
||||
content="When the menace known as the Joker wreaks havoc and chaos on the people of Gotham, "
|
||||
"Batman must accept one of the greatest psychological and physical tests of his "
|
||||
"ability to fight injustice.",
|
||||
memory_type="observation",
|
||||
id="1"
|
||||
|
||||
user_id="1",
|
||||
status="valid",
|
||||
memory_id="bbb456",
|
||||
meta_data={"1": "1"}
|
||||
),
|
||||
MemoryNode(
|
||||
content="An insomniac office worker and a devil-may-care soapmaker form an underground fight club that evolves into something much, much more.",
|
||||
content="An insomniac office worker and a devil-may-care soapmaker form an underground fight "
|
||||
"club that evolves into something much, much more.",
|
||||
memory_type="insights",
|
||||
id="2"
|
||||
|
||||
user_id="2",
|
||||
status="valid",
|
||||
memory_id="ccc789",
|
||||
meta_data={"2": "2"}
|
||||
),
|
||||
MemoryNode(
|
||||
content="A thief who steals corporate secrets through the use of dream-sharing technology is given the inverse task of planting an idea into thed of a C.E.O.",
|
||||
content="A thief who steals corporate secrets through the use of dream-sharing technology "
|
||||
"is given the inverse task of planting an idea into thed of a C.E.O.",
|
||||
memory_type="insights",
|
||||
id="3"
|
||||
user_id="3",
|
||||
status="valid",
|
||||
memory_id="ddd012",
|
||||
meta_data={"3": "3"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="A computer hacker learns from mysterious rebels about the true nature of his reality and his role in the war against its controllers.",
|
||||
content="A computer hacker learns from mysterious rebels about the true nature of his reality "
|
||||
"and his role in the war against its controllers.",
|
||||
memory_type="profile",
|
||||
id="4"
|
||||
user_id="4",
|
||||
status="valid",
|
||||
memory_id="eee345",
|
||||
meta_data={"4": "4"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven deadly sins as his motives.",
|
||||
content="Two detectives, a rookie and a veteran, hunt a serial killer who uses the seven "
|
||||
"deadly sins as his motives.",
|
||||
memory_type="profile",
|
||||
id="5"
|
||||
user_id="5",
|
||||
status="valid",
|
||||
memory_id="fff678",
|
||||
meta_data={"5": "5"},
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="An organized crime dynasty's aging patriarch transfers control of his clandestine empire to his reluctant son.",
|
||||
content="An organized crime dynasty's aging patriarch transfers control of his clandestine "
|
||||
"empire to his reluctant son.",
|
||||
memory_type="insights",
|
||||
id="6"),
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="ggg901",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
),
|
||||
MemoryNode(
|
||||
content="ggggggggg",
|
||||
memory_type="profile",
|
||||
id="6"),
|
||||
|
||||
]
|
||||
# @unittest.skip("tmp")
|
||||
def test_insert(self, ):
|
||||
for node in self.data:
|
||||
self.es_store.insert(node)
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="ggg234",
|
||||
meta_data={"5": "5"}
|
||||
|
||||
|
||||
# @unittest.skip("tmp")
|
||||
def test_retrieve(self, ):
|
||||
|
||||
filter = {
|
||||
"id": ["1", "2", "3"],
|
||||
"memory_type": "insights",
|
||||
),
|
||||
]
|
||||
|
||||
def test_retrieve(self):
|
||||
filter_dict = {
|
||||
"user_id": "6",
|
||||
}
|
||||
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter, top_k=10)
|
||||
for node in self.data:
|
||||
self.es_store.insert(node)
|
||||
self.es_store.insert(MemoryNode(
|
||||
content="xxxxxx",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="valid",
|
||||
memory_id="ggg567",
|
||||
meta_data={"5": "5"}
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
||||
self.es_store.update(MemoryNode(
|
||||
content="test update",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="invalid",
|
||||
memory_id="ggg567"
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
self.es_store.delete(MemoryNode(
|
||||
content="test update",
|
||||
memory_type="profile",
|
||||
user_id="6",
|
||||
status="invalid",
|
||||
memory_id="ggg567"
|
||||
))
|
||||
res = self.es_store.retrieve(query="hacker", filter_dict=filter_dict, top_k=10)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
def tearDown(self):
|
||||
self.es_store.close()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue