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.