--- 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.