mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-13 23:11:03 +00:00
fix test file import problem
This commit is contained in:
parent
70ecdb1d94
commit
2237a6abaf
20 changed files with 175 additions and 178 deletions
0
examples/api/__init__.py
Normal file
0
examples/api/__init__.py
Normal file
|
|
@ -8,7 +8,6 @@ from memoryscope.core.memoryscope import MemoryScope
|
|||
|
||||
|
||||
def cli_job(**kwargs):
|
||||
kwargs["memory_chat_type"] = "cli_chat"
|
||||
MemoryScope(**kwargs).default_memory_chat.run()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -53,7 +53,10 @@ class ApiMemoryChat(BaseMemoryChat):
|
|||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs)
|
||||
self._prompt_handler = PromptHandler(__file__,
|
||||
language=self.context.language,
|
||||
prompt_file="memory_chat_prompt",
|
||||
**self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -68,7 +68,10 @@ class CliMemoryChat(BaseMemoryChat):
|
|||
PromptHandler: An instance of the PromptHandler configured for this CLI session.
|
||||
"""
|
||||
if self._prompt_handler is None:
|
||||
self._prompt_handler = PromptHandler(__file__, prompt_file="memory_chat_prompt", **self.kwargs)
|
||||
self._prompt_handler = PromptHandler(__file__,
|
||||
language=self.context.language,
|
||||
prompt_file="memory_chat_prompt",
|
||||
**self.kwargs)
|
||||
return self._prompt_handler
|
||||
|
||||
def print_logo(self):
|
||||
|
|
|
|||
|
|
@ -12,8 +12,8 @@ class Arguments(object):
|
|||
|
||||
logger_name_time_suffix: str = field(default="%Y%m%d_%H%M%S")
|
||||
|
||||
memory_chat_type: str = field(default="cli_chat", metadata={
|
||||
"help": "cli_chat(Command-line interaction), api_chat(API interface interaction), etc."})
|
||||
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."})
|
||||
|
|
|
|||
|
|
@ -65,14 +65,9 @@ class ConfigManager(object):
|
|||
|
||||
@staticmethod
|
||||
def update_memory_chat_by_arguments(config: dict, arguments: Arguments):
|
||||
if arguments.memory_chat_type == "cli_chat":
|
||||
memory_chat_class = "chat.cli_memory_chat"
|
||||
elif arguments.memory_chat_type == "api_chat":
|
||||
memory_chat_class = "chat.api_memory_chat"
|
||||
else:
|
||||
raise NotImplementedError(f"known memory_chat_type={arguments.memory_chat_type}")
|
||||
memory_chat_class_split = config["class"].split(".")
|
||||
config.update({
|
||||
"class": memory_chat_class,
|
||||
"class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]),
|
||||
"human_name": DEFAULT_HUMAN_NAME[LanguageEnum(arguments.language)],
|
||||
"assistant_name": "AI",
|
||||
})
|
||||
|
|
@ -154,7 +149,7 @@ class ConfigManager(object):
|
|||
def clear_node_all(self, node: str):
|
||||
self.config[node].clear()
|
||||
|
||||
def dump_config(self, file_type: Literal["json", "yaml"], to_stream: bool = True, file_path: Optional[str] = None):
|
||||
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":
|
||||
|
|
@ -162,9 +157,8 @@ class ConfigManager(object):
|
|||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if to_stream:
|
||||
print(content)
|
||||
|
||||
if file_type:
|
||||
if file_path:
|
||||
with open(file_path, "w") as f:
|
||||
f.write(content)
|
||||
|
||||
return content
|
||||
|
|
|
|||
|
|
@ -7,78 +7,78 @@ global:
|
|||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
class: core.chat.cli_memory_chat
|
||||
memory_service: memoryscope_service
|
||||
generation_model: generation_model
|
||||
|
||||
memory_service:
|
||||
memoryscope_service:
|
||||
class: memory.service.memory_scope_service
|
||||
class: core.service.memory_scope_service
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: memory.operation.frontend_operation
|
||||
class: core.operation.frontend_operation
|
||||
workflow: read_message
|
||||
description: "read short memory"
|
||||
|
||||
retrieve_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
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: memory.operation.frontend_operation
|
||||
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: memory.operation.frontend_operation
|
||||
class: core.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
|
||||
class: core.operation.frontend_operation
|
||||
workflow: set_query,retrieve_all_memory,delete_all
|
||||
description: "delete all long-term memory"
|
||||
|
||||
add_memory:
|
||||
class: memory.operation.frontend_operation
|
||||
class: core.operation.frontend_operation
|
||||
workflow: add_memory
|
||||
description: "add a single observation"
|
||||
|
||||
consolidate_memory:
|
||||
class: memory.operation.consolidate_memory_op
|
||||
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: memory.operation.backend_operation
|
||||
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: memory.worker.dummy_worker
|
||||
class: core.worker.dummy_worker
|
||||
generation_model: generation_model
|
||||
embedding_model: embedding_model
|
||||
rank_model: rank_model
|
||||
read_message:
|
||||
class: memory.worker.frontend.read_message_worker
|
||||
class: core.worker.frontend.read_message_worker
|
||||
set_query:
|
||||
class: memory.worker.frontend.set_query_worker
|
||||
class: core.worker.frontend.set_query_worker
|
||||
retrieve_obs_ins:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
class: core.worker.frontend.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_top_k: 100
|
||||
extract_time:
|
||||
class: memory.worker.frontend.extract_time_worker
|
||||
class: core.worker.frontend.extract_time_worker
|
||||
generation_model: generation_model
|
||||
semantic_rank:
|
||||
class: memory.worker.frontend.semantic_rank_worker
|
||||
class: core.worker.frontend.semantic_rank_worker
|
||||
rank_model: rank_model
|
||||
fuse_rerank:
|
||||
class: memory.worker.frontend.fuse_rerank_worker
|
||||
class: core.worker.frontend.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.01
|
||||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
|
|
@ -88,88 +88,88 @@ worker:
|
|||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
retrieve_top_memory:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
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: memory.worker.frontend.print_memory_worker
|
||||
class: core.worker.frontend.print_memory_worker
|
||||
retrieve_all_memory:
|
||||
class: memory.worker.frontend.retrieve_memory_worker
|
||||
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: memory.worker.backend.update_memory_worker
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: delete_memory
|
||||
delete_all:
|
||||
class: memory.worker.backend.update_memory_worker
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: delete_all
|
||||
add_memory:
|
||||
class: memory.worker.backend.update_memory_worker
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: from_query
|
||||
info_filter:
|
||||
class: memory.worker.backend.info_filter_worker
|
||||
class: core.worker.backend.info_filter_worker
|
||||
generation_model: generation_model
|
||||
load_today_memory:
|
||||
class: memory.worker.backend.load_memory_worker
|
||||
class: core.worker.backend.load_memory_worker
|
||||
retrieve_today_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.backend.get_observation_worker
|
||||
class: core.worker.backend.get_observation_worker
|
||||
generation_model: generation_model
|
||||
get_observation_with_time:
|
||||
class: memory.worker.backend.get_observation_with_time_worker
|
||||
class: core.worker.backend.get_observation_with_time_worker
|
||||
generation_model: generation_model
|
||||
contra_repeat:
|
||||
class: memory.worker.backend.contra_repeat_worker
|
||||
class: core.worker.backend.contra_repeat_worker
|
||||
generation_model: generation_model
|
||||
store_memory:
|
||||
class: memory.worker.backend.update_memory_worker
|
||||
class: core.worker.backend.update_memory_worker
|
||||
method: from_memory_key
|
||||
memory_key: all
|
||||
load_obs_and_insight:
|
||||
class: memory.worker.backend.load_memory_worker
|
||||
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: memory.worker.backend.get_reflection_subject_worker
|
||||
class: core.worker.backend.get_reflection_subject_worker
|
||||
generation_model: generation_model
|
||||
reflect_obs_cnt_threshold: 10
|
||||
update_insight:
|
||||
class: memory.worker.backend.update_insight_worker
|
||||
class: core.worker.backend.update_insight_worker
|
||||
generation_model: generation_model
|
||||
rank_model: rank_model
|
||||
long_contra_repeat:
|
||||
class: memory.worker.backend.long_contra_repeat_worker
|
||||
class: core.worker.backend.long_contra_repeat_worker
|
||||
generation_model: generation_model
|
||||
|
||||
model:
|
||||
generation_model:
|
||||
class: models.llama_index_generation_model
|
||||
class: core.models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
max_tokens: 2000
|
||||
embedding_model:
|
||||
class: models.llama_index_embedding_model
|
||||
class: core.models.llama_index_embedding_model
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
rank_model:
|
||||
class: models.llama_index_rank_model
|
||||
class: core.models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
top_n: 500
|
||||
dummy_generation:
|
||||
class: models.dummy_generation_model
|
||||
class: core.models.dummy_generation_model
|
||||
module_name: dummy_generation
|
||||
model_name: dummy_generation_model
|
||||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_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: storage.dummy_monitor
|
||||
class: core.storage.dummy_monitor
|
||||
|
|
@ -82,6 +82,8 @@ class MemoryScope(ConfigManager):
|
|||
if self.context.monitor:
|
||||
self.context.monitor.close()
|
||||
|
||||
self.logger.close()
|
||||
|
||||
def __enter__(self):
|
||||
self.init_context_by_config()
|
||||
|
||||
|
|
|
|||
|
|
@ -16,9 +16,9 @@ class PromptHandler(object):
|
|||
|
||||
def __init__(self,
|
||||
class_path: str,
|
||||
language: LanguageEnum | str,
|
||||
prompt_file: str = "",
|
||||
prompt_dict: dict = None,
|
||||
language_enum: LanguageEnum = LanguageEnum.EN,
|
||||
**kwargs):
|
||||
"""
|
||||
Initializes the PromptHandler with paths to prompt sources and additional keyword arguments.
|
||||
|
|
@ -27,13 +27,13 @@ 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_enum (LanguageEnum): context language.
|
||||
language (LanguageEnum, str): context language.
|
||||
**kwargs: Additional keyword arguments that might be used in prompt handling.
|
||||
"""
|
||||
class_path: Path = Path(class_path)
|
||||
self._class_dir: Path = class_path.parent
|
||||
self._class_name: str = class_path.stem
|
||||
self._language_enum: LanguageEnum = language_enum
|
||||
self._language_enum: LanguageEnum = LanguageEnum(language)
|
||||
self.kwargs = kwargs
|
||||
|
||||
self._prompt_dict: Dict[str, str] = {}
|
||||
|
|
|
|||
|
|
@ -190,7 +190,7 @@ 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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -2,10 +2,10 @@ import sys
|
|||
|
||||
sys.path.append(".") # noqa: E402
|
||||
|
||||
from memoryscope.utils.prompt_handler import PromptHandler
|
||||
from memoryscope.core.utils.prompt_handler import PromptHandler
|
||||
|
||||
if __name__ == "__main__":
|
||||
file_path: str = __file__
|
||||
print(file_path)
|
||||
handler = PromptHandler(__file__, "read_prompt")
|
||||
handler = PromptHandler(__file__, language="cn", prompt_file="read_prompt", )
|
||||
print(handler.prompt_dict)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -1,44 +1,51 @@
|
|||
import datetime
|
||||
import unittest
|
||||
|
||||
from memoryscope.cli import MemoryScope
|
||||
from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \
|
||||
MEMORYSCOPE_CONTEXT
|
||||
from memoryscope.core.config.arguments import Arguments
|
||||
from memoryscope.core.memoryscope import MemoryScope
|
||||
from memoryscope.core.utils.tool_functions import init_instance_by_config
|
||||
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.memory_node import MemoryNode
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class TestWorkersCn(unittest.TestCase):
|
||||
"""Tests for LLIEmbedding"""
|
||||
|
||||
def setUp(self):
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
|
||||
|
||||
ms = MemoryScope()
|
||||
ms.read_config("config/demo_config_cn.yaml")
|
||||
ms.init_global_content_by_config()
|
||||
arguments = Arguments(
|
||||
language="cn",
|
||||
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",
|
||||
)
|
||||
self.ms = MemoryScope(arguments=arguments)
|
||||
config = self.ms.dump_config()
|
||||
self.ms.logger.info(f"config=\n{config}")
|
||||
|
||||
def tearDown(self):
|
||||
self.logger.close()
|
||||
self.ms.close()
|
||||
|
||||
@unittest.skip
|
||||
def test_extract_time(self):
|
||||
name = "extract_time"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
query = "明天我去上海出差"
|
||||
query_timestamp = int(datetime.datetime.now().timestamp())
|
||||
|
|
@ -53,13 +60,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "info_filter"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
|
||||
|
|
@ -79,13 +85,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "info_filter"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗"),
|
||||
|
|
@ -115,13 +120,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "get_observation"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜"),
|
||||
|
|
@ -149,13 +153,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "get_observation"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。"),
|
||||
|
|
@ -181,13 +184,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "get_observation_with_time"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术"),
|
||||
|
|
@ -211,13 +213,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "contra_repeat"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"),
|
||||
|
|
@ -272,13 +273,12 @@ class TestWorkersCn(unittest.TestCase):
|
|||
name = "get_reflection_subject"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
|
||||
|
|
@ -304,18 +304,17 @@ class TestWorkersCn(unittest.TestCase):
|
|||
|
||||
@unittest.skip
|
||||
def test_update_insight_worker(self):
|
||||
reflection_worker = self.test_get_reflection_subject.__wrapped__(self)
|
||||
reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self)
|
||||
|
||||
name = "update_insight"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context=reflection_worker.context,
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="用户喜欢打王者荣耀"),
|
||||
|
|
@ -327,18 +326,17 @@ class TestWorkersCn(unittest.TestCase):
|
|||
result = "\n".join(result)
|
||||
worker.logger.info(f"result.update_insight={result}")
|
||||
|
||||
@unittest.skip
|
||||
# @unittest.skip
|
||||
def test_long_contra_repeat_worker(self):
|
||||
name = "long_contra_repeat"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。"),
|
||||
|
|
|
|||
|
|
@ -1,44 +1,51 @@
|
|||
import datetime
|
||||
import unittest
|
||||
|
||||
from memoryscope.cli import MemoryScope
|
||||
from memoryscope.constants.common_constants import CHAT_MESSAGES, NEW_OBS_NODES, NEW_OBS_WITH_TIME_NODES, \
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES
|
||||
MERGE_OBS_NODES, QUERY_WITH_TS, EXTRACT_TIME_DICT, NOT_REFLECTED_NODES, INSIGHT_NODES, NOT_UPDATED_NODES, \
|
||||
MEMORYSCOPE_CONTEXT
|
||||
from memoryscope.core.config.arguments import Arguments
|
||||
from memoryscope.core.memoryscope import MemoryScope
|
||||
from memoryscope.core.utils.tool_functions import init_instance_by_config
|
||||
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.memory_node import MemoryNode
|
||||
from memoryscope.scheme.message import Message
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.utils.logger import Logger
|
||||
from memoryscope.utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class TestWorkersEn(unittest.TestCase):
|
||||
"""Tests for LLIEmbedding"""
|
||||
|
||||
def setUp(self):
|
||||
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
self.logger: Logger = Logger.get_logger(f"test_worker_{datetime_suffix}", to_stream=True)
|
||||
|
||||
ms = MemoryScope()
|
||||
ms.read_config("config/demo_config_en.yaml")
|
||||
ms.init_global_content_by_config()
|
||||
arguments = Arguments(
|
||||
language="en",
|
||||
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",
|
||||
)
|
||||
self.ms = MemoryScope(arguments=arguments)
|
||||
config = self.ms.dump_config()
|
||||
self.ms.logger.info(f"config=\n{config}")
|
||||
|
||||
def tearDown(self):
|
||||
self.logger.close()
|
||||
self.ms.close()
|
||||
|
||||
@unittest.skip
|
||||
def test_extract_time(self):
|
||||
name = "extract_time"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
query = "I will be on a business trip to Shanghai tomorrow."
|
||||
query_timestamp = int(datetime.datetime.now().timestamp())
|
||||
|
|
@ -53,13 +60,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "info_filter"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."),
|
||||
|
|
@ -80,13 +86,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "info_filter"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value, content="Do you know where the freshest seafood is in Beijing?"),
|
||||
|
|
@ -129,13 +134,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "get_observation"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
# FIXME Does the appearance of 'am' indicate the presence of a time keyword?
|
||||
chat_messages = [
|
||||
|
|
@ -159,13 +163,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "get_observation"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value,
|
||||
|
|
@ -202,13 +205,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "get_observation_with_time"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
chat_messages = [
|
||||
Message(role=MessageRoleEnum.USER.value,
|
||||
|
|
@ -235,13 +237,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "contra_repeat"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"),
|
||||
|
|
@ -294,13 +295,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "get_reflection_subject"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),
|
||||
|
|
@ -327,18 +327,17 @@ class TestWorkersEn(unittest.TestCase):
|
|||
|
||||
@unittest.skip
|
||||
def test_update_insight_worker(self):
|
||||
reflection_worker = self.test_get_reflection_subject.__wrapped__(self)
|
||||
reflection_worker: MemoryBaseWorker = self.test_get_reflection_subject.__wrapped__(self)
|
||||
|
||||
name = "update_insight"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context=reflection_worker.context,
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="Users like to play King of Glory"),
|
||||
|
|
@ -355,13 +354,12 @@ class TestWorkersEn(unittest.TestCase):
|
|||
name = "long_contra_repeat"
|
||||
|
||||
worker: MemoryBaseWorker = init_instance_by_config(
|
||||
config=G_CONTEXT.worker_config[name],
|
||||
suffix_name="worker",
|
||||
config=self.ms.context.worker_conf_dict[name],
|
||||
name=name,
|
||||
is_multi_thread=False,
|
||||
context={},
|
||||
context={MEMORYSCOPE_CONTEXT: self.ms.context},
|
||||
context_lock=None,
|
||||
thread_pool=G_CONTEXT.thread_pool)
|
||||
thread_pool=self.ms.context.thread_pool)
|
||||
|
||||
nodes = [
|
||||
MemoryNode(content="Users are interested in strategy games and looking for new challenges."),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue