|
Download README.md from sallamsaka/M31-Coding-Test: direct link, hf CLI and curl.
- Browser
- Download file 5.13 kB
-
https://huggingface.co/sallamsaka/M31-Coding-Test/resolve/main/README.md
- Command line
-
hf download hf://sallamsaka/M31-Coding-Test/README.md
-
curl -L -o README.md https://huggingface.co/sallamsaka/M31-Coding-Test/resolve/main/README.md
5.13 kB
| 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. | |