| 272 | |
| 273 | # https://arxiv.org/abs/2311.06190 |
| 274 | class 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]) - \ |
nothing calls this directly
no outgoing calls
no test coverage detected