Distributed Stratified Cox Regression

The statistical problem

The Cox proportional hazards model is widely used in medical statistics. Given covariates \(x_i\) for subject \(i\), an event time \(t_i\), and an event indicator \(\delta_i\), the hazard is modeled as

\[ h(t \mid x_i) \;=\; h_0(t) \, \exp(\beta^\top x_i) \]

where \(h_0(t)\) is an unspecified baseline hazard and \(\beta\) is the vector of regression coefficients we want to estimate. The partial log-likelihood depends only on \(\beta\):

\[ \ell(\beta) \;=\; \sum_{i: \delta_i = 1} \left[\, \beta^\top x_i \;-\; \log\!\!\sum_{j \in R_i} \exp(\beta^\top x_j) \,\right] \]

where \(R_i\) is the risk set at time \(t_i\).

For stratified Cox regression — when baseline hazards differ across strata (e.g. across study sites) but the coefficients \(\beta\) are shared — the partial log-likelihood becomes a sum over strata:

\[ \ell(\beta) \;=\; \sum_{s=1}^{S} \ell_s(\beta). \]

Because the log-likelihood is a sum over strata, each site can compute its own term. The master/worker setup used by distcomp, DataSHIELD, and WebDISCO relies on this: a master sends the current \(\beta\) to the sites, each site computes \(\ell_s(\beta)\) on its own data, and the master adds the results. We use the same setup, but the sites encrypt their terms under CKKS, so the master sees only the sum and not the individual terms.

DLBCL lymphoma cohort

We use the diffuse large-B-cell lymphoma (DLBCL) cohort of Rosenwald et al. (2002), the same dataset that Bayle et al. (2025) use to motivate distributed Cox estimation. The table load_dlbcl() returns excludes the five patients with zero follow-up time (following Bayle et al. (2025)), leaving 235 patients with 133 deaths over a median follow-up of 2.8 years. The published outcome predictor combines four gene-expression signatures (germinal-center B cell, lymph node, proliferation, and MHC class II) and the single gene BMP6. We model the hazard as a function of those five variables, stratified by molecular subgroup (GCB, ABC, Type III), and we treat each subgroup as a site. The three sites differ in size (GCB \(n=115\), ABC \(n=71\), Type III \(n=49\), with 54, 49, and 30 deaths respectively); the protocol does not require equal sizes.

import numpy as np
from homomorphepy import load_dlbcl, site_order

df = load_dlbcl()
COVARIATES = ["GCB_sig", "LN_sig", "Prolif_sig", "BMP6", "MHC2_sig"]

def site_frame(name):
    m = (df["Subgroup"] == name).to_numpy()
    return {
        "time":   df["time"].to_numpy(dtype=float)[m],
        "status": df["status"].to_numpy(dtype=int)[m],
        "X":      np.column_stack([df[c].to_numpy(dtype=float)[m]
                                   for c in COVARIATES]),
    }

cox_data = {name: site_frame(name) for name in site_order()}
df.groupby("Subgroup", observed=True).agg(n=("time", "size"),
                                          events=("status", "sum"))
n events
Subgroup
GCB 115 54
ABC 71 49
Type III 49 30

The aggregated fit

If all data were in one place, we would fit the stratified Cox model directly:

from statsmodels.duration.hazard_regression import PHReg

agg_model = PHReg(df["time"], df[COVARIATES], status=df["status"],
                  strata=df["Subgroup"].cat.codes, ties="efron")
agg_fit = agg_model.fit()
print(agg_fit.summary())
agg_loglik = (agg_model.loglike(np.zeros(len(COVARIATES))),
              agg_model.loglike(agg_fit.params))
agg_loglik
                         Results: PHReg
================================================================
Model:                   PH Reg      Num strata:            3   
Dependent variable:      time        Min stratum size:      49  
Ties:                    Efron       Max stratum size:      115 
Sample size:             235         Avg stratum size:      78.3
Num. events:             133                                    
----------------------------------------------------------------
            log HR log HR SE   HR      t    P>|t|  [0.025 0.975]
----------------------------------------------------------------
GCB_sig    -0.2639    0.1194 0.7681 -2.2099 0.0271 0.6078 0.9706
LN_sig     -0.2544    0.0852 0.7754 -2.9871 0.0028 0.6562 0.9163
Prolif_sig  0.3031    0.1498 1.3541  2.0234 0.0430 1.0095 1.8162
BMP6        0.3036    0.1073 1.3548  2.8304 0.0046 1.0979 1.6718
MHC2_sig   -0.3191    0.0941 0.7268 -3.3905 0.0007 0.6043 0.8740
================================================================
Confidence intervals are for the hazard ratios
(np.float64(-516.59855116308), np.float64(-495.22902176698074))

The first log-likelihood is at \(\beta = 0\) (the null model); the second is at the MLE. The goal is to reproduce these estimates without the three sites pooling their data.

The protocol

We use the same master/worker topology as the MLE page: master broadcasts \(\beta\), each worker computes its local Cox partial log-likelihood at \(\beta\), encrypts it under the master’s public key, and returns the encrypted value. The master sums the encrypted contributions homomorphically and decrypts the total. Mathematically nothing changes from the MLE case; only the local computation differs.

The local computation uses PHReg.loglike(), which returns the partial log-likelihood at the supplied \(\beta\) without fitting. PHReg defaults to Breslow’s treatment of ties; we pass ties="efron".

def local_cox_nll(data, beta):
    b = np.asarray(beta, dtype=float).ravel()
    model = PHReg(data["time"], data["X"],
                  status=data["status"], ties="efron")
    return float(-model.loglike(b))

Wiring up the protocol

