lara roth commited on
Commit
238676a
·
1 Parent(s): 661766d

Upload ChebAutoencoder model.

Browse files
Files changed (3) hide show
  1. AutoencoderCheb.py +169 -0
  2. eval.py +7 -0
  3. gae_25_08_2025.pt +3 -0
AutoencoderCheb.py ADDED
@@ -0,0 +1,169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ # import torch.nn.functional as F
4
+ # import torch.optim as optim
5
+ import torch
6
+ # import torchvision.transforms as transforms
7
+ # import numpy as np
8
+ # import os
9
+ # from os import listdir
10
+ from os.path import join
11
+ # import pandas as pd
12
+ # import cv2 as cv
13
+ # from torch.utils.data import DataLoader, Dataset, random_split
14
+ # from PIL import Image, ImageOps
15
+ #from torchmetrics import Accuracy
16
+ from pytorch_lightning import LightningModule
17
+ # from pytorch_lightning.callbacks.early_stopping import EarlyStopping
18
+ # from pytorch_lightning.callbacks.model_checkpoint import ModelCheckpoint
19
+ #from pytorch_lightning.callbacks.progress import TQDMProgressBar
20
+ # from pytorch_lightning.loggers import TensorBoardLogger
21
+ # from torchvision import models
22
+ # import general_functions
23
+ # from data_loader import DatasetCostum
24
+ # import time
25
+ # import optuna
26
+ # from pytorch_lightning.callbacks.progress import TQDMProgressBar
27
+ # from pytorch_lightning.callbacks.early_stopping import EarlyStopping
28
+ # import statistics
29
+ # import panorama.datasets
30
+ # from optuna.integration import PyTorchLightningPruningCallback
31
+ # from graph_attention_layer_or import GATLayerImp3
32
+
33
+ # from pytorch_forecasting import TimeSeriesDataSet
34
+ # from pytorch_forecasting.data.encoders import TorchNormalizer
35
+ # from torch_geometric import EdgeIndex
36
+ # from graph.astgcn import ASTGCN
37
+ # from graph.gconv_gru import GConvGRU
38
+ #import ray
39
+ #from ray import tune
40
+ #from ray.tune.tuner import Tuner, TuneConfig
41
+ #from ray.tune.search.optuna import OptunaSearch
42
+ # import pickle
43
+ # import sqlite3
44
+ # import json
45
+ # from torch.utils.tensorboard import SummaryWriter
46
+ # from tgcn import TGCN2
47
+ from torch_geometric.nn import ChebConv, Sequential
48
+ # from torch.optim.lr_scheduler import ReduceLROnPlateau
49
+ # from GAT_layer import GATLayer, GraphAttentionLayer
50
+
51
+
52
+ class AutoEncoderModel(LightningModule):
53
+ def __init__(self, cuda_true, batch_size):
54
+ super().__init__()
55
+
56
+
57
+ self.epochs, self.conditions = list(), list()
58
+ self.recon_loss_test_step_list = list()
59
+ self.num_step = 0
60
+
61
+ if cuda_true:
62
+ self.dev = "cuda"
63
+ else:
64
+ self.dev = "cpu"
65
+ self.num_nodes = 15
66
+ # self.optimizer_name = hyper_params["optimizer_name"]
67
+
68
+
69
+ # self.tb_writer = summary_writer
70
+
71
+ self.edge_index_att = None
72
+
73
+ self.batch_size = batch_size
74
+ self.criterion = nn.MSELoss(reduction='mean')
75
+
76
+ self.steps = 0
77
+ self.recon_loss_train_step = 0
78
+ self.recon_loss_train_step_list = list()
79
+ self.recon_loss_tain_epoch_list = list()
80
+ self.recon_loss_val_step = 0
81
+ self.recon_loss_val_step_list = list()
82
+ self.recon_loss_val_epoch_list = list()
83
+ self.recon_loss_test_step_list = list()
84
+ self.epoch = 0
85
+
86
+ self.window = 64
87
+ self.automatic_optimization = True
88
+ self.test_target_data, self.test_predict_data = list(), list()
89
+ # self.output_decoder = torch.rand(self.batch_size, self.num_nodes, self.window) # check size
90
+ # self.output_decoder.requires_grad_()
91
+ self.output_first_layer_decoder = torch.rand(self.batch_size, self.num_nodes, self.window*4) # check size
92
+ self.output_first_layer_decoder.requires_grad_()
93
+ self.output_first_layer_decoder.to(self.dev)
94
+
95
+ self.node_num_featues = 5
96
+ self.total_feat = self.node_num_featues * self.window
97
+
98
+
99
+
100
+ ### Original Code
101
+ self.k = 4
102
+ latent_dim = 104
103
+ self.beta = 0.009256865323169841
104
+ self.encoder = Sequential('x, edge_index', [
105
+ (ChebConv(in_channels=self.window*self.node_num_featues, out_channels=self.window*2, K=self.k), 'x, edge_index -> x'),
106
+ nn.ReLU(inplace=True),
107
+ (ChebConv(in_channels=self.window*2, out_channels=self.window*4, K=self.k), 'x, edge_index -> x'),
108
+ nn.ReLU(inplace=True),
109
+ ])
110
+
111
+ self.encoder_2 = Sequential('x, edge_index', [
112
+ (ChebConv(in_channels=self.window, out_channels=self.window*4, K=self.k), 'x, edge_index -> x'),
113
+ nn.ReLU(inplace=True)
114
+ ])
115
+ self.latent = nn.Sequential(
116
+ nn.Flatten(),
117
+ nn.Linear(self.window*4*self.num_nodes, latent_dim),
118
+
119
+ nn.Linear(latent_dim, self.window*4*self.num_nodes),
120
+ nn.Unflatten(-1, (int(self.num_nodes), int(self.window*4)))
121
+ )
122
+
123
+
124
+ self.latent.to(self.dev)
125
+ self.decoder_2 = Sequential('x, edge_index' ,[
126
+ (ChebConv(in_channels=self.window*4, out_channels=self.window, K=self.k), 'x, edge_index -> x'),
127
+ nn.ReLU(inplace=True)
128
+ ])
129
+ self.decoder = Sequential('x, edge_index' ,[
130
+ (ChebConv(in_channels=self.window*4, out_channels=self.window*2, K=self.k), 'x, edge_index -> x'),
131
+ nn.ReLU(inplace=True),
132
+ (ChebConv(in_channels=self.window*2, out_channels=self.window*self.node_num_featues, K=self.k), 'x, edge_index -> x')
133
+ ])
134
+
135
+
136
+
137
+ self.softmax = nn.Softmax()
138
+
139
+
140
+ def forward(self, input_data, edge_indices, adj_matrix):
141
+ self.edge_indices = edge_indices
142
+ self.adj_matrix = adj_matrix
143
+ # print("input data", input_data)
144
+
145
+
146
+
147
+
148
+ 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]))
149
+ input_data_reshaped = input_data_reshaped.to(self.dev)
150
+ output_encoder = self.encoder(input_data_reshaped, self.edge_indices)
151
+ scaled_encoder = torch.mul(output_encoder, self.beta)
152
+ output_latent = self.latent(scaled_encoder)
153
+ output_decoder = self.decoder(output_latent, self.edge_indices)
154
+ self.output_decoder = torch.reshape(input=output_decoder, shape=(output_decoder.shape[0], self.window, self.num_nodes, self.node_num_featues))
155
+
156
+ recon_loss_list = list()
157
+ for i in range(input_data.shape[1]):
158
+ recon_loss_list.append(self.criterion(self.output_decoder[:,i,:,:], input_data[:,i,:,:]).to(self.dev))
159
+ recon_loss = sum(recon_loss_list)/len(recon_loss_list)
160
+
161
+ return recon_loss, self.output_decoder
162
+
163
+
164
+
165
+ def calc_edge_weight(edge_index, adj_matrix):
166
+ edge_weight = torch.rand(edge_index.shape[1])
167
+ for i, element in enumerate(edge_index.T):
168
+ edge_weight[i] = (adj_matrix[element[0]][element[1]] + adj_matrix[element[1]][element[0]])/2.0
169
+ return edge_weight
eval.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from HuggingFaceProjects.FalseDetector.AutoencoderCheb import AutoEncoderModel
3
+ model = AutoEncoderModel(cuda_true=True, batch_size=32)
4
+
5
+ model.load_state_dict(torch.load("/home/roth/git_projects/HuggingFaceProjects/FalseDetector/gae_25_08_2025.pt"))
6
+ print("model", model)
7
+ model.eval()
gae_25_08_2025.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:89fba40a12f9f9c534b6984f877fe3ba7d802054c8f1b4420be6aab2771c0d03
3
+ size 6112122