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.
Get the code and install its dependencies:
git clone https://github.com/Sallamsaka/M31-Coding-Test cd M31-Coding-Test pip install -r requirements.txtCopy 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 cloneRun:
python -m src.reproduce --from-hub sallamsaka/M31-Coding-TestThis 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).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.