MCPcopy Create free account
hub / github.com/IntelligentSystemsLab/ST-EVCDP / FGN

Class FGN

baselines.py:274–418  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

272
273# https://arxiv.org/abs/2311.06190
274class FGN(nn.Module):
275 def __init__(self, pre_length=1, embed_size=64,
276 feature_size=0, seq_length=12, hidden_size=32, hard_thresholding_fraction=1, hidden_size_factor=1, sparsity_threshold=0.01):
277 super().__init__()
278 self.embed_size = embed_size
279 self.hidden_size = hidden_size
280 self.number_frequency = 1
281 self.pre_length = pre_length
282 self.feature_size = feature_size
283 self.seq_length = seq_length
284 self.frequency_size = self.embed_size // self.number_frequency
285 self.hidden_size_factor = hidden_size_factor
286 self.sparsity_threshold = sparsity_threshold
287 self.hard_thresholding_fraction = hard_thresholding_fraction
288 self.scale = 0.02
289 self.embeddings = nn.Parameter(torch.randn(1, self.embed_size))
290
291 self.encoder = nn.Linear(2, 1)
292 self.w1 = nn.Parameter(
293 self.scale * torch.randn(2, self.frequency_size, self.frequency_size * self.hidden_size_factor))
294 self.b1 = nn.Parameter(self.scale * torch.randn(2, self.frequency_size * self.hidden_size_factor))
295 self.w2 = nn.Parameter(
296 self.scale * torch.randn(2, self.frequency_size * self.hidden_size_factor, self.frequency_size))
297 self.b2 = nn.Parameter(self.scale * torch.randn(2, self.frequency_size))
298 self.w3 = nn.Parameter(
299 self.scale * torch.randn(2, self.frequency_size,
300 self.frequency_size * self.hidden_size_factor))
301 self.b3 = nn.Parameter(
302 self.scale * torch.randn(2, self.frequency_size * self.hidden_size_factor))
303 self.embeddings_10 = nn.Parameter(torch.randn(self.seq_length, 8))
304 self.fc = nn.Sequential(
305 nn.Linear(self.embed_size * 8, 64),
306 nn.LeakyReLU(),
307 nn.Linear(64, self.hidden_size),
308 nn.LeakyReLU(),
309 nn.Linear(self.hidden_size, self.pre_length)
310 )
311 self.to('cuda:0')
312
313 def tokenEmb(self, x):
314 x = x.unsqueeze(2)
315 y = self.embeddings
316 return x * y
317
318 # FourierGNN
319 def fourierGC(self, x, B, N, L):
320 o1_real = torch.zeros([B, (N*L)//2 + 1, self.frequency_size * self.hidden_size_factor],
321 device=x.device)
322 o1_imag = torch.zeros([B, (N*L)//2 + 1, self.frequency_size * self.hidden_size_factor],
323 device=x.device)
324 o2_real = torch.zeros(x.shape, device=x.device)
325 o2_imag = torch.zeros(x.shape, device=x.device)
326
327 o3_real = torch.zeros(x.shape, device=x.device)
328 o3_imag = torch.zeros(x.shape, device=x.device)
329
330 o1_real = F.relu(
331 torch.einsum('bli,ii->bli', x.real, self.w1[0]) - \

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected