## -----------------------------------------------------------------------------
suppressPackageStartupMessages({
  library(homomorpheR)
  library(openfhe.R)
})
set.seed(20260428)

p             <- 32L
n_sites       <- 3L
cohort_sizes  <- c(80L, 60L, 100L)
n_anchor      <- 100L
n_phenotypes  <- 5L
top_k         <- 5L

## -----------------------------------------------------------------------------
phenotype_centers <- matrix(rnorm(n_phenotypes * p, sd = 1),
                            n_phenotypes, p)

embed_public <- function(n) {
  labels  <- sample.int(n_phenotypes, n, replace = TRUE)
  centers <- phenotype_centers[labels, , drop = FALSE]
  noise   <- matrix(rnorm(n * p, sd = 0.4), n, p)
  z <- centers + noise
  z <- z / sqrt(rowSums(z^2))            # unit norm
  list(z = z, label = labels)
}

## -----------------------------------------------------------------------------
random_drift <- function(p, beta) {
  ## Free fine-tuning, simulated as B = Q D: a random rotation Q
  ## composed with an anisotropic stretch D = diag(exp(beta g)).
  ## beta = 0 gives an exactly orthogonal (isometric) drift;
  ## beta > 0 is non-isometric, the generic fine-tuned case.
  Q <- qr.Q(qr(matrix(rnorm(p * p), p, p)))
  if (beta == 0) return(Q)
  Q %*% diag(exp(beta * rnorm(p)))
}

embed_private <- function(z_public, B_k) {
  ## Site k's fine-tuned model in column-vector convention:
  ## f_k(x) = B_k · f(x). For a matrix z_public of n row-stacked
  ## vectors, the private embeddings are z_public %*% t(B_k),
  ## then unit-normalized (the drift acts on the sphere; under a
  ## non-isometric B the normalization is a genuine nonlinearity).
  v <- z_public %*% t(B_k)
  v / sqrt(rowSums(v^2))
}

## -----------------------------------------------------------------------------
fit_adapter <- function(Z_priv, Z_pub, mu) {
  ## Compatibility adapter A (private -> public): Z_priv A ~ Z_pub,
  ## with a near-isometry penalty mu * ||A^T A - I||^2.
  p <- ncol(Z_priv)
  if (is.infinite(mu)) {                       # orthogonal Procrustes
    sv <- svd(crossprod(Z_priv, Z_pub))
    return(sv$u %*% t(sv$v))
  }
  A_ls <- solve(crossprod(Z_priv) + 1e-6 * diag(p), crossprod(Z_priv, Z_pub))
  if (mu == 0) return(A_ls)                     # least squares
  fn <- function(par) { A <- matrix(par, p, p)  # near-orthogonal
    sum((Z_priv %*% A - Z_pub)^2) + mu * sum((crossprod(A) - diag(p))^2) }
  gr <- function(par) { A <- matrix(par, p, p)
    as.vector(2 * crossprod(Z_priv, Z_priv %*% A - Z_pub) +
              4 * mu * (A %*% (crossprod(A) - diag(p)))) }
  matrix(optim(as.vector(A_ls), fn, gr, method = "L-BFGS-B",
               control = list(maxit = 400))$par, p, p)
}

## -----------------------------------------------------------------------------
fit_gram <- function(Z_priv, Z_pub) {
  M <- tcrossprod(solve(crossprod(Z_priv) + 1e-6 * diag(ncol(Z_priv)),
                        crossprod(Z_priv, Z_pub)))
  e <- eigen(M, symmetric = TRUE)
  e$vectors %*% (sqrt(pmax(e$values, 0)) * t(e$vectors))   # symmetric M^{1/2}
}

unit_rows <- function(Z) Z / sqrt(rowSums(Z^2))            # for Design-1 folding

## -----------------------------------------------------------------------------
beta_demo <- 0.6
mu_demo   <- Inf      # orthogonal endpoint for the encrypted walk-through

## -----------------------------------------------------------------------------
public_anchor <- embed_public(n_anchor)
public_query  <- embed_public(1L)        # one query for the protocol walk-through

site_cohorts <- lapply(cohort_sizes, embed_public)

## ----eval=RECOMPUTE-----------------------------------------------------------
# cc <- fhe_context(
#   scheme               = "CKKS",
#   multiplicative_depth = 3L,
#   scaling_mod_size     = 45L,
#   batch_size           = p,
#   features             = c(Feature$MULTIPARTY, Feature$KEYSWITCH))
# 
# sites <- lapply(seq_len(n_sites), function(k)
#   make_worker(paste("Site", k), data = NULL,
#               contribution_fn = function(data, theta) NULL))
# 
# master <- make_threshold_master(name = "master",
#                                 crypto_context = cc,
#                                 sites = sites)
# 
# ## Published once, when key generation completes, to every party that
# ## will encrypt -- the sites and the querying party alike. It is public
# ## in full, and it is the last thing anyone needs from the aggregator.
# ## `site_params()` asks a site what it is holding; no aggregator is
# ## involved, and the object has no property a key share could sit in.
# pub <- site_params(sites[[1]])
# pub

## ----eval=RECOMPUTE-----------------------------------------------------------
# rotation_indices <- seq_len(p - 1L)
# joint_pk_tag <- get_key_tag(master@joint_pubkey)
# 
# ## Each of the two functions below runs *at a site*, and touches only
# ## that site's own share. Nothing collects the shares into one list:
# ## a variable holding every `sk` would be exactly the single point of
# ## compromise the n-of-n split exists to remove.
# 
# ## Lead site. Generates its rotation keys under its own share, which
# ## populates the context's automorphism-key registry under that
# ## share's tag, and returns the resulting (public) key map.
# site_rotation_lead <- function(site, indices) {
#   eval_rotate_key_gen(cc, site@state$sk, indices)
#   get_eval_automorphism_key_map(get_key_tag(site@state$sk))
# }
# 
# ## A following site. Contributes its share of the rotation keys to the
# ## running joint map and returns the contribution.
# site_rotation_share <- function(site, running, indices, tag)
#   multi_eval_at_index_key_gen(cc, site@state$sk, running,
#                               index_list = indices, key_tag = tag)
# 
# rot_running <- site_rotation_lead(sites[[1]], rotation_indices)
# 
# ## Sites 2..n contribute their rotation-key shares in turn.
# ## At each step, the running joint share is registered under
# ## the cumulative-pubkey tag — which for our `make_threshold_master`
# ## is the same `joint_pk_tag` at every step (the chain daisy-chains
# ## pubkeys forward but only the final joint pubkey is retained).
# for (k in 2:n_sites) {
#   share_k <- site_rotation_share(sites[[k]], rot_running,
#                                  rotation_indices, joint_pk_tag)
#   rot_running <- multi_add_eval_automorphism_keys(
#     cc, rot_running, share_k, key_tag = joint_pk_tag)
# }
# 
# insert_eval_automorphism_key(rot_running, key_tag = joint_pk_tag)

## ----eval=RECOMPUTE-----------------------------------------------------------
# x <- 1:p
# ct_x      <- encrypt(pub, x)
# ct_rot    <- eval_rotate(ct_x, 3L)
# recovered <- decrypt(master, ct_rot, len = p)
# 
# expected <- x[((seq_len(p) - 1L + 3L) %% p) + 1L]
# stopifnot(max(abs(recovered - expected)) < 1e-6)
# cat("rotation round-trip max error:",
#     sprintf("%.2e\n", max(abs(recovered - expected))))

## -----------------------------------------------------------------------------
B <- lapply(seq_len(n_sites), function(k) random_drift(p, beta_demo))

## Each site's private database, embedded under that site's
## fine-tuned model and unit-normalized on the sphere.
db <- lapply(seq_len(n_sites), function(k) {
  cohort <- site_cohorts[[k]]
  list(z = embed_private(cohort$z, B[[k]]),
       label = cohort$label)
})

## Each site fits its compatibility adapter on the anchor cohort.
A_hat <- lapply(seq_len(n_sites), function(k) {
  Z_pub  <- public_anchor$z
  Z_priv <- embed_private(Z_pub, B[[k]])
  fit_adapter(Z_priv, Z_pub, mu_demo)
})

