Gradient-boosted survival model minimizing the Cox partial-likelihood loss.
GradientBoostingSurvivalAnalysis(
*,
n_estimators=100,
learning_rate=0.1,
max_depth=3,
min_samples_leaf=3,
subsample=1.0,
max_features=None,
random_state=None,
)
Builds an additive risk score by sequentially fitting shallow squared-error regression trees to the negative gradient of the negative Cox partial log-likelihood (the martingale residuals under the current model) and adding a shrunken version of each tree. The result is a flexible, non-linear generalization of the Cox model: strong predictive performance on structured tabular data, without a global proportional-hazards assumption on individual covariates. A Breslow baseline hazard turns the additive log-risk score into per-subject survival and cumulative-hazard functions.
Parameters
n_estimators: int = 100
-
Number of boosting iterations (trees).
learning_rate: float = 0.1
-
Shrinkage applied to each tree’s contribution. Smaller values need more trees but generalize better.
max_depth: int = 3
-
Maximum depth of each regression tree (interaction depth).
min_samples_leaf: int = 3
-
Minimum number of samples in each leaf of a tree.
subsample: float = 1.0
-
Fraction of the training rows sampled (without replacement) to grow each tree. Values below 1.0 give stochastic gradient boosting.
max_features: Any = None
-
Covariates considered per split: None (all, default), "sqrt", "log2", an int, or a float fraction.
random_state: Any = None
-
Seed or
numpy.random.Generator for subsampling and per-node feature selection.
Examples
Fit a gradient-boosted survival model on the bundled lung dataset:
import greenwood as gw
# Load data and build a right-censored response
lung = gw.load_dataset("lung", backend="pandas").dropna(subset=["ph.ecog", "ph.karno"])
y = gw.Surv.right(lung["time"], event=(lung["status"] == 2))
cols = ["age", "sex", "ph.ecog", "ph.karno", "wt.loss"]
# Fit the model and score the first five subjects
gbm = gw.GradientBoostingSurvivalAnalysis(
n_estimators=200, learning_rate=0.05, max_depth=2, random_state=0
).fit(y, lung[cols])
gbm.predict(lung[cols])[:5]
array([1.66691083, 1.09057783, 0.59369757, 1.18981321, 0.58051085])
Methods
|
Name
|
Description
|
|
fit()
|
Fit the gradient-boosted survival model to a right-censored response.
|
|
predict()
|
Predict a risk score, linear predictor, or survival / cumulative-hazard curves.
|
|
variable_importance()
|
Impurity-based variable importance (normalized total squared-error reduction).
|
fit()
Fit the gradient-boosted survival model to a right-censored response.
fit(
surv,
covariates,
*,
data=None,
)
Parameters
surv: Surv
-
A right-censored Surv response (built with Surv.right()).
covariates: Any
-
A dataframe, a 2-D array, or a right-hand-side formula string evaluated against data.
data: Any = None
-
DataFrame used to evaluate a formula string.
Returns
self
-
The fitted estimator, with cached attributes including
trees_, event_times_, feature_importances_, n_features_in_, and feature_names_in_.
predict()
Predict a risk score, linear predictor, or survival / cumulative-hazard curves.
predict(
newdata=None,
*,
type="risk",
times=None,
format=None,
)
type="risk" (default): the relative hazard \exp(F(x)), a 1-D array where higher values indicate higher risk.
type="lp": the additive log-risk score F(x) (the boosted Cox linear predictor).
type="survival": survival S(t \mid x) = \exp(-H_0(t)\,e^{F(x)}) at times.
type="cumulative_hazard": H(t \mid x) = H_0(t)\,e^{F(x)} at times.
Parameters
newdata: Any = None
-
Covariates to predict for. None predicts for the training subjects.
type: str = "risk"
-
One of "risk", "lp", "survival", or "cumulative_hazard".
times: Any = None
-
Times at which to evaluate curves. Defaults to the training event times.
format: str | None = None
-
DataFrame backend for curve output.
Returns
numpy.ndarray or DataFrame
-
A risk / linear-predictor vector, or a curve frame with a
time column and one column per subject.
variable_importance()
Impurity-based variable importance (normalized total squared-error reduction).
variable_importance(
*,
format=None,
)
Each split’s reduction in squared error is credited to its covariate and summed across all trees, then normalized to sum to one. Returns a table sorted by importance.