11
22from collections import namedtuple
3+ import numpy as np
34import torch
45import torch .nn as nn
56import torch .optim as optim
67from torch .utils .data import DataLoader
8+ from tqdm import tqdm # type: ignore
79from typing import List , Tuple
810
911"""To do
@@ -19,46 +21,78 @@ class ModelRunner(object):
1921 # Runner for Autoencoder
2022
2123 def __init__ (self , model : nn .Module , num_epoch :int = 3 , learning_rate :float = 1e-3 ,
22- criterion :nn .Module = nn .MSELoss (), is_autoencoder :bool = False , is_report :bool = True ):
24+ criterion :nn .Module = nn .MSELoss (), is_autoencoder :bool = False ,
25+ is_normalized :bool = False , is_report :bool = True ):
2326 """
2427 Args:
25- model (nn.Module): _description_
28+ model (nn.Module): Model being run
2629 num_epoch (int, optional): Defaults to 3.
2730 learning_rate (float, optional): Defaults to 1e-3.
2831 is_autoencoder (bool, optional): target data is features Defaults to False.
29- is_report (bool, optional): Print text for progress. Defaults to False.
32+ is_normalized (bool, optional): Whether to normalize the input data (divide by std).
33+ Defaults to False.
34+ is_report (bool, optional): Print text for progress.
35+ Defaults to False.
3036 """
3137 self .device = torch .accelerator .current_accelerator ().type if torch .accelerator .is_available () else "cpu" # type: ignore
3238 self .model = model .to (self .device )
3339 self .num_epoch = num_epoch
3440 self .learning_rate = learning_rate
3541 self .criterion = criterion
3642 self .is_autoencoder = is_autoencoder
43+ self .is_normalized = is_normalized
3744 self .is_report = is_report
45+ # Calculated state
46+ self .feature_std_tnsr = torch .tensor ([np .nan ])
47+ self .target_std_tnsr = torch .tensor ([np .nan ])
3848
3949 def train (self , train_loader : DataLoader ) -> RunnerResult :
40- """Train the network."""
50+ """
51+ Train the model.
52+
53+ Args:
54+ train_loader (DataLoader): DataLoader for training data
55+ Returns:
56+ RunnerResult: losses and number of epochs
57+ """
58+ ##
59+ def calculate_std (loader_idx : int ) -> torch .Tensor :
60+ # loader_idx (int): Index into the DataLoader
61+ full_tnsr = torch .cat ([x [loader_idx ] for x in train_loader ])
62+ if self .is_normalized :
63+ return full_tnsr .std (dim = 0 )
64+ else :
65+ return torch .ones (full_tnsr .size ()[1 ])
66+ ##
67+ # Handle normalization adjustments
68+ self .feature_std_tnsr = calculate_std (0 )
69+ self .target_std_tnsr = calculate_std (1 )
70+ if self .is_autoencoder :
71+ self .target_std_tnsr = self .feature_std_tnsr
72+ # Initialize for training
4173 optimizer = optim .Adam (self .model .parameters (), lr = self .learning_rate )
4274 self .model .to (self .device )
43-
4475 self .model .train ()
4576 losses = []
4677 avg_loss = 0.0
47-
48- for epoch in range (self .num_epoch ):
78+ epoch_loss = np .inf
79+ # Training loop
80+ pbar = tqdm (range (self .num_epoch ), desc = f"epochs (loss={ epoch_loss :.4f} )" )
81+ for epoch in pbar :
82+ pbar .set_description_str (f"epochs (loss={ epoch_loss :.4f} )" )
4983 epoch_loss = 0
50- for data in list (train_loader ):
51- data = data .to (self .device , non_blocking = True )
52- feature_tnsr , target_tnsr = data
84+ #for idx, (feature_tnsr, target_tnsr) in list(train_loader):
85+ for (feature_tnsr , target_tnsr ) in train_loader :
5386 if self .is_autoencoder :
5487 # For autoencoder, target is the same as input
5588 target_tnsr = feature_tnsr
56- feature_tnsr = feature_tnsr .view (feature_tnsr .size (0 ), - 1 )
57- # Forward pass
89+ feature_tnsr = feature_tnsr / self .feature_std_tnsr
5890 feature_tnsr = feature_tnsr .to (self .device )
91+ target_tnsr = target_tnsr / self .target_std_tnsr
92+ target_tnsr = target_tnsr .to (self .device )
93+ # Forward pass
5994 prediction_tnsr = self .model (feature_tnsr )
6095 loss = self .criterion (prediction_tnsr , target_tnsr )
61- loss = loss .to (CPU )
6296 # Backward pass
6397 optimizer .zero_grad ()
6498 loss .backward ()
@@ -71,22 +105,35 @@ def train(self, train_loader: DataLoader) -> RunnerResult:
71105 if self .is_report :
72106 print (f'Epoch [{ epoch + 1 } /{ self .num_epoch } ], Loss: { avg_loss :.4f} ' )
73107 #
108+ self .model .to (CPU )
74109 return RunnerResult (losses = losses , num_epochs = self .num_epoch )
75110
76- def evaluate (self , test_loader : DataLoader ) -> RunnerResult :
77- """Evaluate the model on the test set."""
111+ def predict (self , feature_tnsr : torch .Tensor ) -> torch .Tensor :
112+ """Predicts the target for the features.
113+
114+ Args:
115+ feature_tnsr (torch.Tensor): Input features for which to predict targets.
116+ Returns:
117+ torch.Tensor: target predictions
118+ """
119+ self .model .eval ()
120+ feature_tnsr = feature_tnsr / self .feature_std_tnsr
121+ with torch .no_grad ():
122+ prediction_tnsr = self .model (feature_tnsr )
123+ return self .feature_std_tnsr * prediction_tnsr
124+
125+ def assess (self , test_loader : DataLoader ) -> RunnerResult :
126+ """Assess the model on a test dataset."""
78127 self .model .eval ()
79128 test_losses = []
80129 #
81130 with torch .no_grad ():
82- for data in list (test_loader ):
83- data = data .to (self .device , non_blocking = True )
84- feature_tnsr , target_tnsr = data
131+ for (feature_tnsr , target_tnsr ) in list (test_loader ):
85132 if self .is_autoencoder :
86133 # For autoencoder, target is the same as input
87134 target_tnsr = feature_tnsr
88135 feature_tnsr = feature_tnsr .view (feature_tnsr .size (0 ), - 1 )
89- prediction_tnsr = self .model (feature_tnsr )
136+ prediction_tnsr = self .predict (feature_tnsr )
90137 loss = self .criterion (prediction_tnsr , target_tnsr ).to (CPU )
91138 test_losses .append (loss .item ())
92139
@@ -110,5 +157,5 @@ def run(self, train_loader: DataLoader, test_loader: DataLoader)->Tuple[RunnerRe
110157 print ("Training Fully Connected Autoencoder..." )
111158 # Create and train fully connected autoencoder
112159 train_runner_result = self .train (train_loader )
113- test_runner_result = self .evaluate (test_loader )
160+ test_runner_result = self .assess (test_loader )
114161 return train_runner_result , test_runner_result
0 commit comments