| import torch |
| import torch.nn as nn |
|
|
| class HybridTabTransformer(nn.Module): |
| def __init__(self, cat_dims, num_continuous, embed_dim=32, n_heads=4, n_layers=2): |
| super().__init__() |
| |
| |
| self.embeddings = nn.ModuleList([ |
| nn.Embedding(dim, embed_dim) for dim in cat_dims |
| ]) |
| |
| |
| encoder_layer = nn.TransformerEncoderLayer( |
| d_model=embed_dim, |
| nhead=n_heads, |
| dim_feedforward=128, |
| batch_first=True, |
| dropout=0.1 |
| ) |
|
|
| self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_layers) |
| |
| |
| |
| self.combined_dim = (len(cat_dims) * embed_dim) + num_continuous |
| |
| self.classifier = nn.Sequential( |
| nn.Linear(self.combined_dim, 64), |
| nn.BatchNorm1d(64), |
| nn.ReLU(), |
| nn.Dropout(0.2), |
| nn.Linear(64, 1), |
| nn.Sigmoid() |
| ) |
|
|
| def forward(self, x_cat, x_num): |
| |
| |
| |
| |
| embeddings = [emb(x_cat[:, i]) for i, emb in enumerate(self.embeddings)] |
| x = torch.stack(embeddings, dim=1) |
| |
| |
| x = self.transformer(x) |
| |
| |
| x = x.flatten(1) |
| x = torch.cat([x, x_num], dim=1) |
| |
| return self.classifier(x) |