M31-Coding-Test / README.md
sallamsaka's picture
Upload README.md with huggingface_hub
ec18ab9 verified
|
Raw History Blame Contribute Delete
5.13 kB
metadata
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.