MCPcopy Create free account
hub / github.com/thuml/Time-Series-Library / forward

Method forward

models/MSGNet.py:41–78  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

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
81class Model(nn.Module):

Callers

nothing calls this directly

Calls 1

FFT_for_PeriodFunction · 0.70

Tested by

no test coverage detected