## Design-1 deployment: the adapter folded into the database
## offline, giving unit-norm public-compatible vectors that an
## encrypted public-model query scores directly. (At mu = Inf the
## adapter is orthogonal, so folding is norm-preserving and
## Design 1 and Design 2 coincide; they part company at finite mu.)
db_fold <- lapply(seq_len(n_sites), function(k)
  list(z = unit_rows(db[[k]]$z %*% A_hat[[k]]), label = db[[k]]$label))

## Setup diagnostics the master would receive: anchor-reconstruction
## error and the adapter's departure from isometry.
cat("Per-site adapter diagnostics (beta =", beta_demo, ", mu = Inf):\n")
for (k in seq_len(n_sites)) {
  Zr <- embed_private(public_anchor$z, B[[k]])
  recon <- norm(Zr %*% A_hat[[k]] - public_anchor$z, "F")
  aniso <- norm(crossprod(A_hat[[k]]) - diag(p), "F")
  cat(sprintf("  site %d: anchor recon %.2e,  ||A^T A - I||_F %.2e\n",
              k, recon, aniso))
}

## -----------------------------------------------------------------------------
build_diagonals <- function(M, p) {
  ## d_i[j] = M[j, ((j-1 + i) %% p) + 1]   (1-based R indexing)
  lapply(0:(p - 1L), function(i) {
    vapply(seq_len(p),
           function(j) M[j, ((j - 1L + i) %% p) + 1L],
           numeric(1L))
  })
}

encrypted_matvec <- function(ct_q, M, cc, p) {
  diags <- build_diagonals(M, p)
  ct_acc <- NULL
  for (i in 0:(p - 1L)) {
    d_pt <- make_ckks_packed_plaintext(cc, diags[[i + 1L]])
    ct_term <- if (i == 0L) {
      eval_mult(ct_q, d_pt)
    } else {
      eval_mult(eval_rotate(ct_q, i), d_pt)
    }
    ct_acc <- if (is.null(ct_acc)) ct_term else eval_add(ct_acc, ct_term)
  }
  ct_acc
}

## ----eval=RECOMPUTE-----------------------------------------------------------
# q_demo <- public_query$z[1, ]
# ct_q  <- encrypt(pub, q_demo)
# ct_Aq <- encrypted_matvec(ct_q, A_hat[[1]], cc, p)
# Aq_recovered <- decrypt(master, ct_Aq, len = p)
# Aq_expected  <- as.numeric(A_hat[[1]] %*% q_demo)
# cat(sprintf("matvec max error (site 1): %.2e\n",
#             max(abs(Aq_recovered - Aq_expected))))

## -----------------------------------------------------------------------------
slot_sum_reduction <- function(ct, p) {
  ## Standard CKKS log-p reduction. The rotations are cyclic over
  ## the whole batch, and here the batch is exactly p slots wide,
  ## so after the loop *every* slot holds the same value: the full
  ## sum_{j=1}^{p} ct[j]. We read slot 0 by convention.
  step <- p %/% 2L
  while (step >= 1L) {
    ct <- eval_add(ct, eval_rotate(ct, step))
    step <- step %/% 2L
  }
  ct
}

encrypted_inner_product <- function(ct_x, v_plain, cc, p) {
  pt_v <- make_ckks_packed_plaintext(cc, v_plain)
  ct_prod <- eval_mult(ct_x, pt_v)
  slot_sum_reduction(ct_prod, p)
}

## ----eval=RECOMPUTE-----------------------------------------------------------
# v_test <- db[[1]]$z[1, ]
# ct_score <- encrypted_inner_product(ct_Aq, v_test, cc, p)
# score_recovered <- decrypt(master, ct_score, len = 1L)
# score_expected  <- sum(Aq_expected * v_test)
# cat(sprintf("inner-product error (site 1, patient 1): %.2e\n",
#             abs(score_recovered - score_expected)))

## ----eval=RECOMPUTE-----------------------------------------------------------
# ct_score_fold <- encrypted_inner_product(ct_q, db_fold[[1]]$z[1, ], cc, p)
# fold_recovered <- decrypt(master, ct_score_fold, len = 1L)
# fold_expected  <- sum(q_demo * db_fold[[1]]$z[1, ])
# cat(sprintf("Design-1 inner-product error (site 1, patient 1): %.2e\n",
#             abs(fold_recovered - fold_expected)))

## ----eval=RECOMPUTE-----------------------------------------------------------
# all_slots <- decrypt(master, ct_score, len = p)
# cat(sprintf("slots holding the released score: %d of %d (max deviation %.2e)\n",
#             sum(abs(all_slots - score_recovered) < 1e-6), p,
#             max(abs(all_slots - score_recovered))))

## ----eval=RECOMPUTE-----------------------------------------------------------
# make_similarity_site_fn <- function(A_k, db_k, cc, p) {
#   function(ct_q) {
#     ct_Aq <- encrypted_matvec(ct_q, A_k, cc, p)
#     n_local <- nrow(db_k$z)
#     ct_scores <- lapply(seq_len(n_local), function(i) {
#       encrypted_inner_product(ct_Aq, db_k$z[i, ], cc, p)
#     })
#     list(scores = ct_scores,
#          local_index = seq_len(n_local),
#          label = db_k$label)
#   }
# }
# 
# site_fns <- lapply(seq_len(n_sites), function(k) {
#   make_similarity_site_fn(A_hat[[k]], db[[k]], cc, p)
# })

## ----eval=RECOMPUTE-----------------------------------------------------------
# t0 <- proc.time()
# site1_out <- site_fns[[1]](ct_q)
# site1_elapsed <- (proc.time() - t0)[["elapsed"]]
# cat(sprintf("site 1 produced %d encrypted scores in %.2f s\n",
#             length(site1_out$scores), site1_elapsed))

## ----eval=RECOMPUTE-----------------------------------------------------------
# run_similarity_query <- function(ct_q, site_fns, master, top_k) {
#   ## Fan out to every site.
#   site_results <- lapply(seq_along(site_fns), function(k) {
#     out <- site_fns[[k]](ct_q)
#     out$site_id <- k
#     out
#   })
# 
#   ## Threshold-decrypt each per-patient inner product. v1 runs
#   ## one ceremony per patient; a packed variant that fuses
#   ## multiple inner products into a single encrypted value
#   ## (via slot tiling) is a natural extension.
#   rows <- list()
#   for (s in site_results) {
#     for (i in seq_along(s$scores)) {
#       score <- decrypt(master, s$scores[[i]], len = 1L)
#       rows[[length(rows) + 1L]] <- data.frame(
#         site_id     = s$site_id,
#         local_index = s$local_index[i],
#         label       = s$label[i],
#         score       = score)
#     }
#   }
#   scored <- do.call(rbind, rows)
#   scored <- scored[order(-scored$score), ]
#   head(scored, top_k)
# }
# 
# t0 <- proc.time()
# top_result <- run_similarity_query(ct_q, site_fns, master, top_k = top_k)
# elapsed <- (proc.time() - t0)[["elapsed"]]
# cat(sprintf("Top-%d retrieval over %d sites and %d patients in %.1f s\n",
#             top_k, n_sites, sum(cohort_sizes), elapsed))
# cat(sprintf("Query phenotype label: %d\n", public_query$label))
# print(top_result, row.names = FALSE)

