63 lines
2.3 KiB
Python
63 lines
2.3 KiB
Python
import os
|
|
import pickle
|
|
from datetime import datetime
|
|
|
|
from utils.utils import get_config_id
|
|
|
|
|
|
def get_model_config(base_config: dict,
|
|
base_result_dir: str,
|
|
dataset_base_dir: str):
|
|
# get config identifier
|
|
name_for_current_config = base_config["model_name"] + "_" + get_config_id(base_config)
|
|
# date_part = datetime.now().strftime("%Y_%m_%d_%H_%M")
|
|
# name_for_current_config = base_config["model_name"] + "_" + date_part
|
|
model_dir = os.path.abspath(f"{base_result_dir}/{name_for_current_config}")
|
|
if not os.path.exists(model_dir):
|
|
# create config
|
|
os.makedirs(model_dir)
|
|
|
|
# append identifier to config
|
|
config = base_config.copy()
|
|
config["id"] = name_for_current_config
|
|
|
|
# append paths
|
|
config["model_dir"] = model_dir
|
|
config["dataset_dir"] = os.path.join(dataset_base_dir, config["feature_config"]["feature_set_name"])
|
|
config["feature_config"]["dataset_dir"] = config["dataset_dir"]
|
|
|
|
# save model configuration
|
|
with open(f"{model_dir}/model_configuration.pickle", "wb") as f:
|
|
pickle.dump(config, f)
|
|
else:
|
|
# fetch config
|
|
with open(f"{model_dir}/model_configuration.pickle", "rb") as f:
|
|
config = pickle.load(f)
|
|
|
|
# append paths
|
|
config["model_dir"] = model_dir
|
|
config["dataset_dir"] = os.path.join(dataset_base_dir, config["feature_config"]["feature_set_name"])
|
|
config["feature_config"]["dataset_dir"] = config["dataset_dir"]
|
|
return config
|
|
|
|
|
|
def get_model_config_from_file(model_dir: str,
|
|
model_base_dir: str,
|
|
lmdb_base_dir: str):
|
|
# fetch config
|
|
with open(f"{model_dir}/model_configuration.pickle", "rb") as f:
|
|
config = pickle.load(f)
|
|
|
|
# update paths based on base directories
|
|
config["model_dir"] = os.path.abspath(f"{model_base_dir}/{config['id']}")
|
|
config["feature_config"]["dataset_dir"] = os.path.abspath(
|
|
f"{lmdb_base_dir}/{config['feature_config']['feature_set_name']}")
|
|
|
|
return config
|
|
|
|
|
|
def save_model_config(model_dir: str, model_config: dict):
|
|
# save model configuration
|
|
with open(f"{model_dir}/model_configuration.pickle", "wb") as f:
|
|
pickle.dump(model_config, f)
|