(self, in_channels, out_channels, seq_len_q, seq_len_kv, modes, c=64,
k=8, ich=512,
L=0,
base='legendre',
mode_select_method='random',
initializer=None, activation='tanh',
**kwargs)
| 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 | |
| 119 | def forward(self, q, k, v, mask=None): |
| 120 | B, N, H, E = q.shape # (B, N, H, E) torch.Size([3, 768, 8, 2]) |
nothing calls this directly
no test coverage detected