156 lines
6.9 KiB
Python
156 lines
6.9 KiB
Python
import datetime
|
|
import os
|
|
|
|
import torch
|
|
from torch import nn
|
|
from torch.optim import AdamW
|
|
from torch.optim.lr_scheduler import OneCycleLR
|
|
from torch.utils.data import IterableDataset, DataLoader
|
|
|
|
from models.third_party.tft_model import TemporalFusionTransformer
|
|
|
|
|
|
def get_tft_model(model_configuration: dict,
|
|
sample_item: dict,
|
|
device: str) -> nn.Module:
|
|
config_class = create_config_class(model_configuration, sample_item)
|
|
model = TemporalFusionTransformer(config_class)
|
|
model.to(device)
|
|
return model
|
|
|
|
|
|
def create_training_state(model_configuration: dict,
|
|
training_configuration: dict,
|
|
train_dataset: IterableDataset | DataLoader,
|
|
device: str) -> dict:
|
|
training_state = dict()
|
|
|
|
sample_item = next(iter(train_dataset))
|
|
|
|
if model_configuration["model_type"] == "TemporalFusionTransformer":
|
|
training_state["model"] = get_tft_model(model_configuration,
|
|
sample_item,
|
|
device)
|
|
else:
|
|
raise NotImplementedError(f"Model type {model_configuration['type']} not implemented")
|
|
|
|
training_state["optimizer"] = AdamW(training_state["model"].parameters(),
|
|
lr=training_configuration["learning_rate"])
|
|
training_state["scheduler"] = OneCycleLR(training_state["optimizer"],
|
|
max_lr=training_configuration["learning_rate"],
|
|
total_steps=len(train_dataset) *
|
|
training_configuration[
|
|
"epochs"])
|
|
training_state["current_epoch"] = 1
|
|
return training_state
|
|
|
|
|
|
def get_checkpoints(model_configuration: dict, training_configuration: dict) -> list:
|
|
checkpoints_dir = f"{model_configuration['model_dir']}/trainings/{training_configuration['id']}/checkpoints"
|
|
if os.path.exists(checkpoints_dir):
|
|
checkpoints = [os.path.join(checkpoints_dir, f) for f in os.listdir(checkpoints_dir) if
|
|
f.endswith('.pt') and "checkpoint" in f]
|
|
checkpoints.sort(key=os.path.getmtime, reverse=False)
|
|
return checkpoints
|
|
else:
|
|
return []
|
|
|
|
|
|
def load_checkpoint(checkpoint_path: str,
|
|
training_configuration: dict,
|
|
model_configuration: dict,
|
|
test_data_loader: IterableDataset | DataLoader,
|
|
device: str) -> tuple:
|
|
checkpoint = torch.load(checkpoint_path)
|
|
training_state = torch.load(checkpoint_path)
|
|
|
|
sample_item = next(iter(test_data_loader))
|
|
|
|
# load model
|
|
if model_configuration["model_type"] == "TemporalFusionTransformer":
|
|
model = get_tft_model(model_configuration,
|
|
sample_item,
|
|
device)
|
|
model_configuration["model"] = model
|
|
else:
|
|
raise NotImplementedError(f"Model type {model_configuration['model_type']} not implemented")
|
|
|
|
# load optimizer
|
|
optimizer = AdamW(model.parameters(), lr=training_configuration["learning_rate"])
|
|
optimizer.load_state_dict(checkpoint["optimizer"])
|
|
training_state["optimizer"] = optimizer
|
|
|
|
# load scheduler
|
|
scheduler = OneCycleLR(optimizer,
|
|
max_lr=training_configuration["learning_rate"],
|
|
total_steps=len(test_data_loader) * training_configuration["epochs"])
|
|
scheduler.load_state_dict(checkpoint["scheduler"])
|
|
training_state["scheduler"] = scheduler
|
|
|
|
return training_state
|
|
|
|
|
|
def save_checkpoint(training_config: dict,
|
|
training_state: dict,
|
|
model_config: dict) -> None:
|
|
checkpoints_dir = f"{model_config['model_dir']}/trainings/{training_config['id']}/checkpoints"
|
|
checkpoint_id = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
checkpoint_path = os.path.join(checkpoints_dir, f"checkpoint_{checkpoint_id}.pt")
|
|
if not os.path.exists(checkpoints_dir):
|
|
os.makedirs(checkpoints_dir)
|
|
|
|
config_to_save = training_state.copy()
|
|
|
|
# replace training parts with their state dicts
|
|
config_to_save["model"] = config_to_save["model"].state_dict()
|
|
config_to_save["optimizer"] = config_to_save["optimizer"].state_dict()
|
|
config_to_save["scheduler"] = config_to_save["scheduler"].state_dict()
|
|
|
|
# save training state
|
|
torch.save(config_to_save, checkpoint_path)
|
|
|
|
|
|
def create_config_class(config: dict, sample_batch: dict) -> object:
|
|
class ConfigClass:
|
|
def __init__(self):
|
|
# Feature sizes
|
|
self.static_categorical_inp_lens = []
|
|
self.temporal_known_categorical_inp_lens = []
|
|
self.temporal_observed_categorical_inp_lens = []
|
|
|
|
model_parameters = config["model_parameters"]
|
|
|
|
self.example_length = model_parameters["encoder_length"] + model_parameters["decoder_length"]
|
|
self.encoder_length = model_parameters["encoder_length"]
|
|
|
|
self.n_head = model_parameters["attention_heads"]
|
|
self.hidden_size = model_parameters["state_size"]
|
|
self.dropout = model_parameters["dropout"]
|
|
self.attn_dropout = model_parameters["attention_dropout"]
|
|
self.quantiles = model_parameters["output_quantiles"]
|
|
self.use_past_targets = model_parameters["use_past_targets"]
|
|
|
|
#### Derived variables ####
|
|
self.temporal_known_continuous_inp_size = sample_batch["k_cont"].shape[2]
|
|
self.temporal_observed_continuous_inp_size = sample_batch["o_cont"].shape[2]
|
|
self.temporal_target_size = sample_batch["target"].shape[2]
|
|
self.static_continuous_inp_size = sample_batch["s_cont"].shape[2]
|
|
|
|
self.num_static_vars = self.static_continuous_inp_size + len(self.static_categorical_inp_lens)
|
|
self.num_future_vars = self.temporal_known_continuous_inp_size + len(
|
|
self.temporal_known_categorical_inp_lens)
|
|
if self.use_past_targets:
|
|
self.num_historic_vars = self.num_future_vars + self.temporal_observed_continuous_inp_size + self.temporal_target_size + len(
|
|
self.temporal_observed_categorical_inp_lens)
|
|
else:
|
|
self.num_historic_vars = self.num_future_vars + self.temporal_observed_continuous_inp_size + len(
|
|
self.temporal_observed_categorical_inp_lens)
|
|
# self.num_historic_vars = sum([self.num_future_vars,
|
|
# self.temporal_observed_continuous_inp_size,
|
|
# self.temporal_target_size,
|
|
# len(self.temporal_observed_categorical_inp_lens),
|
|
# ])
|
|
self.target_size = self.temporal_target_size
|
|
|
|
return ConfigClass()
|