mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-09 22:31:05 +00:00
add config manager for global configs
This commit is contained in:
parent
6174f84b23
commit
70ecdb1d94
85 changed files with 538 additions and 625 deletions
|
|
@ -1,3 +1,2 @@
|
|||
""" Version of MemoryScope."""
|
||||
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
__version__ = "0.1.0"
|
||||
|
|
|
|||
|
|
@ -1,182 +0,0 @@
|
|||
DEFAULT_GLOBAL_ARGUMENTS = {
|
||||
"language": "en",
|
||||
"thread_pool_max_workers": 5,
|
||||
"logger_name": "memoryscope",
|
||||
"logger_name_time_suffix": "%Y%m%d_%H%M%S"
|
||||
}
|
||||
|
||||
DEFAULT_MEMORY_CHAT_ARGUMENTS = {
|
||||
"cli_memory_chat": {
|
||||
"class": "chat.cli_memory_chat",
|
||||
"memory_service": "memoryscope_service",
|
||||
"generation_model": "generation_model"
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_MEMORY_SERVICE_ARGUMENTS = {
|
||||
"memoryscope_service": {
|
||||
"class": "memory.service.memory_scope_service",
|
||||
"memory_operations": {
|
||||
"read_message": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "read_message",
|
||||
"description": "read short memory"
|
||||
},
|
||||
"retrieve_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank",
|
||||
"description": "retrieve long-term memory"
|
||||
},
|
||||
"list_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_top_memory,print_memory",
|
||||
"description": "read all long-term memory of the user"
|
||||
},
|
||||
"delete_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_all_memory,delete_memory",
|
||||
"description": "delete a single long-term memory"
|
||||
},
|
||||
"delete_all": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "set_query,retrieve_all_memory,delete_all",
|
||||
"description": "delete all long-term memory"
|
||||
},
|
||||
"add_memory": {
|
||||
"class": "memory.operation.frontend_operation",
|
||||
"workflow": "add_memory",
|
||||
"description": "add a single observation"
|
||||
},
|
||||
"consolidate_memory": {
|
||||
"class": "memory.operation.consolidate_memory_op",
|
||||
"workflow": "info_filter,[get_observation|get_observation_with_time|load_today_memory],contra_repeat,"
|
||||
"store_memory",
|
||||
"description": "summary user's observation memory",
|
||||
"interval_time": 1
|
||||
},
|
||||
"reflect_and_reconsolidate": {
|
||||
"class": "memory.operation.backend_operation",
|
||||
"workflow": "load_obs_and_insight,get_reflection_subject,update_insight,long_contra_repeat,"
|
||||
"store_memory",
|
||||
"description": "summary user's insight memory",
|
||||
"interval_time": 15
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_WORKER_ARGUMENTS = {
|
||||
"dummy": {
|
||||
"class": "memory.worker.dummy_worker",
|
||||
"generation_model": "generation_model",
|
||||
"embedding_model": "embedding_model",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"read_message": {
|
||||
"class": "memory.worker.frontend.read_message_worker"
|
||||
},
|
||||
"set_query": {
|
||||
"class": "memory.worker.frontend.set_query_worker"
|
||||
},
|
||||
"retrieve_obs_ins": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 100,
|
||||
"retrieve_ins_top_k": 100
|
||||
},
|
||||
"extract_time": {
|
||||
"class": "memory.worker.frontend.extract_time_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"semantic_rank": {
|
||||
"class": "memory.worker.frontend.semantic_rank_worker",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"fuse_rerank": {
|
||||
"class": "memory.worker.frontend.fuse_rerank_worker",
|
||||
"fuse_score_threshold": 0.01,
|
||||
"fuse_ratio_dict": {
|
||||
"conversation": 0.5,
|
||||
"observation": 1,
|
||||
"obs_customized": 1.2,
|
||||
"insight": 2
|
||||
},
|
||||
"fuse_time_ratio": 2,
|
||||
"fuse_rerank_top_k": 10
|
||||
},
|
||||
"retrieve_top_memory": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 100,
|
||||
"retrieve_ins_top_k": 100,
|
||||
"retrieve_expired_top_k": 100
|
||||
},
|
||||
"print_memory": {
|
||||
"class": "memory.worker.frontend.print_memory_worker"
|
||||
},
|
||||
"retrieve_all_memory": {
|
||||
"class": "memory.worker.frontend.retrieve_memory_worker",
|
||||
"retrieve_obs_top_k": 1000,
|
||||
"retrieve_ins_top_k": 1000,
|
||||
"retrieve_expired_top_k": 1000
|
||||
},
|
||||
"delete_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "delete_memory"
|
||||
},
|
||||
"delete_all": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "delete_all"
|
||||
},
|
||||
"add_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "from_query"
|
||||
},
|
||||
"info_filter": {
|
||||
"class": "memory.worker.backend.info_filter_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"load_today_memory": {
|
||||
"class": "memory.worker.backend.load_memory_worker",
|
||||
"retrieve_today_top_k": 100
|
||||
},
|
||||
"get_observation": {
|
||||
"class": "memory.worker.backend.get_observation_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"get_observation_with_time": {
|
||||
"class": "memory.worker.backend.get_observation_with_time_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"contra_repeat": {
|
||||
"class": "memory.worker.backend.contra_repeat_worker",
|
||||
"generation_model": "generation_model"
|
||||
},
|
||||
"store_memory": {
|
||||
"class": "memory.worker.backend.update_memory_worker",
|
||||
"method": "from_memory_key",
|
||||
"memory_key": "all"
|
||||
},
|
||||
"load_obs_and_insight": {
|
||||
"class": "memory.worker.backend.load_memory_worker",
|
||||
"retrieve_not_reflected_top_k": 100,
|
||||
"retrieve_not_updated_top_k": 100,
|
||||
"retrieve_insight_top_k": 100
|
||||
},
|
||||
"get_reflection_subject": {
|
||||
"class": "memory.worker.backend.get_reflection_subject_worker",
|
||||
"generation_model": "generation_model",
|
||||
"reflect_obs_cnt_threshold": 10
|
||||
},
|
||||
"update_insight": {
|
||||
"class": "memory.worker.backend.update_insight_worker",
|
||||
"generation_model": "generation_model",
|
||||
"rank_model": "rank_model"
|
||||
},
|
||||
"long_contra_repeat": {
|
||||
"class": "memory.worker.backend.long_contra_repeat_worker",
|
||||
"generation_model": "generation_model"
|
||||
}
|
||||
}
|
||||
|
||||
DEFAULT_MONITOR_ARGUMENTS = {
|
||||
"class": "storage.dummy_monitor"
|
||||
}
|
||||
|
|
@ -1,17 +1,15 @@
|
|||
import sys
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.memoryscope import MemoryScope
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import fire
|
||||
|
||||
from memoryscope.core.memoryscope import MemoryScope
|
||||
|
||||
def cli_job(config_path: str):
|
||||
ms = MemoryScope(config_path=config_path)
|
||||
memory_chat: BaseMemoryChat = ms.default_memory_chat
|
||||
memory_chat.run()
|
||||
|
||||
def cli_job(**kwargs):
|
||||
kwargs["memory_chat_type"] = "cli_chat"
|
||||
MemoryScope(**kwargs).default_memory_chat.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ MEMORYSCOPE_CONTEXT = "memoryscope_context"
|
|||
|
||||
RESULT = "result"
|
||||
|
||||
MEMORIES = "memories"
|
||||
|
||||
CHAT_MESSAGES = "chat_messages"
|
||||
|
||||
MEMORY_MANAGER = "memory_manager"
|
||||
|
|
|
|||
|
|
@ -1,14 +1,15 @@
|
|||
from typing import List
|
||||
from typing import List, Optional
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.constants.common_constants import MEMORIES
|
||||
from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME
|
||||
from memoryscope.core.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.models.base_model import BaseModel
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
class ApiMemoryChat(BaseMemoryChat):
|
||||
|
|
@ -96,50 +97,79 @@ class ApiMemoryChat(BaseMemoryChat):
|
|||
self._generation_model = self.context.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def get_new_message(self, query: str, role_name: str = "") -> Message:
|
||||
if not role_name:
|
||||
role_name = self.human_name
|
||||
return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query)
|
||||
|
||||
def get_system_message_with_memory(self, memories: str) -> Message:
|
||||
# Incorporate memory into the system prompt if available
|
||||
system_prompt = self.prompt_handler.system_prompt
|
||||
if memories:
|
||||
memory_prompt = self.prompt_handler.memory_prompt
|
||||
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,
|
||||
role_name: str = "",
|
||||
remember_response: bool = True):
|
||||
|
||||
role_name: Optional[str] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
memory_prompt: Optional[str] = None,
|
||||
extra_memories: Optional[str] = None,
|
||||
add_not_memorized_messages: bool = True,
|
||||
remember_response: bool = True,
|
||||
**kwargs):
|
||||
"""
|
||||
The core function that carries out conversation with memory accepts user queries through query and returns the
|
||||
conversation results through model_response. The retrieved memories are stored in the memories within meta_data.
|
||||
Args:
|
||||
query (str, optional): User's query, includes the user's question.
|
||||
role_name (str, optional): User's role name.
|
||||
system_prompt (str, optional): System prompt. Defaults to the system_prompt in "memory_chat_prompt.yaml".
|
||||
memory_prompt (str, optional): Memory prompt. Defaults to the memory_prompt in "memory_chat_prompt.yaml".
|
||||
extra_memories (str, optional): Manually added user memory in this function.
|
||||
add_not_memorized_messages (bool, optional): whether add not memorized messages to LLM.
|
||||
remember_response (bool, optional): Flag indicating whether to save the AI's response to memory.
|
||||
Defaults to False.
|
||||
Returns:
|
||||
- ModelResponse: In non-streaming mode, returns a complete AI response.
|
||||
- ModelResponseGen: In streaming mode, returns a generator yielding AI response parts.
|
||||
- Memories: To obtain the memory by invoking the method of model_response.meta_data[MEMORIES]
|
||||
"""
|
||||
chat_messages: List[Message] = []
|
||||
|
||||
new_message: Message = self.get_new_message(query=query, role_name=role_name)
|
||||
# prepare query message
|
||||
if not role_name:
|
||||
role_name = self.human_name
|
||||
query_message = Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query)
|
||||
|
||||
# To retrieve memory, prepare the query timestamp and role name by adding new_message.
|
||||
memories: str = self.memory_service.retrieve_memory(query=new_message.content,
|
||||
role_name=new_message.role_name,
|
||||
timestamp=new_message.time_created)
|
||||
# To retrieve memory, prepare the query timestamp and role name by adding query_message.
|
||||
memories: str = self.memory_service.retrieve_memory(query=query_message.content,
|
||||
role_name=query_message.role_name,
|
||||
timestamp=query_message.time_created)
|
||||
|
||||
# format system_message with memories
|
||||
system_message: Message = self.get_system_message_with_memory(memories=memories)
|
||||
system_prompt_list = []
|
||||
if system_prompt:
|
||||
system_prompt_list.append(system_prompt)
|
||||
else:
|
||||
system_prompt_list.append(self.prompt_handler.system_prompt)
|
||||
|
||||
if memories:
|
||||
# add memory prompt
|
||||
if memory_prompt:
|
||||
system_prompt_list.append(memory_prompt)
|
||||
else:
|
||||
system_prompt_list.append(self.prompt_handler.memory_prompt)
|
||||
system_prompt_list.append(memories)
|
||||
|
||||
if extra_memories:
|
||||
system_prompt_list.extend(extra_memories)
|
||||
|
||||
system_prompt_join = "\n".join([x.strip() for x in system_prompt_list])
|
||||
system_message = Message(role=MessageRoleEnum.SYSTEM, content=system_prompt_join)
|
||||
chat_messages.append(system_message)
|
||||
|
||||
# Include past conversation history in the message list
|
||||
history_messages = self.memory_service.read_message()
|
||||
if history_messages:
|
||||
chat_messages.extend(history_messages)
|
||||
if add_not_memorized_messages:
|
||||
history_messages = self.memory_service.read_message()
|
||||
if history_messages:
|
||||
chat_messages.extend(history_messages)
|
||||
|
||||
# Append the current user's message to the conversation context
|
||||
chat_messages.append(new_message)
|
||||
chat_messages.append(query_message)
|
||||
self.logger.info(f"chat_messages={chat_messages}")
|
||||
|
||||
resp = self.generation_model.call(messages=chat_messages, stream=self.stream, **self.generation_model_kwargs)
|
||||
|
||||
if self.stream:
|
||||
assert isinstance(resp, ModelResponseGen)
|
||||
model_response: ModelResponse | None = None
|
||||
for model_response in resp:
|
||||
yield model_response
|
||||
|
|
@ -147,17 +177,18 @@ class ApiMemoryChat(BaseMemoryChat):
|
|||
if remember_response:
|
||||
if model_response and model_response.message:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages([new_message, model_response.message])
|
||||
model_response.meta_data[MEMORIES] = memories
|
||||
self.memory_service.add_messages([query_message, model_response.message])
|
||||
else:
|
||||
self.logger.info("model_response or model_response.message is empty!")
|
||||
self.logger.warning("model_response or model_response.message is empty!")
|
||||
|
||||
else:
|
||||
assert isinstance(resp, ModelResponse)
|
||||
model_response: ModelResponse = resp
|
||||
if remember_response:
|
||||
if model_response and model_response.message:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages([new_message, model_response.message])
|
||||
model_response.meta_data[MEMORIES] = memories
|
||||
self.memory_service.add_messages([query_message, model_response.message])
|
||||
else:
|
||||
self.logger.info("model_response or model_response.message is empty!")
|
||||
self.logger.warning("model_response or model_response.message is empty!")
|
||||
return model_response
|
||||
|
|
@ -1,9 +1,7 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List
|
||||
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
|
|
@ -17,12 +15,14 @@ class BaseMemoryChat(metaclass=ABCMeta):
|
|||
self.kwargs: dict = kwargs
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@abstractmethod
|
||||
def get_new_message(self, query: str, role_name: str = "") -> Message:
|
||||
raise NotImplementedError
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Abstract property to access the memory service.
|
||||
|
||||
@abstractmethod
|
||||
def get_system_message_with_memory(self, memories: str) -> Message:
|
||||
Raises:
|
||||
NotImplementedError: This method should be implemented in a subclass.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
|
|
@ -40,25 +40,6 @@ class BaseMemoryChat(metaclass=ABCMeta):
|
|||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Abstract property to access the memory service.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: This method should be implemented in a subclass.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
self.memory_service.add_messages(messages)
|
||||
|
||||
def start_backend_service(self):
|
||||
self.memory_service.start_backend_service()
|
||||
|
||||
def do_memory_operation(self, operation_name: str, **kwargs):
|
||||
return self.memory_service.do_operation(name=operation_name, **kwargs)
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Abstract method to run the chat system.
|
||||
|
|
@ -4,16 +4,16 @@ from typing import List
|
|||
|
||||
import questionary
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME
|
||||
from memoryscope.core.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.models.base_model import BaseModel
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.core.utils.tool_functions import char_logo
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.utils.tool_functions import char_logo
|
||||
from memoryscope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
class CliMemoryChat(BaseMemoryChat):
|
||||
|
|
@ -122,7 +122,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
self._generation_model = self.context.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def get_new_message(self, query: str, role_name: str = "") -> Message:
|
||||
def get_user_message(self, query: str, role_name: str = "") -> Message:
|
||||
if not role_name:
|
||||
role_name = self.human_name
|
||||
return Message(role=MessageRoleEnum.USER.value, role_name=role_name, content=query)
|
||||
|
|
@ -142,7 +142,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
|
||||
chat_messages: List[Message] = []
|
||||
|
||||
new_message: Message = self.get_new_message(query=query, role_name=role_name)
|
||||
new_message: Message = self.get_user_message(query=query, role_name=role_name)
|
||||
|
||||
# To retrieve memory, prepare the query timestamp and role name by adding new_message.
|
||||
memories: str = self.memory_service.retrieve_memory(query=new_message.content,
|
||||
|
|
@ -168,18 +168,11 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
**self.generation_model_kwargs)
|
||||
|
||||
if self.stream:
|
||||
assert isinstance(resp, ModelResponseGen)
|
||||
model_response: ModelResponse | None = None
|
||||
for model_response in resp:
|
||||
questionary.print(model_response.delta, end="")
|
||||
questionary.print("")
|
||||
|
||||
if remember_response and model_response and model_response.message:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages([new_message, model_response.message])
|
||||
|
||||
else:
|
||||
assert isinstance(resp, ModelResponse)
|
||||
model_response: ModelResponse = resp
|
||||
questionary.print(model_response.message.content)
|
||||
|
||||
|
|
@ -315,7 +308,7 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response
|
||||
self.start_backend_service()
|
||||
self.memory_service.start_backend_service()
|
||||
self.chat_with_memory(query=query)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
|
|
@ -3,7 +3,7 @@ from typing import Literal, Dict
|
|||
|
||||
|
||||
@dataclass
|
||||
class MemoryscopeArguments(object):
|
||||
class Arguments(object):
|
||||
language: Literal["cn", "en"] = field(default="en", metadata={"help": "support en & cn now"})
|
||||
|
||||
thread_pool_max_workers: int = field(default=5, metadata={"help": "thread pool max workers"})
|
||||
|
|
@ -12,12 +12,8 @@ class MemoryscopeArguments(object):
|
|||
|
||||
logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S")
|
||||
|
||||
memory_chat_class: str = field(default="chat.api_memory_chat", metadata={
|
||||
"help": "The memory chat class for dynamic import: chat.cli_memory_chat, chat.api_memory_chat"})
|
||||
|
||||
human_name: str = field(default="user", metadata={"help": "en: user, cn: 用户"})
|
||||
|
||||
assistant_name: str = field(default="AI")
|
||||
memory_chat_type: str = field(default="cli_chat", metadata={
|
||||
"help": "cli_chat(Command-line interaction), api_chat(API interface interaction), etc."})
|
||||
|
||||
consolidate_memory_interval_time: int = field(default=1, metadata={
|
||||
"help": "If you feel that the token consumption is relatively high, please increase the time interval."})
|
||||
|
|
@ -40,7 +36,7 @@ class MemoryscopeArguments(object):
|
|||
embedding_backend: str = field(default="openai_embedding", metadata={
|
||||
"help": "global embedding backend: openai_embedding, dashscope_embedding, etc."})
|
||||
|
||||
embedding_model: str = field(default="gpt-4o", metadata={
|
||||
embedding_model: str = field(default="text-embedding-ada-002", metadata={
|
||||
"help": "global embedding model: text-embedding-ada-002, text-embedding-v2, etc."})
|
||||
|
||||
embedding_params: dict = field(default_factory=lambda: {})
|
||||
|
|
@ -61,6 +57,3 @@ class MemoryscopeArguments(object):
|
|||
|
||||
retrieve_mode: str = field(default="dense", metadata={
|
||||
"help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"})
|
||||
|
||||
hybrid_alpha: float | None = field(default=1.0, metadata={
|
||||
"help": "fuse alpha params used in hybrid mode(not implemented)"})
|
||||
170
memoryscope/core/config/config_manager.py
Normal file
170
memoryscope/core/config/config_manager.py
Normal file
|
|
@ -0,0 +1,170 @@
|
|||
import json
|
||||
from dataclasses import fields
|
||||
from pathlib import Path
|
||||
from typing import Optional, Literal
|
||||
|
||||
import yaml
|
||||
|
||||
from memoryscope.constants.language_constants import DEFAULT_HUMAN_NAME
|
||||
from memoryscope.core.config.arguments import Arguments
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
class ConfigManager(object):
|
||||
|
||||
def __init__(self,
|
||||
config: dict = None,
|
||||
config_path: Optional[str] = None,
|
||||
arguments: Optional[Arguments] = None,
|
||||
demo_config_name: str = "demo_config.yaml",
|
||||
**kwargs):
|
||||
self.config: dict = {}
|
||||
self.kwargs = kwargs
|
||||
|
||||
if config:
|
||||
self.config = config
|
||||
|
||||
elif config_path:
|
||||
self.read_config(config_path)
|
||||
|
||||
else:
|
||||
self.read_demo_config(demo_config_name)
|
||||
|
||||
if arguments:
|
||||
self.update_config_by_arguments(arguments)
|
||||
|
||||
elif kwargs:
|
||||
key_list = [x.name for x in fields(Arguments)]
|
||||
arguments = Arguments(**{k: v for k, v in kwargs.items() if k in key_list})
|
||||
self.update_config_by_arguments(arguments)
|
||||
|
||||
def read_config(self, config_path: str):
|
||||
if config_path.endswith(".yaml"):
|
||||
with open(config_path) as f:
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
|
||||
elif config_path.endswith(".json"):
|
||||
with open(config_path) as f:
|
||||
self.config = json.load(f)
|
||||
|
||||
def read_demo_config(self, demo_config_name: str):
|
||||
file_path = Path(__file__)
|
||||
demo_config_path = (file_path.parent / demo_config_name).__str__()
|
||||
with open(demo_config_path) as f:
|
||||
self.config = yaml.load(f, yaml.FullLoader)
|
||||
|
||||
@staticmethod
|
||||
def update_global_by_arguments(config: dict, arguments: Arguments):
|
||||
config.update({
|
||||
"language": arguments.language,
|
||||
"thread_pool_max_workers": arguments.thread_pool_max_workers,
|
||||
"logger_name": arguments.logger_name,
|
||||
"logger_name_time_suffix": arguments.logger_name_time_suffix,
|
||||
"use_dummy_ranker": arguments.use_dummy_ranker,
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def update_memory_chat_by_arguments(config: dict, arguments: Arguments):
|
||||
if arguments.memory_chat_type == "cli_chat":
|
||||
memory_chat_class = "chat.cli_memory_chat"
|
||||
elif arguments.memory_chat_type == "api_chat":
|
||||
memory_chat_class = "chat.api_memory_chat"
|
||||
else:
|
||||
raise NotImplementedError(f"known memory_chat_type={arguments.memory_chat_type}")
|
||||
config.update({
|
||||
"class": memory_chat_class,
|
||||
"human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)],
|
||||
"assistant_name": "AI",
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def update_memory_service_by_arguments(config: dict, arguments: Arguments):
|
||||
config.update({
|
||||
"human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)],
|
||||
"assistant_name": "AI",
|
||||
})
|
||||
config["memory_operations"]["consolidate_memory"]["interval_time"] = \
|
||||
arguments.consolidate_memory_interval_time
|
||||
config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \
|
||||
arguments.reflect_and_reconsolidate_interval_time
|
||||
|
||||
@staticmethod
|
||||
def update_worker_by_arguments(config: dict, arguments: Arguments):
|
||||
for worker_name, kv_dict in arguments.worker_params.items():
|
||||
if worker_name not in config:
|
||||
continue
|
||||
config[worker_name].update(kv_dict)
|
||||
|
||||
@staticmethod
|
||||
def update_model_by_arguments(config: dict, arguments: Arguments):
|
||||
config["generation_model"].update({
|
||||
"module_name": arguments.generation_backend,
|
||||
"model_name": arguments.generation_model,
|
||||
**arguments.generation_params,
|
||||
})
|
||||
|
||||
config["embedding_model"].update({
|
||||
"module_name": arguments.embedding_backend,
|
||||
"model_name": arguments.embedding_model,
|
||||
**arguments.embedding_params,
|
||||
})
|
||||
|
||||
config["rank_model"].update({
|
||||
"module_name": arguments.rank_backend,
|
||||
"model_name": arguments.rank_model,
|
||||
**arguments.rank_params,
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def update_memory_store_by_arguments(config: dict, arguments: Arguments):
|
||||
config.update({
|
||||
"index_name": arguments.es_index_name,
|
||||
"es_url": arguments.es_url,
|
||||
"retrieve_mode": arguments.retrieve_mode})
|
||||
|
||||
def update_config_by_arguments(self, arguments: Arguments):
|
||||
# prepare global
|
||||
self.update_global_by_arguments(self.config["global"], arguments)
|
||||
|
||||
# prepare memory chat
|
||||
memory_chat_conf_dict = self.config["memory_chat"]
|
||||
memory_chat_config = list(memory_chat_conf_dict.values())[0]
|
||||
self.update_memory_chat_by_arguments(memory_chat_config, arguments)
|
||||
|
||||
# prepare memory service
|
||||
memory_service_conf_dict = self.config["memory_service"]
|
||||
memory_service_config = list(memory_service_conf_dict.values())[0]
|
||||
self.update_memory_service_by_arguments(memory_service_config, arguments)
|
||||
|
||||
# prepare worker
|
||||
self.update_worker_by_arguments(self.config["worker"], arguments)
|
||||
|
||||
# prepare model
|
||||
self.update_model_by_arguments(self.config["model"], arguments)
|
||||
|
||||
# prepare memory store
|
||||
self.update_memory_store_by_arguments(self.config["memory_store"], arguments)
|
||||
|
||||
def add_node_object(self, node: str, name: str, config: dict):
|
||||
self.config[node][name] = config
|
||||
|
||||
def pop_node_object(self, node: str, name: str):
|
||||
return self.config[node].pop(name, None)
|
||||
|
||||
def clear_node_all(self, node: str):
|
||||
self.config[node].clear()
|
||||
|
||||
def dump_config(self, file_type: Literal["json", "yaml"], to_stream: bool = True, file_path: Optional[str] = None):
|
||||
if file_type == "json":
|
||||
content = json.dumps(self.config, indent=2, ensure_ascii=False)
|
||||
elif file_type == "yaml":
|
||||
content = yaml.dump(self.config, indent=2, allow_unicode=True)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if to_stream:
|
||||
print(content)
|
||||
|
||||
if file_type:
|
||||
with open(file_path, "w") as f:
|
||||
f.write(content)
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
global_config:
|
||||
global:
|
||||
language: en
|
||||
thread_pool_max_workers: 5
|
||||
logger_name: memoryscope
|
||||
logger_name_time_suffix: %Y%m%d_%H%M%S
|
||||
use_dummy_ranker: true
|
||||
logger_name_time_suffix: "%Y%m%d_%H%M%S"
|
||||
use_dummy_ranker: false
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
|
|
@ -169,8 +169,7 @@ memory_store:
|
|||
embedding_model: embedding_model
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
retrieve_type: dense
|
||||
hybrid_alpha: 1.0
|
||||
retrieve_mode: dense
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
105
memoryscope/core/memoryscope.py
Normal file
105
memoryscope/core/memoryscope.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
import datetime
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from memoryscope.core.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.core.config.config_manager import ConfigManager
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.core.utils.tool_functions import init_instance_by_config
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
|
||||
|
||||
class MemoryScope(ConfigManager):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self.logger = self._init_logger()
|
||||
|
||||
self.context: MemoryscopeContext = MemoryscopeContext()
|
||||
self.init_context_by_config()
|
||||
|
||||
def _init_logger(self) -> Logger:
|
||||
global_config = self.config["global"]
|
||||
logger_name = global_config["logger_name"]
|
||||
logger_name_time_suffix = global_config["logger_name_time_suffix"]
|
||||
if logger_name_time_suffix:
|
||||
suffix = datetime.datetime.now().strftime(logger_name_time_suffix)
|
||||
logger_name = f"{logger_name}_{suffix}"
|
||||
return Logger.get_logger(logger_name, to_stream=False)
|
||||
|
||||
def init_context_by_config(self):
|
||||
# set global config
|
||||
global_conf = self.config["global"]
|
||||
self.context.language = LanguageEnum(global_conf["language"])
|
||||
self.context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"])
|
||||
self.context.meta_data["use_dummy_ranker"] = global_conf["use_dummy_ranker"]
|
||||
|
||||
# init memory_chat
|
||||
memory_chat_conf_dict = self.config["memory_chat"]
|
||||
if memory_chat_conf_dict:
|
||||
for name, conf in memory_chat_conf_dict.items():
|
||||
self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# set memory_service
|
||||
memory_service_conf_dict = self.config["memory_service"]
|
||||
assert memory_service_conf_dict
|
||||
for name, conf in memory_service_conf_dict.items():
|
||||
self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# init model
|
||||
model_conf_dict = self.config["model"]
|
||||
assert model_conf_dict
|
||||
for name, conf in model_conf_dict.items():
|
||||
self.context.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init memory_store
|
||||
memory_store_conf = self.config["memory_store"]
|
||||
assert memory_store_conf
|
||||
emb_model_name: str = memory_store_conf[ModelEnum.EMBEDDING_MODEL.value]
|
||||
embedding_model = self.context.model_dict[emb_model_name]
|
||||
self.context.memory_store = init_instance_by_config(memory_store_conf, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
monitor_conf = self.config["monitor"]
|
||||
if monitor_conf:
|
||||
self.context.monitor = init_instance_by_config(monitor_conf)
|
||||
|
||||
# set worker config
|
||||
self.context.worker_conf_dict = self.config["worker"]
|
||||
|
||||
def close(self):
|
||||
# wait service to stop
|
||||
for _, service in self.context.memory_service_dict.items():
|
||||
service.stop_backend_service(wait_service_end=True)
|
||||
|
||||
self.context.thread_pool.shutdown()
|
||||
|
||||
self.context.memory_store.close()
|
||||
|
||||
if self.context.monitor:
|
||||
self.context.monitor.close()
|
||||
|
||||
def __enter__(self):
|
||||
self.init_context_by_config()
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.close()
|
||||
|
||||
@property
|
||||
def memory_chat_dict(self):
|
||||
return self.context.memory_chat_dict
|
||||
|
||||
@property
|
||||
def memory_service_dict(self):
|
||||
return self.context.memory_service_dict
|
||||
|
||||
@property
|
||||
def default_memory_chat(self) -> BaseMemoryChat:
|
||||
return list(self.memory_chat_dict.values())[0]
|
||||
|
||||
@property
|
||||
def default_service(self) -> BaseMemoryService:
|
||||
return list(self.memory_service_dict.values())[0]
|
||||
|
|
@ -3,11 +3,11 @@ import time
|
|||
from abc import abstractmethod, ABCMeta
|
||||
from typing import Any
|
||||
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.core.utils.registry import Registry
|
||||
from memoryscope.core.utils.timer import Timer
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.registry import Registry
|
||||
from memoryscope.utils.timer import Timer
|
||||
|
||||
MODEL_REGISTRY = Registry("models")
|
||||
|
||||
|
|
@ -3,9 +3,9 @@ from typing import List
|
|||
|
||||
from llama_index.core.base.llms.types import ChatMessage
|
||||
|
||||
from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
|
@ -2,8 +2,8 @@ from typing import List
|
|||
|
||||
from llama_index.embeddings.dashscope import DashScopeEmbedding
|
||||
|
||||
from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
|
|
@ -3,9 +3,9 @@ from typing import List
|
|||
from llama_index.core.base.llms.types import ChatMessage, ChatResponse, CompletionResponse
|
||||
from llama_index.llms.dashscope import DashScope
|
||||
|
||||
from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
|
@ -4,8 +4,8 @@ from llama_index.core.data_structs import Node
|
|||
from llama_index.core.schema import NodeWithScore
|
||||
from llama_index.postprocessor.dashscope_rerank import DashScopeRerank
|
||||
|
||||
from memoryscope.core.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.models.base_model import BaseModel, MODEL_REGISTRY
|
||||
from memoryscope.scheme.model_response import ModelResponse
|
||||
|
||||
|
||||
|
|
@ -2,10 +2,10 @@ import time
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.common_constants import CHAT_KWARGS, RESULT, CHAT_MESSAGES
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memoryscope.memory.operation.base_workflow import BaseWorkflow
|
||||
from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memoryscope.core.operation.base_workflow import BaseWorkflow
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BackendOperation(BaseWorkflow, BaseOperation):
|
||||
|
|
@ -5,11 +5,11 @@ from itertools import zip_longest
|
|||
from typing import Dict, Any, List
|
||||
|
||||
from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT
|
||||
from memoryscope.memory.worker.base_worker import BaseWorker
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.timer import Timer
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.core.utils.timer import Timer
|
||||
from memoryscope.core.utils.tool_functions import init_instance_by_config
|
||||
from memoryscope.core.worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class BaseWorkflow(object):
|
||||
|
|
@ -1,12 +1,12 @@
|
|||
from memoryscope.constants.common_constants import CHAT_KWARGS, CHAT_MESSAGES, RESULT
|
||||
from memoryscope.core.operation.backend_operation import BackendOperation
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.operation.backend_operation import BackendOperation
|
||||
|
||||
|
||||
class ConsolidateOperation(BackendOperation):
|
||||
class ConsolidateMemoryOp(BackendOperation):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super(ConsolidateOperation, self).__init__(**kwargs)
|
||||
super(ConsolidateMemoryOp, self).__init__(**kwargs)
|
||||
|
||||
self.message_lock = kwargs.get("message_lock", None)
|
||||
self.contextual_msg_min_count: int = kwargs.get("contextual_msg_min_count", 0)
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.common_constants import RESULT, CHAT_MESSAGES, CHAT_KWARGS
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memoryscope.memory.operation.base_workflow import BaseWorkflow
|
||||
from memoryscope.core.operation.base_operation import BaseOperation, OPERATION_TYPE
|
||||
from memoryscope.core.operation.base_workflow import BaseWorkflow
|
||||
from memoryscope.scheme.message import Message
|
||||
|
||||
|
||||
|
|
@ -1,10 +1,10 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List, Dict
|
||||
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.operation.base_operation import BaseOperation
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
|
|
@ -1,10 +1,10 @@
|
|||
import threading
|
||||
from typing import List
|
||||
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.operation.base_operation import BaseOperation
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.tool_functions import init_instance_by_config
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class MemoryScopeService(BaseMemoryService):
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.core.models.base_model import BaseModel
|
||||
from memoryscope.core.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
|
||||
|
||||
class DummyMemoryStore(BaseMemoryStore):
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
from memoryscope.core.storage.base_monitor import BaseMonitor
|
||||
|
||||
|
||||
class DummyMonitor(BaseMonitor):
|
||||
|
|
@ -4,12 +4,14 @@ from typing import Dict, List
|
|||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle
|
||||
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.core.models.base_model import BaseModel
|
||||
from memoryscope.core.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.core.storage.llama_index_sync_elasticsearch import (SyncElasticsearchStore,
|
||||
ESCombinedRetrieveStrategy,
|
||||
_to_elasticsearch_filter,
|
||||
SPECIAL_QUERY)
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore, ESCombinedRetrieveStrategy, \
|
||||
_to_elasticsearch_filter
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
||||
|
|
@ -38,7 +40,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.logger = Logger.get_logger()
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: str = "**--**",
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
filter_dict: Dict[str, List[str]] | Dict[str, str] = None) -> List[MemoryNode]:
|
||||
# if index is not created, return []
|
||||
|
|
@ -53,8 +55,12 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
retriever = self.index.as_retriever(vector_store_kwargs={"es_filter": es_filter, "fields": ['embedding']},
|
||||
similarity_top_k=top_k,
|
||||
sparse_top_k=top_k)
|
||||
|
||||
if not query:
|
||||
query = SPECIAL_QUERY
|
||||
|
||||
if not query and self.emb_dims:
|
||||
query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector())
|
||||
query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector())
|
||||
|
||||
text_nodes = retriever.retrieve(query)
|
||||
if text_nodes and text_nodes[0].embedding:
|
||||
|
|
@ -80,7 +86,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
sparse_top_k=top_k)
|
||||
|
||||
if not query:
|
||||
query = QueryBundle(query_str='**--**', embedding=self.dummy_query_vector())
|
||||
query = SPECIAL_QUERY
|
||||
|
||||
if not query:
|
||||
query = QueryBundle(query_str=SPECIAL_QUERY, embedding=self.dummy_query_vector())
|
||||
|
||||
text_nodes: List[NodeWithScore] = await retriever.aretrieve(query)
|
||||
|
||||
|
|
@ -38,6 +38,8 @@ DISTANCE_STRATEGIES = Literal[
|
|||
"EUCLIDEAN_DISTANCE",
|
||||
]
|
||||
|
||||
SPECIAL_QUERY: str = "**--**"
|
||||
|
||||
|
||||
def get_elasticsearch_client(
|
||||
url: Optional[str] = None,
|
||||
|
|
@ -134,15 +136,15 @@ def _mode_must_match_retrieval_strategy(
|
|||
class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
distance: DistanceMetric = DistanceMetric.COSINE,
|
||||
model_id: Optional[str] = None,
|
||||
retrieve_mode: str = "dense",
|
||||
rrf: Union[bool, Dict[str, Any]] = True,
|
||||
text_field: Optional[str] = "text_field",
|
||||
hybrid_alpha: Optional[float] = None,
|
||||
):
|
||||
self,
|
||||
*,
|
||||
distance: DistanceMetric = DistanceMetric.COSINE,
|
||||
model_id: Optional[str] = None,
|
||||
retrieve_mode: str = "dense",
|
||||
rrf: Union[bool, Dict[str, Any]] = True,
|
||||
text_field: Optional[str] = "text_field",
|
||||
hybrid_alpha: Optional[float] = None,
|
||||
):
|
||||
if retrieve_mode == "dense":
|
||||
self.alpha = 1.0
|
||||
elif retrieve_mode == "sparse":
|
||||
|
|
@ -151,7 +153,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy):
|
|||
elif retrieve_mode == "hybrid":
|
||||
# self.alpha = hybrid_alpha
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
super().__init__(distance=distance, model_id=model_id, hybrid=True, rrf=rrf, text_field=text_field)
|
||||
|
||||
def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]:
|
||||
|
|
@ -159,7 +161,7 @@ class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy):
|
|||
# RRF is used to even the score from the knn query and text query
|
||||
# RRF has two optional parameters: {'rank_constant':int, 'window_size':int}
|
||||
# https://www.elastic.co/guide/en/elasticsearch/reference/current/rrf.html
|
||||
if query == "**--**":
|
||||
if query == SPECIAL_QUERY:
|
||||
query_body = {
|
||||
"query": {
|
||||
"bool": {
|
||||
|
|
@ -699,7 +701,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy)
|
||||
and self.retrieval_strategy.hybrid
|
||||
):
|
||||
# total_rank = sum(top_k_scores)
|
||||
total_rank = sum(top_k_scores)
|
||||
top_k_scores = [rank for rank in top_k_scores]
|
||||
# top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores]
|
||||
# top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores]
|
||||
|
|
@ -3,8 +3,8 @@ import re
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class DatetimeHandler(object):
|
||||
|
|
@ -222,9 +222,9 @@ class DatetimeHandler(object):
|
|||
Returns:
|
||||
dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails.
|
||||
"""
|
||||
func_name = f"extract_date_parts_{language}"
|
||||
func_name = f"extract_date_parts_{language.value}"
|
||||
if not hasattr(cls, func_name):
|
||||
cls.logger.warning(f"language={language} needs to complete extract_date_parts func!")
|
||||
cls.logger.warning(f"language={language.value} needs to complete extract_date_parts func!")
|
||||
return {}
|
||||
return getattr(cls, func_name)(input_string=input_string)
|
||||
|
||||
|
|
@ -272,13 +272,13 @@ class DatetimeHandler(object):
|
|||
|
||||
@classmethod
|
||||
def has_time_word(cls, query: str, language: LanguageEnum) -> bool:
|
||||
func_name = f"has_time_word_{language}"
|
||||
func_name = f"has_time_word_{language.value}"
|
||||
if not hasattr(cls, func_name):
|
||||
cls.logger.warning(f"language={language} needs to complete has_time_word function!")
|
||||
cls.logger.warning(f"language={language.value} needs to complete has_time_word function!")
|
||||
return False
|
||||
|
||||
if language not in DATATIME_WORD_LIST:
|
||||
cls.logger.warning(f"language={language} is missing in DATATIME_WORD_LIST!")
|
||||
cls.logger.warning(f"language={language.value} is missing in DATATIME_WORD_LIST!")
|
||||
return False
|
||||
|
||||
datetime_word_list = DATATIME_WORD_LIST[language]
|
||||
|
|
@ -2,8 +2,8 @@ import re
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import NONE_WORD
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
import time
|
||||
from typing import Literal
|
||||
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
|
||||
TIME_LOG_TYPE = Literal["end", "wrap", "none"]
|
||||
|
||||
|
|
@ -75,7 +75,7 @@ class Timer(object):
|
|||
self.logger.info(f"----- {self.name}.begin -----")
|
||||
return self
|
||||
|
||||
def __exit__(self, *args, **kwargs):
|
||||
def __exit__(self, exc_type, exc_value, exc_tb):
|
||||
"""
|
||||
End timing and print the formatted log.
|
||||
"""
|
||||
|
|
@ -2,11 +2,11 @@ from typing import List
|
|||
|
||||
from memoryscope.constants.common_constants import NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, MERGE_OBS_NODES, TODAY_NODES
|
||||
from memoryscope.constants.language_constants import NONE_WORD, CONTRADICTORY_WORD, CONTAINED_WORD
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
|
@ -2,10 +2,10 @@ from typing import List
|
|||
|
||||
from memoryscope.constants.common_constants import NEW_OBS_WITH_TIME_NODES
|
||||
from memoryscope.constants.language_constants import COLON_WORD
|
||||
from memoryscope.memory.worker.backend.get_observation_worker import GetObservationWorker
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.backend.get_observation_worker import GetObservationWorker
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(GetObservationWorker):
|
||||
|
|
@ -2,14 +2,14 @@ from typing import List
|
|||
|
||||
from memoryscope.constants.common_constants import NEW_OBS_NODES, TIME_INFER
|
||||
from memoryscope.constants.language_constants import REPEATED_WORD, NONE_WORD, COLON_WORD, TIME_INFER_WORD
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
|
|
@ -2,13 +2,13 @@ from typing import List
|
|||
|
||||
from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES
|
||||
from memoryscope.constants.language_constants import COMMA_WORD
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class GetReflectionSubjectWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,11 +1,11 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import COLON_WORD
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,12 +1,12 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.common_constants import NOT_REFLECTED_NODES, NOT_UPDATED_NODES, INSIGHT_NODES, TODAY_NODES
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.timer import timer
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.timer import timer
|
||||
|
||||
|
||||
class LoadMemoryWorker(MemoryBaseWorker):
|
||||
|
|
@ -2,13 +2,13 @@ from typing import List, Dict
|
|||
|
||||
from memoryscope.constants.common_constants import NOT_UPDATED_NODES, MERGE_OBS_NODES
|
||||
from memoryscope.constants.language_constants import NONE_WORD, CONTAINED_WORD, CONTRADICTORY_WORD
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
|
|
@ -3,12 +3,12 @@ from typing import List
|
|||
|
||||
from memoryscope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES
|
||||
from memoryscope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg, cosine_similarity
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg, cosine_similarity
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,10 +1,10 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
|
||||
|
||||
class UpdateMemoryWorker(MemoryBaseWorker):
|
||||
|
|
@ -3,8 +3,8 @@ from abc import ABCMeta, abstractmethod
|
|||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Any, Dict
|
||||
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.timer import Timer
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.core.utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(metaclass=ABCMeta):
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
import datetime
|
||||
|
||||
from memoryscope.constants.common_constants import RESULT, WORKFLOW_NAME, CHAT_KWARGS
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
|
|
@ -3,9 +3,9 @@ from typing import Dict
|
|||
|
||||
from memoryscope.constants.common_constants import QUERY_WITH_TS, EXTRACT_TIME_DICT
|
||||
from memoryscope.constants.language_constants import DATATIME_KEY_MAP
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memoryscope.constants.common_constants import EXTRACT_TIME_DICT, RANKED_MEMORY_NODES, RESULT
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
|
||||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,11 +1,11 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, RESULT
|
||||
from memoryscope.core.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
|
||||
|
||||
class PrintMemoryWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
from memoryscope.constants.common_constants import RESULT
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ReadMessageWorker(MemoryBaseWorker):
|
||||
|
|
@ -1,12 +1,12 @@
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.common_constants import QUERY_WITH_TS, RETRIEVE_MEMORY_NODES
|
||||
from memoryscope.core.utils.timer import timer
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.timer import timer
|
||||
|
||||
|
||||
class RetrieveMemoryWorker(MemoryBaseWorker):
|
||||
|
|
@ -120,6 +120,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
7. Stores the processed memory nodes for further use.
|
||||
"""
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
self.logger.info(f"retrieve memory with query={query}.")
|
||||
self.submit_thread_task(self.retrieve_from_observation, query=query)
|
||||
self.submit_thread_task(self.retrieve_from_insight, query=query)
|
||||
self.submit_thread_task(self.retrieve_expired_memory, query=query)
|
||||
|
|
@ -136,7 +137,7 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
|
|||
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_recall, reverse=True)
|
||||
for node in memory_node_list:
|
||||
node.action_status = ActionStatusEnum.NONE.value
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_rerank} type={node.memory_type} "
|
||||
self.logger.info(f"recall_stage: content={node.content} score={node.score_recall} type={node.memory_type} "
|
||||
f"store_status={node.store_status} action_status={node.action_status}")
|
||||
|
||||
self.memory_manager.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from memoryscope.constants.common_constants import RETRIEVE_MEMORY_NODES, QUERY_WITH_TS, RANKED_MEMORY_NODES
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
|
|
@ -39,22 +39,23 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
for node in memory_node_list:
|
||||
node.score_rank = node.score_recall
|
||||
self.logger.warning("use score_recall instead of score_rank!")
|
||||
return
|
||||
|
||||
# drop repeated
|
||||
memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if n.content.strip()}
|
||||
memory_node_list = list(memory_node_dict.values())
|
||||
else:
|
||||
# drop repeated
|
||||
memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if
|
||||
n.content.strip()}
|
||||
memory_node_list = list(memory_node_dict.values())
|
||||
|
||||
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
|
||||
if not response.status or not response.rank_scores:
|
||||
return
|
||||
response = self.rank_model.call(query=query, documents=[n.content for n in memory_node_list])
|
||||
if not response.status or not response.rank_scores:
|
||||
return
|
||||
|
||||
# set score
|
||||
for idx, score in response.rank_scores.items():
|
||||
if idx >= len(memory_node_list):
|
||||
self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!")
|
||||
continue
|
||||
memory_node_list[idx].score_rank = score
|
||||
# set score
|
||||
for idx, score in response.rank_scores.items():
|
||||
if idx >= len(memory_node_list):
|
||||
self.logger.warning(f"Idx={idx} exceeds the maximum length of rank_scores!")
|
||||
continue
|
||||
memory_node_list[idx].score_rank = score
|
||||
|
||||
# sort by score
|
||||
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
import datetime
|
||||
|
||||
from memoryscope.constants.common_constants import QUERY_WITH_TS
|
||||
from memoryscope.core.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memoryscope.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SetQueryWorker(MemoryBaseWorker):
|
||||
|
|
@ -3,15 +3,15 @@ from typing import List, Dict, Any
|
|||
|
||||
from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORYSCOPE_CONTEXT, \
|
||||
WORKFLOW_NAME, MEMORY_MANAGER
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.models.base_model import BaseModel
|
||||
from memoryscope.core.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.core.storage.base_monitor import BaseMonitor
|
||||
from memoryscope.core.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.core.worker.base_worker import BaseWorker
|
||||
from memoryscope.core.worker.memory_manager import MemoryManager
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.memory.worker.base_worker import BaseWorker
|
||||
from memoryscope.memory.worker.memory_manager import MemoryManager
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
||||
|
|
@ -1,11 +1,11 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from memoryscope.core.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.core.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.enumeration.action_status_enum import ActionStatusEnum
|
||||
from memoryscope.enumeration.store_status_enum import StoreStatusEnum
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class MemoryManager(object):
|
||||
|
|
@ -1,201 +0,0 @@
|
|||
import datetime
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import yaml
|
||||
|
||||
from memoryscope.argument import default_arguments
|
||||
from memoryscope.argument.memoryscope_arguments import MemoryscopeArguments
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.memoryscope_context import MemoryscopeContext
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class MemoryScope(object):
|
||||
|
||||
def __init__(self,
|
||||
arguments: MemoryscopeArguments | None = None,
|
||||
config: dict | None = None,
|
||||
config_path: str = ""):
|
||||
|
||||
self.global_conf: dict = {}
|
||||
self.memory_chat_conf_dict: dict = {}
|
||||
self.memory_service_conf_dict: dict = {}
|
||||
self.worker_conf_dict: dict = {}
|
||||
self.model_conf_dict: dict = {}
|
||||
self.memory_store_conf: dict = {}
|
||||
self.monitor_conf: dict = {}
|
||||
|
||||
self.context: MemoryscopeContext = MemoryscopeContext()
|
||||
|
||||
if arguments:
|
||||
self._init_by_arguments(arguments=arguments)
|
||||
elif config:
|
||||
self._init_by_config(config=config)
|
||||
elif config_path:
|
||||
self._init_by_config_path(config_path=config_path)
|
||||
else:
|
||||
raise RuntimeError("At least one of arguments, config, or file_path must not be empty!")
|
||||
|
||||
self.logger = self._init_logger()
|
||||
|
||||
self._init_context_by_config()
|
||||
|
||||
def _init_by_arguments(self, arguments: MemoryscopeArguments):
|
||||
# prepare global
|
||||
self.global_conf = {
|
||||
"language": arguments.language,
|
||||
"thread_pool_max_workers": arguments.thread_pool_max_workers,
|
||||
"logger_name": arguments.logger_name,
|
||||
"logger_name_time_suffix": arguments.logger_name_time_suffix,
|
||||
"use_dummy_ranker": arguments.use_dummy_ranker,
|
||||
}
|
||||
|
||||
# prepare memory chat
|
||||
self.memory_chat_conf_dict = default_arguments.DEFAULT_MEMORY_CHAT_ARGUMENTS.copy()
|
||||
memory_chat_config = list(self.memory_chat_conf_dict.values())[0]
|
||||
memory_chat_config.update({
|
||||
"class": arguments.memory_chat_class,
|
||||
"human_name": arguments.human_name,
|
||||
"assistant_name": arguments.assistant_name,
|
||||
})
|
||||
|
||||
# prepare memory service
|
||||
self.memory_service_conf_dict = default_arguments.DEFAULT_MEMORY_SERVICE_ARGUMENTS.copy()
|
||||
memory_service_config = list(self.memory_service_conf_dict.values())[0]
|
||||
memory_service_config.update({
|
||||
"human_name": arguments.human_name,
|
||||
"assistant_name": arguments.assistant_name,
|
||||
})
|
||||
memory_service_config["memory_operations"]["consolidate_memory"]["interval_time"] = \
|
||||
arguments.consolidate_memory_interval_time
|
||||
memory_service_config["memory_operations"]["reflect_and_reconsolidate"]["interval_time"] = \
|
||||
arguments.reflect_and_reconsolidate_interval_time
|
||||
|
||||
# prepare memory service
|
||||
self.worker_conf_dict = default_arguments.DEFAULT_WORKER_ARGUMENTS.copy()
|
||||
if arguments.worker_params:
|
||||
for worker_name, kv_dict in arguments.worker_params.items():
|
||||
if worker_name not in self.worker_conf_dict:
|
||||
continue
|
||||
self.worker_conf_dict[worker_name].update(kv_dict)
|
||||
|
||||
# prepare models
|
||||
self.model_conf_dict = {
|
||||
"generation_model": {
|
||||
"class": "models.llama_index_generation_model",
|
||||
"module_name": arguments.generation_backend,
|
||||
"model_name": arguments.generation_model,
|
||||
**arguments.generation_params,
|
||||
},
|
||||
"embedding_model": {
|
||||
"class": "models.llama_index_embedding_model",
|
||||
"module_name": arguments.embedding_backend,
|
||||
"model_name": arguments.embedding_model,
|
||||
**arguments.embedding_params,
|
||||
},
|
||||
"rank_model": {
|
||||
"class": "models.llama_index_rank_model",
|
||||
"module_name": arguments.rank_backend,
|
||||
"model_name": arguments.rank_model,
|
||||
**arguments.rank_params,
|
||||
},
|
||||
}
|
||||
|
||||
# prepare memory store
|
||||
self.memory_store_conf = {
|
||||
"class": "storage.llama_index_es_memory_store",
|
||||
"embedding_model": "embedding_model",
|
||||
"index_name": arguments.es_index_name,
|
||||
"es_url": arguments.es_url,
|
||||
"retrieve_mode": arguments.retrieve_mode,
|
||||
"hybrid_alpha": arguments.hybrid_alpha,
|
||||
}
|
||||
|
||||
self.monitor_conf = default_arguments.DEFAULT_MONITOR_ARGUMENTS.copy()
|
||||
|
||||
def _init_by_config(self, config: dict):
|
||||
self.global_conf = config["global_config"]
|
||||
self.memory_service_conf_dict = config["memory_service"]
|
||||
self.worker_conf_dict = config["worker"]
|
||||
self.model_conf_dict = config["model"]
|
||||
self.memory_store_conf = config["memory_store"]
|
||||
|
||||
# not necessary
|
||||
self.memory_chat_conf_dict = config.get("memory_chat")
|
||||
self.monitor_conf = config.get("monitor")
|
||||
|
||||
def _init_by_config_path(self, config_path: str):
|
||||
with open(config_path) as f:
|
||||
if config_path.endswith("yaml"):
|
||||
config = yaml.load(f, yaml.FullLoader)
|
||||
elif config_path.endswith("json"):
|
||||
config = json.load(f)
|
||||
else:
|
||||
raise RuntimeError("not supported config file type!")
|
||||
return self._init_by_config(config)
|
||||
|
||||
def _init_logger(self) -> Logger:
|
||||
logger_name = self.global_conf.get("logger_name")
|
||||
assert logger_name, "logger_name is empty!"
|
||||
logger_name_time_suffix = self.global_conf.get("logger_name_time_suffix")
|
||||
if logger_name_time_suffix:
|
||||
suffix = datetime.datetime.now().strftime(logger_name_time_suffix)
|
||||
logger_name = f"{logger_name}_{suffix}"
|
||||
return Logger.get_logger(logger_name, to_stream=False)
|
||||
|
||||
def _init_context_by_config(self):
|
||||
# set global config
|
||||
self.context.language = LanguageEnum(self.global_conf["language"])
|
||||
self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"])
|
||||
self.context.meta_data["use_dummy_ranker"] = self.global_conf["use_dummy_ranker"]
|
||||
|
||||
# init memory_chat
|
||||
if self.memory_chat_conf_dict:
|
||||
for name, conf in self.memory_chat_conf_dict.items():
|
||||
self.context.memory_chat_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# set memory_service
|
||||
assert self.memory_service_conf_dict
|
||||
for name, conf in self.memory_service_conf_dict.items():
|
||||
self.context.memory_service_dict[name] = init_instance_by_config(conf, name=name, context=self.context)
|
||||
|
||||
# init models
|
||||
assert self.model_conf_dict
|
||||
for name, conf in self.model_conf_dict.items():
|
||||
self.context.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
assert self.memory_store_conf
|
||||
emb_model_name: str = self.memory_store_conf[ModelEnum.EMBEDDING_MODEL.value]
|
||||
embedding_model = self.context.model_dict[emb_model_name]
|
||||
self.context.memory_store = init_instance_by_config(self.memory_store_conf, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
if self.monitor_conf:
|
||||
self.context.monitor = init_instance_by_config(self.monitor_conf)
|
||||
|
||||
# set worker config
|
||||
self.context.worker_config = self.worker_conf_dict
|
||||
|
||||
def close(self):
|
||||
# wait service to stop
|
||||
for _, service in self.context.memory_service_dict.items():
|
||||
service.stop_backend_service(wait_service_end=True)
|
||||
self.context.memory_store.close()
|
||||
self.context.thread_pool.shutdown()
|
||||
|
||||
if self.context.monitor:
|
||||
self.context.monitor.close()
|
||||
|
||||
@property
|
||||
def default_memory_chat(self) -> BaseMemoryChat:
|
||||
return list(self.context.memory_chat_dict.values())[0]
|
||||
|
||||
@property
|
||||
def default_service(self) -> BaseMemoryService:
|
||||
return list(self.context.memory_service_dict.values())[0]
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
from memoryscope.cli import MemoryScope
|
||||
from memoryscope.scheme.message import Message
|
||||
|
||||
ms = MemoryScope().load_config("config/demo_config_no_stream.yaml")
|
||||
ms = MemoryScope().read_config("config/demo_config_no_stream.yaml")
|
||||
memory_service = ms.default_service
|
||||
memory_chat = ms.default_chat_handle
|
||||
|
||||
|
|
|
|||
14
tests/other/test_cli.py
Normal file
14
tests/other/test_cli.py
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
import fire
|
||||
|
||||
|
||||
class CLI:
|
||||
def run(self, **kwargs):
|
||||
"""
|
||||
打印传入的 kwargs
|
||||
"""
|
||||
for key, value in kwargs.items():
|
||||
print(f"{key}: {value}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
fire.Fire(CLI().run)
|
||||
|
|
@ -164,7 +164,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
meta_data={"5": "5"},
|
||||
timestamp=13
|
||||
))
|
||||
|
||||
|
||||
def test_retrieve(self):
|
||||
filter_dict = {
|
||||
"timestamp": 12,
|
||||
|
|
@ -172,12 +172,11 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
# "score_rank": 0,
|
||||
}
|
||||
|
||||
|
||||
res = self.es_store.retrieve_memories(query="hacker", filter_dict=filter_dict, top_k=15)
|
||||
print(len(res))
|
||||
print(res)
|
||||
|
||||
def test_retrieve_wo_query(self,):
|
||||
def test_retrieve_wo_query(self, ):
|
||||
filter_dict = {
|
||||
"memory_id": "bbb456",
|
||||
}
|
||||
|
|
@ -185,6 +184,5 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
|||
print(len(res))
|
||||
print(res)
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
self.es_store.close()
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class TestWorkersCn(unittest.TestCase):
|
|||
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
|
||||
|
||||
ms = MemoryScope()
|
||||
ms.load_config("config/demo_config_cn.yaml")
|
||||
ms.read_config("config/demo_config_cn.yaml")
|
||||
ms.init_global_content_by_config()
|
||||
|
||||
def tearDown(self):
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ class TestWorkersEn(unittest.TestCase):
|
|||
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
|
||||
|
||||
ms = MemoryScope()
|
||||
ms.load_config("config/demo_config_en.yaml")
|
||||
ms.read_config("config/demo_config_en.yaml")
|
||||
ms.init_global_content_by_config()
|
||||
|
||||
def tearDown(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue