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