(self, q, k, v, mask=None)
| 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]) |
| 121 | _, S, _, _ = k.shape # (B, S, H, E) torch.Size([3, 96, 8, 2]) |
| 122 | |
| 123 | q = q.view(q.shape[0], q.shape[1], -1) |
| 124 | k = k.view(k.shape[0], k.shape[1], -1) |
| 125 | v = v.view(v.shape[0], v.shape[1], -1) |
| 126 | q = self.Lq(q) |
| 127 | q = q.view(q.shape[0], q.shape[1], self.c, self.k) |
| 128 | k = self.Lk(k) |
| 129 | k = k.view(k.shape[0], k.shape[1], self.c, self.k) |
| 130 | v = self.Lv(v) |
| 131 | v = v.view(v.shape[0], v.shape[1], self.c, self.k) |
| 132 | |
| 133 | if N > S: |
| 134 | zeros = torch.zeros_like(q[:, :(N - S), :]).float() |
| 135 | v = torch.cat([v, zeros], dim=1) |
| 136 | k = torch.cat([k, zeros], dim=1) |
| 137 | else: |
| 138 | v = v[:, :N, :, :] |
| 139 | k = k[:, :N, :, :] |
| 140 | |
| 141 | ns = math.floor(np.log2(N)) |
| 142 | nl = pow(2, math.ceil(np.log2(N))) |
| 143 | extra_q = q[:, 0:nl - N, :, :] |
| 144 | extra_k = k[:, 0:nl - N, :, :] |
| 145 | extra_v = v[:, 0:nl - N, :, :] |
| 146 | q = torch.cat([q, extra_q], 1) |
| 147 | k = torch.cat([k, extra_k], 1) |
| 148 | v = torch.cat([v, extra_v], 1) |
| 149 | |
| 150 | Ud_q = torch.jit.annotate(List[Tuple[Tensor]], []) |
| 151 | Ud_k = torch.jit.annotate(List[Tuple[Tensor]], []) |
| 152 | Ud_v = torch.jit.annotate(List[Tuple[Tensor]], []) |
| 153 | |
| 154 | Us_q = torch.jit.annotate(List[Tensor], []) |
| 155 | Us_k = torch.jit.annotate(List[Tensor], []) |
| 156 | Us_v = torch.jit.annotate(List[Tensor], []) |
| 157 | |
| 158 | Ud = torch.jit.annotate(List[Tensor], []) |
| 159 | Us = torch.jit.annotate(List[Tensor], []) |
| 160 | |
| 161 | # decompose |
| 162 | for i in range(ns - self.L): |
| 163 | # print('q shape',q.shape) |
| 164 | d, q = self.wavelet_transform(q) |
| 165 | Ud_q += [tuple([d, q])] |
| 166 | Us_q += [d] |
| 167 | for i in range(ns - self.L): |
| 168 | d, k = self.wavelet_transform(k) |
| 169 | Ud_k += [tuple([d, k])] |
| 170 | Us_k += [d] |
| 171 | for i in range(ns - self.L): |
| 172 | d, v = self.wavelet_transform(v) |
| 173 | Ud_v += [tuple([d, v])] |
| 174 | Us_v += [d] |
| 175 | for i in range(ns - self.L): |
| 176 | dk, sk = Ud_k[i], Us_k[i] |
nothing calls this directly
no test coverage detected