---
title: "Distributed Stratified Cox Regression"
author: "Balasubramanian Narasimhan"
date: '`r Sys.Date()`'
bibliography: homomorphing.bib
output:
  html_document:
  fig_caption: yes
  theme: cerulean
  toc: yes
  toc_depth: 2
vignette: >
  %\VignetteIndexEntry{Distributed Stratified Cox Regression}
  %\VignetteEngine{knitr::rmarkdown}
  \usepackage[utf8]{inputenc}
---

```{r ktab, echo=FALSE}
## Tables through kableExtra. Math written as $...$ in headers, cells
## and captions becomes \( ... \) and `code` becomes <code>, because
## pandoc does not process math inside a raw HTML table.
ktab <- function(x, ..., col.names = names(x), caption = NULL) {
    tex <- function(s)
        gsub("`([^`]*)`", "<code>\\1</code>",
             gsub("\\$([^$]+)\\$", "\\\\(\\1\\\\)", s))
    chr <- vapply(x, is.character, logical(1))
    x[chr] <- lapply(x[chr], tex)
    tab <- knitr::kable(x, format = "html", escape = FALSE,
                        col.names = tex(col.names),
                        caption = if (!is.null(caption)) tex(caption), ...)
    kableExtra::kable_styling(tab, bootstrap_options = c("striped", "condensed"),
                              full_width = TRUE)
}
```

```{r echo=F, purl=FALSE}
knitr::opts_chunk$set(
  message = FALSE,
  warning = FALSE,
  error = FALSE,
  tidy = FALSE,
  cache = FALSE
)
## The encrypted chunks are gated `eval = RECOMPUTE` and are shown but
## not run when the vignette builds; their results are read from
## data(cox_results), produced by data-raw/cox_results.R from these
## same chunks. HOMOMORPHER_RECOMPUTE=true runs them for real.
if (!exists("RECOMPUTE"))
    RECOMPUTE <- isTRUE(as.logical(Sys.getenv("HOMOMORPHER_RECOMPUTE", "FALSE")))
data(cox_results, package = "homomorpheR")
```

## 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

```{r}
#| echo: false
suppressPackageStartupMessages(library(homomorpheR))
data(DLBCL)
.sg  <- c("GCB", "ABC", "Type III")
.coh <- list(
  n      = nrow(DLBCL),
  deaths = sum(DLBCL$status),
  mfu    = format(round(median(DLBCL$time), 1), nsmall = 1),
  sgn    = vapply(.sg, function(s) sum(DLBCL$Subgroup == s), 1L),
  sgd    = vapply(.sg, function(s) sum(DLBCL$status[DLBCL$Subgroup == s]), 1L))
```

We use the diffuse large-B-cell lymphoma (DLBCL) cohort of
@Rosenwald2002, the same dataset that @BayleFanLou2025 use to
motivate distributed Cox estimation. The `data(DLBCL)` table
shipped with homomorpheR excludes the five patients with zero
follow-up time (following @BayleFanLou2025), leaving
`r .coh$n` patients with `r .coh$deaths` deaths over a median
follow-up of `r .coh$mfu` 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=`r .coh$sgn["GCB"]`$, ABC $n=`r .coh$sgn["ABC"]`$, Type III
$n=`r .coh$sgn["Type III"]`$, with `r .coh$sgd["GCB"]`,
`r .coh$sgd["ABC"]`, and `r .coh$sgd["Type III"]` deaths
respectively); the protocol does not require equal sizes.

```{r}
suppressPackageStartupMessages(library(survival))
library(homomorpheR)
data(DLBCL)

cox_data <- split(
  DLBCL[, c("time", "status", "GCB_sig", "LN_sig",
            "Prolif_sig", "BMP6", "MHC2_sig", "Subgroup")],
  DLBCL$Subgroup)
sapply(cox_data, function(df) c(n = nrow(df), events = sum(df$status)))
```

## The aggregated fit

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

```{r}
agg_model <- coxph(Surv(time, status) ~ GCB_sig + LN_sig +
                       Prolif_sig + BMP6 + MHC2_sig +
                       strata(Subgroup),
                   data = DLBCL)
agg_model
agg_model$loglik
```

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 vignette: 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 a feature of `coxph()`: with
`iter.max = 0` it returns the partial log-likelihood at the
supplied `init` without taking any Newton-Raphson steps.

```{r}
cph_control <- replace(coxph.control(), "iter.max", 0)

