File size: 4,159 Bytes
238676a
 
 
 
 
faeee34
238676a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
faeee34
238676a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
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