diff --git a/examples/api/agentscope_example.py b/examples/api/agentscope_example.py index f4dd59e6..0f8b06cc 100644 --- a/examples/api/agentscope_example.py +++ b/examples/api/agentscope_example.py @@ -40,7 +40,6 @@ def main(): language="cn", human_name="User", assistant_name="AI", - logger_to_screen=False, memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", generation_model="qwen-max", diff --git a/examples/api/autogen_example.py b/examples/api/autogen_example.py index ba0cef60..b260678d 100644 --- a/examples/api/autogen_example.py +++ b/examples/api/autogen_example.py @@ -56,7 +56,6 @@ def main(): language="cn", human_name="User", assistant_name="AI", - logger_to_screen=False, memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", generation_model="qwen-max", diff --git a/examples/api/chat_example.py b/examples/api/chat_example.py index 763f9f78..7c1e7077 100644 --- a/examples/api/chat_example.py +++ b/examples/api/chat_example.py @@ -6,7 +6,6 @@ arguments = Arguments( language="cn", human_name="User", assistant_name="AI", - logger_to_screen=False, memory_chat_class="api_memory_chat", generation_backend="dashscope_generation", generation_model="qwen2-72b-instruct", diff --git a/examples/api/simple_usages_cn.ipynb b/examples/api/simple_usages_cn.ipynb index 059ad530..61b092d8 100644 --- a/examples/api/simple_usages_cn.ipynb +++ b/examples/api/simple_usages_cn.ipynb @@ -43,7 +43,6 @@ " language=\"cn\",\n", " human_name=\"用户\",\n", " assistant_name=\"AI\",\n", - " logger_to_screen=False,\n", " memory_chat_class=\"api_memory_chat\",\n", " generation_backend=\"dashscope_generation\",\n", " generation_model=\"qwen2-72b-instruct\",\n", diff --git a/examples/api/simple_usages_en.ipynb b/examples/api/simple_usages_en.ipynb index 056ef1d0..29cd896e 100644 --- a/examples/api/simple_usages_en.ipynb +++ b/examples/api/simple_usages_en.ipynb @@ -42,7 +42,6 @@ " language=\"en\",\n", " human_name=\"User\",\n", " assistant_name=\"AI\",\n", - " logger_to_screen=False,\n", " memory_chat_class=\"api_memory_chat\",\n", " generation_backend=\"dashscope_generation\",\n", " generation_model=\"qwen2-72b-instruct\",\n", diff --git a/memoryscope/core/config/arguments.py b/memoryscope/core/config/arguments.py index e13c6ec5..0646d6e3 100644 --- a/memoryscope/core/config/arguments.py +++ b/memoryscope/core/config/arguments.py @@ -7,12 +7,6 @@ class Arguments(object): 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."}) @@ -68,15 +62,19 @@ class Arguments(object): enable_ranker: bool = field(default=False, 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."}) + "substitute. However, the ranking effectiveness will be somewhat compromised.", + "map_yaml": "global->enable_ranker"}) enable_today_contra_repeat: bool = field(default=True, metadata={ "help": "Whether enable conflict resolution and deduplication for the day? " - "Note that enabling this will increase token consumption."}) + "Note that enabling this will increase token consumption.", + "map_yaml": "global->enable_today_contra_repeat"}) enable_long_contra_repeat: bool = field(default=False, metadata={ "help": "Whether to enable long-term conflict resolution and deduplication. " - "Note that enabling this will increase token consumption."}) + "Note that enabling this will increase token consumption.", + "map_yaml": "global->enable_long_contra_repeat"}) output_memory_max_count: int = field(default=20, metadata={ - "help": "The maximum number of memories retrieved during memory recall."}) + "help": "The maximum number of memories retrieved during memory recall.", + "map_yaml": "global->output_memory_max_count"}) diff --git a/memoryscope/core/config/config_manager.py b/memoryscope/core/config/config_manager.py index da4152ac..829b7265 100644 --- a/memoryscope/core/config/config_manager.py +++ b/memoryscope/core/config/config_manager.py @@ -1,4 +1,5 @@ import json +import os from dataclasses import fields from datetime import datetime from pathlib import Path @@ -13,47 +14,40 @@ from memoryscope.core.utils.logger import Logger class ConfigManager(object): def __init__(self, - config: dict = None, config_path: Optional[str] = None, arguments: Optional[Arguments] = None, demo_config_name: str = "demo_config_zh.yaml", **kwargs): self.config: dict = {} self.kwargs = kwargs + self.logger = Logger.get_logger("memoryscope") - if config: - self.config = config - self.logger = self._init_logger() - self.logger.info("init by config mode:") + if not (config_path or kwargs or arguments): + raise RuntimeError("can not init config manager without kwargs or --config_path!") - elif config_path: + if 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__}") + self.read_config((Path(__file__).parent / demo_config_name).__str__()) - 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}") + kwargs = {k: v for k, v in kwargs.items() if k in [x.name for x in fields(Arguments)]} + kwargs_padding = {x.name: None for x in fields(Arguments) if x.name not in kwargs} + kwargs.update(kwargs_padding) - else: - raise RuntimeError("can not init config manager without kwargs!") + # (high) when there are environment variables, read them and merge into kwargs + kwargs_from_env = {x.name:os.environ.get(x.name, None) for x in fields(Arguments) if os.environ.get(x.name, None) is not None} + kwargs.update(kwargs_from_env) + + # generate argument dataclass + if not arguments: + arguments = Arguments(**kwargs) + else: + # (highest) when arguments is passed into the memoryscope + arguments = arguments + + self.update_config_by_arguments(arguments) self.logger.info("\n" + self.dump_config()) - def _init_logger(self) -> Logger: - global_config = self.config["global"] - logger_name = global_config["logger_name"] - 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: @@ -63,42 +57,45 @@ class ConfigManager(object): 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_ignore_none(config, new_config_dict): + update_dict = {k:v for k, v in new_config_dict.items() if v is not None} + config.update(update_dict) + return @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, - "enable_ranker": arguments.enable_ranker, - "enable_today_contra_repeat": arguments.enable_today_contra_repeat, - "enable_long_contra_repeat": arguments.enable_long_contra_repeat, - "output_memory_max_count": arguments.output_memory_max_count, - }) + ConfigManager.update_ignore_none( + config, + { + "language": arguments.language, + "thread_pool_max_workers": arguments.thread_pool_max_workers, + "enable_ranker": arguments.enable_ranker, + "enable_today_contra_repeat": arguments.enable_today_contra_repeat, + "enable_long_contra_repeat": arguments.enable_long_contra_repeat, + "output_memory_max_count": arguments.output_memory_max_count, + } + ) @staticmethod def update_memory_chat_by_arguments(config: dict, arguments: Arguments): - memory_chat_class_split = config["class"].split(".") - stream = arguments.chat_stream - if stream is None: - stream = arguments.memory_chat_class in ["cli_memory_chat", ] - config.update({ - "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), - "stream": stream, - }) + if arguments.memory_chat_class is not None: + memory_chat_class_split = config["class"].split(".") + stream = arguments.chat_stream + if stream is None: + stream = arguments.memory_chat_class in ["cli_memory_chat", ] + config.update( + { + "class": ".".join(memory_chat_class_split[:-1] + [arguments.memory_chat_class]), + "stream": stream, + } + ) @staticmethod def update_memory_service_by_arguments(config: dict, arguments: Arguments): - config.update({ - "human_name": arguments.human_name if arguments.human_name else "", - "assistant_name": arguments.assistant_name if arguments.assistant_name else "", + ConfigManager.update_ignore_none(config, { + "human_name": arguments.human_name, + "assistant_name": arguments.assistant_name, }) if arguments.consolidate_memory_interval_time is not None: config["memory_operations"]["consolidate_memory"]["interval_time"] = \ @@ -110,37 +107,48 @@ class ConfigManager(object): @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) + if arguments.worker_params is not None: + 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({ + ConfigManager.update_ignore_none(config["generation_model"], { "module_name": arguments.generation_backend, "model_name": arguments.generation_model, - **arguments.generation_params, }) + if isinstance(arguments.generation_params, dict): + ConfigManager.update_ignore_none(config["generation_model"], { + **arguments.generation_params, + }) - config["embedding_model"].update({ + ConfigManager.update_ignore_none(config["embedding_model"], { "module_name": arguments.embedding_backend, "model_name": arguments.embedding_model, - **arguments.embedding_params, }) + if isinstance(arguments.embedding_params, dict): + ConfigManager.update_ignore_none(config["embedding_model"], { + **arguments.embedding_params, + }) - config["rank_model"].update({ + ConfigManager.update_ignore_none(config["rank_model"], { "module_name": arguments.rank_backend, "model_name": arguments.rank_model, - **arguments.rank_params, }) + if isinstance(arguments.rank_params, dict): + ConfigManager.update_ignore_none(config["rank_model"], { + **arguments.rank_params, + }) @staticmethod def update_memory_store_by_arguments(config: dict, arguments: Arguments): - config.update({ + ConfigManager.update_ignore_none(config, { "index_name": arguments.es_index_name, "es_url": arguments.es_url, - "retrieve_mode": arguments.retrieve_mode}) + "retrieve_mode": arguments.retrieve_mode} + ) def update_config_by_arguments(self, arguments: Arguments): # prepare global diff --git a/memoryscope/core/config/demo_config.yaml b/memoryscope/core/config/demo_config.yaml index 28b25cac..f110235b 100644 --- a/memoryscope/core/config/demo_config.yaml +++ b/memoryscope/core/config/demo_config.yaml @@ -1,9 +1,6 @@ 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 enable_ranker: false enable_today_contra_repeat: true enable_long_contra_repeat: false diff --git a/memoryscope/core/config/demo_config_zh.yaml b/memoryscope/core/config/demo_config_zh.yaml index b1ecdba0..22fa2ea6 100644 --- a/memoryscope/core/config/demo_config_zh.yaml +++ b/memoryscope/core/config/demo_config_zh.yaml @@ -1,9 +1,6 @@ global: language: cn thread_pool_max_workers: 5 - logger_name: memoryscope - logger_name_time_suffix: "%Y%m%d_%H%M%S" - logger_to_screen: false enable_ranker: true enable_today_contra_repeat: true enable_long_contra_repeat: false diff --git a/memoryscope/core/storage/llama_index_sync_elasticsearch.py b/memoryscope/core/storage/llama_index_sync_elasticsearch.py index a5d835ac..bcd675cc 100644 --- a/memoryscope/core/storage/llama_index_sync_elasticsearch.py +++ b/memoryscope/core/storage/llama_index_sync_elasticsearch.py @@ -633,7 +633,10 @@ class SyncElasticsearchStore(BasePydanticVectorStore): return q_res def sync_delete_all(self): - self._store.client.delete_by_query(index=[self.index_name], body={"query": {"match_all": {}}}) + try: + self._store.client.delete_by_query(index=[self.index_name], body={"query": {"match_all": {}}}) + except: # elasticsearch.NotFoundError + pass def sync_search_all(self): search_res = self._store.client.search(index=[self.index_name], body={"query": {"match_all": {}}}) diff --git a/tests/worker/test_workers_cn.py b/tests/worker/test_workers_cn.py index da3c3baa..86fa9580 100644 --- a/tests/worker/test_workers_cn.py +++ b/tests/worker/test_workers_cn.py @@ -19,7 +19,6 @@ class TestWorkersCn(unittest.TestCase): def setUp(self): self.arguments = Arguments( language="cn", - logger_to_screen=True, human_name="用户", assistant_name="AI", memory_chat_class="api_memory_chat",