Repository navigation
Add XGBoostModel (gradient-boosted trees) with exact TreeSHAP - #1278
Merged
Merged
Conversation
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]>
jhnwu3
approved these changes
Oct 8, 2026
jhnwu3
left a comment
Collaborator
There was a problem hiding this comment.
Ohh, always wanted one of these models for PyHealth!
Lgtm.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 smallGradientBoostedTreeModelbase so a LightGBM variant can follow (LightGBM itself is out of scope).Usage
Design choices
input_schemaorder (model.feature_keys). Per-field column ranges are inmodel.feature_layout, names inmodel.feature_names.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 aValueErrornames the field. Other PyHealth processors raise aValueErrornaming the field.sequence/nested_sequence/deep_nested_sequenceare never flattened. They raise aValueErrornaming the field, unless you opt in withbag_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.XGBClassifier(binary:logistic); multiclass →multi:softprob; regression →XGBRegressor. Multilabel → one booster per label, so each label gets its ownscale_pos_weightand early-stopping round; xgboost's native multi-output takes only one weight.logitis the booster margin (output_margin=True) andy_probis the booster's own probability, with no clipping or inversion.losscomes fromget_loss_function()on the logits, and tensors are onself.device.scale_pos_weight: a number, one per label,"balanced"(the name used in #1267'sset_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_roundsandeval_metric(e.g.aucpr) plusval_data, predictions and SHAP usebest_iteration.missing=np.nanby 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.get_extra_state/set_extra_statestore the raw UBJ bytes of each booster as uint8 tensors, plus the feature layout as JSON, so it loads undertorch.load(weights_only=True). On load, the fields, processors, mode and vocabulary widths are checked, and a mismatch raises aValueError.Trainer.trainTypeErrorpointing tomodel.fit(...)for any model that setsfit_outside_trainer = True. This is the only Trainer change. I chose raising over silently callingfit, becausefitwould ignoreepochs, the optimizer and the other arguments.pred_contribs=True:model.explain(**batch)gives per-field columns plus bias, andmodel.mean_abs_shap(loader)gives a global ranking. The newpyhealth.interpret.methods.TreeSHAPreturns attributions in each field's input shape; each code's value is split over the positions holding it.pip install "pyhealth[xgboost]"(xgboost>=2.0). If it's missing, constructing the model raises anImportErrorwith an install hint. Importingpyhealth.modelsdoesn't need xgboost.n × columns × 4bytes); this is documented. QuantileDMatrix / DataIter isn't used, which keeps exact parity withXGBClassifier.Validation (
tests/core/test_xgboost_model.py, 21 tests, ~11 s)XGBoostModelwas compared againstxgboost.XGBClassifierwithn_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=0on the same matrix. Test probabilities fromTrainer.inferenceare exactly equal (assert_array_equal).SimpleImputer(median, keep_empty_features=True) + StandardScaler. The matrices and the probabilities are exactly equal.n_jobs=1and atn_jobs=4.Trainer.inference/evaluatewith the right shapes and sensible metrics.Trainer.trainraises.sequenceraises an error naming the field;nested_sequence_floatsraises too.drop_last=Trueloaders are rejected.best_iterationand the probabilities equalXGBClassifierwith the same eval set (aucpr), and they differ from using all trees.save_ckptthenload_ckptinto a fresh model gives identicaly_probandlogitin all 4 modes (one with early stopping). A different field set raises, and so does a different vocabulary size.shap.TreeExplainer(rtol 1e-4), with the same top-10 ranking. This check skips if shap isn't installed.tools/check_pr_rules.pypasses.Example:
examples/mortality_prediction/mortality_mimic3_xgboost.py, on synthetic MIMIC-III (2,196 samples, patient split 70/10/20, test set):The synthetic data has little real signal, so this shows the workflow, not clinical performance.
Notes
libomp.dyliband XGBoost loads Homebrew's, and two copies in one process can crash multithreaded fits. The model and install docs describe the fix (make torch'slibomp.dyliba symlink to Homebrew's, or usen_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.testenvironment doesn't install xgboost, so these tests currently skip in CI. Addingxgboostto[tool.pixi.feature.test.pypi-dependencies]would run them, but needs apixi.lockrefresh. I can do that in this PR or a follow-up if you'd like.🤖 Generated with Claude Code