| 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) |
| 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 |
| |
|
|
| |
|
|
| |
| 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 |