(self, x)
| 151 | self.head = M.Linear(embed_dim, 1) |
| 152 | |
| 153 | def forward(self, x): |
| 154 | x = self.proj(x) |
| 155 | x = F.flatten(x, 2).transpose(0, 2, 1) |
| 156 | |
| 157 | x = self.extra(x) |
| 158 | x = x.mean(axis=1) |
| 159 | x = self.head(x) |
| 160 | x = x.sum() |
| 161 | return x |
| 162 | |
| 163 | |
| 164 | def test_ViTmode_trace_train(): |