(self, x)
| 39 | |
| 40 | |
| 41 | def forward(self, x): |
| 42 | B, T, N = x.size() |
| 43 | scale_list, scale_weight = FFT_for_Period(x, self.k) |
| 44 | res = [] |
| 45 | for i in range(self.k): |
| 46 | scale = scale_list[i] |
| 47 | #Gconv |
| 48 | x = self.gconv[i](x) |
| 49 | # paddng |
| 50 | if (self.seq_len) % scale != 0: |
| 51 | length = (((self.seq_len) // scale) + 1) * scale |
| 52 | padding = torch.zeros([x.shape[0], (length - (self.seq_len)), x.shape[2]]).to(x.device) |
| 53 | out = torch.cat([x, padding], dim=1) |
| 54 | else: |
| 55 | length = self.seq_len |
| 56 | out = x |
| 57 | out = out.reshape(B, length // scale, scale, N) |
| 58 | |
| 59 | #for Mul-attetion |
| 60 | out = out.reshape(-1 , scale , N) |
| 61 | out = self.norm(self.att0(out)) |
| 62 | out = self.gelu(out) |
| 63 | out = out.reshape(B, -1 , scale , N).reshape(B ,-1 ,N) |
| 64 | # #for simpleVIT |
| 65 | # out = self.att(out.permute(0, 3, 1, 2).contiguous()) #return |
| 66 | # out = out.permute(0, 2, 3, 1).reshape(B, -1 ,N) |
| 67 | |
| 68 | out = out[:, :self.seq_len, :] |
| 69 | res.append(out) |
| 70 | |
| 71 | res = torch.stack(res, dim=-1) |
| 72 | # adaptive aggregation |
| 73 | scale_weight = F.softmax(scale_weight, dim=1) |
| 74 | scale_weight = scale_weight.unsqueeze(1).unsqueeze(1).repeat(1, T, N, 1) |
| 75 | res = torch.sum(res * scale_weight, -1) |
| 76 | # residual connection |
| 77 | res = res + x |
| 78 | return res |
| 79 | |
| 80 | |
| 81 | class Model(nn.Module): |
nothing calls this directly
no test coverage detected