mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
[features] add api features for memory scope
This commit is contained in:
commit
3889b28133
103 changed files with 1947 additions and 1676 deletions
|
|
@ -1,192 +0,0 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
|
||||
logger_config:
|
||||
logger_name: memoryscope
|
||||
logger_suffix: time
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
memory_service: memory_scope_service
|
||||
generation_model: dashscope_generation
|
||||
|
||||
memory_service:
|
||||
memory_scope_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"
|
||||
|
||||
summary_observation_memory:
|
||||
class: memory.operation.summary_observation_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
|
||||
|
||||
summary_insight_memory:
|
||||
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
|
||||
|
||||
worker:
|
||||
dummy:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
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: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
semantic_rank:
|
||||
class: memory.worker.frontend.semantic_rank_worker
|
||||
rank_model: dashscope_rank
|
||||
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.0
|
||||
fuse_time_ratio: 2.0
|
||||
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: 100
|
||||
retrieve_ins_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
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: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
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: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
get_observation_with_time:
|
||||
class: memory.worker.backend.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
contra_repeat:
|
||||
class: memory.worker.backend.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
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: dashscope_generation
|
||||
reflect_obs_cnt_threshold: 10
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
update_insight:
|
||||
class: memory.worker.backend.update_insight_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
rank_model: dashscope_rank
|
||||
long_contra_repeat:
|
||||
class: memory.worker.backend.long_contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_kwargs:
|
||||
top_k: 1
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
max_tokens: 2000
|
||||
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
|
||||
top_n: 10
|
||||
dummy_generation:
|
||||
class: models.dummy_generation_model
|
||||
module_name: dummy_generation
|
||||
model_name: dummy_generation_model
|
||||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
use_hybrid: true
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
|
|
@ -1,148 +0,0 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
stream: false
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.service.chat_memory_service
|
||||
history_msg_count: 32
|
||||
contextual_msg_count: 6
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: memory.operation.read_message
|
||||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: set_query,retrieve_memory2,print_memory
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 5
|
||||
# summary_memory:
|
||||
# class: memory.operation.summary_memory
|
||||
# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
# description: "summary observation memories of the user"
|
||||
# interval_time: 60
|
||||
|
||||
worker:
|
||||
dummy:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
set_query:
|
||||
class: memory.worker.read.set_query_worker
|
||||
retrieve_memory1:
|
||||
class: memory.worker.read.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_expired_top_k: 0
|
||||
extract_time:
|
||||
class: memory.worker.read.extract_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
semantic_rank:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
fuse_rerank:
|
||||
class: memory.worker.read.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.1
|
||||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
observation: 1
|
||||
obs_customized: 1.2
|
||||
insight: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
retrieve_memory2:
|
||||
class: memory.worker.read.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
print_memory:
|
||||
class: memory.worker.read.print_memory_worker
|
||||
info_filter:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
info_filter_msg_max_size: 200
|
||||
generation_model_top_k: 1
|
||||
load_memory1:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 0
|
||||
retrieve_not_updated_top_k: 0
|
||||
retrieve_insight_top_k: 0
|
||||
today_obs_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.write.get_observation_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
get_observation_with_time:
|
||||
class: memory.worker.write.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
contra_repeat:
|
||||
class: memory.worker.write.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
retrieve_top_k: 30
|
||||
contra_repeat_max_count: 50
|
||||
store_memory:
|
||||
class: memory.worker.write.store_memory_worker
|
||||
store_key: all
|
||||
load_memory2:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 100
|
||||
retrieve_not_updated_top_k: 100
|
||||
retrieve_insight_top_k: 100
|
||||
today_obs_top_k: 0
|
||||
get_reflection_subject:
|
||||
class: memory.worker.summary.get_reflection_subject_worker
|
||||
retrieve_top_k: 100
|
||||
reflect_obs_cnt_threshold: 32
|
||||
generation_model_top_k: 1
|
||||
update_insight:
|
||||
class: memory.worker.summary.update_insight_worker
|
||||
update_insight_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
update_insight_max_thread: 10
|
||||
long_contra_repeat:
|
||||
class: memory.worker.summary.long_contra_repeat_worker
|
||||
long_contra_repeat_top_k: 2
|
||||
long_contra_repeat_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
|
||||
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
|
||||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
use_hybrid: false
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
|
|
@ -1,83 +0,0 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
dash_scope_apikey:
|
||||
open_ai_apikey:
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat # select class
|
||||
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 # select class
|
||||
history_msg_count: 32
|
||||
contextual_msg_count: 6
|
||||
read_memory_key: read_memory
|
||||
memory_operations:
|
||||
read_message: # define operation
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow # select workflow
|
||||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: dummy_workflow
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 60
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: dummy_workflow
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 300
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model # select class
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
dashscope_embedding:
|
||||
class: models.llama_index_embedding_model # select class
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
class: models.llama_index_rank_model # select class
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
|
||||
vector_store:
|
||||
class: storage.dummy_vector_store # select class
|
||||
embedding_model: dashscope_embedding
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor # select class
|
||||
|
||||
worker:
|
||||
dummy_workflow:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
fuse_rerank_worker:
|
||||
class: memory.worker.read.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.1
|
||||
fuse_ratio_dict:
|
||||
observation: 1
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
|
||||
66
examples/api/chat_example.py
Normal file
66
examples/api/chat_example.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
import sys
|
||||
|
||||
sys.path.append(".")
|
||||
|
||||
from memoryscope import MemoryScope, Arguments
|
||||
|
||||
arguments = Arguments(
|
||||
language="cn",
|
||||
logger_to_screen=False,
|
||||
memory_chat_class="api_memory_chat",
|
||||
generation_backend="dashscope_generation",
|
||||
generation_model="qwen-max",
|
||||
embedding_backend="dashscope_embedding",
|
||||
embedding_model="text-embedding-v2",
|
||||
use_dummy_ranker=False,
|
||||
rank_backend="dashscope_rank",
|
||||
rank_model="gte-rerank",
|
||||
)
|
||||
|
||||
|
||||
def chat_example1():
|
||||
with MemoryScope(config_path="memoryscope/core/config/demo_config.yaml") as ms:
|
||||
memory_chat = ms.default_memory_chat
|
||||
response = memory_chat.chat_with_memory(query="我的爱好是弹琴。")
|
||||
print("回答:\n" + response.message.content)
|
||||
print("记忆:\n" + response.meta_data["memories"])
|
||||
|
||||
|
||||
def chat_example2():
|
||||
with MemoryScope(arguments=arguments) as ms:
|
||||
memory_chat = ms.default_memory_chat
|
||||
|
||||
response = memory_chat.chat_with_memory(query="我的爱好是弹琴。")
|
||||
print("回答:\n" + response.message.content)
|
||||
print("记忆:\n" + response.meta_data["memories"])
|
||||
|
||||
|
||||
def chat_example3():
|
||||
with MemoryScope(arguments=arguments) as ms:
|
||||
memory_chat = ms.default_memory_chat
|
||||
|
||||
response = memory_chat.chat_with_memory(
|
||||
query="你知道我消灭了几个星系吗?",
|
||||
extra_memories="用户是一只哥斯拉,这是真是的,不是虚构的。\n前天用户消灭了星系0x4be?\n大前天用户消灭了星系0x5be?")
|
||||
print("回答:\n" + response.message.content)
|
||||
|
||||
|
||||
def chat_example4():
|
||||
with MemoryScope(arguments=arguments) as ms:
|
||||
memory_chat = ms.default_memory_chat
|
||||
|
||||
response = memory_chat.chat_with_memory(query="我的爱好是弹琴。")
|
||||
print("回答1:\n" + response.message.content)
|
||||
memory_chat.memory_service.consolidate_memory()
|
||||
|
||||
response = memory_chat.chat_with_memory(query="你知道我的乐器爱好是什么?",
|
||||
history_message_strategy=None)
|
||||
print("回答2:\n" + response.message.content)
|
||||
print("记忆2:\n" + response.meta_data["memories"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
chat_example1()
|
||||
# chat_example2()
|
||||
# chat_example3()
|
||||
# chat_example4()
|
||||
1
examples/cli/dash_cli_cn1.sh
Normal file
1
examples/cli/dash_cli_cn1.sh
Normal file
|
|
@ -0,0 +1 @@
|
|||
python memoryscope/cli.py -config_path=memoryscope/core/config/demo_config.yaml
|
||||
10
examples/cli/dash_cli_cn2.sh
Normal file
10
examples/cli/dash_cli_cn2.sh
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
python memoryscope/cli.py \
|
||||
-language="cn" \
|
||||
-memory_chat_class="cli_memory_chat" \
|
||||
-generation_backend="dashscope_generation" \
|
||||
-generation_model="qwen-max" \
|
||||
-embedding_backend="dashscope_embedding" \
|
||||
-embedding_model="text-embedding-v2" \
|
||||
-use_dummy_ranker=False \
|
||||
-rank_backend="dashscope_rank" \
|
||||
-rank_model="gte-rerank"
|
||||
|
|
@ -1,3 +1,5 @@
|
|||
""" Version of MemoryScope."""
|
||||
from memoryscope.core.config.arguments import Arguments
|
||||
from memoryscope.core.memoryscope import MemoryScope
|
||||
|
||||
__version__ = "0.1.0-alpha.1"
|
||||
""" Version of MemoryScope."""
|
||||
__version__ = "0.1.0"
|
||||
|
|
|
|||
|
|
@ -1,43 +0,0 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
"""
|
||||
An abstract base class representing a chat system integrated with memory services.
|
||||
It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self, query: str):
|
||||
"""
|
||||
Initiates a chat interaction using the memory service, with the provided query as input.
|
||||
|
||||
Args:
|
||||
query (str): The user's query or message to start the chat.
|
||||
|
||||
Returns:
|
||||
This method should return the chat response generated after processing the query
|
||||
with the associated memory context. The actual return type and content are defined by the implementing
|
||||
subclass.
|
||||
"""
|
||||
|
||||
@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 run(self):
|
||||
"""
|
||||
Abstract method to run the chat system.
|
||||
|
||||
This method should contain the logic to initiate and manage the chat process,
|
||||
utilizing the memory service as needed. It must be implemented by subclasses.
|
||||
"""
|
||||
pass
|
||||
|
|
@ -1,346 +0,0 @@
|
|||
import os
|
||||
import time
|
||||
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.enumeration.message_role_enum import MessageRoleEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.utils.tool_functions import char_logo
|
||||
|
||||
|
||||
class CliMemoryChat(BaseMemoryChat):
|
||||
"""
|
||||
Command-line interface for chatting with an AI that integrates memory functionality.
|
||||
Allows users to interact, manage chat history, adjust streaming settings, and view commands' help.
|
||||
"""
|
||||
USER_COMMANDS = {
|
||||
"exit": "Exit the CLI.",
|
||||
"clear": "Clear the command history.",
|
||||
"help": "Display available CLI commands and their descriptions.",
|
||||
"stream": "Toggle between getting streamed responses from the model."
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
memory_service: str,
|
||||
generation_model: str,
|
||||
stream: bool = True,
|
||||
human_name: str = DEFAULT_HUMAN_NAME[G_CONTEXT.language],
|
||||
assistant_name: str = "AI",
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the CLI chat instance with specified services, models, and personalized settings.
|
||||
|
||||
Args:
|
||||
memory_service (str | BaseMemoryService): The memory service to be used for storing conversation history.
|
||||
generation_model (str | BaseModel): The model responsible for generating AI responses.
|
||||
stream (bool, optional): Flag indicating whether responses should be streamed. Defaults to True.
|
||||
human_name (str, optional): The name assigned to the human user. Defaults to a language-specific user.
|
||||
assistant_name (str, optional): The name of the AI assistant. Defaults to "AI".
|
||||
**kwargs: Additional keyword arguments for flexibility or future extensions.
|
||||
|
||||
Side Effects:
|
||||
- Updates global context with human and AI names.
|
||||
- Initializes logging for the instance.
|
||||
"""
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.generation_model_kwargs: dict = kwargs.get("generation_model_kwargs", {})
|
||||
|
||||
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._prompt_handler: PromptHandler | None = None
|
||||
G_CONTEXT.meta_data.update({
|
||||
"human_name": human_name,
|
||||
"assistant_name": assistant_name,
|
||||
})
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
"""
|
||||
Lazy initialization property for the prompt handler.
|
||||
|
||||
This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed.
|
||||
It uses the current file's path and additional keyword arguments for configuration.
|
||||
|
||||
Returns:
|
||||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def print_logo(self):
|
||||
"""
|
||||
Prints the logo of the CLI application to the console.
|
||||
|
||||
The logo is composed of multiple lines, which are iterated through
|
||||
and printed one by one to provide a visual identity for the chat interface.
|
||||
"""
|
||||
for line in self._logo:
|
||||
print(line)
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Property to access the memory service. If the service is initially set as a string,
|
||||
it will be looked up in the memory service dictionary of global context, initialized,
|
||||
and then returned as an instance of `BaseMemoryService`. Ensures the memory service
|
||||
is properly started before use.
|
||||
|
||||
Returns:
|
||||
BaseMemoryService: An active memory service instance.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of memory service is not found in the memory service dictionary of global context.
|
||||
"""
|
||||
if isinstance(self._memory_service, str):
|
||||
if self._memory_service not in G_CONTEXT.memory_service_dict:
|
||||
raise ValueError("Missing declaration of memory_service in yaml configuration: " + self._memory_service)
|
||||
self._memory_service = G_CONTEXT.memory_service_dict[self._memory_service]
|
||||
self._memory_service.init_service()
|
||||
self._memory_service.start_backend_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
"""
|
||||
Property to get the generation model. If the model is set as a string, it will be resolved from the global
|
||||
context's model dictionary.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of generation model is not found in the model dictionary of global context .
|
||||
|
||||
Returns:
|
||||
BaseModel: An actual generation model instance.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
if self._generation_model not in G_CONTEXT.model_dict:
|
||||
raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}")
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def chat_with_memory(self, query: str, remember_response: bool = False) -> ModelResponse | ModelResponseGen:
|
||||
"""
|
||||
Engages in a conversation with the AI model, utilizing conversation memory.
|
||||
The function sends the user's query, incorporates conversation history and memory,
|
||||
and optionally remembers the AI's response based on the user's preference.
|
||||
|
||||
Args:
|
||||
query (str): The user's input or query for the AI.
|
||||
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.
|
||||
|
||||
Side Effects:
|
||||
- Updates the conversation memory with the query of user and (optionally) the response of AI.
|
||||
- Retrieves and includes historical messages and memory content in the context of conversation.
|
||||
"""
|
||||
new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name=self.human_name, content=query)
|
||||
self.memory_service.add_messages(new_message)
|
||||
|
||||
messages: List[Message] = []
|
||||
|
||||
# Incorporate memory into the system prompt if available
|
||||
system_prompt = self.prompt_handler.system_prompt
|
||||
memories: str = self.memory_service.retrieve_memory()
|
||||
if memories:
|
||||
memory_prompt = self.prompt_handler.memory_prompt
|
||||
system_prompt = "\n".join([x.strip() for x in [system_prompt, memory_prompt, memories]])
|
||||
messages.append(Message(role=MessageRoleEnum.SYSTEM, content=system_prompt))
|
||||
|
||||
# Include past conversation history in the message list
|
||||
history_messages = self.memory_service.read_message()
|
||||
if history_messages:
|
||||
messages.extend(history_messages)
|
||||
|
||||
# Append the current user's message to the conversation context
|
||||
messages.append(new_message)
|
||||
self.logger.info(f"messages={messages}")
|
||||
|
||||
# Invoke the Language Model with the constructed message context, respecting streaming setting
|
||||
generated = self.generation_model.call(messages=messages, stream=self.stream, **self.generation_model_kwargs)
|
||||
|
||||
# In non-streaming interactions, explicitly save the AI's reply to memory if instructed
|
||||
if remember_response:
|
||||
assert not self.stream # Ensure we're not in streaming mode when remembering responses
|
||||
generated.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages(generated.message)
|
||||
|
||||
# Return the AI's response directly or as a generator based on the streaming mode
|
||||
return generated
|
||||
|
||||
@staticmethod
|
||||
def parse_query_command(query: str):
|
||||
"""
|
||||
Parses the user's input query command, separating it into the command and its associated keyword arguments.
|
||||
|
||||
Args:
|
||||
query (str): The raw input string from the user which includes the command and its arguments.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword
|
||||
arguments as the second element.
|
||||
"""
|
||||
query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command
|
||||
command = query_split[0] # Extract the command
|
||||
args = query_split[1:] # Extract the arguments following the command
|
||||
kwargs = {} # Initialize dictionary to hold keyword arguments
|
||||
|
||||
for arg in args:
|
||||
# Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original)
|
||||
if not args:
|
||||
continue
|
||||
arg_split = arg.split("=") # Split argument into key-value pair
|
||||
if len(arg_split) >= 2: # Ensure there's both a key and value
|
||||
k = arg_split[0] # Extract key
|
||||
v = arg_split[1] # Extract value
|
||||
if k and v: # Only add to kwargs if both key and value are non-empty
|
||||
kwargs[k] = v
|
||||
|
||||
return command, kwargs # Return the parsed command and keyword arguments
|
||||
|
||||
def process_commands(self, query: str) -> bool:
|
||||
"""
|
||||
Parses and executes commands from user input in the CLI chat interface.
|
||||
Supports operations like exiting, clearing screen, showing help, toggling stream mode,
|
||||
executing predefined memory operations, and handling unknown commands.
|
||||
|
||||
Args:
|
||||
query (str): The user's input command string.
|
||||
|
||||
Returns:
|
||||
bool: Indicates whether to continue running the CLI after processing the command.
|
||||
"""
|
||||
continue_run = True
|
||||
command, kwargs = self.parse_query_command(query)
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
if command == "exit":
|
||||
self.memory_service.stop_backend_service()
|
||||
continue_run = False
|
||||
|
||||
elif command == "clear":
|
||||
os.system("clear")
|
||||
|
||||
elif command == "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 command == "stream":
|
||||
self.stream = not self.stream
|
||||
questionary.print(f"set stream: {self.stream}")
|
||||
|
||||
elif command in self.memory_service.op_description_dict:
|
||||
refresh_time = kwargs.pop("refresh_time", "")
|
||||
if refresh_time and refresh_time.isdigit():
|
||||
refresh_time = int(refresh_time)
|
||||
self.memory_service.stop_backend_service()
|
||||
while True:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
os.system("clear")
|
||||
self.print_logo()
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
time.sleep(refresh_time)
|
||||
|
||||
else:
|
||||
result = self.memory_service.do_operation(op_name=command, **kwargs)
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
||||
else:
|
||||
questionary.print(f"Unknown command={command} received.")
|
||||
|
||||
return continue_run
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Runs the CLI chat loop, which handles user input, processes commands,
|
||||
communicates with the AI model, manages conversation memory, and controls
|
||||
the chat session including streaming responses, command execution, and error handling.
|
||||
|
||||
The loop continues until the user explicitly chooses to exit.
|
||||
"""
|
||||
self.print_logo()
|
||||
self.USER_COMMANDS.update(self.memory_service.op_description_dict)
|
||||
|
||||
while True:
|
||||
try:
|
||||
query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask()
|
||||
if not query:
|
||||
continue
|
||||
|
||||
query: str = query.strip()
|
||||
|
||||
# Handle special commands prefixed with '/'
|
||||
if query.startswith("/"):
|
||||
if self.process_commands(query=query):
|
||||
continue
|
||||
else:
|
||||
break
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response, with support for streaming
|
||||
self.memory_service.start_backend_service()
|
||||
if self.stream:
|
||||
model_response = None
|
||||
for model_response in self.chat_with_memory(query=query):
|
||||
questionary.print(model_response.delta, end="")
|
||||
questionary.print("")
|
||||
|
||||
else:
|
||||
model_response = self.chat_with_memory(query=query)
|
||||
questionary.print(model_response.message.content)
|
||||
|
||||
# Append AI's response to the conversation memory
|
||||
model_response.message.role_name = self.assistant_name
|
||||
self.memory_service.add_messages(model_response.message)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
# Handle user interruption and confirm exit
|
||||
questionary.print("User interrupt occurred.")
|
||||
is_exit = questionary.confirm("Continue exit?").unsafe_ask()
|
||||
if is_exit:
|
||||
self.memory_service.stop_backend_service()
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
# Log and handle any unanticipated exceptions
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.")
|
||||
continue
|
||||
|
|
@ -1,105 +1,17 @@
|
|||
import datetime
|
||||
import sys
|
||||
|
||||
import questionary
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
import fire
|
||||
import yaml
|
||||
import atexit
|
||||
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.enumeration.model_enum import ModelEnum
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.timer import timer
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config, camelcase_to_underscore
|
||||
from memoryscope.core.memoryscope import MemoryScope
|
||||
|
||||
|
||||
class MemoryScope(object):
|
||||
|
||||
def __init__(self):
|
||||
self.config: Dict[str, Any] = {}
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
class_name = camelcase_to_underscore(self.__class__.__name__)
|
||||
self.logger: Logger = Logger.get_logger(f"{class_name}_{datetime_suffix}", to_stream=False)
|
||||
|
||||
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!")
|
||||
self.init_global_content_by_config()
|
||||
atexit.register(self.shutdown) # register clean up function
|
||||
return self
|
||||
|
||||
@staticmethod
|
||||
def shutdown():
|
||||
questionary.print('Gracefully executing the shutdown function...')
|
||||
G_CONTEXT.memory_store.close()
|
||||
G_CONTEXT.monitor.close()
|
||||
G_CONTEXT.thread_pool.shutdown()
|
||||
|
||||
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):
|
||||
# set global config
|
||||
self.set_global_config()
|
||||
|
||||
# init memory_chat
|
||||
for name, conf in self.config["memory_chat"].items():
|
||||
G_CONTEXT.memory_chat_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# set memory_service
|
||||
for name, conf in self.config["memory_service"].items():
|
||||
G_CONTEXT.memory_service_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init models
|
||||
for name, conf in self.config["models"].items():
|
||||
G_CONTEXT.model_dict[name] = init_instance_by_config(conf, name=name)
|
||||
|
||||
# init vector_store
|
||||
if "memory_store" not in self.config:
|
||||
raise RuntimeError("memory_store config is required!")
|
||||
memory_store_config = self.config["memory_store"]
|
||||
embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
|
||||
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model)
|
||||
|
||||
# init monitor
|
||||
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
|
||||
|
||||
# set worker config
|
||||
G_CONTEXT.worker_config = self.config["worker"]
|
||||
|
||||
@property
|
||||
def default_chat_handle(self):
|
||||
return list(G_CONTEXT.memory_chat_dict.values())[0]
|
||||
|
||||
@property
|
||||
def default_service(self):
|
||||
return self.default_chat_handle.memory_service
|
||||
|
||||
|
||||
class CliJob(MemoryScope):
|
||||
|
||||
def run(self, config: str):
|
||||
self.load_config(config)
|
||||
self.init_global_content_by_config()
|
||||
self.default_chat_handle.run()
|
||||
def cli_job(**kwargs):
|
||||
with MemoryScope(**kwargs) as ms:
|
||||
memory_chat = ms.default_memory_chat
|
||||
memory_chat.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
cli_job = CliJob()
|
||||
fire.Fire(cli_job.run)
|
||||
fire.Fire(cli_job)
|
||||
|
|
|
|||
|
|
@ -5,11 +5,15 @@
|
|||
|
||||
WORKFLOW_NAME = "workflow_name"
|
||||
|
||||
MEMORYSCOPE_CONTEXT = "memoryscope_context"
|
||||
|
||||
RESULT = "result"
|
||||
|
||||
MEMORIES = "memories"
|
||||
|
||||
CHAT_MESSAGES = "chat_messages"
|
||||
|
||||
MEMORY_HANDLER = "memory_handler"
|
||||
MEMORY_MANAGER = "memory_manager"
|
||||
|
||||
CHAT_KWARGS = "chat_kwargs"
|
||||
|
||||
|
|
|
|||
217
memoryscope/core/chat/api_memory_chat.py
Normal file
217
memoryscope/core/chat/api_memory_chat.py
Normal file
|
|
@ -0,0 +1,217 @@
|
|||
from typing import List, Optional, Literal
|
||||
|
||||
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.scheme.message import Message
|
||||
from memoryscope.scheme.model_response import ModelResponse, ModelResponseGen
|
||||
|
||||
|
||||
class ApiMemoryChat(BaseMemoryChat):
|
||||
|
||||
def __init__(self,
|
||||
memory_service: str,
|
||||
generation_model: str,
|
||||
context: MemoryscopeContext,
|
||||
stream: bool = False,
|
||||
human_name: str = None,
|
||||
assistant_name: str = None,
|
||||
**kwargs):
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self._memory_service: BaseMemoryService | str = memory_service
|
||||
self._generation_model: BaseModel | str = generation_model
|
||||
self.context: MemoryscopeContext = context
|
||||
self.stream: bool = stream
|
||||
self.generation_model_kwargs: dict = kwargs.pop("generation_model_kwargs", {})
|
||||
|
||||
self.human_name: str = human_name
|
||||
if not self.human_name:
|
||||
self.human_name = DEFAULT_HUMAN_NAME[self.context.language]
|
||||
self.context.meta_data["human_name"] = self.human_name
|
||||
|
||||
self.assistant_name: str = assistant_name
|
||||
if not self.assistant_name:
|
||||
self.assistant_name = "AI"
|
||||
self.context.meta_data["assistant_name"] = self.assistant_name
|
||||
|
||||
self._prompt_handler: PromptHandler | None = None
|
||||
|
||||
@property
|
||||
def prompt_handler(self) -> PromptHandler:
|
||||
"""
|
||||
Lazy initialization property for the prompt handler.
|
||||
|
||||
This property ensures that the `_prompt_handler` attribute is only instantiated when it is first accessed.
|
||||
It uses the current file's path and additional keyword arguments for configuration.
|
||||
|
||||
Returns:
|
||||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__,
|
||||
language=self.context.language,
|
||||
prompt_file="memory_chat_prompt",
|
||||
**self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
@property
|
||||
def memory_service(self) -> BaseMemoryService:
|
||||
"""
|
||||
Property to access the memory service. If the service is initially set as a string,
|
||||
it will be looked up in the memory service dictionary of context, initialized,
|
||||
and then returned as an instance of `BaseMemoryService`. Ensures the memory service
|
||||
is properly started before use.
|
||||
|
||||
Returns:
|
||||
BaseMemoryService: An active memory service instance.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of memory service is not found in the memory service dictionary of context.
|
||||
"""
|
||||
if isinstance(self._memory_service, str):
|
||||
if self._memory_service not in self.context.memory_service_dict:
|
||||
raise ValueError(f"Missing declaration of memory_service in context: {self._memory_service}")
|
||||
|
||||
self._memory_service: BaseMemoryService = self.context.memory_service_dict[self._memory_service]
|
||||
# init service & update kwargs
|
||||
self._memory_service.init_service()
|
||||
return self._memory_service
|
||||
|
||||
@property
|
||||
def generation_model(self) -> BaseModel:
|
||||
"""
|
||||
Property to get the generation model. If the model is set as a string, it will be resolved from the global
|
||||
context's model dictionary.
|
||||
|
||||
Raises:
|
||||
ValueError: If the declaration of generation model is not found in the model dictionary of context .
|
||||
|
||||
Returns:
|
||||
BaseModel: An actual generation model instance.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
if self._generation_model not in self.context.model_dict:
|
||||
raise ValueError(f"Missing declaration of generation model in yaml config: {self._generation_model}")
|
||||
self._generation_model = self.context.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
def iter_response(self,
|
||||
remember_response: bool,
|
||||
resp: ModelResponseGen,
|
||||
memories: str,
|
||||
query_message: Message) -> ModelResponseGen:
|
||||
|
||||
model_response: ModelResponse | None = None
|
||||
for model_response in resp:
|
||||
yield model_response
|
||||
|
||||
if remember_response:
|
||||
if model_response and model_response.message:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
model_response.meta_data[MEMORIES] = memories
|
||||
self.memory_service.add_messages([query_message, model_response.message])
|
||||
else:
|
||||
self.logger.warning("model_response or model_response.message is empty!")
|
||||
|
||||
def chat_with_memory(self,
|
||||
query: str,
|
||||
role_name: Optional[str] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
memory_prompt: Optional[str] = None,
|
||||
extra_memories: Optional[str] = None,
|
||||
history_message_strategy: Literal["auto", None] | int = "auto",
|
||||
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.
|
||||
history_message_strategy ("auto", None, int):
|
||||
- If it is set to "auto", the history messages in the conversation will retain those that have not
|
||||
yet been summarized. Default to "auto".
|
||||
- If it is set to None, no conversation history will be saved.
|
||||
- If it is set to an integer value "n", the most recent "n" messages will be retained.
|
||||
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] = []
|
||||
|
||||
# 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 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_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
|
||||
if history_message_strategy:
|
||||
history_messages = []
|
||||
|
||||
if history_message_strategy == "auto":
|
||||
history_messages = self.memory_service.read_message()
|
||||
|
||||
elif isinstance(history_message_strategy, int):
|
||||
history_messages = self.memory_service.chat_messages[-history_message_strategy:]
|
||||
|
||||
if history_messages:
|
||||
chat_messages.extend(history_messages)
|
||||
|
||||
# Append the current user's message to the conversation context
|
||||
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:
|
||||
return self.iter_response(remember_response, resp, memories, query_message)
|
||||
|
||||
else:
|
||||
model_response: ModelResponse = resp
|
||||
if remember_response:
|
||||
if model_response and model_response.message:
|
||||
model_response.message.role_name = self.assistant_name
|
||||
model_response.meta_data[MEMORIES] = memories
|
||||
self.memory_service.add_messages([query_message, model_response.message])
|
||||
else:
|
||||
self.logger.warning("model_response or model_response.message is empty!")
|
||||
return model_response
|
||||
68
memoryscope/core/chat/base_memory_chat.py
Normal file
68
memoryscope/core/chat/base_memory_chat.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import Optional, Literal
|
||||
|
||||
from memoryscope.core.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryChat(metaclass=ABCMeta):
|
||||
"""
|
||||
An abstract base class representing a chat system integrated with memory services.
|
||||
It outlines the method to initiate a chat session leveraging memory data, which concrete subclasses must implement.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs: dict = kwargs
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@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
|
||||
|
||||
@abstractmethod
|
||||
def chat_with_memory(self,
|
||||
query: str,
|
||||
role_name: Optional[str] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
memory_prompt: Optional[str] = None,
|
||||
extra_memories: Optional[str] = None,
|
||||
history_message_strategy: Literal["auto", None] | int = "auto",
|
||||
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.
|
||||
history_message_strategy ("auto", None, int):
|
||||
- If it is set to "auto", the history messages in the conversation will retain those that have not
|
||||
yet been summarized. Default to "auto".
|
||||
- If it is set to None, no conversation history will be saved.
|
||||
- If it is set to an integer value "n", the most recent "n" messages will be retained.
|
||||
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]
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Abstract method to run the chat system.
|
||||
|
||||
This method should contain the logic to initiate and manage the chat process,
|
||||
utilizing the memory service as needed. It must be implemented by subclasses.
|
||||
"""
|
||||
pass
|
||||
206
memoryscope/core/chat/cli_memory_chat.py
Normal file
206
memoryscope/core/chat/cli_memory_chat.py
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
import os
|
||||
import time
|
||||
from typing import Optional, Literal
|
||||
|
||||
import questionary
|
||||
|
||||
from memoryscope.core.chat.api_memory_chat import ApiMemoryChat
|
||||
from memoryscope.core.utils.tool_functions import char_logo
|
||||
|
||||
|
||||
class CliMemoryChat(ApiMemoryChat):
|
||||
"""
|
||||
Command-line interface for chatting with an AI that integrates memory functionality.
|
||||
Allows users to interact, manage chat history, adjust streaming settings, and view commands' help.
|
||||
"""
|
||||
USER_COMMANDS = {
|
||||
"exit": "Exit the CLI.",
|
||||
"clear": "Clear the command history.",
|
||||
"help": "Display available CLI commands and their descriptions.",
|
||||
"stream": "Toggle between getting streamed responses from the model."
|
||||
}
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._logo = char_logo("MemoryScope")
|
||||
|
||||
def print_logo(self):
|
||||
"""
|
||||
Prints the logo of the CLI application to the console.
|
||||
|
||||
The logo is composed of multiple lines, which are iterated through
|
||||
and printed one by one to provide a visual identity for the chat interface.
|
||||
"""
|
||||
for line in self._logo:
|
||||
print(line)
|
||||
|
||||
def chat_with_memory(self,
|
||||
query: str,
|
||||
role_name: Optional[str] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
memory_prompt: Optional[str] = None,
|
||||
extra_memories: Optional[str] = None,
|
||||
history_message_strategy: Literal["auto", None] | int = "auto",
|
||||
remember_response: bool = True,
|
||||
**kwargs):
|
||||
resp = super().chat_with_memory(query=query,
|
||||
role_name=role_name,
|
||||
system_prompt=system_prompt,
|
||||
memory_prompt=memory_prompt,
|
||||
extra_memories=extra_memories,
|
||||
history_message_strategy=history_message_strategy,
|
||||
remember_response=remember_response,
|
||||
**kwargs)
|
||||
|
||||
if self.stream:
|
||||
for _resp in resp:
|
||||
questionary.print(_resp.delta, end="")
|
||||
questionary.print("")
|
||||
else:
|
||||
questionary.print(resp.message.content)
|
||||
|
||||
@staticmethod
|
||||
def parse_query_command(query: str):
|
||||
"""
|
||||
Parses the user's input query command, separating it into the command and its associated keyword arguments.
|
||||
|
||||
Args:
|
||||
query (str): The raw input string from the user which includes the command and its arguments.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing the command (str) as the first element and a dictionary (kwargs) of keyword
|
||||
arguments as the second element.
|
||||
"""
|
||||
query_split = query.lstrip("/").lower().split(" ") # Split and preprocess the input command
|
||||
command = query_split[0] # Extract the command
|
||||
args = query_split[1:] # Extract the arguments following the command
|
||||
kwargs = {} # Initialize dictionary to hold keyword arguments
|
||||
|
||||
for arg in args:
|
||||
# Skip if no arguments exist (unnecessary check due to prior assignment, but retained as per original)
|
||||
if not args:
|
||||
continue
|
||||
arg_split = arg.split("=") # Split argument into key-value pair
|
||||
if len(arg_split) >= 2: # Ensure there's both a key and value
|
||||
k = arg_split[0] # Extract key
|
||||
v = arg_split[1] # Extract value
|
||||
if k and v: # Only add to kwargs if both key and value are non-empty
|
||||
kwargs[k] = v
|
||||
|
||||
return command, kwargs # Return the parsed command and keyword arguments
|
||||
|
||||
def process_commands(self, query: str) -> bool:
|
||||
"""
|
||||
Parses and executes commands from user input in the CLI chat interface.
|
||||
Supports operations like exiting, clearing screen, showing help, toggling stream mode,
|
||||
executing predefined memory operations, and handling unknown commands.
|
||||
|
||||
Args:
|
||||
query (str): The user's input command string.
|
||||
|
||||
Returns:
|
||||
bool: Indicates whether to continue running the CLI after processing the command.
|
||||
"""
|
||||
continue_run = True
|
||||
command, kwargs = self.parse_query_command(query)
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
if command == "exit":
|
||||
self.memory_service.stop_backend_service()
|
||||
continue_run = False
|
||||
|
||||
elif command == "clear":
|
||||
os.system("clear")
|
||||
|
||||
elif command == "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 command == "stream":
|
||||
self.stream = not self.stream
|
||||
questionary.print(f"set stream: {self.stream}")
|
||||
|
||||
elif command in self.memory_service.op_description_dict:
|
||||
refresh_time = kwargs.pop("refresh_time", "")
|
||||
if refresh_time and refresh_time.isdigit():
|
||||
refresh_time = int(refresh_time)
|
||||
self.memory_service.stop_backend_service()
|
||||
while True:
|
||||
result = self.memory_service.do_operation(name=command, **kwargs)
|
||||
os.system("clear")
|
||||
self.print_logo()
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
time.sleep(refresh_time)
|
||||
|
||||
else:
|
||||
result = self.memory_service.do_operation(name=command, **kwargs)
|
||||
if result:
|
||||
if isinstance(result, list):
|
||||
result = "\n".join([str(x) for x in result])
|
||||
questionary.print(result)
|
||||
else:
|
||||
questionary.print(f"command={command} result is empty! kwargs={kwargs}")
|
||||
|
||||
else:
|
||||
questionary.print(f"Unknown command={command} received.")
|
||||
|
||||
return continue_run
|
||||
|
||||
def run(self):
|
||||
"""
|
||||
Runs the CLI chat loop, which handles user input, processes commands,
|
||||
communicates with the AI model, manages conversation memory, and controls
|
||||
the chat session including streaming responses, command execution, and error handling.
|
||||
|
||||
The loop continues until the user explicitly chooses to exit.
|
||||
"""
|
||||
self.print_logo()
|
||||
self.USER_COMMANDS.update(self.memory_service.op_description_dict)
|
||||
|
||||
while True:
|
||||
try:
|
||||
query = questionary.text(message=f"{self.human_name}:", multiline=False, qmark=">").unsafe_ask()
|
||||
if not query:
|
||||
continue
|
||||
|
||||
query: str = query.strip()
|
||||
|
||||
# Handle special commands prefixed with '/'
|
||||
if query.startswith("/"):
|
||||
if self.process_commands(query=query):
|
||||
continue
|
||||
else:
|
||||
break
|
||||
|
||||
# Print prompt for AI's response
|
||||
questionary.print("> ", end="", style="fg:yellow")
|
||||
questionary.print(f"{self.assistant_name}: ", end="", style="bold")
|
||||
|
||||
# Fetch and display AI's response
|
||||
self.memory_service.start_backend_service()
|
||||
self.chat_with_memory(query=query)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
# Handle user interruption and confirm exit
|
||||
questionary.print("User interrupt occurred.")
|
||||
is_exit = questionary.confirm("Continue exit?").unsafe_ask()
|
||||
if is_exit:
|
||||
self.memory_service.stop_backend_service()
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
# Log and handle any unanticipated exceptions
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
self.logger.exception(f"An exception occurred when running cli memory chat. args={e.args}.")
|
||||
continue
|
||||
61
memoryscope/core/config/arguments.py
Normal file
61
memoryscope/core/config/arguments.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import Literal, Dict
|
||||
|
||||
|
||||
@dataclass
|
||||
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"})
|
||||
|
||||
logger_name: str = field(default="memoryscope")
|
||||
|
||||
logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S")
|
||||
|
||||
logger_to_screen: bool = field(default=False, metadata={"help": "If false, it does not print to the screen."})
|
||||
|
||||
memory_chat_class: str = field(default="cli_memory_chat", metadata={
|
||||
"help": "cli_memory_chat(Command-line interaction), api_memory_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."})
|
||||
|
||||
reflect_and_reconsolidate_interval_time: int = field(default=15, metadata={
|
||||
"help": "If you feel that the token consumption is relatively high, please increase the time interval."})
|
||||
|
||||
worker_params: Dict[str, dict] = field(default_factory=lambda: {}, metadata={
|
||||
"help": "dict format: worker_name -> param_key -> param_value"})
|
||||
|
||||
generation_backend: str = field(default="openai_generation", metadata={
|
||||
"help": "global generation backend: openai_generation, dashscope_generation, etc."})
|
||||
|
||||
generation_model: str = field(default="gpt-4o", metadata={
|
||||
"help": "global generation model: gpt-4o, gpt-4, qwen-max, etc."})
|
||||
|
||||
generation_params: dict = field(default_factory=lambda: {}, metadata={
|
||||
"help": "global generation params: max_tokens, top_p, temperature, etc."})
|
||||
|
||||
embedding_backend: str = field(default="openai_embedding", metadata={
|
||||
"help": "global embedding backend: openai_embedding, dashscope_embedding, etc."})
|
||||
|
||||
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: {})
|
||||
|
||||
use_dummy_ranker: bool = field(default=True, metadata={
|
||||
"help": "If a semantic ranking model is not available, MemoryScope will use cosine similarity scoring as a "
|
||||
"substitute. However, the ranking effectiveness will be somewhat compromised."})
|
||||
|
||||
rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."})
|
||||
|
||||
rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."})
|
||||
|
||||
rank_params: dict = field(default_factory=lambda: {})
|
||||
|
||||
es_index_name: str = field(default="memory_index")
|
||||
|
||||
es_url: str = field(default="http://localhost:9200")
|
||||
|
||||
retrieve_mode: str = field(default="dense", metadata={
|
||||
"help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"})
|
||||
189
memoryscope/core/config/config_manager.py
Normal file
189
memoryscope/core/config/config_manager.py
Normal file
|
|
@ -0,0 +1,189 @@
|
|||
import json
|
||||
from dataclasses import fields
|
||||
from datetime import datetime
|
||||
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.core.utils.logger import Logger
|
||||
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
|
||||
self.logger = self._init_logger()
|
||||
self.logger.info("init by config mode:")
|
||||
|
||||
elif config_path:
|
||||
self.read_config(config_path)
|
||||
self.logger = self._init_logger()
|
||||
self.logger.info("init by config_path mode:")
|
||||
|
||||
else:
|
||||
self.read_demo_config(demo_config_name)
|
||||
if arguments:
|
||||
self.update_config_by_arguments(arguments)
|
||||
self.logger = self._init_logger()
|
||||
self.logger.info(f"init by arguments mode: {arguments.__dict__}")
|
||||
|
||||
elif kwargs:
|
||||
kwargs = {k: v for k, v in kwargs.items() if k in [x.name for x in fields(Arguments)]}
|
||||
arguments = Arguments(**kwargs)
|
||||
self.update_config_by_arguments(arguments)
|
||||
self.logger = self._init_logger()
|
||||
self.logger.info(f"init by kwargs mode: {kwargs}")
|
||||
|
||||
else:
|
||||
raise RuntimeError("can not init config manager without kwargs!")
|
||||
self.logger.info(self.dump_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.now().strftime(logger_name_time_suffix)
|
||||
logger_name = f"{logger_name}_{suffix}"
|
||||
return Logger.get_logger(logger_name, to_stream=global_config["logger_to_screen"])
|
||||
|
||||
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,
|
||||
"logger_to_screen": arguments.logger_to_screen,
|
||||
"use_dummy_ranker": arguments.use_dummy_ranker,
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def update_memory_chat_by_arguments(config: dict, arguments: Arguments):
|
||||
memory_chat_class_split = config["class"].split(".")
|
||||
stream = arguments.memory_chat_class in ["cli_memory_chat", ]
|
||||
config.update({
|
||||
"class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]),
|
||||
"human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)],
|
||||
"assistant_name": "AI",
|
||||
"stream": stream,
|
||||
})
|
||||
|
||||
@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"] = "yaml", file_path: Optional[str] = None) -> str:
|
||||
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 file_path:
|
||||
with open(file_path, "w") as f:
|
||||
f.write(content)
|
||||
|
||||
return content
|
||||
176
memoryscope/core/config/demo_config.yaml
Normal file
176
memoryscope/core/config/demo_config.yaml
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
global:
|
||||
language: en
|
||||
thread_pool_max_workers: 5
|
||||
logger_name: memoryscope
|
||||
logger_name_time_suffix: "%Y%m%d_%H%M%S"
|
||||
logger_to_screen: false
|
||||
use_dummy_ranker: false
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: core.chat.cli_memory_chat
|
||||
memory_service: memoryscope_service
|
||||
generation_model: generation_model
|
||||
|
||||
memory_service:
|
||||
memoryscope_service:
|
||||
class: core.service.memory_scope_service
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: read_message
|
||||
description: "read short memory"
|
||||
|
||||
retrieve_memory:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: set_query,[extract_time|retrieve_obs_ins,semantic_rank],fuse_rerank
|
||||
description: "retrieve long-term memory"
|
||||
|
||||
list_memory:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: set_query,retrieve_top_memory,print_memory
|
||||
description: "read all long-term memory of the user"
|
||||
|
||||
delete_memory:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: set_query,retrieve_all_memory,delete_memory
|
||||
description: "delete a single long-term memory"
|
||||
|
||||
delete_all:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: set_query,retrieve_all_memory,delete_all
|
||||
description: "delete all long-term memory"
|
||||
|
||||
add_memory:
|
||||
class: core.operation.frontend_operation
|
||||
workflow: add_memory
|
||||
description: "add a single observation"
|
||||
|
||||
consolidate_memory:
|
||||
class: core.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: core.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
|
||||
|
||||
worker:
|
||||
dummy:
|
||||
class: core.worker.dummy_worker
|
||||
generation_model: generation_model
|
||||
embedding_model: embedding_model
|
||||
rank_model: rank_model
|
||||
read_message:
|
||||
class: core.worker.frontend.read_message_worker
|
||||
set_query:
|
||||
class: core.worker.frontend.set_query_worker
|
||||
retrieve_obs_ins:
|
||||
class: core.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_top_k: 100
|
||||
extract_time:
|
||||
class: core.worker.frontend.extract_time_worker
|
||||
generation_model: generation_model
|
||||
semantic_rank:
|
||||
class: core.worker.frontend.semantic_rank_worker
|
||||
rank_model: rank_model
|
||||
fuse_rerank:
|
||||
class: core.worker.frontend.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.01
|
||||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
observation: 1
|
||||
obs_customized: 1.2
|
||||
insight: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 20
|
||||
retrieve_top_memory:
|
||||
class: core.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
print_memory:
|
||||
class: core.worker.frontend.print_memory_worker
|
||||
retrieve_all_memory:
|
||||
class: core.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 1000
|
||||
retrieve_ins_top_k: 1000
|
||||
retrieve_expired_top_k: 1000
|
||||
delete_memory:
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: delete_memory
|
||||
delete_all:
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: delete_all
|
||||
add_memory:
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: from_query
|
||||
info_filter:
|
||||
class: core.worker.backend.info_filter_worker
|
||||
generation_model: generation_model
|
||||
load_today_memory:
|
||||
class: core.worker.backend.load_memory_worker
|
||||
retrieve_today_top_k: 100
|
||||
get_observation:
|
||||
class: core.worker.backend.get_observation_worker
|
||||
generation_model: generation_model
|
||||
get_observation_with_time:
|
||||
class: core.worker.backend.get_observation_with_time_worker
|
||||
generation_model: generation_model
|
||||
contra_repeat:
|
||||
class: core.worker.backend.contra_repeat_worker
|
||||
generation_model: generation_model
|
||||
store_memory:
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: from_memory_key
|
||||
memory_key: all
|
||||
load_obs_and_insight:
|
||||
class: core.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: core.worker.backend.get_reflection_subject_worker
|
||||
generation_model: generation_model
|
||||
reflect_obs_cnt_threshold: 10
|
||||
update_insight:
|
||||
class: core.worker.backend.update_insight_worker
|
||||
generation_model: generation_model
|
||||
rank_model: rank_model
|
||||
long_contra_repeat:
|
||||
class: core.worker.backend.long_contra_repeat_worker
|
||||
generation_model: generation_model
|
||||
|
||||
model:
|
||||
generation_model:
|
||||
class: core.models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
max_tokens: 2000
|
||||
embedding_model:
|
||||
class: core.models.llama_index_embedding_model
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
rank_model:
|
||||
class: core.models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
top_n: 500
|
||||
dummy_generation:
|
||||
class: core.models.dummy_generation_model
|
||||
module_name: dummy_generation
|
||||
model_name: dummy_generation_model
|
||||
|
||||
memory_store:
|
||||
class: core.storage.llama_index_es_memory_store
|
||||
embedding_model: embedding_model
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
retrieve_mode: dense
|
||||
|
||||
monitor:
|
||||
class: core.storage.dummy_monitor
|
||||
94
memoryscope/core/memoryscope.py
Normal file
94
memoryscope/core/memoryscope.py
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
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.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.context: MemoryscopeContext = MemoryscopeContext()
|
||||
self.init_context_by_config()
|
||||
|
||||
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()
|
||||
|
||||
self.logger.close()
|
||||
|
||||
def __enter__(self):
|
||||
self.init_context_by_config()
|
||||
return self
|
||||
|
||||
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]
|
||||
29
memoryscope/core/memoryscope_context.py
Normal file
29
memoryscope/core/memoryscope_context.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class MemoryscopeContext(object):
|
||||
"""
|
||||
The context class archives all configs utilized by store, monitor, services and workers.
|
||||
"""
|
||||
|
||||
language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
thread_pool: ThreadPoolExecutor | None = None
|
||||
|
||||
memory_store = None
|
||||
|
||||
monitor = None
|
||||
|
||||
memory_chat_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_chat"})
|
||||
|
||||
memory_service_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> memory_service"})
|
||||
|
||||
model_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> model"})
|
||||
|
||||
worker_conf_dict: dict = field(default_factory=lambda: {}, metadata={"help": "name -> worker_conf"})
|
||||
|
||||
meta_data: dict = field(default_factory=lambda: {})
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -19,38 +19,37 @@ class DummyGenerationModel(BaseModel):
|
|||
"""
|
||||
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
|
||||
|
||||
class DummyModel:
|
||||
"""
|
||||
An inner class representing the dummy model placeholder.
|
||||
"""
|
||||
pass
|
||||
MODEL_REGISTRY.register("dummy_generation", object)
|
||||
|
||||
MODEL_REGISTRY.register("dummy_generation", DummyModel)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
def before_call(self, model_response: ModelResponse, **kwargs):
|
||||
"""
|
||||
Prepares the input data before making a call to the model's generate function.
|
||||
Accepts either a 'prompt' or a list of 'messages'. If both are provided or missing,
|
||||
a RuntimeError is raised. Transforms the input into a standardized format for processing.
|
||||
Prepares the input data before making a call to the language model.
|
||||
It accepts either a 'prompt' directly or a list of 'messages'.
|
||||
If 'prompt' is provided, it sets the data accordingly.
|
||||
If 'messages' are provided, it constructs a list of ChatMessage objects from the list.
|
||||
Raises an error if neither 'prompt' nor 'messages' are supplied.
|
||||
|
||||
Args:
|
||||
**kwargs: Arbitrary keyword arguments including 'prompt' or 'messages'.
|
||||
|
||||
model_response: model_response
|
||||
**kwargs: Arbitrary keyword arguments including 'prompt' and 'messages'.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If neither 'prompt' nor 'messages' is provided, or both are provided.
|
||||
RuntimeError: When both 'prompt' and 'messages' inputs are not provided.
|
||||
"""
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Message] | List[dict] = kwargs.pop("messages", [])
|
||||
|
||||
if prompt:
|
||||
self.data = {"prompt": prompt}
|
||||
data = {"prompt": prompt}
|
||||
elif messages:
|
||||
if isinstance(messages[0], dict):
|
||||
self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]}
|
||||
data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]}
|
||||
else:
|
||||
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
|
||||
data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
|
||||
else:
|
||||
raise RuntimeError("Both 'prompt' and 'messages' are empty!")
|
||||
raise RuntimeError("prompt and messages are both empty!")
|
||||
data.update(**kwargs)
|
||||
model_response.meta_data["data"] = data
|
||||
|
||||
def after_call(self,
|
||||
model_response: ModelResponse,
|
||||
|
|
@ -79,43 +78,16 @@ class DummyGenerationModel(BaseModel):
|
|||
for delta in call_result:
|
||||
model_response.message.content += delta
|
||||
model_response.delta = delta
|
||||
time.sleep(0.1) # ⭐ Introduce a delay to simulate streaming
|
||||
time.sleep(0.1)
|
||||
yield model_response
|
||||
|
||||
return gen()
|
||||
else:
|
||||
model_response.message.content = "".join(call_result) # ⭐ Concatenate results for non-streaming
|
||||
model_response.message.content = "".join(call_result)
|
||||
return model_response
|
||||
|
||||
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
|
||||
"""
|
||||
Generates a dummy response based on the input data, supporting both immediate
|
||||
and streamed response types.
|
||||
def _call(self, model_response: ModelResponse, stream: bool = False, **kwargs):
|
||||
return model_response
|
||||
|
||||
Args:
|
||||
stream (bool, optional): If True, indicates the response should be generated
|
||||
in a streaming manner. Defaults to False.
|
||||
**kwargs: Additional keyword arguments not used in this dummy implementation.
|
||||
|
||||
Returns:
|
||||
Union[ModelResponse, ModelResponseGen]: A dummy response object or a generator
|
||||
object capable of streaming responses.
|
||||
"""
|
||||
assert "prompt" in self.data or "messages" in self.data
|
||||
results = ModelResponse(m_type=self.m_type)
|
||||
return results
|
||||
|
||||
async def _async_call(self, **kwargs) -> ModelResponse:
|
||||
"""
|
||||
Asynchronous version of `_call`, providing the same functionality but designed
|
||||
to be used in asynchronous contexts.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments not used in this dummy implementation.
|
||||
|
||||
Returns:
|
||||
ModelResponse: A dummy response object suitable for asynchronous use.
|
||||
"""
|
||||
assert "prompt" in self.data or "messages" in self.data
|
||||
results = ModelResponse(m_type=self.m_type)
|
||||
return results
|
||||
async def _async_call(self, model_response: ModelResponse, **kwargs):
|
||||
return model_response
|
||||
|
|
@ -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,11 +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.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BackendOperation(BaseWorkflow, BaseOperation):
|
||||
|
|
@ -30,7 +29,7 @@ class BackendOperation(BaseWorkflow, BaseOperation):
|
|||
|
||||
self._operation_status_run: bool = False
|
||||
self._loop_switch: bool = False
|
||||
self._run_thread = None
|
||||
self._backend_task = None
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
|
|
@ -54,17 +53,14 @@ class BackendOperation(BaseWorkflow, BaseOperation):
|
|||
Returns:
|
||||
Any: The result obtained after executing the workflow.
|
||||
"""
|
||||
self.context.clear()
|
||||
|
||||
# Add additional arguments to the context
|
||||
kwargs.update(**self.kwargs)
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
|
||||
# Include the most recent messages in the operation context
|
||||
self.context[CHAT_MESSAGES] = self.chat_messages
|
||||
# prepare kwargs
|
||||
workflow_kwargs = {
|
||||
CHAT_MESSAGES: self.chat_messages,
|
||||
CHAT_KWARGS: {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
# Execute the workflow with the prepared context
|
||||
self.run_workflow()
|
||||
self.run_workflow(**workflow_kwargs)
|
||||
|
||||
# Retrieve the result from the context after workflow execution
|
||||
return self.context.get(RESULT)
|
||||
|
|
@ -107,17 +103,24 @@ class BackendOperation(BaseWorkflow, BaseOperation):
|
|||
if self._loop_switch:
|
||||
self.run_operation()
|
||||
|
||||
def run_operation_backend(self):
|
||||
def start_operation_backend(self):
|
||||
"""
|
||||
Initiates the background operation loop if it's not already running.
|
||||
Sets the _loop_switch to True and submits the _loop_operation to a thread from the global thread pool.
|
||||
"""
|
||||
if not self._loop_switch:
|
||||
self._loop_switch = True
|
||||
self._run_thread = G_CONTEXT.thread_pool.submit(self._loop_operation)
|
||||
self._backend_task = self.thread_pool.submit(self._loop_operation)
|
||||
self.logger.info(f"start operation={self.name}...")
|
||||
|
||||
def stop_operation_backend(self):
|
||||
def stop_operation_backend(self, wait_task_end: bool = False):
|
||||
"""
|
||||
Stops the background operation loop by setting the _loop_switch to False.
|
||||
"""
|
||||
self._loop_switch = False
|
||||
if self._backend_task:
|
||||
if wait_task_end:
|
||||
self._backend_task.result()
|
||||
self.logger.info(f"stop operation={self.name}...")
|
||||
else:
|
||||
self.logger.info(f"send stop signal to operation={self.name}...")
|
||||
|
|
@ -12,7 +12,6 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
operation_type (OPERATION_TYPE): Specifies the type of operation, defaulting to "frontend".
|
||||
name (str): The name of the operation.
|
||||
description (str): A description of the operation.
|
||||
kwargs (dict): Additional keyword arguments for operation configuration.
|
||||
"""
|
||||
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
|
@ -51,14 +50,14 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def run_operation_backend(self):
|
||||
def start_operation_backend(self):
|
||||
"""
|
||||
Placeholder method for running an operation specific to the backend.
|
||||
Intended to be overridden by subclasses if backend operations are required.
|
||||
"""
|
||||
pass
|
||||
|
||||
def stop_operation_backend(self):
|
||||
def stop_operation_backend(self, wait_task_end: bool = False):
|
||||
"""
|
||||
Placeholder method to stop any ongoing backend operations.
|
||||
Should be implemented in subclasses where backend operations are managed.
|
||||
|
|
@ -4,25 +4,26 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
|
|||
from itertools import zip_longest
|
||||
from typing import Dict, Any, List
|
||||
|
||||
from memoryscope.constants.common_constants import WORKFLOW_NAME
|
||||
from memoryscope.memory.worker.base_worker import BaseWorker
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.timer import Timer
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
from memoryscope.constants.common_constants import WORKFLOW_NAME, MEMORYSCOPE_CONTEXT
|
||||
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):
|
||||
|
||||
def __init__(self,
|
||||
name: str,
|
||||
memoryscope_context: MemoryscopeContext,
|
||||
workflow: str = "",
|
||||
thread_pool: ThreadPoolExecutor = G_CONTEXT.thread_pool,
|
||||
**kwargs):
|
||||
|
||||
self.name: str = name
|
||||
self.memoryscope_context: MemoryscopeContext = memoryscope_context
|
||||
self.thread_pool: ThreadPoolExecutor = self.memoryscope_context.thread_pool
|
||||
self.workflow: str = workflow
|
||||
self.thread_pool: ThreadPoolExecutor = thread_pool
|
||||
self.kwargs = kwargs
|
||||
|
||||
self.workflow_worker_list: List[List[List[str]]] = []
|
||||
|
|
@ -128,17 +129,16 @@ class BaseWorkflow(object):
|
|||
This method modifies `self.worker_dict` in-place, replacing the keys with actual worker instances.
|
||||
"""
|
||||
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!")
|
||||
if name not in self.memoryscope_context.worker_conf_dict:
|
||||
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",
|
||||
config=self.memoryscope_context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=is_backend or self.worker_dict[name],
|
||||
context=self.context,
|
||||
context_lock=self.context_lock,
|
||||
thread_pool=G_CONTEXT.thread_pool,
|
||||
thread_pool=self.thread_pool,
|
||||
**kwargs)
|
||||
|
||||
def _run_sub_workflow(self, worker_list: List[str]) -> bool:
|
||||
|
|
@ -150,7 +150,7 @@ class BaseWorkflow(object):
|
|||
return False
|
||||
return True
|
||||
|
||||
def run_workflow(self):
|
||||
def run_workflow(self, **kwargs):
|
||||
"""
|
||||
Executes the workflow by orchestrating the steps defined in `self.workflow_worker_list`.
|
||||
This method supports both sequential and parallel execution of sub-workflows based on the structure
|
||||
|
|
@ -159,9 +159,18 @@ class BaseWorkflow(object):
|
|||
If a workflow part consists of a single item, it is executed sequentially. For parts with multiple items,
|
||||
they are submitted for parallel execution using a thread pool. The workflow will stop if any sub-workflow
|
||||
returns False.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments to be passed to context.
|
||||
"""
|
||||
with Timer(f"workflow.{self.name}", time_log_type="wrap"):
|
||||
self.context[WORKFLOW_NAME] = self.name
|
||||
self.context.clear()
|
||||
|
||||
self.context.update({
|
||||
WORKFLOW_NAME: self.name,
|
||||
MEMORYSCOPE_CONTEXT: self.memoryscope_context,
|
||||
**kwargs,
|
||||
})
|
||||
|
||||
# Iterate over each part of the workflow
|
||||
for workflow_part in self.workflow_worker_list:
|
||||
|
|
@ -174,7 +183,7 @@ class BaseWorkflow(object):
|
|||
t_list = []
|
||||
# Submit tasks to the thread pool
|
||||
for sub_workflow in workflow_part:
|
||||
t_list.append(G_CONTEXT.thread_pool.submit(self._run_sub_workflow, sub_workflow))
|
||||
t_list.append(self.thread_pool.submit(self._run_sub_workflow, sub_workflow))
|
||||
|
||||
# Check results; if any task returns False, stop the workflow
|
||||
flag = True
|
||||
|
|
@ -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 SummaryObservationOp(BackendOperation):
|
||||
class ConsolidateMemoryOp(BackendOperation):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super(SummaryObservationOp, 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)
|
||||
|
|
@ -43,17 +43,14 @@ class SummaryObservationOp(BackendOperation):
|
|||
f"contextual_msg_min_count({self.contextual_msg_min_count}), skip.")
|
||||
return
|
||||
|
||||
self.context.clear()
|
||||
|
||||
# Add additional arguments to the context
|
||||
kwargs.update(**self.kwargs)
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
|
||||
# Include the most recent messages in the operation context
|
||||
self.context[CHAT_MESSAGES] = chat_messages
|
||||
# prepare kwargs
|
||||
workflow_kwargs = {
|
||||
CHAT_MESSAGES: chat_messages,
|
||||
CHAT_KWARGS: {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
# Execute the workflow with the prepared context
|
||||
self.run_workflow()
|
||||
self.run_workflow(**workflow_kwargs)
|
||||
|
||||
# Retrieve the result from the context after workflow execution
|
||||
result = self.context.get(RESULT)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
@ -39,17 +39,15 @@ class FrontendOperation(BaseWorkflow, BaseOperation):
|
|||
Returns:
|
||||
Any: The result obtained from executing the workflow.
|
||||
"""
|
||||
self.context.clear()
|
||||
|
||||
# Include the most recent messages in the operation context
|
||||
self.context[CHAT_MESSAGES] = self.chat_messages
|
||||
|
||||
# Add additional arguments to the context
|
||||
kwargs.update(**self.kwargs)
|
||||
self.context[CHAT_KWARGS] = kwargs
|
||||
# prepare kwargs
|
||||
workflow_kwargs = {
|
||||
CHAT_MESSAGES: self.chat_messages,
|
||||
CHAT_KWARGS: {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
# Execute the workflow with the prepared context
|
||||
self.run_workflow()
|
||||
self.run_workflow(**workflow_kwargs)
|
||||
|
||||
# Retrieve the result from the context after workflow execution
|
||||
return self.context.get(RESULT)
|
||||
82
memoryscope/core/service/base_memory_service.py
Normal file
82
memoryscope/core/service/base_memory_service.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List, Dict
|
||||
|
||||
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
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
"""
|
||||
An abstract base class for managing memory operations within a multithreaded context.
|
||||
It sets up the infrastructure for operation handling, message storage, and synchronization,
|
||||
along with logging capabilities and customizable configurations.
|
||||
"""
|
||||
|
||||
def __init__(self, memory_operations: Dict[str, dict], context: MemoryscopeContext, **kwargs):
|
||||
"""
|
||||
Initializes the BaseMemoryService with operation definitions, keys for memory access,
|
||||
and additional keyword arguments for flexibility.
|
||||
|
||||
Args:
|
||||
memory_operations (Dict[str, dict]): A dictionary defining available memory operations.
|
||||
**kwargs: Additional parameters to customize service behavior.
|
||||
"""
|
||||
self.memory_operations_conf: Dict[str, dict] = memory_operations
|
||||
self.context: MemoryscopeContext = context
|
||||
self.kwargs = kwargs
|
||||
|
||||
self._operation_dict: Dict[str, BaseOperation] = {}
|
||||
self.chat_messages: List[Message] = []
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
@property
|
||||
def op_description_dict(self) -> Dict[str, str]:
|
||||
"""
|
||||
Property to retrieve a dictionary mapping operation keys to their descriptions.
|
||||
Returns:
|
||||
Dict[str, str]: A dictionary where keys are operation identifiers and values are their descriptions.
|
||||
"""
|
||||
return {k: v.description for k, v in self._operation_dict.items()}
|
||||
|
||||
@abstractmethod
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def register_operation(self, name: str, operation_config: dict, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def init_service(self, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_backend_service(self, name: str = None):
|
||||
pass
|
||||
|
||||
def stop_backend_service(self, wait_service_end: bool = False):
|
||||
pass
|
||||
|
||||
def do_operation(self, name: str, **kwargs):
|
||||
"""
|
||||
Executes a specific operation by its name with provided keyword arguments.
|
||||
|
||||
Args:
|
||||
name (str): The name of the operation to execute.
|
||||
**kwargs: Keyword arguments for the operation's execution.
|
||||
|
||||
Returns:
|
||||
The result of the operation execution, if any. Otherwise, None.
|
||||
|
||||
Raises:
|
||||
Warning: If the operation name is not initialized in `_operation_dict`.
|
||||
"""
|
||||
if name not in self._operation_dict:
|
||||
self.logger.warning(f"operation={name} is not registered!")
|
||||
return
|
||||
return self._operation_dict[name].run_operation(**kwargs)
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
assert name in self._operation_dict, f"operation={name} is not registered!"
|
||||
return lambda **kwargs: self.do_operation(name=name, **kwargs)
|
||||
|
|
@ -1,9 +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):
|
||||
|
|
@ -11,6 +12,8 @@ class MemoryScopeService(BaseMemoryService):
|
|||
history_msg_count: int = 100,
|
||||
contextual_msg_max_count: int = 20,
|
||||
contextual_msg_min_count: int = 0,
|
||||
human_name: str = None,
|
||||
assistant_name: str = None,
|
||||
**kwargs):
|
||||
"""
|
||||
init function.
|
||||
|
|
@ -20,13 +23,21 @@ class MemoryScopeService(BaseMemoryService):
|
|||
it will not be included in the context to prevent token overflow.
|
||||
contextual_msg_min_count (int): The minimum context length in a conversation. If it is shorter than this
|
||||
length, no conversation summary will be made and no long-term memory will be generated.
|
||||
kwargs (dict): other kwargs
|
||||
human_name (str): human name.
|
||||
assistant_name (str): assistant name.
|
||||
kwargs (dict): other kwargs.
|
||||
"""
|
||||
super().__init__(**kwargs)
|
||||
self.history_msg_count: int = history_msg_count
|
||||
self.contextual_msg_max_count: int = contextual_msg_max_count
|
||||
self.contextual_msg_min_count: int = contextual_msg_min_count
|
||||
assert history_msg_count >= contextual_msg_max_count >= contextual_msg_min_count
|
||||
if human_name:
|
||||
self.context.meta_data["human_name"] = human_name
|
||||
if assistant_name:
|
||||
self.context.meta_data["assistant_name"] = assistant_name
|
||||
|
||||
self.message_lock = threading.Lock()
|
||||
|
||||
def add_messages(self, messages: List[Message] | Message):
|
||||
"""
|
||||
|
|
@ -54,60 +65,45 @@ class MemoryScopeService(BaseMemoryService):
|
|||
for _ in range(gap_size):
|
||||
self.chat_messages.pop(0)
|
||||
|
||||
def do_operation(self, op_name: str, **kwargs):
|
||||
"""
|
||||
Executes a specific operation by its name with provided keyword arguments.
|
||||
|
||||
Args:
|
||||
op_name (str): The name of the operation to execute.
|
||||
**kwargs: Keyword arguments for the operation's execution.
|
||||
|
||||
Returns:
|
||||
The result of the operation execution, if any. Otherwise, None.
|
||||
|
||||
Raises:
|
||||
Warning: If the operation name is not initialized in `_operation_dict`.
|
||||
"""
|
||||
if op_name not in self._operation_dict:
|
||||
self.logger.warning(f"op_name={op_name} is not inited!") # Warn if operation not initialized
|
||||
def register_operation(self, name: str, operation_config: dict, **kwargs):
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"op_name={name} is registered before!")
|
||||
return
|
||||
return self._operation_dict[op_name].run_operation(**kwargs) # Execute the operation
|
||||
|
||||
operation: BaseOperation = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
memoryscope_context=self.context,
|
||||
contextual_msg_max_count=self.contextual_msg_max_count,
|
||||
contextual_msg_min_count=self.contextual_msg_min_count)
|
||||
|
||||
# Initialize workflow for each operation
|
||||
operation.init_workflow(**kwargs)
|
||||
self._operation_dict[name] = operation
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
||||
def init_service(self, **kwargs):
|
||||
for name, operation_config in self.memory_operations.items():
|
||||
if name in self._operation_dict:
|
||||
self.logger.warning(f"memory operation={name} is repeated!")
|
||||
continue
|
||||
for name, operation_config in self.memory_operations_conf.items():
|
||||
self.register_operation(name, operation_config, **kwargs)
|
||||
|
||||
# ⭐ Initialize operation instance by its config
|
||||
operation: BaseOperation = init_instance_by_config(
|
||||
config=operation_config,
|
||||
name=name,
|
||||
chat_messages=self.chat_messages,
|
||||
message_lock=self.message_lock,
|
||||
contextual_msg_max_count=self.contextual_msg_max_count,
|
||||
contextual_msg_min_count=self.contextual_msg_min_count)
|
||||
operation.init_workflow(**kwargs) # Initialize workflow for each operation
|
||||
|
||||
self._operation_dict[name] = operation
|
||||
self.logger.info(f"service={self.__class__.__name__} init operation={name}")
|
||||
|
||||
def start_backend_service(self):
|
||||
def start_backend_service(self, name: str = None):
|
||||
"""
|
||||
Start all backend operations.
|
||||
"""
|
||||
for _, operation in self._operation_dict.items():
|
||||
if operation.operation_type == "backend":
|
||||
# Run backend operations
|
||||
operation.run_operation_backend()
|
||||
self.logger.info(f"start operation={operation.name}...")
|
||||
for op_name, operation in self._operation_dict.items():
|
||||
if name:
|
||||
if op_name == name:
|
||||
operation.start_operation_backend()
|
||||
else:
|
||||
if operation.operation_type == "backend":
|
||||
operation.start_operation_backend()
|
||||
|
||||
def stop_backend_service(self):
|
||||
def stop_backend_service(self, wait_service_end: bool = False):
|
||||
"""
|
||||
Stops all backend operations that are currently running.
|
||||
"""
|
||||
for _, operation in self._operation_dict.items():
|
||||
if operation.operation_type == "backend":
|
||||
# Stop backend operations
|
||||
operation.stop_operation_backend()
|
||||
self.logger.info(f"stop operation={operation.name}...")
|
||||
operation.stop_operation_backend(wait_task_end=wait_service_end)
|
||||
|
|
@ -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):
|
||||
|
|
@ -12,6 +12,17 @@ class DummyMemoryStore(BaseMemoryStore):
|
|||
semantic retrieval. Actual storage operations are not implemented.
|
||||
"""
|
||||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
"""
|
||||
Initializes the DummyMemoryStore with an embedding model and additional keyword arguments.
|
||||
|
||||
Args:
|
||||
embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval.
|
||||
**kwargs: Additional keyword arguments for configuration or future expansion.
|
||||
"""
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def retrieve_memories(self,
|
||||
query: str = "",
|
||||
top_k: int = 3,
|
||||
|
|
@ -24,17 +35,6 @@ class DummyMemoryStore(BaseMemoryStore):
|
|||
filter_dict: Dict[str, List[str]] = None) -> List[MemoryNode]:
|
||||
pass
|
||||
|
||||
def __init__(self, embedding_model: BaseModel, **kwargs):
|
||||
"""
|
||||
Initializes the DummyMemoryStore with an embedding model and additional keyword arguments.
|
||||
|
||||
Args:
|
||||
embedding_model (BaseModel): The model used to embed data for potential similarity-based retrieval.
|
||||
**kwargs: Additional keyword arguments for configuration or future expansion.
|
||||
"""
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
self.kwargs = kwargs
|
||||
|
||||
def batch_insert(self, nodes: List[MemoryNode]):
|
||||
pass
|
||||
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
from memoryscope.core.storage.base_monitor import BaseMonitor
|
||||
|
||||
|
||||
class DummyMonitor(BaseMonitor):
|
||||
|
|
@ -1,15 +1,17 @@
|
|||
import random
|
||||
from typing import Dict, List, Optional
|
||||
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):
|
||||
|
|
@ -19,15 +21,15 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
index_name: str,
|
||||
es_url: str,
|
||||
retrieve_mode: str = "dense",
|
||||
hybrid_alpha: float = None,
|
||||
hybrid_alpha: float = None,
|
||||
**kwargs):
|
||||
self.emb_dims = None
|
||||
self.index_name = index_name
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
retrieval_strategy = ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, hybrid_alpha=hybrid_alpha)
|
||||
self.es_store = SyncElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode,
|
||||
hybrid_alpha=hybrid_alpha),
|
||||
retrieval_strategy=retrieval_strategy,
|
||||
**kwargs)
|
||||
|
||||
# TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
@ -144,7 +153,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
embedding=embedding,
|
||||
metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"}))
|
||||
metadata=memory_node.model_dump(exclude={"content",
|
||||
"vector",
|
||||
"score_recall",
|
||||
"score_rank",
|
||||
"score_rerank"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
|
|
@ -38,6 +38,8 @@ DISTANCE_STRATEGIES = Literal[
|
|||
"EUCLIDEAN_DISTANCE",
|
||||
]
|
||||
|
||||
SPECIAL_QUERY: str = "**--**"
|
||||
|
||||
|
||||
def get_elasticsearch_client(
|
||||
url: Optional[str] = None,
|
||||
|
|
@ -133,17 +135,16 @@ 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":
|
||||
|
|
@ -152,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]:
|
||||
|
|
@ -160,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": {
|
||||
|
|
@ -268,7 +269,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
|
|||
}
|
||||
}
|
||||
)
|
||||
result['bool'].update({"should": operands}) # ⭐ Add 'should' clause for OR logic
|
||||
result['bool'].update({"should": operands}) # Add 'should' clause for OR logic
|
||||
result['bool'].update({"minimum_should_match": 1}) # Ensure at least one 'should' match
|
||||
else:
|
||||
key_str = f"metadata.{key}.keyword" if isinstance(value, str) else f"metadata.{key}"
|
||||
|
|
@ -613,7 +614,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
] = None,
|
||||
es_filter: Optional[List[Dict]] = None,
|
||||
fields: List[str] = [],
|
||||
**kwargs: Any,
|
||||
) -> VectorStoreQueryResult:
|
||||
"""
|
||||
Asynchronously queries the Elasticsearch index for the top k most similar nodes
|
||||
|
|
@ -626,6 +626,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
A custom function to modify the Elasticsearch query body. Defaults to None.
|
||||
es_filter (List[Dict], optional): Additional filters to apply during the query.
|
||||
If filters are present in the query, these filters will not be used. Defaults to None.
|
||||
fields (List[str], optional): .
|
||||
|
||||
Returns:
|
||||
VectorStoreQueryResult: The result of the query, including nodes, their IDs,
|
||||
|
|
@ -1,9 +1,10 @@
|
|||
import datetime
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import WEEKDAYS, DATATIME_WORD_LIST, MONTH_DICT
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
class DatetimeHandler(object):
|
||||
|
|
@ -40,7 +41,7 @@ class DatetimeHandler(object):
|
|||
|
||||
self._dt_info_dict: dict | None = None
|
||||
|
||||
def _parse_dt_info(self):
|
||||
def _parse_dt_info(self, language: LanguageEnum):
|
||||
"""
|
||||
Parses the datetime object (_dt) into a dictionary containing detailed date and time components,
|
||||
including language-specific weekday representation.
|
||||
|
|
@ -52,17 +53,16 @@ class DatetimeHandler(object):
|
|||
"""
|
||||
return {
|
||||
"year": self._dt.year,
|
||||
"month": MONTH_DICT[G_CONTEXT.language][self._dt.month - 1],
|
||||
"month": MONTH_DICT[language][self._dt.month - 1],
|
||||
"day": self._dt.day,
|
||||
"hour": self._dt.hour,
|
||||
"minute": self._dt.minute,
|
||||
"second": self._dt.second,
|
||||
"week": self._dt.isocalendar().week,
|
||||
"weekday": WEEKDAYS[G_CONTEXT.language][self._dt.isocalendar().weekday - 1],
|
||||
"weekday": WEEKDAYS[language][self._dt.isocalendar().weekday - 1],
|
||||
}
|
||||
|
||||
@property
|
||||
def dt_info_dict(self):
|
||||
def get_dt_info_dict(self, language: LanguageEnum):
|
||||
"""
|
||||
Property method to get the dictionary containing parsed datetime information.
|
||||
If None, initialize using `_parse_dt_info`.
|
||||
|
|
@ -71,7 +71,7 @@ class DatetimeHandler(object):
|
|||
dict: A dictionary with parsed datetime information.
|
||||
"""
|
||||
if self._dt_info_dict is None:
|
||||
self._dt_info_dict = self._parse_dt_info()
|
||||
self._dt_info_dict = self._parse_dt_info(language=language)
|
||||
return self._dt_info_dict
|
||||
|
||||
@classmethod
|
||||
|
|
@ -207,7 +207,7 @@ class DatetimeHandler(object):
|
|||
return date_info
|
||||
|
||||
@classmethod
|
||||
def extract_date_parts(cls, input_string: str) -> dict:
|
||||
def extract_date_parts(cls, input_string: str, language: LanguageEnum) -> dict:
|
||||
"""
|
||||
Extracts various date components from the input string based on the current language context.
|
||||
|
||||
|
|
@ -217,48 +217,51 @@ class DatetimeHandler(object):
|
|||
|
||||
Args:
|
||||
input_string (str): The string containing date information to be parsed.
|
||||
language (str): current language.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing extracted date components, or an empty dictionary if parsing fails.
|
||||
"""
|
||||
func_name = f"extract_date_parts_{G_CONTEXT.language.value}"
|
||||
func_name = f"extract_date_parts_{language.value}"
|
||||
if not hasattr(cls, func_name):
|
||||
cls.logger.warning(f"language={G_CONTEXT.language.value} 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)
|
||||
|
||||
@classmethod
|
||||
def has_time_word_cn(cls, query: str) -> bool:
|
||||
def has_time_word_cn(cls, query: str, datetime_word_list: List[str]) -> bool:
|
||||
"""
|
||||
Check if the input query contains any datetime-related words based on the cn language context.
|
||||
|
||||
Args:
|
||||
query (str): The input string to check for datetime-related words.
|
||||
datetime_word_list (list[str]): datetime keywords
|
||||
|
||||
Returns:
|
||||
bool: True if the query contains at least one datetime-related word, False otherwise.
|
||||
"""
|
||||
contain_datetime = False
|
||||
# TODO use re
|
||||
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
|
||||
for datetime_word in datetime_word_list:
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
return contain_datetime
|
||||
|
||||
@classmethod
|
||||
def has_time_word_en(cls, query: str) -> bool:
|
||||
def has_time_word_en(cls, query: str, datetime_word_list: List[str]) -> bool:
|
||||
"""
|
||||
Check if the input query contains any datetime-related words based on the en language context.
|
||||
|
||||
Args:
|
||||
query (str): The input string to check for datetime-related words.
|
||||
datetime_word_list (list[str]): datetime keywords
|
||||
|
||||
Returns:
|
||||
bool: True if the query contains at least one datetime-related word, False otherwise.
|
||||
"""
|
||||
contain_datetime = False
|
||||
for datetime_word in DATATIME_WORD_LIST[G_CONTEXT.language]:
|
||||
for datetime_word in datetime_word_list:
|
||||
datetime_word = datetime_word.lower()
|
||||
# TODO fix strip
|
||||
if datetime_word in [x.strip().lower().strip(",").strip(".").strip("?").strip(":")
|
||||
|
|
@ -268,12 +271,18 @@ class DatetimeHandler(object):
|
|||
return contain_datetime
|
||||
|
||||
@classmethod
|
||||
def has_time_word(cls, query: str) -> bool:
|
||||
func_name = f"has_time_word_{G_CONTEXT.language.value}"
|
||||
def has_time_word(cls, query: str, language: LanguageEnum) -> bool:
|
||||
func_name = f"has_time_word_{language.value}"
|
||||
if not hasattr(cls, func_name):
|
||||
cls.logger.warning(f"language={G_CONTEXT.language.value} needs to complete has_time_word func!")
|
||||
cls.logger.warning(f"language={language.value} needs to complete has_time_word function!")
|
||||
return False
|
||||
return getattr(cls, func_name)(query=query)
|
||||
|
||||
if language not 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]
|
||||
return getattr(cls, func_name)(query=query, datetime_word_list=datetime_word_list)
|
||||
|
||||
def datetime_format(self, dt_format: str = "%Y%m%d") -> str:
|
||||
"""
|
||||
|
|
@ -287,17 +296,18 @@ class DatetimeHandler(object):
|
|||
"""
|
||||
return self._dt.strftime(dt_format)
|
||||
|
||||
def string_format(self, string_format: str) -> str:
|
||||
def string_format(self, string_format: str, language: LanguageEnum) -> str:
|
||||
"""
|
||||
Format the datetime information stored in the instance using a custom string format.
|
||||
|
||||
Args:
|
||||
string_format (str): A format string where placeholders are keys from `dt_info_dict`.
|
||||
language (str): current language.
|
||||
|
||||
Returns:
|
||||
str: A formatted datetime string.
|
||||
"""
|
||||
return string_format.format(**self.dt_info_dict)
|
||||
return string_format.format(**self.get_dt_info_dict(language=language))
|
||||
|
||||
@property
|
||||
def timestamp(self) -> int:
|
||||
|
|
@ -26,7 +26,7 @@ class Logger(logging.Logger):
|
|||
max_bytes: int = 1024 * 1024 * 1024,
|
||||
backup_count: int = 10):
|
||||
"""
|
||||
Initializes the Logger instance, setting up handlers for console and/or file logging based on provided parameters.
|
||||
Initializes the Logger instance, setting up handlers for console and file logging based on provided parameters.
|
||||
|
||||
Args:
|
||||
name (str): Identifier for the logger.
|
||||
|
|
@ -105,7 +105,8 @@ class Logger(logging.Logger):
|
|||
by the handlers are freed properly.
|
||||
"""
|
||||
for handler in self.handlers:
|
||||
handler.close() # ⭐ Close each handler to release resources
|
||||
# Close each handler to release resources
|
||||
handler.close()
|
||||
|
||||
def clear(self):
|
||||
"""
|
||||
|
|
@ -1,19 +1,25 @@
|
|||
import json
|
||||
import os.path
|
||||
from pathlib import Path
|
||||
from typing import Dict
|
||||
|
||||
import yaml
|
||||
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
class PromptHandler(object):
|
||||
"""
|
||||
The `PromptHandler` class manages prompt messages by loading them from YAML or JSON files and dictionaries,
|
||||
supporting language selection based on a global context, and providing dictionary-like access to the prompt messages.
|
||||
supporting language selection based on a context, and providing dictionary-like access to the prompt messages.
|
||||
"""
|
||||
|
||||
def __init__(self, class_path: str, prompt_file: str = "", prompt_dict: dict = None, **kwargs):
|
||||
def __init__(self,
|
||||
class_path: str,
|
||||
language: LanguageEnum | str,
|
||||
prompt_file: str = "",
|
||||
prompt_dict: dict = None,
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the PromptHandler with paths to prompt sources and additional keyword arguments.
|
||||
|
||||
|
|
@ -21,30 +27,32 @@ class PromptHandler(object):
|
|||
class_path (str): The path to the class where prompts are utilized.
|
||||
prompt_file (str, optional): The path to an external file containing prompts. Defaults to "".
|
||||
prompt_dict (dict, optional): A dictionary directly containing prompt definitions. Defaults to None.
|
||||
language (LanguageEnum, str): context language.
|
||||
**kwargs: Additional keyword arguments that might be used in prompt handling.
|
||||
"""
|
||||
self._class_path: str = class_path
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
class_path: Path = Path(class_path)
|
||||
self._class_dir: Path = class_path.parent
|
||||
self._class_name: str = class_path.stem
|
||||
self._language_enum: LanguageEnum = LanguageEnum(language)
|
||||
self.kwargs = kwargs
|
||||
|
||||
file_path = self._class_path.strip(".py")
|
||||
|
||||
self.add_prompt_file(file_path)
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
|
||||
self.add_prompt_file((self._class_dir / self._class_name).__str__(), raise_exception=False)
|
||||
if prompt_file:
|
||||
self.add_prompt_file(prompt_file)
|
||||
|
||||
self.add_prompt_file((self._class_dir / prompt_file).__str__())
|
||||
if prompt_dict:
|
||||
self.add_prompt_dict(prompt_dict)
|
||||
|
||||
@staticmethod
|
||||
def file_path_completion(file_path: str) -> str:
|
||||
def file_path_completion(file_path: str, raise_exception: bool = True) -> str:
|
||||
"""
|
||||
Attempts to complete the given file path by appending either a `.yaml` or `.json` extension
|
||||
based on the existence of the respective file. If neither exists, an exception is raised.
|
||||
|
||||
Args:
|
||||
file_path (str): The base path of the file to be completed.
|
||||
raise_exception (bool): If the file cannot be found, report an error.
|
||||
|
||||
Returns:
|
||||
str: The completed file path with the appropriate extension.
|
||||
|
|
@ -61,9 +69,10 @@ class PromptHandler(object):
|
|||
if os.path.exists(f"{file_path}.json"):
|
||||
return f"{file_path}.json"
|
||||
|
||||
raise RuntimeError(f"{file_path}/yaml/json is not exists!")
|
||||
if raise_exception:
|
||||
raise RuntimeError(f"{file_path}/yaml/json is not exists!")
|
||||
|
||||
def add_prompt_file(self, file_path: str):
|
||||
def add_prompt_file(self, file_path: str, raise_exception: bool = True):
|
||||
"""
|
||||
Adds prompt messages from a YAML or JSON file to the internal dictionary.
|
||||
|
||||
|
|
@ -72,8 +81,11 @@ class PromptHandler(object):
|
|||
|
||||
Args:
|
||||
file_path (str): The path to the YAML or JSON file containing the prompts.
|
||||
raise_exception (bool): If the file cannot be found, report an error.
|
||||
"""
|
||||
file_path = self.file_path_completion(file_path)
|
||||
file_path = self.file_path_completion(file_path, raise_exception=raise_exception)
|
||||
if not file_path:
|
||||
return
|
||||
|
||||
prompt_dict = {}
|
||||
|
||||
|
|
@ -102,9 +114,9 @@ class PromptHandler(object):
|
|||
RuntimeError: If a prompt message for the current language is not found.
|
||||
"""
|
||||
for key, language_dict in prompt_dict.items():
|
||||
prompts = language_dict.get(G_CONTEXT.language)
|
||||
prompts = language_dict.get(self._language_enum.value)
|
||||
if not prompts:
|
||||
raise RuntimeError(f"{key}.prompt.{G_CONTEXT.language} is empty!")
|
||||
raise RuntimeError(f"{key}.prompt.{self._language_enum.value} is empty!")
|
||||
self._prompt_dict[key] = prompts.strip()
|
||||
|
||||
@property
|
||||
|
|
@ -12,7 +12,8 @@ class Registry(object):
|
|||
|
||||
Attributes:
|
||||
name (str): The name of the registry.
|
||||
module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are the modules themselves.
|
||||
module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are
|
||||
the modules themselves.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str):
|
||||
|
|
@ -31,7 +32,7 @@ class Registry(object):
|
|||
|
||||
Args:
|
||||
module_name (str): The name of module to be registered.
|
||||
modules (List[Any] | Dict[str, Any]): The module to be registered.
|
||||
module (List[Any] | Dict[str, Any]): The module to be registered.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the input is already registered.
|
||||
|
|
@ -46,7 +47,8 @@ class Registry(object):
|
|||
|
||||
def batch_register(self, modules: List[Any] | Dict[str, Any]):
|
||||
"""
|
||||
Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary mapping names to modules.
|
||||
Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary
|
||||
mapping names to modules.
|
||||
|
||||
Args:
|
||||
modules (List[Any] | Dict[str, Any]): A list of modules or a dictionary mapping module names to the modules.
|
||||
|
|
@ -2,28 +2,22 @@ import re
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import NONE_WORD
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
"""
|
||||
The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract specific
|
||||
The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract
|
||||
patterns from the text and filter out unnecessary information, while also logging the processing steps and outcomes.
|
||||
"""
|
||||
|
||||
pattern_v1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets
|
||||
|
||||
def __init__(self, response_text: str, logger_prefix: str = ""):
|
||||
"""
|
||||
Initializes the `ResponseTextParser` instance with the provided response text and sets up a logger.
|
||||
|
||||
Args:
|
||||
response_text (str): The raw response text that needs to be parsed and processed.
|
||||
"""
|
||||
PATTERN_V1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets
|
||||
|
||||
def __init__(self, response_text: str, language: LanguageEnum, logger_prefix: str = ""):
|
||||
# Strips leading and trailing whitespace from the response text
|
||||
self.response_text: str = response_text.strip()
|
||||
self.language: LanguageEnum = language
|
||||
|
||||
# The prefix of log. Defaults to "".
|
||||
self.logger_prefix: str = logger_prefix
|
||||
|
|
@ -43,7 +37,7 @@ class ResponseTextParser(object):
|
|||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
matches = [match.group(1) for match in self.pattern_v1.finditer(line)]
|
||||
matches = [match.group(1) for match in self.PATTERN_V1.finditer(line)]
|
||||
if matches:
|
||||
result.append(matches)
|
||||
self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
|
|
@ -51,18 +45,15 @@ class ResponseTextParser(object):
|
|||
|
||||
def parse_v2(self) -> List[str]:
|
||||
"""
|
||||
Extract lines which contain NONE_WORD in Chinese or English.
|
||||
Extract lines which contain NONE_WORD.
|
||||
|
||||
Args:
|
||||
prefix (str): The prefix of log. Defaults to "".
|
||||
|
||||
Returns:
|
||||
Contents match the specific patterns.
|
||||
"""
|
||||
result = []
|
||||
for line in self.response_text.split("\n"):
|
||||
line = line.strip()
|
||||
if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language):
|
||||
if not line or line.lower() == NONE_WORD.get(self.language):
|
||||
continue
|
||||
result.append(line)
|
||||
self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
|
|
@ -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"]
|
||||
|
||||
|
|
@ -26,7 +26,7 @@ class Timer(object):
|
|||
Args:
|
||||
name (str): The log name.
|
||||
time_log_type (str): The log type. Defaults to 'End'.
|
||||
use_ms (bool): Use 'ms' as the time scale or not. Defaults to True.
|
||||
use_ms (bool): Use 'ms' as the timescale or not. Defaults to True.
|
||||
stack_level (int): The stack level of log. Defaults to 2.
|
||||
float_precision (int): The precision of cost time. Defaults to 4.
|
||||
|
||||
|
|
@ -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.
|
||||
"""
|
||||
|
|
@ -6,6 +6,7 @@ from copy import deepcopy
|
|||
from importlib import import_module
|
||||
from typing import List
|
||||
|
||||
import numpy as np
|
||||
import pyfiglet
|
||||
from termcolor import colored
|
||||
|
||||
|
|
@ -18,7 +19,7 @@ ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey",
|
|||
|
||||
def underscore_to_camelcase(name: str, is_first_title: bool = True) -> str:
|
||||
"""
|
||||
Converts a underscore_notation string to CamelCase.
|
||||
Converts an underscore_notation string to CamelCase.
|
||||
|
||||
Args:
|
||||
name (str): The underscore_notation string to be converted.
|
||||
|
|
@ -47,10 +48,7 @@ def camelcase_to_underscore(name: str) -> str:
|
|||
return re.sub(r'(?<!^)(?=[A-Z])', '_', name).lower()
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict,
|
||||
default_class_path: str = "memoryscope",
|
||||
suffix_name: str = "",
|
||||
**kwargs):
|
||||
def init_instance_by_config(config: dict, default_class_dir: str = "memoryscope", **kwargs):
|
||||
"""
|
||||
Initialize an instance of a class specified in the configuration dictionary.
|
||||
|
||||
|
|
@ -62,12 +60,9 @@ def init_instance_by_config(config: dict,
|
|||
Args:
|
||||
config (dict): A dictionary containing the configuration, including
|
||||
the 'class' key that specifies the class's module path.
|
||||
default_class_path (str, optional): The default module path prefix
|
||||
default_class_dir (str, optional): The default module path prefix
|
||||
to use if not explicitly defined in
|
||||
'config'. Defaults to "memory_scope".
|
||||
suffix_name (str, optional): A string to append to the class name,
|
||||
ensuring the final class name ends with it.
|
||||
Defaults to "".
|
||||
**kwargs: Additional keyword arguments to pass to the class constructor.
|
||||
|
||||
Returns:
|
||||
|
|
@ -80,21 +75,22 @@ def init_instance_by_config(config: dict,
|
|||
raise RuntimeError("empty class path!")
|
||||
user_defined: bool = config_copy.pop("user_defined", False)
|
||||
|
||||
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
|
||||
class_path_list = []
|
||||
if not user_defined and default_class_dir and not origin_class_path.startswith(default_class_dir):
|
||||
class_path_list.append(default_class_dir)
|
||||
|
||||
class_paths = []
|
||||
if not user_defined and 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))
|
||||
class_path_split = origin_class_path.split(".")
|
||||
class_file_name: str = class_path_split[-1]
|
||||
|
||||
cls_name = underscore_to_camelcase(class_name)
|
||||
class_name = underscore_to_camelcase(class_file_name)
|
||||
if class_name == class_file_name:
|
||||
class_path_list.extend(class_path_split[-1:])
|
||||
else:
|
||||
class_path_list.extend(class_path_split)
|
||||
|
||||
module = import_module(".".join(class_path_list))
|
||||
config_copy.update(kwargs)
|
||||
return getattr(module, cls_name)(**config_copy)
|
||||
return getattr(module, class_name)(**config_copy)
|
||||
|
||||
|
||||
def prompt_to_msg(system_prompt: str,
|
||||
|
|
@ -193,3 +189,21 @@ def contains_keyword(text, keywords) -> bool:
|
|||
escaped_keywords = map(re.escape, keywords)
|
||||
pattern = re.compile('|'.join(escaped_keywords), re.IGNORECASE)
|
||||
return pattern.search(text) is not None
|
||||
|
||||
|
||||
def cosine_similarity(query: List[float], documents: List[List[float]]):
|
||||
query = np.array(query)
|
||||
documents = np.array(documents)
|
||||
|
||||
query_norm = np.linalg.norm(query)
|
||||
if query_norm == 0:
|
||||
raise ValueError("Query vector norm is zero, which will result in a division by zero")
|
||||
|
||||
documents_norm = np.linalg.norm(documents, axis=1)
|
||||
if np.any(documents_norm == 0):
|
||||
raise ValueError("One of the document vectors has zero norm, which will result in a division by zero")
|
||||
|
||||
dot_product = np.dot(documents, query)
|
||||
|
||||
cosine_similarities = dot_product / (query_norm * documents_norm)
|
||||
return cosine_similarities.tolist()
|
||||
|
|
@ -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):
|
||||
|
|
@ -43,13 +43,13 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
6. Updates the status of nodes accordingly.
|
||||
7. Persists the changes back to memory storage.
|
||||
"""
|
||||
all_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
|
||||
all_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories([NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES])
|
||||
if not all_obs_nodes:
|
||||
self.logger.info("all_obs_nodes is empty!")
|
||||
# self.continue_run = False
|
||||
return
|
||||
|
||||
today_obs_nodes: List[MemoryNode] = self.memory_handler.get_memories(TODAY_NODES)
|
||||
today_obs_nodes: List[MemoryNode] = self.memory_manager.get_memories(TODAY_NODES)
|
||||
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
|
|
@ -80,7 +80,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
idx_merge_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.logger.warning("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
|
@ -121,4 +121,4 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
merge_obs_nodes.append(node)
|
||||
|
||||
# save context
|
||||
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)
|
||||
self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes, log_repeat=False)
|
||||
|
|
@ -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):
|
||||
|
|
@ -26,7 +26,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
filter_messages = []
|
||||
for msg in self.chat_messages:
|
||||
# Checks if the message content has any time reference words
|
||||
if DatetimeHandler.has_time_word(query=msg.content):
|
||||
if DatetimeHandler.has_time_word(query=msg.content, language=self.language):
|
||||
filter_messages.append(msg)
|
||||
return filter_messages
|
||||
|
||||
|
|
@ -49,7 +49,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
for i, msg in enumerate(filter_messages):
|
||||
# Create a DatetimeHandler instance for each message's timestamp and format it
|
||||
dt_handler = DatetimeHandler(dt=msg.time_created)
|
||||
dt = dt_handler.string_format(self.prompt_handler.time_string_format)
|
||||
dt = dt_handler.string_format(string_format=self.prompt_handler.time_string_format, language=self.language)
|
||||
# Append formatted timestamp-query pairs to the user_query_list
|
||||
user_query_list.append(f"{i + 1} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
|
||||
|
||||
|
|
@ -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):
|
||||
|
|
@ -42,11 +42,11 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
MemoryTypeEnum.CONVERSATION.value: message.content,
|
||||
TIME_INFER: time_infer,
|
||||
"keywords": keywords,
|
||||
**{k: str(v) for k, v in dt_handler.dt_info_dict.items()},
|
||||
**{k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()},
|
||||
}
|
||||
|
||||
if time_infer:
|
||||
dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer)
|
||||
dt_info_dict = DatetimeHandler.extract_date_parts(input_string=time_infer, language=self.language)
|
||||
meta_data.update({f"event_{k}": str(v) for k, v in dt_info_dict.items()})
|
||||
obs_content = (f"{obs_content} ({self.get_language_value(TIME_INFER_WORD)}"
|
||||
f"{self.get_language_value(COLON_WORD)} {time_infer})")
|
||||
|
|
@ -68,7 +68,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
"""
|
||||
filter_messages = []
|
||||
for msg in self.chat_messages:
|
||||
if not DatetimeHandler.has_time_word(query=msg.content):
|
||||
if not DatetimeHandler.has_time_word(query=msg.content, language=self.language):
|
||||
filter_messages.append(msg)
|
||||
|
||||
self.logger.info(f"after filter_messages.size from {len(self.chat_messages)} to {len(filter_messages)}")
|
||||
|
|
@ -139,7 +139,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# Parses the generated text to extract observation indices, times, contents, and keywords
|
||||
idx_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
idx_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.logger.warning("idx_obs_list is empty!")
|
||||
return
|
||||
|
|
@ -184,4 +184,4 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
keywords=keywords))
|
||||
|
||||
# Stores the extracted and structured observations in the conversation memory
|
||||
self.memory_handler.set_memories(self.OBS_STORE_KEY, new_obs_nodes)
|
||||
self.memory_manager.set_memories(self.OBS_STORE_KEY, new_obs_nodes)
|
||||
|
|
@ -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):
|
||||
|
|
@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
"""
|
||||
dt_handler = DatetimeHandler()
|
||||
# Prepare metadata with current datetime info
|
||||
meta_data = {k: str(v) for k, v in dt_handler.dt_info_dict.items()}
|
||||
meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}
|
||||
|
||||
return MemoryNode(user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
|
|
@ -58,8 +58,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
- Parsing the model's responses for new insight keys.
|
||||
- Creating new insight nodes and updating the memory status accordingly.
|
||||
"""
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_REFLECTED_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES)
|
||||
|
||||
# Count unaudited nodes
|
||||
not_reflected_count = len(not_reflected_nodes)
|
||||
|
|
@ -101,10 +101,11 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# Parse LLM response for new insight keys and update memory
|
||||
new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2()
|
||||
new_insight_keys = ResponseTextParser(response.message.content, self.language,
|
||||
self.__class__.__name__).parse_v2()
|
||||
if new_insight_keys:
|
||||
for insight_key in new_insight_keys:
|
||||
self.memory_handler.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
|
||||
self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
|
||||
|
||||
# Mark unaudited nodes as reflected
|
||||
for node in not_reflected_nodes:
|
||||
|
|
@ -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):
|
||||
|
|
@ -76,7 +76,7 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
info_score_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}")
|
||||
|
||||
|
|
@ -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):
|
||||
|
|
@ -33,7 +33,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_reflected_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
self.memory_manager.set_memories(NOT_REFLECTED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_not_updated_memory(self):
|
||||
|
|
@ -52,7 +52,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_not_updated_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
self.memory_manager.set_memories(NOT_UPDATED_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_insight_memory(self):
|
||||
|
|
@ -70,7 +70,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
}
|
||||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_insight_top_k,
|
||||
filter_dict=filter_dict)
|
||||
self.memory_handler.set_memories(INSIGHT_NODES, nodes)
|
||||
self.memory_manager.set_memories(INSIGHT_NODES, nodes)
|
||||
|
||||
@timer
|
||||
def retrieve_today_memory(self, dt: str):
|
||||
|
|
@ -93,7 +93,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
|
|||
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(top_k=self.retrieve_today_top_k,
|
||||
filter_dict=filter_dict)
|
||||
|
||||
self.memory_handler.set_memories(TODAY_NODES, nodes)
|
||||
self.memory_manager.set_memories(TODAY_NODES, nodes)
|
||||
|
||||
def _run(self):
|
||||
"""
|
||||
|
|
@ -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):
|
||||
|
|
@ -63,7 +63,7 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
|
||||
The process helps in maintaining conversation coherence by resolving contradictions and redundancies.
|
||||
"""
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES)
|
||||
for node in not_updated_nodes:
|
||||
self.submit_thread_task(fn=self.retrieve_similar_content, node=node)
|
||||
|
||||
|
|
@ -111,7 +111,8 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# Parses the model's response text to identify updates for memory nodes
|
||||
idx_obs_info_list = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v1()
|
||||
idx_obs_info_list = ResponseTextParser(response.message.content, self.language,
|
||||
self.__class__.__name__).parse_v1()
|
||||
if len(idx_obs_info_list) <= 0:
|
||||
self.logger.warning("idx_obs_info_list is empty!")
|
||||
return
|
||||
|
|
@ -157,4 +158,4 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
f"action_status={node.action_status}")
|
||||
|
||||
# save context
|
||||
self.memory_handler.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
|
||||
self.memory_manager.set_memories(MERGE_OBS_NODES, merge_obs_nodes)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
|
@ -27,13 +27,15 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
obs_nodes: List[MemoryNode],
|
||||
use_dummy_ranker: bool) -> (MemoryNode, List[MemoryNode], float):
|
||||
"""
|
||||
Filters observed nodes based on their relevance to a given insight node using a ranking model.
|
||||
|
||||
Args:
|
||||
insight_node (MemoryNode): The insight node used as the basis for filtering.
|
||||
obs_nodes (List[MemoryNode]): A list of observed nodes to be filtered.
|
||||
use_dummy_ranker (bool): Global parameters, whether to use rank model or not.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing:
|
||||
|
|
@ -53,24 +55,48 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.warning("obs_nodes is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Call the ranking model to get scores for each observed node's content against the insight key
|
||||
documents = [x.content for x in obs_nodes]
|
||||
self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}")
|
||||
response = self.rank_model.call(query=insight_node.key, documents=documents)
|
||||
if not response.status:
|
||||
return insight_node, filtered_nodes, max_score
|
||||
if use_dummy_ranker:
|
||||
if not insight_node.key_vector:
|
||||
key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results
|
||||
if not key_vector:
|
||||
self.logger.warning(f"embedding call {insight_node.key} failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Iterate over the ranked scores to filter nodes
|
||||
for index, score in response.rank_scores.items():
|
||||
node = obs_nodes[index]
|
||||
# Determine if the node should be kept based on the threshold
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
insight_node.key_vector = key_vector
|
||||
|
||||
score_recall_list = cosine_similarity(insight_node.key_vector, [x.vector for x in obs_nodes])
|
||||
assert len(score_recall_list) == len(obs_nodes), \
|
||||
f"size is not as excepted. {len(score_recall_list)} v.s. {len(obs_nodes)}"
|
||||
|
||||
for score, node in zip(score_recall_list, obs_nodes):
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
else:
|
||||
# Call the ranking model to get scores for each observed node's content against the insight key
|
||||
documents = [x.content for x in obs_nodes]
|
||||
self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}")
|
||||
response = self.rank_model.call(query=insight_node.key, documents=documents)
|
||||
if not response.status:
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Iterate over the ranked scores to filter nodes
|
||||
for index, score in response.rank_scores.items():
|
||||
node = obs_nodes[index]
|
||||
# Determine if the node should be kept based on the threshold
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
# Warn if no nodes were filtered
|
||||
if not filtered_nodes:
|
||||
|
|
@ -95,7 +121,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}"
|
||||
insight_node.content = content
|
||||
insight_node.value = insight_value
|
||||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
|
||||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()})
|
||||
insight_node.timestamp = dt_handler.timestamp
|
||||
insight_node.dt = dt_handler.datetime_format()
|
||||
if insight_node.action_status == ActionStatusEnum.NONE.value:
|
||||
|
|
@ -136,7 +162,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
if not response.status or not response.message.content:
|
||||
return insight_node
|
||||
|
||||
insight_value_list = ResponseTextParser(response.message.content,
|
||||
insight_value_list = ResponseTextParser(response.message.content, self.language,
|
||||
f"update_{insight_node.key}").parse_v1()
|
||||
if not insight_value_list:
|
||||
self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!")
|
||||
|
|
@ -175,26 +201,30 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
6. Gather the results of all update tasks.
|
||||
7. Mark processed nodes as updated in memory.
|
||||
"""
|
||||
insight_nodes: List[MemoryNode] = self.memory_handler.get_memories(INSIGHT_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_handler.get_memories(NOT_UPDATED_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_handler.get_memories(keys=[NOT_REFLECTED_NODES,
|
||||
insight_nodes: List[MemoryNode] = self.memory_manager.get_memories(INSIGHT_NODES)
|
||||
not_updated_nodes: List[MemoryNode] = self.memory_manager.get_memories(NOT_UPDATED_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.memory_manager.get_memories(keys=[NOT_REFLECTED_NODES,
|
||||
NOT_UPDATED_NODES])
|
||||
|
||||
if not insight_nodes:
|
||||
self.logger.warning("insight_nodes is empty, stopping processing.")
|
||||
return
|
||||
|
||||
use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"]
|
||||
|
||||
# Process active insight nodes with corresponding not updated nodes
|
||||
for node in insight_nodes:
|
||||
time.sleep(1)
|
||||
if node.action_status == ActionStatusEnum.NEW.value:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
obs_nodes=not_reflected_nodes)
|
||||
obs_nodes=not_reflected_nodes,
|
||||
use_dummy_ranker=use_dummy_ranker)
|
||||
else:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
obs_nodes=not_updated_nodes)
|
||||
obs_nodes=not_updated_nodes,
|
||||
use_dummy_ranker=use_dummy_ranker)
|
||||
|
||||
# select top n
|
||||
result_list = []
|
||||
|
|
@ -216,12 +246,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# delete empty nodes
|
||||
empty_nodes = [n for n in insight_nodes if not n.content.strip()]
|
||||
self.memory_handler.delete_memories(empty_nodes)
|
||||
self.memory_manager.delete_memories(empty_nodes)
|
||||
|
||||
for node in not_updated_nodes:
|
||||
node.obs_updated = 1
|
||||
node.action_status = ActionStatusEnum.MODIFIED
|
||||
|
||||
# for node in not_reflected_nodes:
|
||||
# node.obs_updated = 1
|
||||
# node.action_status = ActionStatusEnum.MODIFIED
|
||||
|
|
@ -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):
|
||||
|
|
@ -46,7 +46,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
|
|||
if not self.memory_key:
|
||||
return
|
||||
|
||||
return self.memory_handler.get_memories(keys=self.memory_key)
|
||||
return self.memory_manager.get_memories(keys=self.memory_key)
|
||||
|
||||
def delete_all(self):
|
||||
"""
|
||||
|
|
@ -55,7 +55,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
|
|||
Returns:
|
||||
List[MemoryNode]: A list of all MemoryNode objects marked for deletion.
|
||||
"""
|
||||
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
|
||||
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
|
||||
for node in nodes:
|
||||
node.action_status = ActionStatusEnum.DELETE.value
|
||||
self.logger.info(f"delete_all.size={len(nodes)}")
|
||||
|
|
@ -74,7 +74,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
i = 0
|
||||
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
|
||||
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
|
||||
for node in nodes:
|
||||
if node.content == query:
|
||||
i += 1
|
||||
|
|
@ -88,7 +88,7 @@ class UpdateMemoryWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
i = 0
|
||||
nodes: List[MemoryNode] = self.memory_handler.get_memories(keys="all")
|
||||
nodes: List[MemoryNode] = self.memory_manager.get_memories(keys="all")
|
||||
for node in nodes:
|
||||
if node.memory_id == memory_id:
|
||||
i += 1
|
||||
|
|
@ -109,4 +109,4 @@ class UpdateMemoryWorker(MemoryBaseWorker):
|
|||
if not hasattr(self, method):
|
||||
self.logger.info(f"method={method} is missing!")
|
||||
return
|
||||
self.memory_handler.update_memories(nodes=getattr(self, method)())
|
||||
self.memory_manager.update_memories(nodes=getattr(self, method)())
|
||||
|
|
@ -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):
|
||||
0
memoryscope/core/worker/frontend/__init__.py
Normal file
0
memoryscope/core/worker/frontend/__init__.py
Normal file
|
|
@ -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):
|
||||
|
|
@ -33,13 +33,14 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
query, query_timestamp = self.get_context(QUERY_WITH_TS)
|
||||
|
||||
# Identify if the query contains datetime keywords
|
||||
contain_datetime = DatetimeHandler.has_time_word(query)
|
||||
contain_datetime = DatetimeHandler.has_time_word(query, self.language)
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# Prepare the prompt with necessary contextual details
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format)
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format,
|
||||
self.language)
|
||||
system_prompt = self.prompt_handler.extract_time_system
|
||||
few_shot = self.prompt_handler.extract_time_few_shot
|
||||
user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str)
|
||||
|
|
@ -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):
|
||||
|
|
@ -62,7 +62,7 @@ class FuseRerankWorker(MemoryBaseWorker):
|
|||
"""
|
||||
# Parse input parameters from the worker's context
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RANKED_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RANKED_MEMORY_NODES)
|
||||
|
||||
# Check if memory nodes are available; warn and return if not
|
||||
if not memory_node_list:
|
||||
|
|
@ -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):
|
||||
|
|
@ -22,7 +22,7 @@ class PrintMemoryWorker(MemoryBaseWorker):
|
|||
3. Set the formatted string back into the worker's context
|
||||
"""
|
||||
# get long-term memory
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list = sorted(memory_node_list, key=lambda x: x.timestamp, reverse=True)
|
||||
|
||||
observation_memory_list: List[str] = []
|
||||
|
|
@ -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_handler.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)
|
||||
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
|
||||
|
||||
|
||||
|
|
@ -29,25 +29,33 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
"""
|
||||
# query
|
||||
query, _ = self.get_context(QUERY_WITH_TS)
|
||||
memory_node_list: List[MemoryNode] = self.memory_handler.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
memory_node_list: List[MemoryNode] = self.memory_manager.get_memories(RETRIEVE_MEMORY_NODES)
|
||||
if not memory_node_list:
|
||||
self.logger.warning("Retrieve memory nodes is empty!")
|
||||
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())
|
||||
use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"]
|
||||
if use_dummy_ranker:
|
||||
for node in memory_node_list:
|
||||
node.score_rank = node.score_recall
|
||||
self.logger.warning("use score_recall instead of score_rank!")
|
||||
|
||||
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
|
||||
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())
|
||||
|
||||
# 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
|
||||
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
|
||||
|
||||
# sort by score
|
||||
memory_node_list = sorted(memory_node_list, key=lambda n: n.score_rank, reverse=True)
|
||||
|
|
@ -58,4 +66,4 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"Rank stage: Content={node.content}, Score={node.score_rank}")
|
||||
|
||||
# save ranked nodes back to memory
|
||||
self.memory_handler.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False)
|
||||
self.memory_manager.set_memories(RANKED_MEMORY_NODES, memory_node_list, log_repeat=False)
|
||||
|
|
@ -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):
|
||||
|
|
@ -22,22 +22,39 @@ class SetQueryWorker(MemoryBaseWorker):
|
|||
along with its creation timestamp.
|
||||
"""
|
||||
query = "" # Default query value
|
||||
query_timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
|
||||
timestamp = int(datetime.datetime.now().timestamp()) # Current timestamp as default
|
||||
|
||||
if "query" in self.chat_kwargs:
|
||||
# Check if a specific 'query' has been provided via chat kwargs
|
||||
# set query if exists
|
||||
query = self.chat_kwargs["query"]
|
||||
if not query:
|
||||
query = ""
|
||||
query = query.strip()
|
||||
|
||||
# set ts if exists
|
||||
_timestamp = self.chat_kwargs.get("timestamp")
|
||||
if _timestamp and isinstance(_timestamp, int):
|
||||
timestamp = _timestamp
|
||||
|
||||
# check role_name
|
||||
role_name = self.chat_kwargs.get("role_name")
|
||||
if role_name:
|
||||
assert role_name == self.target_name, \
|
||||
f"role_name={role_name} is not supported in human/assistant memory workflow!"
|
||||
|
||||
elif self.chat_messages:
|
||||
# If no explicit query is given, use the content of the latest chat message
|
||||
chat_messages = [msg for msg in self.chat_messages if msg.role == MessageRoleEnum.USER.value]
|
||||
if chat_messages:
|
||||
message = chat_messages[-1]
|
||||
query = message.content
|
||||
query_timestamp = message.time_created
|
||||
timestamp = message.time_created
|
||||
|
||||
# check role_name
|
||||
role_name = message.role_name
|
||||
if role_name:
|
||||
assert role_name == self.target_name, \
|
||||
f"role_name={role_name} is not supported in human/assistant memory workflow!"
|
||||
|
||||
# Store the determined query and its timestamp in the context
|
||||
self.set_context(QUERY_WITH_TS, (query, query_timestamp))
|
||||
self.set_context(QUERY_WITH_TS, (query, timestamp))
|
||||
|
|
@ -1,15 +1,17 @@
|
|||
from abc import ABCMeta
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from memoryscope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS, MEMORY_HANDLER
|
||||
from memoryscope.memory.worker.base_worker import BaseWorker
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
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.scheme.message import Message
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.memory_handler import MemoryHandler
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
|
||||
|
||||
class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
||||
|
|
@ -77,6 +79,18 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
"""
|
||||
return self.get_context(CHAT_KWARGS)
|
||||
|
||||
@property
|
||||
def workflow_name(self) -> str:
|
||||
return self.get_context(WORKFLOW_NAME)
|
||||
|
||||
@property
|
||||
def memoryscope_context(self) -> MemoryscopeContext:
|
||||
return self.get_context(MEMORYSCOPE_CONTEXT)
|
||||
|
||||
@property
|
||||
def language(self) -> LanguageEnum:
|
||||
return self.memoryscope_context.language
|
||||
|
||||
@property
|
||||
def embedding_model(self) -> BaseModel:
|
||||
"""
|
||||
|
|
@ -87,8 +101,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The embedding model used for converting text into vector representations.
|
||||
"""
|
||||
if isinstance(self._embedding_model, str):
|
||||
self._embedding_model = G_CONTEXT.model_dict[self._embedding_model]
|
||||
# ⭐ Retrieve the actual model instance when the attribute is a string reference
|
||||
self._embedding_model = self.memoryscope_context.model_dict[self._embedding_model]
|
||||
return self._embedding_model
|
||||
|
||||
@property
|
||||
|
|
@ -101,8 +114,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The model used for text generation.
|
||||
"""
|
||||
if isinstance(self._generation_model, str):
|
||||
self._generation_model = G_CONTEXT.model_dict[self._generation_model]
|
||||
# ⭐ Retrieve the model instance if currently a string reference
|
||||
self._generation_model = self.memoryscope_context.model_dict[self._generation_model]
|
||||
return self._generation_model
|
||||
|
||||
@property
|
||||
|
|
@ -115,7 +127,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseModel: The rank model instance used for ranking tasks.
|
||||
"""
|
||||
if isinstance(self._rank_model, str):
|
||||
self._rank_model = G_CONTEXT.model_dict[self._rank_model] # Fetch model instance if string reference
|
||||
self._rank_model = self.memoryscope_context.model_dict[self._rank_model]
|
||||
return self._rank_model
|
||||
|
||||
@property
|
||||
|
|
@ -128,7 +140,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseMemoryStore: The memory store instance used for inserting, updating, retrieving and deleting operations.
|
||||
"""
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
self._memory_store = self.memoryscope_context.memory_store
|
||||
return self._memory_store
|
||||
|
||||
@property
|
||||
|
|
@ -141,7 +153,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
BaseMonitor: The monitoring component instance.
|
||||
"""
|
||||
if self._monitor is None:
|
||||
self._monitor = G_CONTEXT.monitor
|
||||
self._monitor = self.memoryscope_context.monitor
|
||||
return self._monitor
|
||||
|
||||
@property
|
||||
|
|
@ -154,7 +166,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
str: The name of the assistant.
|
||||
"""
|
||||
if self._user_name is None:
|
||||
self._user_name = G_CONTEXT.meta_data["assistant_name"]
|
||||
self._user_name = self.memoryscope_context.meta_data["assistant_name"]
|
||||
return self._user_name
|
||||
|
||||
@property
|
||||
|
|
@ -166,7 +178,7 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
str: The readable name of the human.
|
||||
"""
|
||||
if self._target_name is None:
|
||||
self._target_name = G_CONTEXT.meta_data["human_name"]
|
||||
self._target_name = self.memoryscope_context.meta_data["human_name"]
|
||||
return self._target_name
|
||||
|
||||
@property
|
||||
|
|
@ -178,23 +190,22 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
PromptHandler: An instance of PromptHandler initialized with specific file path and keyword arguments.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(self.FILE_PATH, **self.kwargs)
|
||||
self._prompt_handler = PromptHandler(self.FILE_PATH, language=self.language, **self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
@property
|
||||
def memory_handler(self) -> MemoryHandler:
|
||||
def memory_manager(self) -> MemoryManager:
|
||||
"""
|
||||
Lazily initializes and returns the MemoryHandler instance.
|
||||
|
||||
Returns:
|
||||
MemoryHandler: An instance of MemoryHandler.
|
||||
"""
|
||||
if not self.has_content(MEMORY_HANDLER):
|
||||
self.set_context(MEMORY_HANDLER, MemoryHandler()) # Initialize the memory handler if not present
|
||||
return self.get_context(MEMORY_HANDLER)
|
||||
if not self.has_content(MEMORY_MANAGER):
|
||||
self.set_context(MEMORY_MANAGER, MemoryManager(self.memoryscope_context))
|
||||
return self.get_context(MEMORY_MANAGER)
|
||||
|
||||
@staticmethod
|
||||
def get_language_value(languages: dict | List[dict]) -> Any | List[Any]:
|
||||
def get_language_value(self, languages: dict | List[dict]) -> Any | List[Any]:
|
||||
"""
|
||||
Retrieves the value(s) corresponding to the current language context.
|
||||
|
||||
|
|
@ -205,5 +216,5 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
|
|||
Any | list[Any]: The value or list of values matching the current language setting.
|
||||
"""
|
||||
if isinstance(languages, list):
|
||||
return [x[G_CONTEXT.language] for x in languages]
|
||||
return languages[G_CONTEXT.language]
|
||||
return [x[self.language] for x in languages]
|
||||
return languages[self.language]
|
||||
|
|
@ -1,22 +1,21 @@
|
|||
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.scheme.memory_node import MemoryNode
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class MemoryHandler(object):
|
||||
class MemoryManager(object):
|
||||
"""
|
||||
The `MemoryHandler` class manages memory nodes with memory store.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initializes the MemoryHandler.
|
||||
"""
|
||||
def __init__(self, memoryscope_context: MemoryscopeContext):
|
||||
self.memoryscope_context: MemoryscopeContext = memoryscope_context
|
||||
|
||||
self._memory_store: BaseMemoryStore | None = None
|
||||
|
||||
# dict: memory_id -> MemoryNode
|
||||
|
|
@ -36,7 +35,7 @@ class MemoryHandler(object):
|
|||
BaseMemoryStore: The memory store instance associated with this worker.
|
||||
"""
|
||||
if self._memory_store is None:
|
||||
self._memory_store = G_CONTEXT.memory_store
|
||||
self._memory_store = self.memoryscope_context.memory_store
|
||||
return self._memory_store
|
||||
|
||||
def clear(self):
|
||||
|
|
@ -7,8 +7,9 @@ class ModelEnum(str, Enum):
|
|||
|
||||
Members:
|
||||
GENERATION_MODEL: Represents a model responsible for generating content.
|
||||
EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a numerical form suitable for machine learning tasks.
|
||||
RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance or importance.
|
||||
EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a
|
||||
numerical form suitable for machine learning tasks.
|
||||
RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance.
|
||||
"""
|
||||
GENERATION_MODEL = "generation_model"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,106 +0,0 @@
|
|||
import threading
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from typing import List, Dict
|
||||
|
||||
from memoryscope.memory.operation.base_operation import BaseOperation
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class BaseMemoryService(metaclass=ABCMeta):
|
||||
"""
|
||||
An abstract base class for managing memory operations within a multi-threaded context.
|
||||
It sets up the infrastructure for operation handling, message storage, and synchronization,
|
||||
along with logging capabilities and customizable configurations.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
memory_operations: Dict[str, dict],
|
||||
retrieve_memory_key: str = "retrieve_memory",
|
||||
read_message_key: str = "read_message",
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the BaseMemoryService with operation definitions, keys for memory access,
|
||||
and additional keyword arguments for flexibility.
|
||||
|
||||
Args:
|
||||
memory_operations (Dict[str, dict]): A dictionary defining available memory operations.
|
||||
retrieve_memory_key (str): The key indicating a retrieve memory operation. Defaults to "retrieve_memory".
|
||||
read_message_key (str): The key for reading messages. Defaults to "read_message".
|
||||
**kwargs: Additional parameters to customize service behavior.
|
||||
"""
|
||||
self.memory_operations: Dict[str, dict] = memory_operations
|
||||
self.retrieve_memory_key: str = retrieve_memory_key
|
||||
self.read_message_key: str = read_message_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 add_messages(self, messages: List[Message] | Message):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def do_operation(self, op_name: str, **kwargs):
|
||||
"""
|
||||
Abstract method defining the interface for executing a specific operation by its name.
|
||||
This method must be implemented by subclasses to provide the actual operation logic.
|
||||
|
||||
Args:
|
||||
op_name (str): The name identifying the operation to be performed.
|
||||
**kwargs: Additional keyword arguments required for the operation execution.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: This exception is raised when the method is not overridden in a subclass.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def op_description_dict(self) -> Dict[str, str]:
|
||||
"""
|
||||
Property to retrieve a dictionary mapping operation keys to their descriptions.
|
||||
Lazily initializes the dictionary on first access.
|
||||
|
||||
Returns:
|
||||
Dict[str, str]: A dictionary where keys are operation identifiers and values are their descriptions.
|
||||
"""
|
||||
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 retrieve_memory(self):
|
||||
"""
|
||||
Executes the operation associated with retrieved memory.
|
||||
Asserts that the operation for retrieved memory has been initialized.
|
||||
|
||||
Returns:
|
||||
Any: The result of the retrieved memory operation.
|
||||
"""
|
||||
assert self.retrieve_memory_key in self._operation_dict, f"op={self.retrieve_memory_key} is not inited!"
|
||||
return self.do_operation(self.retrieve_memory_key)
|
||||
|
||||
def read_message(self):
|
||||
"""
|
||||
Executes the operation associated with reading messages.
|
||||
Asserts that the operation for reading messages has been initialized.
|
||||
|
||||
Returns:
|
||||
Any: The result of the read message operation.
|
||||
"""
|
||||
assert self.read_message_key in self._operation_dict, f"op={self.read_message_key} is not inited!"
|
||||
return self.do_operation(self.read_message_key)
|
||||
|
||||
@abstractmethod
|
||||
def init_service(self, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_backend_service(self):
|
||||
pass
|
||||
|
||||
def stop_backend_service(self):
|
||||
pass
|
||||
|
|
@ -23,6 +23,8 @@ class MemoryNode(BaseModel):
|
|||
|
||||
key: str = Field("", description="memory key")
|
||||
|
||||
key_vector: List[float] = Field([], description="memory key embedding result")
|
||||
|
||||
value: str = Field("", description="memory value")
|
||||
|
||||
score_recall: float = Field(0, description="embedding similarity score used in recall stage")
|
||||
|
|
@ -37,7 +39,7 @@ class MemoryNode(BaseModel):
|
|||
|
||||
store_status: str = Field("valid", description="store_status: valid / expired")
|
||||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
vector: List[float] = Field([], description="content embedding result")
|
||||
|
||||
timestamp: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()),
|
||||
description="timestamp of the memory node")
|
||||
|
|
|
|||
|
|
@ -1,33 +0,0 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, Any
|
||||
|
||||
from memoryscope.chat.base_memory_chat import BaseMemoryChat
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.memory.service.base_memory_service import BaseMemoryService
|
||||
from memoryscope.models.base_model import BaseModel
|
||||
from memoryscope.storage.base_memory_store import BaseMemoryStore
|
||||
from memoryscope.storage.base_monitor import BaseMonitor
|
||||
|
||||
|
||||
class GlobalContext(object):
|
||||
"""
|
||||
The GlobalContext class archives all configs utilized by store, monitor, services and workers.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
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.memory_store: BaseMemoryStore | None = None
|
||||
self.monitor: BaseMonitor | None = None
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.language: LanguageEnum = LanguageEnum.EN
|
||||
|
||||
self.meta_data: Dict[str, Any] = {}
|
||||
|
||||
|
||||
G_CONTEXT = GlobalContext()
|
||||
|
|
@ -5,8 +5,8 @@ sys.path.append(".") # noqa: E402
|
|||
import asyncio
|
||||
import unittest
|
||||
|
||||
from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
|
||||
|
||||
class TestLLIEmbedding(unittest.TestCase):
|
||||
|
|
|
|||
|
|
@ -6,8 +6,8 @@ import unittest
|
|||
import time
|
||||
import asyncio
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.models.llama_index_generation_model import LlamaIndexGenerationModel
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.core.models.llama_index_generation_model import LlamaIndexGenerationModel
|
||||
from memoryscope.core.utils.logger import Logger
|
||||
|
||||
|
||||
class TestLLILLM(unittest.TestCase):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import unittest
|
||||
|
||||
from memoryscope.models.llama_index_rank_model import LlamaIndexRankModel
|
||||
from memoryscope.core.models.llama_index_rank_model import LlamaIndexRankModel
|
||||
|
||||
|
||||
class TestLLIReRank(unittest.TestCase):
|
||||
|
|
|
|||
|
|
@ -1,18 +0,0 @@
|
|||
from memoryscope.cli import MemoryScope
|
||||
from memoryscope.scheme.message import Message
|
||||
|
||||
ms = MemoryScope().load_config("config/demo_config_no_stream.yaml")
|
||||
memory_service = ms.default_service
|
||||
memory_chat = ms.default_chat_handle
|
||||
|
||||
# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。")
|
||||
# memory_service.add_messages(new_message)
|
||||
|
||||
res: Message = memory_chat.chat_with_memory(query="我的爱好是弹琴。", remember_response=True)
|
||||
print(res.message.content)
|
||||
|
||||
res: Message = memory_chat.chat_with_memory(query="昨天弹出一个光粒,消灭了星系0x4be。", remember_response=True)
|
||||
print(res.message.content)
|
||||
|
||||
res: Message = memory_chat.chat_with_memory(query="今天弹出一个二向箔,消灭了星系0xa2e。", remember_response=True)
|
||||
print(res.message.content)
|
||||
3
tests/other/read_prompt.yaml
Normal file
3
tests/other/read_prompt.yaml
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
a:
|
||||
cn: c
|
||||
en: e
|
||||
11
tests/other/read_yaml.py
Normal file
11
tests/other/read_yaml.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
import sys
|
||||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
from memoryscope.core.utils.prompt_handler import PromptHandler
|
||||
|
||||
if __name__ == "__main__":
|
||||
file_path: str = __file__
|
||||
print(file_path)
|
||||
handler = PromptHandler(__file__, language="cn", prompt_file="read_prompt", )
|
||||
print(handler.prompt_dict)
|
||||
15
tests/other/test_attr.py
Normal file
15
tests/other/test_attr.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
class MyClass:
|
||||
def __init__(self):
|
||||
self.existing_attribute = "I exist"
|
||||
|
||||
def do(self, name: str, **kwargs):
|
||||
print("do %s %s" % (name, kwargs))
|
||||
|
||||
def __getattr__(self, name):
|
||||
return lambda **kwargs: self.do(name, **kwargs)
|
||||
|
||||
|
||||
# 创建类的实例
|
||||
obj = MyClass()
|
||||
|
||||
obj.haha(a=1, b=2)
|
||||
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)
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
import unittest
|
||||
|
||||
from memoryscope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memoryscope.core.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
|
||||
from memoryscope.core.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
|
||||
|
||||
|
||||
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue