(
self,
lighten,
dropout: float = 0.1,
activation: str = "relu",
norm_first: bool = False,
)
| 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) |
nothing calls this directly
no outgoing calls
no test coverage detected