mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
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:
parent
2c99afaab0
commit
8ffd4175ee
11 changed files with 86 additions and 89 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {}}})
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue