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