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

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

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

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

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

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

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

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

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

## ----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.")

