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__() # Embedding layer for each categorical feature - uses Encoder outputs feed into nn.Embedding self.embeddings = nn.ModuleList([ nn.Embedding(dim, embed_dim) for dim in cat_dims ]) # Transformer Encoder Layer - treats categorical embeddings as a sequence of tokens 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) # Final MLP Head (MLP can classify non-linearly separable classes) # Concat: (Number of Cats * embed_dim) + Number of Continuous 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() # binary classification ) def forward(self, x_cat, x_num): # x_cat: [batch, len(cat_cols)] # x_num: [batch, len(num_cols)] # Embed each category and stack: [batch, num_cats, embed_dim] embeddings = [emb(x_cat[:, i]) for i, emb in enumerate(self.embeddings)] x = torch.stack(embeddings, dim=1) # Apply Transformer Attention across the "tokens" (features) x = self.transformer(x) # Flatten categorical tokens and concatenate with numerical features x = x.flatten(1) x = torch.cat([x, x_num], dim=1) return self.classifier(x)