| 22 | |
| 23 | # ########## fourier layer ############# |
| 24 | class 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 #################### |
| 79 | class FourierCrossAttention(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected