Skip to content

Add XGBoostModel (gradient-boosted trees) with exact TreeSHAP - #1278

Merged
jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:feat/xgboost-model
Oct 8, 2026
Merged

jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:feat/xgboost-model

Conversation

@solarsys

@solarsys solarsys commented Oct 8, 2026

Copy link
Copy Markdown
Collaborator

Why

PyHealth has no gradient-boosted tree model; the only tree model is CaliForest. In a downstream pediatric-asthma early-warning project (8,156 landmark samples, 92 engineered tabular features, ~15% positives), XGBoost on the engineered features was the best model on held-out patients (PR-AUC 0.545, AUROC 0.834). It beat all nine PyHealth models tried. Because of that, the project had to run XGBoost outside PyHealth for its main model.

This PR adds XGBoostModel, built on a small GradientBoostedTreeModel base so a LightGBM variant can follow (LightGBM itself is out of scope).

Usage

model = XGBoostModel(samples, bag_of_codes=True, n_estimators=400, max_depth=4,
                     learning_rate=0.05, scale_pos_weight="balanced")
model.fit(train_loader, val_loader)          # once, outside Trainer.train
trainer = Trainer(model=model)
trainer.evaluate(test_loader)                # unchanged Trainer / metrics
trainer.save_ckpt("xgb.ckpt")                # fitted trees round-trip
model.mean_abs_shap(test_loader)             # global TreeSHAP ranking
TreeSHAP(model).attribute(**batch)           # per-field, like other interpreters

Design choices

