| 4 | |
| 5 | |
| 6 | class ITransformer(pl.LightningModule): |
| 7 | def __init__( |
| 8 | self, |
| 9 | lighten, |
| 10 | dropout: float = 0.1, |
| 11 | activation: str = "relu", |
| 12 | norm_first: bool = False, |
| 13 | ): |
| 14 | super().__init__() |
| 15 | self.name = "itransformer" |
| 16 | if lighten: |
| 17 | self.name += "-lighten" |
| 18 | |
| 19 | d_model = 64 if not lighten else 32 |
| 20 | dim_feedforward = 256 if not lighten else 128 |
| 21 | nhead = 8 if not lighten else 4 |
| 22 | num_layers = 2 if not lighten else 1 |
| 23 | |
| 24 | self.embed = nn.Linear(100, d_model, bias=False) |
| 25 | layer_norm_eps: float = 1e-5 |
| 26 | encoder_layer = nn.TransformerEncoderLayer( |
| 27 | d_model=d_model, |
| 28 | nhead=nhead, |
| 29 | dim_feedforward=dim_feedforward, |
| 30 | dropout=dropout, |
| 31 | activation=activation, |
| 32 | layer_norm_eps=layer_norm_eps, |
| 33 | norm_first=norm_first, |
| 34 | batch_first=True, |
| 35 | ) |
| 36 | encoder_norm = nn.LayerNorm(d_model, eps=layer_norm_eps) |
| 37 | self.transformer_encoder = nn.TransformerEncoder( |
| 38 | encoder_layer, num_layers=num_layers, norm=encoder_norm |
| 39 | ) |
| 40 | self.cat_head = nn.Linear(d_model, 3) |
| 41 | |
| 42 | def forward(self, x): |
| 43 | x = x.squeeze(1) |
| 44 | # transpose |
| 45 | x = x.permute(0, 2, 1) |
| 46 | x = self.embed(x) |
| 47 | |
| 48 | # transformer encoder |
| 49 | x = self.transformer_encoder(x) |
| 50 | |
| 51 | # mean pool for classification |
| 52 | x = torch.mean(x, dim=1) |
| 53 | |
| 54 | logits = self.cat_head(x) |
| 55 | return logits |