## ----eval=RECOMPUTE-----------------------------------------------------------
# plaintext_top_k <- function(q, site_data, A_list, top_k) {
#   rows <- list()
#   for (k in seq_along(site_data)) {
#     Aq  <- as.numeric(A_list[[k]] %*% q)
#     s_k <- site_data[[k]]
#     scores <- as.numeric(s_k$z %*% Aq)
#     for (i in seq_along(scores)) {
#       rows[[length(rows) + 1L]] <- data.frame(
#         site_id     = k,
#         local_index = i,
#         label       = s_k$label[i],
#         score       = scores[i])
#     }
#   }
#   scored <- do.call(rbind, rows)
#   scored <- scored[order(-scored$score), ]
#   head(scored, top_k)
# }
# 
# plain_top <- plaintext_top_k(q_demo, db, A_hat, top_k)
# 
# ## Compare encrypted-domain top-k against the cleartext reference
# ## by joining on (site_id, local_index).
# compare <- merge(top_result, plain_top,
#                  by = c("site_id", "local_index"),
#                  suffixes = c("_enc", "_plain"))
# score_err <- max(abs(compare$score_enc - compare$score_plain))
# cat(sprintf("Top-%d encrypted vs cleartext score max error: %.2e\n",
#             top_k, score_err))
# 
# ## Whether the encrypted-domain top-k contains the same
# ## (site_id, local_index) pairs as the cleartext reference.
# enc_set   <- paste(top_result$site_id, top_result$local_index, sep = ":")
# plain_set <- paste(plain_top$site_id,  plain_top$local_index,  sep = ":")
# cat(sprintf("Top-%d set match: %d of %d\n",
#             top_k, length(intersect(enc_set, plain_set)), top_k))

## ----eval=RECOMPUTE, echo=FALSE-----------------------------------------------
# similarity_results <- list(
#     pub_print     = capture.output(print(pub)),
#     rot_err       = max(abs(recovered - expected)),
#     matvec_err    = max(abs(Aq_recovered - Aq_expected)),
#     ip_err        = abs(score_recovered - score_expected),
#     fold_err      = abs(fold_recovered - fold_expected),
#     slots_same    = sum(abs(all_slots - score_recovered) < 1e-6),
#     slots_dev     = max(abs(all_slots - score_recovered)),
#     site1_n       = length(site1_out$scores),
#     site1_elapsed = site1_elapsed,
#     query_elapsed = elapsed,
#     top_result    = top_result,
#     score_err     = score_err,
#     set_match     = length(intersect(enc_set, plain_set)))

## -----------------------------------------------------------------------------
centers_u <- phenotype_centers / sqrt(rowSums(phenotype_centers^2))
embed_cfg <- function(n, sep = 0.85, sd = 0.40) {
  labels <- sample.int(n_phenotypes, n, replace = TRUE)
  z <- sep * centers_u[labels, , drop = FALSE] +
       matrix(rnorm(n * p, sd = sd), n, p)
  list(z = unit_rows(z), label = labels)
}

recall_at_k <- function(query_label, top_rows, k = top_k)
  sum(top_rows$label == query_label) / k

## Federated recall of a query population under one deployment.
## design 1: fold A into the db (unit) and score <q, z>;
## design 2: apply A to the query and score <Aq, h> against raw db.
fed_recall <- function(q_pop, db_list, A_list, design) {
  mean(vapply(seq_len(nrow(q_pop$z)), function(qi) {
    q <- q_pop$z[qi, ]
    parts <- lapply(seq_along(db_list), function(k) {
      s <- db_list[[k]]
      sc <- if (design == 1L) as.numeric(unit_rows(s$z %*% A_list[[k]]) %*% q)
            else              as.numeric(s$z %*% as.numeric(A_list[[k]] %*% q))
      data.frame(label = s$label, score = sc)
    })
    scored <- do.call(rbind, parts)
    recall_at_k(q_pop$label[qi],
                head(scored[order(-scored$score), ], top_k))
  }, numeric(1)))
}

## A world at non-isometry beta with an anchor cohort of size na.
make_world <- function(beta, na) {
  anchor  <- embed_cfg(na)
  cohorts <- lapply(cohort_sizes, embed_cfg)
  Bs <- lapply(seq_len(n_sites), function(k) random_drift(p, beta))
  list(anchor  = anchor,
       pub_db  = cohorts,                          # public embeddings (ideal ref)
       priv_db = lapply(seq_len(n_sites), function(k)
         list(z = embed_private(cohorts[[k]]$z, Bs[[k]]),
              label = cohorts[[k]]$label)),
       Hanchor = lapply(seq_len(n_sites), function(k)
         embed_private(anchor$z, Bs[[k]])))
}