local_cox_nll <- function(data, beta) {
    fit <- tryCatch(
        coxph(Surv(time, status) ~ GCB_sig + LN_sig + Prolif_sig +
                  BMP6 + MHC2_sig,
              data    = data,
              init    = beta,
              control = cph_control),
        error = function(e) NULL)
    if (is.null(fit)) NA_real_ else -fit$loglik[1]
}
```

`tryCatch` returns `NA_real_` if the local fit fails at an extreme
$\beta$. `master_aggregate()` passes that `NA` back to the
optimizer, which then tries a shorter step.

## Wiring up the protocol

The summed negative log-likelihood on this cohort is
`r round(-agg_model$loglik[1])` at $\beta = 0$ and
`r round(-agg_model$loglik[2])` 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.

```{r eval=RECOMPUTE}
cc <- openfhe.R::fhe_context("CKKS",
                           multiplicative_depth = 1L,
                           scaling_mod_size     = 59L,
                           first_mod_size       = 60L,
                           batch_size           = 8L)
keys <- openfhe.R::key_gen(cc)

worker_gcb <- make_worker(name = "GCB",      data = cox_data[["GCB"]],
                          contribution_fn = local_cox_nll)
worker_abc <- make_worker(name = "ABC",      data = cox_data[["ABC"]],
                          contribution_fn = local_cox_nll)
worker_t3  <- make_worker(name = "Type III", data = cox_data[["Type III"]],
                          contribution_fn = local_cox_nll)
master     <- make_ckks_master("Master", crypto_context = cc, keypair = keys)
set_workers(master, list(worker_gcb, worker_abc, worker_t3))
```

## Iterative MLE through the encrypted protocol

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

```{r eval=RECOMPUTE}
library(stats4)

encrypted_nLL <- function(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig) {
    master_aggregate(master, c(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig))
}

fit <- mle(encrypted_nLL,
           start   = list(GCB_sig = 0, LN_sig = 0, Prolif_sig = 0,
                          BMP6    = 0, MHC2_sig = 0),
           method  = "BFGS",
           control = list(reltol = 1e-7))
summary(fit)
logLik(fit)
```

```{r eval=RECOMPUTE, echo=FALSE}
cox_results <- list(coef   = summary(fit)@coef,
                    loglik = as.numeric(logLik(fit)),
                    counts = fit@details$counts)
```

```{r echo=FALSE, purl=FALSE, eval=!RECOMPUTE}
cox_results$coef
cat(sprintf("'log Lik.' %f (df=5)\n", cox_results$loglik))
```

## Comparison with the cleartext fit

To check the encrypted fit, we run the identical `mle()` 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.

```{r}
library(stats4)

plain_nLL <- function(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig) {
    beta <- c(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig)
    sum(vapply(cox_data, local_cox_nll, numeric(1), beta = beta))
}

fit_plain <- mle(plain_nLL,
                 start   = list(GCB_sig = 0, LN_sig = 0, Prolif_sig = 0,
                                BMP6    = 0, MHC2_sig = 0),
                 method  = "BFGS",
                 control = list(reltol = 1e-7))
```

```{r echo = FALSE}
mle_coefs   <- cox_results$coef[, "Estimate"]
plain_coefs <- coef(fit_plain)[names(mle_coefs)]
enc_diff    <- abs(mle_coefs - plain_coefs)
tex_sci <- function(x) {
    e <- floor(log10(abs(x)))
    sprintf("$%.2f \\times 10^{%d}$", x / 10^e, as.integer(e))
}
comparison <- data.frame(
    Coefficient = names(mle_coefs),
    encrypted   = unname(mle_coefs),
    cleartext   = unname(plain_coefs),
    abs_diff    = tex_sci(unname(enc_diff))
)
ktab(comparison, digits = 7, row.names = FALSE, align = "lrrr",
             col.names = c("Coefficient", "$\\hat\\beta$, `mle()` encrypted",
                           "$\\hat\\beta$, `mle()` cleartext",
                           "$\\lvert \\text{difference} \\rvert$"),
             caption = "Single-decrypter CKKS DLBCL Cox against the same objective evaluated in the clear.")
```

The encrypted fit agrees with the cleartext fit to within
`r formatC(max(enc_diff), format = "e", digits = 2)` in every
coefficient.

## What just happened

`stats4::mle()` ran its usual BFGS 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. `mle()` was not modified, and
its result matches the cleartext fit of the same objective.

## What this demonstrates

1. **R optimizers work unchanged.** Any routine that takes the
   objective as a function, such as `mle()`, `optim()`, or
   `nlm()`, 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
   (`vignette("mle")`) works unchanged for survival analysis. Only
   the local computation changes (`coxph` instead of `dpois`).
3. **CKKS encrypts real numbers directly.** The earlier Paillier
   code had to split each value into integer and fractional parts
   and approximate the fractional part as a fraction with
   denominator $2^{256}$. CKKS needs no such encoding.
4. **The master/worker classes are reusable.** The same exported
   `Site` / `Master` classes and `master_aggregate()` runner
   drive both this Cox vignette and the Poisson MLE vignette —
   only the per-worker `contribution_fn` 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 `mle()` 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. `vignette("cox-threshold")` 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.
