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

## ----cvxr-libs, eval=RECOMPUTE------------------------------------------------
# library(survival)
# library(CVXR)
# library(openfhe.R)
# library(homomorpheR)

## ----cvxr-data, eval=RECOMPUTE------------------------------------------------
# data(DLBCL,     package = "homomorpheR")  # 235 patients: survival + signatures
# data(DLBCL_gex, package = "homomorpheR")  # 235 patients x 6416 Lymphochip probes
# dlbcl <- DLBCL
# dlbcl$Subgroup <- factor(dlbcl$Subgroup, levels = c("GCB","ABC","Type III"))
# stopifnot(identical(as.character(dlbcl$ID), rownames(DLBCL_gex)))  # rows aligned
# 
# sites_raw <- lapply(levels(dlbcl$Subgroup), function(s) {
#     idx <- which(dlbcl$Subgroup == s)
#     list(name = s, X = DLBCL_gex[idx, ], time = dlbcl$time[idx],
#          status = dlbcl$status[idx])
# })
# names(sites_raw) <- levels(dlbcl$Subgroup)
# N_sites <- length(sites_raw)
# N_total <- sum(vapply(sites_raw, function(s) nrow(s$X), 1L))
# P_raw   <- ncol(DLBCL_gex)

## ----cvxr-standardize-plain, eval=RECOMPUTE-----------------------------------
# pool_plain <- function(sites, n_total) {
#     s <- Reduce(`+`, lapply(sites, function(s) colSums(s$X)))
#     q <- Reduce(`+`, lapply(sites, function(s) colSums(s$X^2)))
#     mu     <- s / n_total
#     sigma2 <- pmax(q / n_total - mu^2, .Machine$double.eps)
#     list(mu = mu, sigma = sqrt(sigma2))
# }
# ## Each site standardizes its own rows with the pooled moments.
# standardize <- function(sites, pool) lapply(sites, function(s) list(
#     name = s$name,
#     X    = sweep(sweep(s$X, 2, pool$mu, "-"), 2, pool$sigma, "/"),
#     time = s$time, status = s$status))
# pool      <- pool_plain(sites_raw, N_total)
# sites_std <- standardize(sites_raw, pool)

## ----cvxr-screen-plain, eval=RECOMPUTE----------------------------------------
# K <- 100L
# 
# score_info_at_zero <- function(X, time, status) {
#     p <- ncol(X)
#     ord <- order(time)
#     X_o <- X[ord, , drop = FALSE]; stat_o <- status[ord]; time_o <- time[ord]
#     n <- nrow(X_o); U <- numeric(p); I <- numeric(p)
#     ## Risk set at event i: every row with time >= t_i, tied rows included.
#     first <- match(time_o, time_o)
#     for (i in seq_len(n)) {
#         if (stat_o[i] == 1L) {
#             risk <- X_o[first[i]:n, , drop = FALSE]
#             mu_R <- colMeans(risk)
#             U <- U + (X_o[i, ] - mu_R)
#             I <- I + colSums(sweep(risk, 2, mu_R, "-")^2) / nrow(risk)
#         }
#     }
#     list(U = U, I = I)
# }
# 
# screen_plain <- function(sites, K) {
#     UI <- lapply(sites, function(s) score_info_at_zero(s$X, s$time, s$status))
#     U  <- Reduce(`+`, lapply(UI, `[[`, "U"))
#     I  <- Reduce(`+`, lapply(UI, `[[`, "I"))
#     Z  <- U / sqrt(pmax(I, .Machine$double.eps))
#     order(abs(Z), decreasing = TRUE)[seq_len(K)]
# }
# ## Each site keeps the screened probes of its own rows.
# keep_probes <- function(sites, idx) lapply(sites, function(s) list(
#     name = s$name, X = s$X[, idx],
#     time = s$time, status = s$status))
# top_idx  <- screen_plain(sites_std, K)
# sites_KS <- keep_probes(sites_std, top_idx)
# sigma_K  <- pool$sigma[top_idx]

