ChebAutoencoder / AutoencoderCheb.py
lara roth
Update AI model
faeee34
Raw
History Blame Contribute Delete
4.16 kB
import torch
import torch.nn as nn
import torch
from pytorch_lightning import LightningModule
from torch_geometric.nn import ChebConv, Sequential
class AutoEncoderModel(LightningModule):
def __init__(self, cuda_true, batch_size):
super().__init__()
self.epochs, self.conditions = list(), list()
self.recon_loss_test_step_list = list()
self.num_step = 0
if cuda_true:
self.dev = "cuda"
else:
self.dev = "cpu"
self.num_nodes = 15
self.edge_index_att = None
self.batch_size = batch_size
self.criterion = nn.MSELoss(reduction='mean')
self.window = 64
self.automatic_optimization = True
self.test_target_data, self.test_predict_data = list(), list()
self.output_first_layer_decoder = torch.rand(self.batch_size, self.num_nodes, self.window*4) # check size
self.output_first_layer_decoder.requires_grad_()
self.output_first_layer_decoder.to(self.dev)
self.node_num_featues = 5
self.total_feat = self.node_num_featues * self.window
self.k = 4
latent_dim = 104
self.beta = 0.009256865323169841
self.encoder = Sequential('x, edge_index', [
(ChebConv(in_channels=self.window*self.node_num_featues, out_channels=self.window*2, K=self.k), 'x, edge_index -> x'),
nn.ReLU(inplace=True),
(ChebConv(in_channels=self.window*2, out_channels=self.window*4, K=self.k), 'x, edge_index -> x'),
nn.ReLU(inplace=True),
])
self.encoder_2 = Sequential('x, edge_index', [
(ChebConv(in_channels=self.window, out_channels=self.window*4, K=self.k), 'x, edge_index -> x'),
nn.ReLU(inplace=True)
])
self.latent = nn.Sequential(
nn.Flatten(),
nn.Linear(self.window*4*self.num_nodes, latent_dim),
nn.Linear(latent_dim, self.window*4*self.num_nodes),
nn.Unflatten(-1, (int(self.num_nodes), int(self.window*4)))
)
self.latent.to(self.dev)
self.decoder_2 = Sequential('x, edge_index' ,[
(ChebConv(in_channels=self.window*4, out_channels=self.window, K=self.k), 'x, edge_index -> x'),
nn.ReLU(inplace=True)
])
self.decoder = Sequential('x, edge_index' ,[
(ChebConv(in_channels=self.window*4, out_channels=self.window*2, K=self.k), 'x, edge_index -> x'),
nn.ReLU(inplace=True),
(ChebConv(in_channels=self.window*2, out_channels=self.window*self.node_num_featues, K=self.k), 'x, edge_index -> x')
])
self.softmax = nn.Softmax()
def forward(self, input_data, edge_indices, adj_matrix):
self.edge_indices = edge_indices
self.adj_matrix = adj_matrix
# print("input data", input_data)
input_data_reshaped = torch.reshape(input=input_data, shape=(input_data.shape[0], input_data.shape[2], input_data.shape[3] * input_data.shape[1]))
input_data_reshaped = input_data_reshaped.to(self.dev)
output_encoder = self.encoder(input_data_reshaped, self.edge_indices)
scaled_encoder = torch.mul(output_encoder, self.beta)
output_latent = self.latent(scaled_encoder)
output_decoder = self.decoder(output_latent, self.edge_indices)
self.output_decoder = torch.reshape(input=output_decoder, shape=(output_decoder.shape[0], self.window, self.num_nodes, self.node_num_featues))
recon_loss_list = list()
for i in range(input_data.shape[1]):
recon_loss_list.append(self.criterion(self.output_decoder[:,i,:,:], input_data[:,i,:,:]).to(self.dev))
recon_loss = sum(recon_loss_list)/len(recon_loss_list)
return recon_loss, self.output_decoder
def calc_edge_weight(edge_index, adj_matrix):
edge_weight = torch.rand(edge_index.shape[1])
for i, element in enumerate(edge_index.T):
edge_weight[i] = (adj_matrix[element[0]][element[1]] + adj_matrix[element[1]][element[0]])/2.0
return edge_weight