Files
2025-09-10 10:37:55 +02:00

380 lines
14 KiB
Python

import math
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, get_collated_batch_for_key
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)
# update paths
training_config["training_dir"] = training_dir
training_config["model_dir"] = model_config["model_dir"]
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 load_model_for_usage(model_configuration: dict,
training_configuration: dict,
sample_id: str) -> tuple[nn.Module, dict]:
"""
Load the model associated with the given model configuration for the usage on the current node, given its training configuration
If the model does not exist, create it
Args:
model_configuration: configuration of the model
training_configuration: training configuration
sample_id: id of a sample data item
Returns:
model: model
training_config: training configuration with updated variables such as batch size
"""
# get computation rank
if torch.distributed.is_initialized():
local_rank = torch.distributed.get_rank()
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
else:
local_rank = 0
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dataset_dir = model_configuration["feature_config"]["dataset_dir"]
lmdb_env = lmdb.open(dataset_dir, readonly=True)
try:
load_fn = model_configuration["model_load_fn"]
model = load_fn(model_configuration=model_configuration,
training_configuration=training_configuration,
sample_key=sample_id,
device=device,
lmdb_env=lmdb_env)
except FileNotFoundError:
model_creation_fn = model_configuration["model_creation_fn"]
model = model_creation_fn(
model_configuration=model_configuration,
sample_key=sample_id,
lmdb_env=lmdb_env
)
# estimate batch size
sample_batch = get_collated_batch_for_key(sample_id, model_configuration, lmdb_env=lmdb_env)
if sample_batch is None or len(sample_batch) == 0:
raise ValueError(
f"Rank {local_rank}: Could not get sample input for key {sample_id} in dataset {dataset_dir}. "
f"Please check if the key exists in the LMDB dataset.")
sample_input = sample_batch[0][0] # get the first item in the batch
input_shape = sample_input.shape
# send model to proper device
model.to(device)
batch_size = find_max_batch_size(model,
input_shape,
device=device)
# update max lr in training configuration, use square root scaling law based on initial batch size
training_configuration["learning_parameters"]["learning_rate"] = training_configuration["learning_parameters"][
"max_lr"] * math.sqrt(
batch_size / training_configuration["batch_size"])
# update batch size in training configuration
training_configuration["batch_size"] = batch_size
# close lmdb env
lmdb_env.close()
return model, training_configuration
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
def find_max_batch_size(model,
input_shape,
device,
dtype=torch.float32,
max_batch=2048):
"""
Find the maximum batch size that can be processed by the model without running out of memory.
Model must already be moved to the appropriate device (e.g., GPU).
Args:
model: model to test
input_shape: shape of the input tensor (excluding batch size)
device: device to work on
dtype: data type of the input tensor (default: torch.float32)
max_batch: maximum batch size to test (default: 2048)
Returns:
last_successful: the largest batch size that did not cause an out of memory error
"""
batch_size = 1
last_successful = 1
while last_successful < max_batch:
try:
inputs = torch.randn((batch_size, *input_shape), device=device, dtype=dtype)
outputs = model(inputs)
loss = outputs.sum()
loss.backward() # simulate training
torch.cuda.empty_cache()
last_successful = batch_size
batch_size *= 2
except RuntimeError as e:
if 'out of memory' in str(e):
torch.cuda.empty_cache()
break
else:
raise e
return last_successful