I_list  <- replicate(n_sites, diag(p), simplify = FALSE)
n_rep   <- 3L
n_q     <- 40L
mu_grid <- c(0, 0.1, 1, 10, Inf)

## --- mu sweep at fixed beta, ample anchor: fidelity + posture + Gram ---
mu_sweep <- function(beta, na) {
  tab <- 0; ideal <- 0; unaligned <- 0; gram <- 0
  for (r in seq_len(n_rep)) {
    w  <- make_world(beta, na)
    qp <- embed_cfg(n_q)
    ideal     <- ideal     + fed_recall(qp, w$pub_db,  I_list, 2L)
    unaligned <- unaligned + fed_recall(qp, w$priv_db, I_list, 2L)
    Ag   <- lapply(seq_len(n_sites), function(k) fit_gram(w$Hanchor[[k]], w$anchor$z))
    gram <- gram + fed_recall(qp, w$priv_db, Ag, 1L)
    rows <- lapply(mu_grid, function(mu) {
      A <- lapply(seq_len(n_sites), function(k)
        fit_adapter(w$Hanchor[[k]], w$anchor$z, mu))
      data.frame(mu = mu,
                 d1 = fed_recall(qp, w$priv_db, A, 1L),
                 d2 = fed_recall(qp, w$priv_db, A, 2L),
                 aniso = mean(vapply(A, function(Ak)
                   norm(crossprod(Ak) - diag(p), "F"), numeric(1))))
    })
    tab <- tab + do.call(rbind, rows)
  }
  tab <- tab / n_rep
  tab$mu <- mu_grid
  list(tab = tab, ideal = ideal / n_rep,
       unaligned = unaligned / n_rep, gram = gram / n_rep)
}

mu_main   <- mu_sweep(beta = 0.6, na = 100L)   # ample calibration
mu_scarce <- mu_sweep(beta = 0.6, na = 24L)    # anchor < p = 32

## --- beta sweep, ample anchor: Procrustes vs near-orthogonal vs LS ---
## Common random numbers: within a replicate the public data and the
## per-site drift directions are fixed and only the stretch magnitude
## beta varies, so the no-drift ideal is a single drift-independent
## reference rather than a wandering curve.
beta_grid <- c(0, 0.3, 0.6, 1.0)
beta_acc <- 0; ideal_b <- 0
for (r in seq_len(n_rep)) {
  anchor  <- embed_cfg(n_anchor)
  cohorts <- lapply(cohort_sizes, embed_cfg)
  qp      <- embed_cfg(n_q)
  Qg <- lapply(seq_len(n_sites), function(k)
    list(Q = qr.Q(qr(matrix(rnorm(p * p), p, p))), g = rnorm(p)))
  ideal_b  <- ideal_b + fed_recall(qp, cohorts, I_list, 2L)
  beta_acc <- beta_acc + do.call(rbind, lapply(beta_grid, function(b) {
    Bs   <- lapply(Qg, function(qg)
      if (b == 0) qg$Q else qg$Q %*% diag(exp(b * qg$g)))
    priv <- lapply(seq_len(n_sites), function(k)
      list(z = embed_private(cohorts[[k]]$z, Bs[[k]]),
           label = cohorts[[k]]$label))
    Hanc <- lapply(seq_len(n_sites), function(k)
      embed_private(anchor$z, Bs[[k]]))
    fitb <- function(mu) lapply(seq_len(n_sites), function(k)
      fit_adapter(Hanc[[k]], anchor$z, mu))
    data.frame(beta = b,
               LS   = fed_recall(qp, priv, fitb(0), 1L),
               near = fed_recall(qp, priv, fitb(1), 1L),
               Proc = fed_recall(qp, priv, fitb(Inf), 2L))
  }))
}
beta_tab <- beta_acc / n_rep; beta_tab$beta <- beta_grid
ideal_b  <- ideal_b / n_rep

cat("mu sweep (beta=0.6, anchor=100):  ideal=",
    sprintf("%.3f", mu_main$ideal), " unaligned=",
    sprintf("%.3f", mu_main$unaligned), " Gram=",
    sprintf("%.3f\n", mu_main$gram), sep = "")
