Fix: ensure encoding=utf-8 for json

This commit is contained in:
litagin02
2024-01-01 09:26:36 +09:00
parent e9db74ac6d
commit f2e7c18aa9
5 changed files with 22 additions and 14 deletions

View File

@@ -39,10 +39,12 @@ def initialize(model_name, batch_size, epochs, save_every_steps, bf16_run):
logger.info("Step 1: start initialization...")
dataset_path, _, train_path, val_path, config_path = get_path(model_name)
if os.path.isfile(config_path):
config = json.load(open(config_path, "r", encoding="utf-8"))
with open(config_path, "r", encoding="utf-8") as f:
config = json.load(f)
else:
# Use default config
config = json.load(open("configs/config.json", "r", encoding="utf-8"))
with open("configs/config.json", "r", encoding="utf-8") as f:
config = json.load(f)
config["model_name"] = model_name
config["data"]["training_files"] = train_path
config["data"]["validation_files"] = val_path