1D Multiwavelet Cross Attention layer.
| 59 | |
| 60 | |
| 61 | class MultiWaveletCross(nn.Module): |
| 62 | """ |
| 63 | 1D Multiwavelet Cross Attention layer. |
| 64 | """ |
| 65 | def __init__(self, in_channels, out_channels, seq_len_q, seq_len_kv, modes, c=64, |
| 66 | k=8, ich=512, |
| 67 | L=0, |
| 68 | base='legendre', |
| 69 | mode_select_method='random', |
| 70 | initializer=None, activation='tanh', |
| 71 | **kwargs): |
| 72 | super(MultiWaveletCross, self).__init__() |
| 73 | print('base', base) |
| 74 | |
| 75 | self.c = c |
| 76 | self.k = k |
| 77 | self.L = L |
| 78 | H0, H1, G0, G1, PHI0, PHI1 = get_filter(base, k) |
| 79 | H0r = H0 @ PHI0 |
| 80 | G0r = G0 @ PHI0 |
| 81 | H1r = H1 @ PHI1 |
| 82 | G1r = G1 @ PHI1 |
| 83 | |
| 84 | H0r[np.abs(H0r) < 1e-8] = 0 |
| 85 | H1r[np.abs(H1r) < 1e-8] = 0 |
| 86 | G0r[np.abs(G0r) < 1e-8] = 0 |
| 87 | G1r[np.abs(G1r) < 1e-8] = 0 |
| 88 | self.max_item = 3 |
| 89 | |
| 90 | self.attn1 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q, |
| 91 | seq_len_kv=seq_len_kv, modes=modes, activation=activation, |
| 92 | mode_select_method=mode_select_method) |
| 93 | self.attn2 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q, |
| 94 | seq_len_kv=seq_len_kv, modes=modes, activation=activation, |
| 95 | mode_select_method=mode_select_method) |
| 96 | self.attn3 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q, |
| 97 | seq_len_kv=seq_len_kv, modes=modes, activation=activation, |
| 98 | mode_select_method=mode_select_method) |
| 99 | self.attn4 = FourierCrossAttentionW(in_channels=in_channels, out_channels=out_channels, seq_len_q=seq_len_q, |
| 100 | seq_len_kv=seq_len_kv, modes=modes, activation=activation, |
| 101 | mode_select_method=mode_select_method) |
| 102 | self.T0 = nn.Linear(k, k) |
| 103 | self.register_buffer('ec_s', torch.Tensor( |
| 104 | np.concatenate((H0.T, H1.T), axis=0))) |
| 105 | self.register_buffer('ec_d', torch.Tensor( |
| 106 | np.concatenate((G0.T, G1.T), axis=0))) |
| 107 | |
| 108 | self.register_buffer('rc_e', torch.Tensor( |
| 109 | np.concatenate((H0r, G0r), axis=0))) |
| 110 | self.register_buffer('rc_o', torch.Tensor( |
| 111 | np.concatenate((H1r, G1r), axis=0))) |
| 112 | |
| 113 | self.Lk = nn.Linear(ich, c * k) |
| 114 | self.Lq = nn.Linear(ich, c * k) |
| 115 | self.Lv = nn.Linear(ich, c * k) |
| 116 | self.out = nn.Linear(c * k, ich) |
| 117 | self.modes1 = modes |
| 118 |