## ----cvxr-centralized, eval=RECOMPUTE-----------------------------------------
# LAMBDA <- 5
# build_cox_breslow_nll <- function(beta_var, X_s, time_s, status_s) {
#     ord <- order(time_s)
#     X_o <- X_s[ord, , drop = FALSE]; stat_o <- status_s[ord]
#     time_o <- time_s[ord]
#     n <- nrow(X_o); eta_o <- X_o %*% beta_var
#     ## Risk set at event i: every row with time >= t_i, tied rows included.
#     first <- match(time_o, time_o)
#     terms <- list()
#     for (i in seq_len(n)) {
#         if (stat_o[i] == 1L) {
#             terms[[length(terms) + 1L]] <-
#                 log_sum_exp(eta_o[first[i]:n, 1]) - eta_o[i, 1]
#         }
#     }
#     Reduce(`+`, terms)
# }
# 
# beta_var <- Variable(K, name = "beta")
# nll_per  <- lapply(sites_KS, function(s)
#     build_cox_breslow_nll(beta_var, s$X, s$time, s$status))
# agg_prob <- Problem(Minimize(Reduce(`+`, nll_per) +
#                                  LAMBDA * p_norm(beta_var, 1)))
# suppressMessages(suppressWarnings(
#     psolve(agg_prob, solver = "CLARABEL", verbose = FALSE)))
# agg_beta <- as.numeric(value(beta_var))

## ----cvxr-local, eval=RECOMPUTE-----------------------------------------------
# RHO <- 50
# 
# build_local <- function(X_k, time_k, status_k, rho) {
#     p <- ncol(X_k)
#     x  <- Variable(p); zp <- Parameter(p); up <- Parameter(p)
#     nll <- build_cox_breslow_nll(x, X_k, time_k, status_k)
#     aug <- (rho / 2) * sum_squares(x - zp + up)
#     list(prob = Problem(Minimize(nll + aug)), x = x, zp = zp, up = up)
# }
# sites_problem <- lapply(sites_KS, function(s)
#     build_local(s$X, s$time, s$status, RHO))

## ----cvxr-admm, eval=RECOMPUTE------------------------------------------------
# MAX_ITER <- 200L; TOL <- 5e-3
# soft_threshold <- function(v, tau) sign(v) * pmax(abs(v) - tau, 0)
# 
# run_admm <- function(sites_problem, consensus) {
#     site_x <- replicate(N_sites, rep(0, K), simplify = FALSE)
#     site_u <- replicate(N_sites, rep(0, K), simplify = FALSE)
#     z_curr <- rep(0, K); trajectory <- list()
#     for (iter in seq_len(MAX_ITER)) {
#         for (i in seq_len(N_sites)) {
#             value(sites_problem[[i]]$zp) <- z_curr
#             value(sites_problem[[i]]$up) <- site_u[[i]]
#             suppressMessages(suppressWarnings(
#                 psolve(sites_problem[[i]]$prob, solver = "CLARABEL",
#                        verbose = FALSE)))
#             site_x[[i]] <- as.numeric(value(sites_problem[[i]]$x))
#         }
#         w_avg  <- consensus(site_x, site_u)
#         z_new  <- soft_threshold(w_avg, LAMBDA / (N_sites * RHO))
#         site_u <- Map(function(u, x) u + (x - z_new), site_u, site_x)
#         primal <- sqrt(mean(vapply(site_x, function(x) sum((x - z_new)^2), 0)))
#         dual   <- RHO * sqrt(sum((z_new - z_curr)^2))
#         z_curr <- z_new; trajectory[[iter]] <- z_new
#         if (primal < TOL && dual < TOL) break
#     }
#     list(z = z_curr, trajectory = trajectory)
# }
# 
# plain_consensus <- function(site_x, site_u)
#     Reduce(`+`, Map(`+`, site_x, site_u)) / length(site_x)
# 
# ref        <- run_admm(sites_problem, plain_consensus)
# z_ref      <- ref$z
# n_iter_ref <- length(ref$trajectory)

## ----cvxr-context, eval=RECOMPUTE---------------------------------------------
# cc <- fhe_context("CKKS",
#                   multiplicative_depth = 1L,
#                   scaling_mod_size     = 59L,
#                   first_mod_size       = 60L,
#                   batch_size           = 8192L,
#                   features             = c(Feature$MULTIPARTY))
# 
# key_sites <- lapply(levels(dlbcl$Subgroup), function(s)
#     make_worker(s, data = NULL, contribution_fn = function(data, theta) NULL))
# 
# master <- make_threshold_master("Aggregator",
#                                 crypto_context = cc, sites = key_sites)