The summed negative log-likelihood on this cohort is 517 at \(\beta = 0\) and 495 at the MLE. CKKS represents values of this size at the default scaling parameters. We raise scaling_mod_size from 50 to 59 for extra precision and set first_mod_size = 60, the library default, explicitly.

from homomorphepy import fhe_context, make_worker, make_ckks_master

cc = fhe_context("CKKS", multiplicative_depth=1, scaling_mod_size=59,
                 first_mod_size=60, batch_size=8)
keys = cc.KeyGen()

workers = [make_worker(name, data, local_cox_nll)
           for name, data in cox_data.items()]
master = make_ckks_master("Master", cc, keys)
master.set_workers(workers)
<CKKSMaster Master>

Iterative MLE through the encrypted protocol

We hand scipy.optimize.minimize() a function that looks like a standard multivariate negative log-likelihood. Each call drives one master/worker round and returns a single decrypted scalar.

from scipy.optimize import minimize

def encrypted_nll(beta):
    value = master.aggregate(np.asarray(beta, dtype=float))
    return 1e12 if (value is None or np.isnan(value)) else float(value)

fit = minimize(encrypted_nll, x0=np.zeros(len(COVARIATES)),
               method="L-BFGS-B", options={"ftol": 1e-9})
beta_hat = np.asarray(fit.x)

Standard errors come from the Hessian of the same encrypted objective at the fit, by central differences. Each of those evaluations is another protocol round.

import pandas as pd
from homomorphepy.examples.cox import fd_hessian

se = np.sqrt(np.diag(np.linalg.inv(fd_hessian(encrypted_nll, beta_hat))))
pd.DataFrame({"estimate": beta_hat, "std. error": se}, index=COVARIATES)
estimate std. error
GCB_sig -0.263872 0.119404
LN_sig -0.254359 0.085152
Prolif_sig 0.303125 0.149813
BMP6 0.303637 0.107278
MHC2_sig -0.319147 0.094129
float(-fit.fun)
-495.22902176698545

Comparison with the cleartext fit

To check the encrypted fit, we run the identical objective a second time with the encrypted aggregation replaced by an ordinary sum of the three sites’ cleartext values: same likelihood, same optimizer, same starting values and tolerance, no encryption.

def cleartext_nll(beta):
    return sum(local_cox_nll(d, beta) for d in cox_data.values())

fit_plain = minimize(cleartext_nll, x0=np.zeros(len(COVARIATES)),
                     method="L-BFGS-B", options={"ftol": 1e-9})
beta_plain = np.asarray(fit_plain.x)
minimize(), encrypted minimize(), cleartext absolute difference
coefficient
GCB_sig -0.263871 -0.263872 7.16e-08
LN_sig -0.254359 -0.254359 3.89e-08
Prolif_sig 0.303125 0.303126 1.31e-07
BMP6 0.303637 0.303638 1.20e-07
MHC2_sig -0.319147 -0.319147 1.45e-07

The encrypted fit agrees with the cleartext fit to within 1.45e-07 in every coefficient.

assert enc_diff.max() < 1e-4

What just happened

scipy.optimize.minimize() ran its usual L-BFGS-B iterations. Each time it asked for the negative log-likelihood at a point in \(\mathbb{R}^5\), the function ran one CKKS master/worker round across the three sites and returned one decrypted number. minimize() was not modified, and its result matches the cleartext fit of the same objective.

The same run, in one call

from homomorphepy.examples import cox

result = cox.run(backend="ckks")
result.loglik_encrypted, result.loglik_centralized, result.max_abs_difference
(-495.2290217669947, -495.22902176698074, 5.579673819089415e-07)

What this demonstrates

  1. Python optimizers work unchanged. Any routine that takes the objective as a function, such as scipy.optimize.minimize(), can be given one that computes its value through the encrypted protocol.
  2. Stratified Cox regression decomposes additively across strata, so the master/worker scheme used for Poisson MLE (MLE page) works unchanged for survival analysis. Only the local computation changes (PHReg instead of the Poisson log-likelihood).
  3. CKKS encrypts real numbers directly. An additive scheme such as Paillier would need each value split into integer and fractional parts, with the fractional part approximated as a fraction with a fixed denominator. CKKS needs no such encoding.
  4. The master/worker classes are reusable. The same exported Site / Master classes and aggregate() runner drive both this Cox page and the Poisson MLE page — only the per-worker local function differs.

Caveats and extensions

  • Performance: each function evaluation requires three CKKS encryptions, three encrypted additions, and one decryption. CKKS encrypt/decrypt dominates the wall-clock cost.
  • Information leakage: the master sees the value of the joint log-likelihood at each \(\beta\). That is less than the individual contributions, and it is what minimize() needs. Hiding it as well would require running Newton-Raphson on encrypted values, which CKKS allows but which is considerably more complex.
  • Threshold key generation: in a real deployment the secret key would be split across the sites (n-of-n threshold), so that no single party, the master included, can decrypt intermediate values on its own. The Cox threshold page adds this.
  • Beyond Cox: the same protocol applies to any model whose log-likelihood is a sum over data partitions, such as generalized linear models, mixed-effects models with site-specific random effects, and frailty survival models. Only the local likelihood evaluation changes.

References

Bayle, Pierre, Jianqing Fan, and Zhipeng Lou. 2025. “Communication-Efficient Distributed Estimation and Inference for Cox’s Model.” Journal of the American Statistical Association 120 (551): 1736–46. https://doi.org/10.1080/01621459.2025.2516820.
Rosenwald, Andreas, George Wright, Wing C. Chan, et al. 2002. “The Use of Molecular Profiling to Predict Survival After Chemotherapy for Diffuse Large-B-Cell Lymphoma.” New England Journal of Medicine 346 (25): 1937–47. https://doi.org/10.1056/NEJMoa012914.