fixes
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user