diff --git a/examples/api/agentscope_example.py b/examples/api/agentscope_example.py index 563ae2d0..f4dd59e6 100644 --- a/examples/api/agentscope_example.py +++ b/examples/api/agentscope_example.py @@ -1,11 +1,12 @@ from typing import Optional, Union, Sequence + import agentscope -import sys -import os from agentscope.agents import AgentBase, UserAgent from agentscope.message import Msg + from memoryscope import MemoryScope, Arguments + class MemoryScopeAgent(AgentBase): def __init__(self, name: str, arguments: Arguments, **kwargs) -> None: # Disable AgentScope memory and use MemoryScope memory instead @@ -13,7 +14,6 @@ class MemoryScopeAgent(AgentBase): # Create a memory client in MemoryScope self.memory_scope = MemoryScope(arguments=arguments) - self.memory_scope.init_context_by_config() self.memory_chat = self.memory_scope.default_memory_chat def reply(self, x: Optional[Union[Msg, Sequence[Msg]]] = None) -> Msg: @@ -46,7 +46,6 @@ def main(): 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" ) @@ -79,4 +78,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/examples/api/autogen_example.py b/examples/api/autogen_example.py index 326de416..ba0cef60 100644 --- a/examples/api/autogen_example.py +++ b/examples/api/autogen_example.py @@ -1,7 +1,7 @@ -import os -import sys -from typing import Optional, Union, Sequence, Literal, Dict, List, Any, Tuple -from autogen import Agent, ConversableAgent, UserProxyAgent, config_list_from_json +from typing import Optional, Union, Literal, Dict, List, Any, Tuple + +from autogen import Agent, ConversableAgent, UserProxyAgent + from memoryscope import MemoryScope, Arguments @@ -25,7 +25,6 @@ class MemoryScopeAgent(ConversableAgent): # Create a memory client in MemoryScope self.memory_scope = MemoryScope(arguments=arguments) - self.memory_scope.init_context_by_config() self.memory_chat = self.memory_scope.default_memory_chat self.register_reply([Agent, None], MemoryScopeAgent.generate_reply_with_memory,remove_other_reply_funcs=True) @@ -63,7 +62,6 @@ def main(): 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" ) @@ -79,4 +77,4 @@ def main(): if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index f409f411..20651ef8 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -38,7 +38,7 @@ class Arguments(object): generation_backend: str = field(default="openai_generation", metadata={ "help": "global generation backend: openai_generation, dashscope_generation, etc."}) - generation_model: str = field(default="gpt-4o-mini", metadata={ + generation_model: str = field(default="gpt-4o", metadata={ "help": "global generation model: gpt-4o, gpt-4o-mini, gpt-4-turbo, qwen-max, etc."}) generation_params: dict = field(default_factory=lambda: {}, metadata={ diff --git a/memoryscope/core/memoryscope.py b/memoryscope/core/memoryscope.py index d673d0ea..26a8bcde 100644 --- a/memoryscope/core/memoryscope.py +++ b/memoryscope/core/memoryscope.py @@ -14,15 +14,15 @@ class MemoryScope(ConfigManager): def __init__(self, **kwargs): super().__init__(**kwargs) - self.context: MemoryscopeContext = MemoryscopeContext() - self.init_context_by_config() + self._context: MemoryscopeContext = MemoryscopeContext() + self._init_context_by_config() - def init_context_by_config(self): + 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.update({ + self._context.language = LanguageEnum(global_conf["language"]) + self._context.thread_pool = ThreadPoolExecutor(max_workers=global_conf["thread_pool_max_workers"]) + self._context.meta_data.update({ "enable_ranker": global_conf["enable_ranker"], "enable_today_contra_repeat": global_conf["enable_today_contra_repeat"], "enable_long_contra_repeat": global_conf["enable_long_contra_repeat"], @@ -42,63 +42,68 @@ class MemoryScope(ConfigManager): 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) + 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) + 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) + 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) + 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) + self._context.monitor = init_instance_by_config(monitor_conf) # set worker config - self.context.worker_conf_dict = self.config["worker"] + self._context.worker_conf_dict = self.config["worker"] def close(self): # wait service to stop - for _, service in self.context.memory_service_dict.items(): + for _, service in self._context.memory_service_dict.items(): service.stop_backend_service(wait_service=True) - self.context.thread_pool.shutdown() + self._context.thread_pool.shutdown() - self.context.memory_store.close() + self._context.memory_store.close() - if self.context.monitor: - self.context.monitor.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): + if exc_type is not None: + self.logger.warning(f"An exception occurred: {exc_type.__name__}: {exc_val}\n{exc_tb}") self.close() + @property + def content(self): + return self._context + @property def memory_chat_dict(self): - return self.context.memory_chat_dict + return self._context.memory_chat_dict @property def memory_service_dict(self): - return self.context.memory_service_dict + return self._context.memory_service_dict @property def default_memory_chat(self) -> BaseMemoryChat: diff --git a/memoryscope/core/storage/llama_index_es_memory_store.py b/memoryscope/core/storage/llama_index_es_memory_store.py index 478e81be..5593b558 100644 --- a/memoryscope/core/storage/llama_index_es_memory_store.py +++ b/memoryscope/core/storage/llama_index_es_memory_store.py @@ -104,14 +104,15 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore): return [self._text_node_2_memory_node(n) for n in text_nodes] def batch_insert(self, nodes: List[MemoryNode]): - # TODO batch insert - for node in nodes: - self.insert(node) + self.index.insert_nodes([self._memory_node_2_text_node(node) for node in nodes]) def batch_update(self, nodes: List[MemoryNode], update_embedding: bool = True): - # TODO batch_update - for node in nodes: - self.update(node, update_embedding=update_embedding) + if update_embedding: + for node in nodes: + node.vector = [] + + self.batch_delete(nodes) + self.batch_insert(nodes) def batch_delete(self, nodes: List[MemoryNode]): # TODO batch_delete diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index b7f23f2d..da3c3baa 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -43,12 +43,12 @@ class TestWorkersCn(unittest.TestCase): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) query = "明天我去上海出差" query_timestamp = int(datetime.datetime.now().timestamp()) @@ -63,12 +63,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name), @@ -97,12 +97,12 @@ class TestWorkersCn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="你知道北京哪里的海鲜最新鲜吗", @@ -145,12 +145,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="我爱吃川菜", role_name=self.arguments.human_name), @@ -179,12 +179,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="有没有推荐的策略游戏?最近想找新的挑战。", role_name=self.arguments.human_name), @@ -210,12 +210,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="去年我们一起合作了因果推断技术", role_name=self.arguments.human_name), @@ -239,12 +239,12 @@ class TestWorkersCn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="用户在美团干活"), @@ -299,12 +299,12 @@ class TestWorkersCn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name), @@ -335,12 +335,12 @@ class TestWorkersCn(unittest.TestCase): name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="用户喜欢打王者荣耀", role_name=self.arguments.human_name), @@ -357,12 +357,12 @@ class TestWorkersCn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context, TARGET_NAME: self.arguments.human_name}, + context={MEMORYSCOPE_CONTEXT: self.ms._context, TARGET_NAME: self.arguments.human_name}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="用户对策略游戏感兴趣,寻找新挑战。", role_name=self.arguments.human_name), diff --git a/tests/worker/test_workers_en.py b/tests/worker/test_workers_en.py index 92f81cd3..3dd46442 100644 --- a/tests/worker/test_workers_en.py +++ b/tests/worker/test_workers_en.py @@ -41,12 +41,12 @@ class TestWorkersEn(unittest.TestCase): name = "extract_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.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()) @@ -61,12 +61,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, content="I love to eat Sichuan cuisine."), @@ -87,12 +87,12 @@ class TestWorkersEn(unittest.TestCase): name = "info_filter" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.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?"), @@ -135,12 +135,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.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 = [ @@ -164,12 +164,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -206,12 +206,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_observation_with_time" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) chat_messages = [ Message(role=MessageRoleEnum.USER.value, @@ -238,12 +238,12 @@ class TestWorkersEn(unittest.TestCase): name = "contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(user_name="AI", target_name="用户", content="User is working in Meituan"), @@ -296,12 +296,12 @@ class TestWorkersEn(unittest.TestCase): name = "get_reflection_subject" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."), @@ -333,12 +333,12 @@ class TestWorkersEn(unittest.TestCase): name = "update_insight" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, context=reflection_worker.context, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="Users like to play King of Glory"), @@ -355,12 +355,12 @@ class TestWorkersEn(unittest.TestCase): name = "long_contra_repeat" worker: MemoryBaseWorker = init_instance_by_config( - config=self.ms.context.worker_conf_dict[name], + config=self.ms._context.worker_conf_dict[name], name=name, is_multi_thread=False, - context={MEMORYSCOPE_CONTEXT: self.ms.context}, + context={MEMORYSCOPE_CONTEXT: self.ms._context}, context_lock=None, - thread_pool=self.ms.context.thread_pool) + thread_pool=self.ms._context.thread_pool) nodes = [ MemoryNode(content="Users are interested in strategy games and looking for new challenges."),