(self, x)
| 91 | self.cat_head = nn.Linear(d_model, 3) |
| 92 | |
| 93 | def forward(self, x): |
| 94 | x = self.embed(x.squeeze(1)) |
| 95 | |
| 96 | embed_pos = self.embed_positions(x.shape) |
| 97 | |
| 98 | # transformer encoder |
| 99 | x = self.transformer_encoder(x + embed_pos) |
| 100 | |
| 101 | # mean pool for classification |
| 102 | x = torch.mean(x, dim=1) |
| 103 | |
| 104 | logits = self.cat_head(x) |
| 105 | return logits |