262 lines
9.7 KiB
Python
262 lines
9.7 KiB
Python
import os
|
|
import pickle
|
|
import random
|
|
from datetime import datetime
|
|
|
|
import numpy as np
|
|
import torch
|
|
from torch import nn
|
|
from torch.nn import init
|
|
from sklearn.model_selection import train_test_split
|
|
import lmdb
|
|
from bson import ObjectId
|
|
from tqdm import tqdm
|
|
|
|
from utils.lmdb_utils import get_lmdb_keys
|
|
from utils.utils import get_config_id
|
|
from utils.data_utils import produce_window_batches
|
|
|
|
from vsm_datascience_common.cycle_database_connection.db_utils import get_cycles_collection
|
|
|
|
|
|
def get_training_config(base_config: dict, model_config: dict):
|
|
try:
|
|
if base_config["model_class"] != model_config["model_class"]:
|
|
raise Exception("Model type differ in model config and training config")
|
|
except KeyError:
|
|
raise Exception("Model class missing in training or model config")
|
|
|
|
# training_config_id = get_config_id(base_config)
|
|
# hash = training_config_id[-5:]
|
|
|
|
timestamp = datetime.now().strftime("%Y%m%d-%H%M")
|
|
training_config_id = f"{timestamp}"
|
|
|
|
training_dir = os.path.abspath(f"{model_config['model_dir']}/trainings/{training_config_id}")
|
|
|
|
# append identifier to config
|
|
training_config = base_config.copy()
|
|
training_config["id"] = training_config_id
|
|
training_config["training_dir"] = training_dir
|
|
|
|
if not os.path.isdir(training_dir):
|
|
os.makedirs(training_dir)
|
|
# save training configuration
|
|
with open(f"{training_dir}/training_configuration.pickle", "wb") as f:
|
|
pickle.dump(training_config, f)
|
|
else:
|
|
# fetch config
|
|
with open(f"{training_dir}/training_configuration.pickle", "rb") as f:
|
|
training_config = pickle.load(f)
|
|
|
|
return training_config
|
|
|
|
|
|
def get_training_config_from_file(training_dir: str,
|
|
base_model_dir: str,
|
|
model_configuration: dict) -> dict:
|
|
# fetch config
|
|
with open(f"{training_dir}/training_configuration.pickle", "rb") as f:
|
|
training_config = pickle.load(f)
|
|
|
|
# update paths based on base directories
|
|
training_config["training_dir"] = os.path.join(os.path.abspath(base_model_dir),
|
|
model_configuration["id"],
|
|
"trainings",
|
|
training_config["id"])
|
|
|
|
return training_config
|
|
|
|
|
|
def weight_init(m):
|
|
"""
|
|
Usage:
|
|
model = Model()
|
|
model.apply(weight_init)
|
|
"""
|
|
if isinstance(m, nn.Conv1d):
|
|
init.normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.Conv2d):
|
|
init.xavier_normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.Conv3d):
|
|
init.xavier_normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.ConvTranspose1d):
|
|
init.normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.ConvTranspose2d):
|
|
init.xavier_normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.ConvTranspose3d):
|
|
init.xavier_normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.BatchNorm1d):
|
|
init.normal_(m.weight.data, mean=1, std=0.02)
|
|
init.constant_(m.bias.data, 0)
|
|
elif isinstance(m, nn.BatchNorm2d):
|
|
init.normal_(m.weight.data, mean=1, std=0.02)
|
|
init.constant_(m.bias.data, 0)
|
|
elif isinstance(m, nn.BatchNorm3d):
|
|
init.normal_(m.weight.data, mean=1, std=0.02)
|
|
init.constant_(m.bias.data, 0)
|
|
elif isinstance(m, nn.Linear):
|
|
init.xavier_normal_(m.weight.data)
|
|
if m.bias is not None:
|
|
init.normal_(m.bias.data)
|
|
elif isinstance(m, nn.LSTM):
|
|
for param in m.parameters():
|
|
if len(param.shape) >= 2:
|
|
init.orthogonal_(param.data)
|
|
else:
|
|
init.normal_(param.data)
|
|
elif isinstance(m, nn.LSTMCell):
|
|
for param in m.parameters():
|
|
if len(param.shape) >= 2:
|
|
init.orthogonal_(param.data)
|
|
else:
|
|
init.normal_(param.data)
|
|
elif isinstance(m, nn.GRU):
|
|
for param in m.parameters():
|
|
if len(param.shape) >= 2:
|
|
init.orthogonal_(param.data)
|
|
else:
|
|
init.normal_(param.data)
|
|
for names in m._all_weights:
|
|
for name in filter(lambda n: "bias" in n, names):
|
|
bias = getattr(m, name)
|
|
n = bias.size(0)
|
|
bias.data[:n // 3].fill_(-1.)
|
|
elif isinstance(m, nn.GRUCell):
|
|
for param in m.parameters():
|
|
if len(param.shape) >= 2:
|
|
init.orthogonal_(param.data)
|
|
else:
|
|
init.normal_(param.data)
|
|
|
|
|
|
def collate(batch_items: list) -> dict:
|
|
batch = dict()
|
|
for key in batch_items[0].keys():
|
|
if key in ["combination_id", "time_index"]:
|
|
continue
|
|
else:
|
|
if batch_items[0][key] is None:
|
|
batch[key] = None
|
|
else:
|
|
batch[key] = np.stack([item[key] for item in batch_items])
|
|
|
|
for key in batch.keys():
|
|
if batch[key] is not None:
|
|
batch[key] = torch.tensor(batch[key], dtype=torch.float32)
|
|
|
|
return batch
|
|
|
|
|
|
def get_splits_by_user(input_keys: list, train_size: float, val_size: float, test_size: float):
|
|
if train_size + val_size + test_size != 1:
|
|
raise ValueError("Train, val and test sizes must sum to 1")
|
|
|
|
if len(input_keys) == 0:
|
|
raise ValueError("Input keys list is empty")
|
|
|
|
items_by_use = dict()
|
|
for input_key in tqdm(input_keys):
|
|
try:
|
|
user_id = get_cycles_collection().find_one({"_id": ObjectId(input_key)})["user_id"]
|
|
if user_id not in items_by_use:
|
|
items_by_use[user_id] = []
|
|
items_by_use[user_id].append(input_key)
|
|
except Exception:
|
|
print(f"Error getting user id for key {input_key}")
|
|
continue
|
|
|
|
user_ids = list(items_by_use.keys())
|
|
random.shuffle(user_ids)
|
|
train_users, temp_users = train_test_split(user_ids, train_size=train_size, test_size=test_size + val_size)
|
|
# compute relative test size, as it must be relative to the remaining users
|
|
relative_test_size = test_size / (1 - train_size)
|
|
val_users, test_users = train_test_split(temp_users, test_size=relative_test_size)
|
|
|
|
train_keys = []
|
|
val_keys = []
|
|
test_keys = []
|
|
for user_id in train_users:
|
|
train_keys.extend(items_by_use[user_id])
|
|
for user_id in val_users:
|
|
val_keys.extend(items_by_use[user_id])
|
|
for user_id in test_users:
|
|
test_keys.extend(items_by_use[user_id])
|
|
|
|
return train_keys, val_keys, test_keys
|
|
|
|
|
|
def save_splits(train_keys: list, val_keys: list, test_keys: list, base_dir: str):
|
|
if not os.path.exists(base_dir):
|
|
os.makedirs(base_dir)
|
|
|
|
with open(f"{base_dir}/train_keys.pickle", "wb") as f:
|
|
pickle.dump(train_keys, f)
|
|
with open(f"{base_dir}/val_keys.pickle", "wb") as f:
|
|
pickle.dump(val_keys, f)
|
|
with open(f"{base_dir}/test_keys.pickle", "wb") as f:
|
|
pickle.dump(test_keys, f)
|
|
|
|
|
|
def load_splits(base_dir: str):
|
|
with open(f"{base_dir}/train_keys.pickle", "rb") as f:
|
|
train_keys = pickle.load(f)
|
|
with open(f"{base_dir}/val_keys.pickle", "rb") as f:
|
|
val_keys = pickle.load(f)
|
|
with open(f"{base_dir}/test_keys.pickle", "rb") as f:
|
|
test_keys = pickle.load(f)
|
|
|
|
return train_keys, val_keys, test_keys
|
|
|
|
|
|
def get_data_ids(model_configuration: dict,
|
|
training_configuration: dict,
|
|
env_path: str,
|
|
limit: int = None) -> tuple:
|
|
env = lmdb.open(f"{env_path}", readonly=True)
|
|
if os.path.exists(f"{model_configuration['feature_config']['dataset_dir']}/train_keys.pickle"):
|
|
train_ids, val_ids, test_ids = load_splits(model_configuration["feature_config"]["dataset_dir"])
|
|
else:
|
|
lmdb_keys = get_lmdb_keys(env, limit)
|
|
# train_ids, val_ids, test_ids = get_splits_by_user(lmdb_keys,
|
|
# training_configuration["train_size"],
|
|
# training_configuration["val_size"],
|
|
# training_configuration["test_size"])
|
|
train_ids, temp_ids = train_test_split(lmdb_keys,
|
|
train_size=training_configuration["train_size"],
|
|
test_size=training_configuration["val_size"] + training_configuration[
|
|
"test_size"])
|
|
# compute relative test size, as it must be relative to the remaining users
|
|
relative_test_size = training_configuration["test_size"] / (1 - training_configuration["train_size"])
|
|
val_ids, test_ids = train_test_split(temp_ids,
|
|
test_size=relative_test_size)
|
|
# save splits to file
|
|
save_splits(train_ids, val_ids, test_ids, model_configuration["feature_config"]["dataset_dir"])
|
|
|
|
if limit is not None:
|
|
train_size = training_configuration["train_size"]
|
|
val_size = training_configuration["val_size"]
|
|
test_size = training_configuration["test_size"]
|
|
|
|
train_lim = int(limit * train_size)
|
|
val_lim = int(limit * val_size)
|
|
test_lim = int(limit * test_size)
|
|
|
|
train_ids = train_ids[:train_lim]
|
|
val_ids = val_ids[:val_lim]
|
|
test_ids = test_ids[:test_lim]
|
|
|
|
return train_ids, val_ids, test_ids
|