Download AutoencoderCheb.py from EnerTEF/ChebAutoencoder: direct link, hf CLI and curl.
- Browser
- Download file 4.16 kB
-
https://huggingface.co/EnerTEF/ChebAutoencoder/resolve/main/AutoencoderCheb.py
- Command line
-
hf download hf://EnerTEF/ChebAutoencoder/AutoencoderCheb.py
-
curl -L -o AutoencoderCheb.py https://huggingface.co/EnerTEF/ChebAutoencoder/resolve/main/AutoencoderCheb.py
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 |