MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / FourierBlock

Class FourierBlock

layers/FourierCorrelation.py:24–76  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

22
23# ########## fourier layer #############
24class FourierBlock(nn.Module):
25 def __init__(self, in_channels, out_channels, n_heads, seq_len, modes=0, mode_select_method='random'):
26 super(FourierBlock, self).__init__()
27 print('fourier enhanced block used!')
28 """
29 1D Fourier block. It performs representation learning on frequency domain,
30 it does FFT, linear transform, and Inverse FFT.
31 """
32 # get modes on frequency domain
33 self.index = get_frequency_modes(seq_len, modes=modes, mode_select_method=mode_select_method)
34 print('modes={}, index={}'.format(modes, self.index))
35
36 self.n_heads = n_heads
37 self.scale = (1 / (in_channels * out_channels))
38 self.weights1 = nn.Parameter(
39 self.scale * torch.rand(self.n_heads, in_channels // self.n_heads, out_channels // self.n_heads,
40 len(self.index), dtype=torch.float))
41 self.weights2 = nn.Parameter(
42 self.scale * torch.rand(self.n_heads, in_channels // self.n_heads, out_channels // self.n_heads,
43 len(self.index), dtype=torch.float))
44
45 # Complex multiplication
46 def compl_mul1d(self, order, x, weights):
47 x_flag = True
48 w_flag = True
49 if not torch.is_complex(x):
50 x_flag = False
51 x = torch.complex(x, torch.zeros_like(x).to(x.device))
52 if not torch.is_complex(weights):
53 w_flag = False
54 weights = torch.complex(weights, torch.zeros_like(weights).to(weights.device))
55 if x_flag or w_flag:
56 return torch.complex(torch.einsum(order, x.real, weights.real) - torch.einsum(order, x.imag, weights.imag),
57 torch.einsum(order, x.real, weights.imag) + torch.einsum(order, x.imag, weights.real))
58 else:
59 return torch.einsum(order, x.real, weights.real)
60
61 def forward(self, q, k, v, mask):
62 # size = [B, L, H, E]
63 B, L, H, E = q.shape
64 x = q.permute(0, 2, 3, 1)
65 # Compute Fourier coefficients
66 x_ft = torch.fft.rfft(x, dim=-1)
67 # Perform Fourier neural operations
68 out_ft = torch.zeros(B, H, E, L // 2 + 1, device=x.device, dtype=torch.cfloat)
69 for wi, i in enumerate(self.index):
70 if i >= x_ft.shape[3] or wi >= out_ft.shape[3]:
71 continue
72 out_ft[:, :, :, wi] = self.compl_mul1d("bhi,hio->bho", x_ft[:, :, :, i],
73 torch.complex(self.weights1, self.weights2)[:, :, :, wi])
74 # Return to time domain
75 x = torch.fft.irfft(out_ft, n=x.size(-1))
76 return (x, None)
77
78# ########## Fourier Cross Former ####################
79class FourierCrossAttention(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected