fixes
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user