-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patheval_model.py
More file actions
87 lines (66 loc) · 3.57 KB
/
Copy patheval_model.py
File metadata and controls
87 lines (66 loc) · 3.57 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
import hydra
import logging
import torch
from src.timetensor.dataset import fetch_training_data, get_sizes, apply_stats
from src.timetensor.models import load_model
from src.timetensor.pipeline import get_losses, load_learner
from src.timetensor.visu import plot_weights
from src.timetensor.utils import get_dirs, set_seed
from tqdm import tqdm
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from src.timetensor.pipeline import launch_eval, launch_example
from src.timetensor.utils import symlog
import warnings
warnings.simplefilter(action='ignore', category=FutureWarning)
@hydra.main(version_base=None, config_path="configs", config_name="config")
def run(cfg):
logger = logging.getLogger(__name__)
logger.info("=====Running eval script=====")
#configs
data_path = cfg.data.path
lags, horizon = int(cfg.task.lags), int(cfg.task.horizon)
criterion_name = cfg.training.loss
criterion, eval_losses = get_losses(criterion_name, complete_evaluation=cfg.training.complete_evaluation)
model_name, norm_name = cfg.model.name, cfg.normalization.name
if norm_name == "None":
norm_name = None
kwargs = {**(cfg.normalization.configs or {}), **(cfg.model.configs or {})}
verbose, seed = cfg.misc.verbose, cfg.misc.seed
output_dir, save_name = cfg.misc.output_dir, cfg.misc.save_name,
save_name, save_dir = get_dirs(output_dir, save_name, model_name, norm_name, criterion_name, cfg.data.subsets.sizes)
if verbose:
logger.info(f"Fetched main configs, save directory : {save_dir}")
logger.info(f"Model {model_name}, norm {norm_name}, criterion {criterion_name}, kwargs {kwargs}")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
set_seed(seed)
#data
loaders_dict, stats_dict, nodes_stats_dict = fetch_training_data(
data_path, cfg.data.splits, cfg.data.subsets, cfg.training.bs, lags, horizon,
clusters=cfg.data.clustering.clusters, seed=seed, random_eval=cfg.training.random_eval, do_nodes=False)
if cfg.data.normalize:
apply_stats(loaders_dict, stats_dict)
shape, shape_str, batch_str = get_sizes(loaders_dict, str_info=True)
if verbose:
logger.info("Fetched dataloaders")
logger.info(shape_str)
logger.info(batch_str)
#model
model = load_model(model_name, shape, norm_name, cfg.training.init, cfg.training.freeze_core, cfg.model.constants, cfg.model.residuals, stats_dict, nodes_stats_dict, device=="cpu", logger, **kwargs)
learner = load_learner(model, norm_name, criterion, cfg.training.lr, eval_losses, device)
logger.info("--Model eval--")
launch_eval(learner, loaders_dict, stats_dict, eval_losses, save_dir, save_name, cfg.training.complete_evaluation, results_dir=output_dir, mode="Test", denormalize=cfg.data.normalize, runs=cfg.training.eval_runs)
launch_example(data_path, model, lags, horizon, device, save_dir, save_name)
#weights
plot_weights(model, save_dir + "plots/", save_name)
if (norm_name is not None) and (("revin" in norm_name) or ("mIN" in norm_name and "cmIN" not in norm_name)):
params = {"beta": model.beta.data.detach().cpu().numpy()[0][0][0], "alpha": model.alpha.data.detach().cpu().numpy()[0][0][0]}
logger.info(f"Final modulations: {params}")
elif (norm_name is not None and "cmIN" in norm_name):
params = {f"beta_{k}": value.data.detach().cpu().numpy()[0][0][0] for k,value in enumerate(model.betas)}
logger.info(f"Final modulations: {params}")
logger.info('End of script\n')
if __name__ == "__main__":
run()