print(round(mu_main$tab, 3), row.names = FALSE)
cat("\nbeta sweep (anchor=100, common random numbers):  ideal=",
    sprintf("%.3f\n", ideal_b), sep = "")
print(round(beta_tab, 3), row.names = FALSE)
cat("\nmu sweep at scarce anchor=24:\n")
print(round(mu_scarce$tab[c("mu", "d1", "d2")], 3), row.names = FALSE)

## ----recall-figure, fig.width=9, fig.height=7, fig.align="center"-------------
op <- par(mfrow = c(2, 2), mar = c(4.2, 4.2, 2.4, 1.0))
xi <- seq_along(mu_grid)
mulab <- function(m) ifelse(is.infinite(m), "Inf", formatC(m, format = "g"))

## (1) recall vs mu: Design 1 vs Design 2, with Gram / ideal / floor
plot(xi, mu_main$tab$d1, type = "b", pch = 19, ylim = c(0, 1), xaxt = "n",
     xlab = expression(mu ~ "(0 = LS    ->    Inf = Procrustes)"),
     ylab = sprintf("recall@%d", top_k), main = "Recall vs mu (beta = 0.6)")
axis(1, xi, mulab(mu_grid))
lines(xi, mu_main$tab$d2, type = "b", pch = 1, lty = 2, col = "firebrick")
abline(h = mu_main$ideal, lty = 3, col = "darkgreen")
abline(h = mu_main$unaligned, lty = 3, col = "grey60")
abline(h = mu_main$gram, lty = 4, lwd = 2, col = "orange")
legend("right", bty = "n", cex = 0.75,
       legend = c("Design 1 (fold)", "Design 2 (matvec)", "Gram-PSD",
                  "ideal", "unaligned"),
       col = c("black", "firebrick", "orange", "darkgreen", "grey60"),
       lty = c(1, 2, 4, 3, 3), pch = c(19, 1, NA, NA, NA))

## (2) recall vs beta: Procrustes vs near-orthogonal vs LS
plot(beta_tab$beta, beta_tab$Proc, type = "b", pch = 19, ylim = c(0, 1),
     xlab = expression(beta ~ "(non-isometry)"), ylab = sprintf("recall@%d", top_k),
     main = "Procrustes vs relaxed adapters")
lines(beta_tab$beta, beta_tab$near, type = "b", pch = 1, lty = 2, col = "blue")
lines(beta_tab$beta, beta_tab$LS, type = "b", pch = 2, lty = 3, col = "firebrick")
abline(h = ideal_b, lty = 3, col = "darkgreen")
legend("bottomleft", bty = "n", cex = 0.75,
       legend = c("Procrustes (mu=Inf)", "near-orth (mu=1)", "LS (mu=0)", "ideal"),
       col = c("black", "blue", "firebrick", "darkgreen"),
       lty = c(1, 2, 3, 3), pch = c(19, 1, 2, NA))

## (3) recall vs mu at scarce anchor (< p): the regularization sweet spot
plot(xi, mu_scarce$tab$d1, type = "b", pch = 19, ylim = c(0, 1), xaxt = "n",
     xlab = expression(mu), ylab = sprintf("recall@%d (Design 1)", top_k),
     main = "Scarce anchor (n = 24 < p = 32)")
axis(1, xi, mulab(mu_grid))
lines(xi, mu_scarce$tab$d2, type = "b", pch = 1, lty = 2, col = "firebrick")
legend("bottomleft", bty = "n", cex = 0.75,
       legend = c("Design 1 (fold)", "Design 2 (matvec)"),
       col = c("black", "firebrick"), lty = c(1, 2), pch = c(19, 1))

## (4) departure from isometry vs mu (the Design-2 / matvec budget)
plot(xi, mu_main$tab$aniso + 1e-12, type = "b", pch = 19, log = "y", xaxt = "n",
     xlab = expression(mu), ylab = expression("||" * A^T * A - I * "||"[F]),
     main = "Departure from isometry")
axis(1, xi, mulab(mu_grid))
par(op)

