feat: support config_path / kwargs / env simutaneously

support config_path / kwargs / env simultaneously
Link: https://code.alibaba-inc.com/OpenRepo/MemoryScope/codereview/18119750
* feat: support config_path / kwargs / env simutaneously

* amend arguments err in api mode
This commit is contained in:
fuqingxu.fqx 2024-08-28 11:59:13 +08:00
parent 2c99afaab0
commit 8ffd4175ee
11 changed files with 86 additions and 89 deletions

View file

@ -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",

View file

@ -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",

View file

@ -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",

View file

@ -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",

View file

@ -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",

View file

@ -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"})

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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": {}}})

View file

@ -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",