MCPcopy Create free account
hub / github.com/Anoise/WTFlib / forward

Method forward

LDPS_Graph/layers/MultiWaveletCorrelation.py:119–189  ·  view source on GitHub ↗
(self, q, k, v, mask=None)

Source from the content-addressed store, hash-verified

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]

Callers

nothing calls this directly

Calls 2

wavelet_transformMethod · 0.95
evenOddMethod · 0.95

Tested by

no test coverage detected