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