## ----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(survival))
library(homomorpheR)
data(DLBCL)

cox_data <- split(
  DLBCL[, c("time", "status", "GCB_sig", "LN_sig",
            "Prolif_sig", "BMP6", "MHC2_sig", "Subgroup")],
  DLBCL$Subgroup)

## -----------------------------------------------------------------------------
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,
#                            features             = c(openfhe.R::Feature$MULTIPARTY))

## ----eval=RECOMPUTE-----------------------------------------------------------
# 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_threshold_master("Aggregator",
#                                 crypto_context = cc,
#                                 sites = list(worker_gcb, worker_abc, worker_t3))

## ----eval=RECOMPUTE-----------------------------------------------------------
# share_check <- c(master_holds_shares = "secret_keys" %in% names(S7::props(master)),
#                  gcb_holds_own_share = !is.null(worker_gcb@state$sk))
# share_check

## ----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_threshold_results <- list(coef        = summary(fit)@coef,
#                               loglik      = as.numeric(logLik(fit)),
#                               counts      = fit@details$counts,
#                               share_check = share_check)

## -----------------------------------------------------------------------------
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-------------------------------------------------------------
thr_coefs   <- cox_threshold_results$coef[, "Estimate"]
plain_coefs <- coef(fit_plain)[names(thr_coefs)]
thr_diff    <- abs(thr_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(thr_coefs),
    threshold   = unname(thr_coefs),
    cleartext   = unname(plain_coefs),
    abs_diff    = tex_sci(unname(thr_diff))
)
ktab(comparison, digits = 7, row.names = FALSE, align = "lrrr",
             col.names = c("Coefficient", "$\\hat\\beta$, `mle()` threshold",
                           "$\\hat\\beta$, `mle()` cleartext",
                           "$\\lvert \\text{difference} \\rvert$"),
             caption = "Threshold-CKKS DLBCL Cox against the same objective evaluated in the clear.")

