MCPcopy Create free account
hub / github.com/FinancialComputingUCL/LOBFrame / ITransformer

Class ITransformer

models/iTransformer/itransformer.py:6–55  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected