Views
No views yet
| Architecture | FT-Transformer (d_token=256, 6 layers, 8 heads, 4.77M params) |
| Task | Binary classification — hypotension during session (yes/no) |
| Dataset | Mendeley 10.17632/7kmtsmsgfw.1 — 97,640 sessions, 758 patients |
| Training split | 73,034 train / 9,960 val / 14,646 test |
| Metric | Value |
|---|---|
| ROC-AUC | 0.7841 |
| Average Precision | 0.5849 |
| F1 (threshold=0.5) | 0.5557 |
| Hypotension recall | ~75% |
1import torch
2import torch.nn as nn
3import math
4
5class FeatureTokenizer(nn.Module):
6 def __init__(self, n_num, cat_cardinalities, d_token):
7 super().__init__()
8 self.num_weight = nn.Parameter(torch.empty(n_num, d_token))
9 self.num_bias = nn.Parameter(torch.empty(n_num, d_token))
10 self.cat_embs = nn.ModuleList([nn.Embedding(c+1, d_token) for c in cat_cardinalities])
11 nn.init.kaiming_uniform_(self.num_weight, a=math.sqrt(5))
12 nn.init.zeros_(self.num_bias)
13 def forward(self, x_num, x_cat):
14 x_n = x_num.unsqueeze(-1) * self.num_weight + self.num_bias
15 x_c = torch.stack([e(x_cat[:, i]) for i, e in enumerate(self.cat_embs)], dim=1)
16 return torch.cat([x_n, x_c], dim=1)
17
18class FTTransformer(nn.Module):
19 def __init__(self, n_num, cat_cardinalities, d_token=256, n_heads=8, n_layers=6, dropout=0.1):
20 super().__init__()
21 self.tokenizer = FeatureTokenizer(n_num, cat_cardinalities, d_token)
22 self.cls_token = nn.Parameter(torch.zeros(1, 1, d_token))
23 enc = nn.TransformerEncoderLayer(d_model=d_token, nhead=n_heads,
24 dim_feedforward=4*d_token, dropout=dropout, batch_first=True,
25 norm_first=True, activation='gelu')
26 self.encoder = nn.TransformerEncoder(enc, num_layers=n_layers)
27 self.head = nn.Sequential(nn.LayerNorm(d_token), nn.Linear(d_token, 1))
28 def forward(self, x_num, x_cat):
29 x = self.tokenizer(x_num, x_cat)
30 cls = self.cls_token.expand(x.size(0), -1, -1)
31 x = self.encoder(torch.cat([cls, x], dim=1))
32 return self.head(x[:, 0]).squeeze(-1)
33
34# Load
35ckpt = torch.load("hemodialysis_ft.pt", map_location="cpu")
36model = FTTransformer(n_num=20, cat_cardinalities=ckpt["cat_cardinalities"])
37model.load_state_dict(ckpt["model_state"])
38model.eval()
39
40# Inference (x_num normalized with ckpt scaler_mean/scaler_scale)
41prob = torch.sigmoid(model(x_num, x_cat)).item()
42print(f"Hypotension probability: {prob:.1%}")