## ----echo=F-------------------------------------------------------------------
knitr::opts_chunk$set(
  message = FALSE,
  warning = FALSE,
  error = FALSE,
  tidy = FALSE,
  cache = FALSE
)
## Inline numbers in scientific notation as HTML (knitr's default
## prints a bare "5.8^{-6}" in HTML output).
sci <- function(x, digits = 2) {
  e <- floor(log10(abs(x)))
  if (x == 0 || e >= -2) return(formatC(x, format = "fg", digits = digits))
  sprintf("%s&nbsp;&times;&nbsp;10<sup>%d</sup>",
          formatC(x / 10^e, format = "f", digits = digits - 1), e)
}

## -----------------------------------------------------------------------------
set.seed(123)
n <- 500
age       <- rnorm(n, 55, 10)
biomarker <- rnorm(n, 0, 1)
prob      <- plogis(-2 + 0.03 * age + 0.8 * biomarker)
outcome   <- rbinom(n, 1, prob)

model <- glm(outcome ~ age + biomarker, family = binomial)
beta  <- coef(model)
cat("Coefficients (intercept, age, biomarker):", round(beta, 4), "\n")

## -----------------------------------------------------------------------------
library(openfhe.R)

cc <- fhe_context("CKKS",
                  multiplicative_depth = 8L,
                  scaling_mod_size     = 50L,
                  batch_size           = 16L,
                  features             = c(Feature$ADVANCEDSHE))
keys <- key_gen(cc, eval_mult = TRUE)

new_age <- c(45, 52, 60, 38, 70, 55, 48, 63,
             41, 57, 66, 44, 72, 50, 59, 35)
new_bm  <- c(-0.5,  0.3, 1.2, -1.0, 0.8, 0.1, -0.3, 1.5,
             -0.8,  0.6, 0.9, -0.4, 1.1, 0.0,  0.7, -1.2)

ct_age <- encrypt(keys@public, make_ckks_packed_plaintext(cc, new_age), cc = cc)
ct_bm  <- encrypt(keys@public, make_ckks_packed_plaintext(cc, new_bm),  cc = cc)

## -----------------------------------------------------------------------------
ct_eta <- ct_age * beta[2]
ct_eta <- ct_eta + ct_bm * beta[3]
ct_eta <- ct_eta + beta[1]

## -----------------------------------------------------------------------------
ct_prob <- eval_logistic(ct_eta, a = -4, b = 4, degree = 16)

## -----------------------------------------------------------------------------
result <- decrypt(ct_prob, keys@secret, cc = cc)
set_length(result, 16L)
encrypted_probs <- get_real_packed_value(result)[1:16]

cleartext_probs <- plogis(beta[1] + beta[2] * new_age + beta[3] * new_bm)
max_err <- max(abs(encrypted_probs - cleartext_probs))
cat(sprintf("Max absolute error vs cleartext sigmoid: %.2e\n", max_err))

