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