## ----cvxr-pool-encrypt, eval=RECOMPUTE----------------------------------------
# ## Site side. Each site forms its own moments and encrypts them under
# ## the joint key; what leaves is encrypted. Encrypting at the
# ## aggregator instead would mean handing it the per-site column sums in
# ## the clear first, which is the disclosure this round exists to avoid.
# site_moments <- function(site, s)
#     list(sum   = encrypt(site, colSums(s$X)),
#          sumsq = encrypt(site, colSums(s$X^2)))
# 
# ## Aggregator side. It receives only the encrypted moments, adds them,
# ## and decrypts the totals.
# pool_encrypted <- function(master, parts, n_total, p_raw) {
#     pooled_sum   <- decrypt(
#         master, Reduce(`+`, lapply(parts, `[[`, "sum")),   len = p_raw)
#     pooled_sumsq <- decrypt(
#         master, Reduce(`+`, lapply(parts, `[[`, "sumsq")), len = p_raw)
#     mu     <- pooled_sum / n_total
#     sigma2 <- pmax(pooled_sumsq / n_total - mu^2, .Machine$double.eps)
#     list(mu = mu, sigma = sqrt(sigma2))
# }
# moment_parts <- Map(site_moments, key_sites, sites_raw)   # at the sites
# fhe_pool     <- pool_encrypted(master, moment_parts, N_total, P_raw)
# pool_agree   <- list(mu    = max(abs(fhe_pool$mu    - pool$mu)),
#                      sigma = max(abs(fhe_pool$sigma - pool$sigma)))
# 
# ## The decrypted moments go back to the sites, which standardize their
# ## own rows with them.
# sites_std_fhe <- standardize(sites_raw, fhe_pool)

## ----cvxr-screen-encrypt, eval=RECOMPUTE--------------------------------------
# ## Site side: compute the score and information at beta = 0 on the
# ## site's own rows, and encrypt both before returning them.
# site_score_info <- function(site, s) {
#     z <- score_info_at_zero(s$X, s$time, s$status)
#     list(U = encrypt(site, z$U), I = encrypt(site, z$I))
# }
# 
# ## Aggregator side: it receives only the encrypted (U, I), adds them,
# ## and decrypts the totals.
# screen_encrypted <- function(master, UI, p_raw, K) {
#     U   <- decrypt(master, Reduce(`+`, lapply(UI, `[[`, "U")),
#                    len = p_raw)
#     I   <- decrypt(master, Reduce(`+`, lapply(UI, `[[`, "I")),
#                    len = p_raw)
#     Z   <- U / sqrt(pmax(I, .Machine$double.eps))
#     order(abs(Z), decreasing = TRUE)[seq_len(K)]
# }
# UI_parts <- Map(site_score_info, key_sites, sites_std_fhe)   # at the sites
# fhe_top  <- screen_encrypted(master, UI_parts, P_raw, K)
# stopifnot(identical(fhe_top, top_idx))   # same probes, same order, as in the clear
# 
# ## The decrypted screen goes back to the sites, which keep those probes.
# sites_KS_fhe <- keep_probes(sites_std_fhe, fhe_top)

## ----cvxr-consensus, eval=RECOMPUTE-------------------------------------------
# ## Site side: site k forms x_k + u_k and encrypts it. What leaves the
# ## site is encrypted.
# site_consensus_term <- function(site, x_k, u_k)
#     encrypt(site, x_k + u_k)
# 
# ## Aggregator side: it receives only the encrypted terms, adds them,
# ## scales by 1/N, and decrypts the average.
# aggregate_consensus <- function(master, cts, K)
#     decrypt(master, Reduce(`+`, cts) * (1 / length(cts)), len = K)
# 
# ## One consensus round: the site step at every site, then the
# ## aggregator step.
# encrypted_consensus <- function(site_x, site_u)
#     aggregate_consensus(master,
#                         Map(site_consensus_term, key_sites, site_x, site_u),
#                         K)
# 
# ## The local problems, on the design the encrypted rounds produced.
# sites_problem_fhe <- lapply(sites_KS_fhe, function(s)
#     build_local(s$X, s$time, s$status, RHO))
# fhe        <- run_admm(sites_problem_fhe, encrypted_consensus)
# z_enc      <- fhe$z
# trajectory <- fhe$trajectory
# n_iter_enc <- length(trajectory)

## ----cvxr-assemble, eval=RECOMPUTE, echo=FALSE--------------------------------
# ## (not shown in the rendered vignette; purled into the regeneration
# ## script) collect the primitives the vignette, manuscript, and data
# ## object share. Light derived quantities are recomputed downstream.
# cvxr_consensus <- list(
#     params     = list(K = K, LAMBDA = LAMBDA, RHO = RHO,
#                       MAX_ITER = MAX_ITER, TOL = TOL),
#     top_idx    = top_idx,
#     sigma_K    = sigma_K,
#     agg_beta   = agg_beta,
#     z_ref      = z_ref,
#     z_enc      = z_enc,
#     trajectory = trajectory,
#     n_iter_ref = n_iter_ref,
#     n_iter_enc = n_iter_enc,
#     pool_agree = pool_agree,
#     screen_match = identical(fhe_top, top_idx))

