Files
temperature-based-fertility…/code/new_realtime/models/tft_utils.py
T
Alex Blank 50cf43b9fe added code
2025-05-19 11:11:04 +02:00

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()