diff --git a/lifelines/fitters/coxph_fitter.py b/lifelines/fitters/coxph_fitter.py index 4f90c848e..315e3081c 100644 --- a/lifelines/fitters/coxph_fitter.py +++ b/lifelines/fitters/coxph_fitter.py @@ -287,6 +287,7 @@ def fit( """ self.strata = utils._to_list_or_singleton(utils.coalesce(strata, self.strata)) + self.formula = formula self._model = self._fit_model( df, duration_col, @@ -858,6 +859,7 @@ def compute_followup_hazard_ratios(self, training_df: DataFrame, followup_times: weights_col=self.weights_col, cluster_col=self.cluster_col, entry_col=self.entry_col, + formula=self.formula, ) results[t] = model.hazard_ratios_ return DataFrame(results).T diff --git a/lifelines/tests/test_estimation.py b/lifelines/tests/test_estimation.py index 00c366086..fcfc16ed1 100644 --- a/lifelines/tests/test_estimation.py +++ b/lifelines/tests/test_estimation.py @@ -3248,6 +3248,14 @@ def test_compute_followup_hazard_ratios(self, cph, cph_spline, rossi): cph_spline.fit(rossi, "week", "arrest") cph_spline.compute_followup_hazard_ratios(rossi, [15, 25, 35, 45]) + def test_compute_followup_hazard_ratios_preserves_formula(self, cph, rossi): + cph.fit(rossi, "week", "arrest", formula="fin + age + race") + + result = cph.compute_followup_hazard_ratios(rossi, [rossi["week"].max()]) + + assert_index_equal(result.columns, cph.hazard_ratios_.index) + assert_series_equal(result.iloc[0], cph.hazard_ratios_, check_names=False) + def test_model_can_accept_null_covariates(self, cph, rossi): cph.fit(rossi[["week", "arrest"]], "week", "arrest") assert True