Item Choice
Column order input_schema order (model.feature_keys). Per-field column ranges are in model.feature_layout, names in model.feature_names.
Inputs Accepted: tensor, timeseries (fixed length), multi_hot, nested_multihot (summed over visits), label processors used as inputs, and custom numeric processors. Dense fields must keep the same per-sample shape in every batch; otherwise a ValueError names the field. Other PyHealth processors raise a ValueError naming the field.
Padded sequences sequence / nested_sequence / deep_nested_sequence are never flattened. They raise a ValueError naming the field, unless you opt in with bag_of_codes=True. That encodes them as per-sample counts over the vocabulary, excluding <pad>/<unk>, so the width is the vocabulary size minus 2, whatever the batch padding. With this, standard MIMIC tasks work.
Modes binary → XGBClassifier(binary:logistic); multiclass → multi:softprob; regression → XGBRegressor. Multilabel → one booster per label, so each label gets its own scale_pos_weight and early-stopping round; xgboost's native multi-output takes only one weight.
Outputs logit is the booster margin (output_margin=True) and y_prob is the booster's own probability, with no clipping or inversion. loss comes from get_loss_function() on the logits, and tensors are on self.device.
Imbalance scale_pos_weight: a number, one per label, "balanced" (the name used in #1267's set_pos_weight, which isn't merged yet) or alias "auto", computed as negatives/positives on the training labels. The docs say it inflates probabilities, so recalibrate on a patient-grouped split (#1272 / Platt scaling). The model does not calibrate silently.
Early stopping Off by default. With early_stopping_rounds and eval_metric (e.g. aucpr) plus val_data, predictions and SHAP use best_iteration.
Missing values missing=np.nan by default, so XGBoost learns where missing values go at each split. The docs note that this differs from median imputation and recommend float64 statistics cast to float32 once.
Persistence get_extra_state / set_extra_state store the raw UBJ bytes of each booster as uint8 tensors, plus the feature layout as JSON, so it loads under torch.load(weights_only=True). On load, the fields, processors, mode and vocabulary widths are checked, and a mismatch raises a ValueError.
Trainer.train It raises a TypeError pointing to model.fit(...) for any model that sets fit_outside_trainer = True. This is the only Trainer change. I chose raising over silently calling fit, because fit would ignore epochs, the optimizer and the other arguments.
Interpretation Exact TreeSHAP from pred_contribs=True: model.explain(**batch) gives per-field columns plus bias, and model.mean_abs_shap(loader) gives a global ranking. The new pyhealth.interpret.methods.TreeSHAP returns attributions in each field's input shape; each code's value is split over the positions holding it.
Dependency Optional extra: pip install "pyhealth[xgboost]" (xgboost>=2.0). If it's missing, constructing the model raises an ImportError with an install hint. Importing pyhealth.models doesn't need xgboost.
Memory The training matrix is dense float32 (n × columns × 4 bytes); this is documented. QuantileDMatrix / DataIter isn't used, which keeps exact parity with XGBClassifier.

Validation (tests/core/test_xgboost_model.py, 21 tests, ~11 s)

  • Parity, native missing values: bit-identical. I built a 92-dim tensor feature with ~10% NaN, one all-NaN column, one constant column and 15% positives, split by patient. XGBoostModel was compared against xgboost.XGBClassifier with n_estimators=400, max_depth=4, learning_rate=0.05, subsample=0.8, colsample_bytree=0.8, tree_method="hist", scale_pos_weight=auto, random_state=0 on the same matrix. Test probabilities from Trainer.inference are exactly equal (assert_array_equal).
  • Parity after median impute + z-score: bit-identical. A processor fitted on training samples only (float64 statistics, cast once; all-NaN → 0; zero variance → scale 1) was compared against SimpleImputer(median, keep_empty_features=True) + StandardScaler. The matrices and the probabilities are exactly equal.
  • Determinism: two fits give identical predictions at n_jobs=1 and at n_jobs=4.
  • Modes: binary, multiclass, multilabel and regression each fit and run Trainer.inference / evaluate with the right shapes and sensible metrics. Trainer.train raises.
  • Inputs:
    • Each fixed-width processor type works, and the column ranges and names are checked.
    • sequence raises an error naming the field; nested_sequence_floats raises too.
    • A tensor whose width varies between batches raises.
    • Bag-of-codes gives identical matrices with batch size 1 and 32 (different padding), and the counts match a hand computation.
    • drop_last=True loaders are rejected.
  • Early stopping: best_iteration and the probabilities equal XGBClassifier with the same eval set (aucpr), and they differ from using all trees.
  • Persistence: save_ckpt then load_ckpt into a fresh model gives identical y_prob and logit in all 4 modes (one with early stopping). A different field set raises, and so does a different vocabulary size.
  • SHAP:
    • Contributions plus bias equal the margin (max error ~5e-7) in all modes.
    • Mean |SHAP| matches shap.TreeExplainer (rtol 1e-4), with the same top-10 ranking. This check skips if shap isn't installed.
    • Interpreter attributions have the input shapes, sum correctly over code positions, and are 0 on padding.
  • Without xgboost: 20 tests skip cleanly, and the install-hint test passes.
  • Full core suite (fresh cache): 1,444 tests, OK. Docstring examples pass and tools/check_pr_rules.py passes.

Example: examples/mortality_prediction/mortality_mimic3_xgboost.py, on synthetic MIMIC-III (2,196 samples, patient split 70/10/20, test set):

model PR-AUC AUROC
XGBoost (bag of codes) 0.143 0.687
LogisticRegression 0.085 0.621
MLP 0.059 0.675

The synthetic data has little real signal, so this shows the workflow, not clinical performance.

Notes

  • macOS / OpenMP: torch wheels bundle libomp.dylib and XGBoost loads Homebrew's, and two copies in one process can crash multithreaded fits. The model and install docs describe the fix (make torch's libomp.dylib a symlink to Homebrew's, or use n_jobs=1). The multithreaded determinism subtest detects two loaded runtimes on macOS and skips instead of crashing. CI runs on Ubuntu only, so it's unaffected.
  • CI coverage: the pixi test environment doesn't install xgboost, so these tests currently skip in CI. Adding xgboost to [tool.pixi.feature.test.pypi-dependencies] would run them, but needs a pixi.lock refresh. I can do that in this PR or a follow-up if you'd like.
  • CaliForest: it flattens padded sequence fields (batch-dependent width) and builds logits by inverting clipped probabilities. Neither pattern is used here. I'll file the flattening issue separately rather than change CaliForest in this PR.
  • Downstream parity: I'll reproduce the asthma project's held-out result (PR-AUC 0.5449, AUROC 0.8335) and its SHAP ranking with this wrapper, on that project's exact split and features.

🤖 Generated with Claude Code

PyHealth had no gradient-boosted tree model, the standard strong baseline
for tabular and bag-of-codes EHR prediction. XGBoostModel is built on a
small GradientBoostedTreeModel base so a LightGBM variant can follow.

- Fit once with model.fit(train, val=None); forward() returns
  {loss, y_prob, y_true, logit}, so Trainer.inference/evaluate and the
  metrics work unchanged. Trainer.train raises a TypeError pointing to
  model.fit for models that set fit_outside_trainer.
- Inputs: tensor, timeseries (fixed length), multi_hot, nested_multihot
  (summed over visits), label processors as inputs, and custom numeric
  processors, concatenated in input_schema order with recorded column
  ranges and names. Padded code sequences raise an error naming the field,
  or with bag_of_codes=True become per-sample counts over the vocabulary
  (excluding <pad>/<unk>), so the width does not depend on batch padding.
- Modes: binary (XGBClassifier), multiclass (multi:softprob), multilabel
  (one booster per label), regression (XGBRegressor). logit is the booster
  margin, y_prob the booster's probabilities; no clipping or inversion.
- scale_pos_weight: number, per label, or "balanced"/"auto" from the
  training labels. Optional early stopping on val data with eval_metric.
- NaN passes through as missing. Checkpoints store the boosters as uint8
  tensors via get/set_extra_state (loads with weights_only=True) together
  with the feature layout, which is checked on load.
- explain() and mean_abs_shap() give exact TreeSHAP from pred_contribs;
  pyhealth.interpret.methods.TreeSHAP maps them onto input fields.
- xgboost is an optional extra: pip install "pyhealth[xgboost]".

Co-Authored-By: Claude Opus 5.5 <[email protected]>
@solarsys
solarsys requested a review from jhnwu3 October 8, 2026 19:41

@jhnwu3 jhnwu3 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ohh, always wanted one of these models for PyHealth!

Lgtm.

@jhnwu3
jhnwu3 merged commit 2db29c7 into sunlabuiuc:master Oct 8, 2026
2 checks passed
@solarsys
solarsys deleted the feat/xgboost-model branch October 9, 2026 04:30
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants