This commit is contained in:
Alex Blank
2025-05-19 13:59:16 +02:00
parent 426f4d6963
commit c6defa2065
196 changed files with 18625 additions and 1 deletions
+261
View File
@@ -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