File size: 5,132 Bytes
33da14d f0f4988 33da14d 3f8aca1 37d7f7d 3f8aca1 33da14d 37d7f7d 33da14d 37d7f7d 33da14d 37d7f7d 33da14d 145b1d2 33da14d ec18ab9 37d7f7d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | ---
license: mit
tags:
- tabular-classification
- healthcare
- synthetic-data
- ehr
library_name: pytorch
---
# Patient Timeline Forecasting — ens_lr_gbdt_tx
Predicts which of 40 target conditions are **newly** diagnosed in the five
years after a patient's anchor date, from structured Synthea EHR events
recorded strictly before it.
Trained for the M31 research-intern take-home, on synthetic data only.
## Model
`sigmoid((logit(LR) + logit(GBDT) + logit(transformer)) / 3)`, per-label Platt
calibrated (fitted out-of-fold). The transformer is 1 layer / width 64 / 2 heads,
reads one token per event with a Time2Vec encoding of time before the anchor **and
of the patient's age at that event**, and fuses the 3,320 tabular features at the
readout; 3 seeds averaged in logit space. Files: `model_lr.joblib`,
`model_gbdt.joblib`, `model_transformer.joblib` (the transformer's output matrix),
`transformer_seed300.pt` / `transformer_seed301.pt` / `transformer_seed302.pt`
(its weights, one per seed), `platt_params.npz` (the 40 calibration maps),
`submission_manifest.joblib` (how the three models are combined), `vocab.json`
(the transformer's vocabulary).
## Task
- 3,514 Synthea patients, split 2,791 train / 365 validation / 358 test.
- Anchor = last recorded encounter minus five calendar years, floored to midnight.
- Label = **first-ever** diagnosis of the condition falls in `[anchor, anchor+5y]`.
A patient already diagnosed before the anchor is *prevalent*: structurally
negative, and reported separately rather than counted as an ordinary negative.
## Results
The test set's outcomes are withheld, so performance is estimated by 5-fold
cross-validation over the 2,791 training patients: each fold is predicted by
models trained on the other four (the transformer uses the provided validation
set only to choose its stopping epoch). Mean ± SD over the 5 test folds:
| model | macro AUROC | macro AP |
|---|---|---|
| logistic regression | 0.7595 ± 0.0056 | 0.2077 ± 0.0192 |
| gradient-boosted trees | 0.7640 ± 0.0081 | 0.2260 ± 0.0185 |
| transformer (3 seeds) | 0.7661 ± 0.0065 | 0.2203 ± 0.0157 |
| **submitted: blend of the three, calibrated** | **0.7755 ± 0.0028** | **0.2435 ± 0.0220** |
(Calibration in the last row is cross-fitted: fitted on four folds, scored on the fifth.)
## What the model is actually learning
The anchor is defined from each patient's last encounter, so for the **42.9%**
of training patients with a death date, the outcome window is exactly their
last five years of life — 100% of those deaths fall within 30 days of the
window end. A substantial part of the achievable signal is therefore
"is this record about to end", which is a property of how the task was
constructed rather than of clinical prediction. The same rule generated the
test anchors, so this is not leakage, but it does bound how the results should
be interpreted.
## Leakage controls
`DEATHDATE`, `HEALTHCARE_EXPENSES` and `HEALTHCARE_COVERAGE` are refused at
load time. `STOP`-derived durations are excluded: the organisers blanked
post-anchor stops in the test split, so such a feature would both leak and
shift. Every fitted statistic — vocabulary, quantile edges, scalers — is fitted
on training patients only. 152 automated checks cover this, including a grep
test that no module outside the time utility parses a timestamp.
## Reproducing the predictions
The files on this page are enough to rebuild the submitted `predictions.csv`
exactly, without any training. The patient data is not included here: you need
the dataset provided with the M31 take-home.
**You need** Python 3.13, git, and that dataset.
1. Get the code and install its dependencies:
```
git clone https://github.com/Sallamsaka/M31-Coding-Test
cd M31-Coding-Test
pip install -r requirements.txt
```
2. Copy the provided data into that folder, so that it looks like this:
```
M31-Coding-Test/
train_val/ provided CSV tables, train and validation patients
test/ provided CSV tables, test patients
patient_splits.csv
target_conditions.csv
test_anchors.csv
src/, outputs/, ... already there from the clone
```
3. Run:
```
python -m src.reproduce --from-hub sallamsaka/M31-Coding-Test
```
This downloads the model files from this page, builds each test patient's
features and event sequence, runs logistic regression, the boosted trees and
the three transformer seeds, averages and calibrates them, and writes
`outputs/predictions_reproduced.csv` (358 patients, 40 conditions).
4. Check the last line it prints. The clone already contains the submitted
`outputs/predictions.csv`, and the script compares against it:
```
max |reproduced - outputs/predictions.csv| = 7.40e-08 (MATCH)
```
Anything below 1e-5 prints `MATCH`; the remaining difference is floating-point
rounding.
Retraining everything from scratch is described in the GitHub README.
|