import numpy as np
import pandas as pd
from homomorphepy.examples import similarity as sim
rng = np.random.default_rng(20260428)
p = 32
n_sites = 3
cohort_sizes = (80, 60, 100)
n_anchor = 100
n_phenotypes = 5
top_k = 5Federated Cosine Similarity with Site-Private Fine-Tuned Models
Introduction
Similarity search across silos is another common problem. A clinician at one hospital sees an unusual case and wants to ask “do other hospitals in our network have similar patients?” — without sending the patient out, without the other hospitals revealing which of their patients were checked, and without any single party (including a coordinating master) able to decrypt a similarity score on its own.
Modern foundation models map patient data (radiology reports, chest X-rays, pathology slides) into a fixed-dimensional vector space where geometric proximity tracks clinical similarity. The natural primitive is cosine similarity between such embeddings, which reduces to an inner product on unit-norm vectors, which is exactly what threshold CKKS computes efficiently.
The wrinkle is that hospitals do not all use the same model. We assume each hospital starts from a public foundation model and fine-tunes it on its own data for its own purposes — it does not constrain the result to be an isometry of the public model. The fine-tuned embeddings therefore differ from the public baseline by a generically non-isometric transformation of the geometry, and a query in the public-model space is no longer directly comparable to a database vector in a hospital’s private space.
One way to close the gap is with a per-site compatibility adapter \(A_k\), fit post hoc on a public anchor cohort. The adapter lives on a one-parameter family indexed by a near-isometry penalty \(\mu\): orthogonal Procrustes at one end (\(\mu\to\infty\)), unconstrained least squares at the other (\(\mu = 0\)), near-orthogonal maps in between. It is deployed in either of two ways — applied to the encrypted query (the homomorphic matrix–vector multiply, valid as a cosine near the orthogonal end), or folded into the database offline so the encrypted query reduces to a plain inner product. Threshold decryption recovers every candidate score at the master, which ranks them and returns the top-k matches; no party can decrypt any intermediate value unilaterally.
The setup
Three sites — labeled \(S_1, S_2, S_3\) — each hold a cohort of patients. Each cohort is embedded with the site’s privately fine-tuned model into \(\mathbf{R}^p\) for some fixed dimension \(p\) (the architecture is shared across sites; only the weights differ).
The query enters the system from a clinician who does not have access to any site’s private model. They embed their candidate patient with the public foundation model, producing \(q \in \mathbf{R}^p\). The query is encrypted under a joint threshold key whose secret-key material is shared \(n\)-of-\(n\) across the three sites.
The retrieval target is the top-\(k\) patients across all sites whose private-model embeddings are most cosine-similar to the query, returned as a list of \((\text{site\_id}, \text{local\_patient\_id}, \text{score})\) tuples.
The difficulty is that \(q\) lives in the public-model geometry while each site’s database lives in the site’s private-model geometry. We close the gap with a per-site compatibility adapter \(A_k\) fit on a public anchor cohort. \(A_k\) is held privately by site \(k\) and never leaves; it is applied to encrypted queries to bring them into agreement with site \(k\)’s geometry before the inner-product step (or, equivalently, folded into the site’s database offline).
Threat model
Three sites and one untrusted aggregator (the master):
- Sites \(S_1, S_2, S_3\) each hold private patient embeddings, a private compatibility adapter \(A_k\), and a secret-key share \(\mathit{sk}_k\). They are honest-but-curious among themselves and toward the master. Each site sees the master’s encrypted query but cannot decrypt it; each site sees its own database in the clear (it is its own data); each site never sees other sites’ databases or scores.
- Master holds no secret-key material. It receives the encrypted query from the querier, broadcasts it to the sites, collects their encrypted scores, orchestrates the n-of-n partial-decryption ceremony, sorts the results, and returns the top-\(k\) to the querier. A curious or compromised master cannot decrypt anything by itself.
- Querier sees only the top-\(k\) result tuples. Repeated adaptive queries leak structural information about the database in the same Hyrum-style sense as any retrieval system; this leakage is acknowledged but not eliminated.
What the master sees, by stage:
- The encrypted query \(\mathit{ct}_q\). Not decryptable alone.
- Per-site encrypted scores, one per candidate patient, with the site id, local index, and label in the clear alongside.
- Partial decryptions from each site. Not decryptable individually.
- After fusion, every candidate score in the clear. It sorts these and returns the top-\(k\) to the querier.
The master decrypts the whole score vector, not just the \(k\) it releases, and it does so because this implementation ranks in the clear. Nothing can do better without encrypted comparison: the global top-\(k\) is not knowable until enough scores are known to rank them. So the master’s view is one row per candidate patient in the whole federation — site, local index, label, and score — and the querier’s view is the top-\(k\) of it.
The site-private compatibility adapters \(\{A_k\}\) never leave their sites. Each is applied to an encrypted vector while it stays encrypted — an unencrypted matrix times an encrypted vector — or else folded into the site’s database offline, so the adapter never appears in the clear at any party other than its owner.
The protocol
The protocol has three phases. The setup phase runs once, offline; the query phase and result phase run for each query.
Setup phase.
- The three sites jointly generate a CKKS threshold key pair (\(n\)-of-\(n\)). Joint rotation keys for slot rotations \(1, 2, \ldots, p-1\) are also generated by an \(n\)-of-\(n\) ceremony.
- A public anchor cohort of patients with public-model embeddings is published.
- Each site embeds the anchor cohort with its private fine-tuned model and fits its \(p \times p\) compatibility adapter \(A_k\) along the \(\mu\) family (orthogonal Procrustes at \(\mu\to\infty\), least squares at \(\mu = 0\)).
- Each site stores its local cohort, embedded with the private model and normalized to unit length — and, for the fold (Design 1) deployment, the public-compatible images of those vectors.
Query phase. The querier embeds their candidate patient with the public model, normalizes to unit length, encrypts into a single slot-packed encrypted vector under the joint public key, and broadcasts that to the master.
For each site \(k\):
- The site receives the encrypted query \(\mathit{ct}_q\).
- The site applies its private \(A_k\) to the encrypted query via diagonal-encoded matrix–vector multiplication, producing \(\mathit{ct}_{A_k q}\) (Design 2). Near the orthogonal end of the family the result is unit-norm, so the score is a genuine cosine; alternatively the site skips this step and scores the query directly against its folded public-compatible database (Design 1).
- The site computes inner products against each of its local embeddings \(v_{k,i}\) via slot-wise multiplication followed by log-\(p\) rotation-and-add slot summation, producing one encrypted scalar per local patient.
- The site returns its encrypted scores to the master, alongside its local patient indices in the clear.
Result phase.
- The master concatenates encrypted scores across all sites. In v1, local indices are kept in the clear alongside the encrypted scores; encrypted top-\(k\) selection in CKKS is feasible but adds depth and complexity orthogonal to this page’s pedagogical aim.
- The master collects partial decryption shares from all three sites for the encrypted scores. Fusion yields the scores in the clear; the master sorts them there and returns the top-\(k\) to the querier.
Synthetic data
The page uses fully synthetic data so it renders quickly. The phenotype-mixture model below is the simplest setup that gives a meaningful retrieval ground truth.
A patient is one of n_phenotypes clinical phenotypes. Each phenotype is a Gaussian cluster in \(\mathbf{R}^p\). Phenotype labels are the retrieval ground truth: a query of phenotype \(j\) should be matched to database patients of phenotype \(j\).
phenotype_centers = rng.normal(size=(n_phenotypes, p))
def unit_rows(Z):
return Z / np.linalg.norm(Z, axis=1, keepdims=True)
def embed_public(n):
labels = rng.integers(0, n_phenotypes, size=n)
centers = phenotype_centers[labels]
noise = rng.normal(scale=0.4, size=(n, p))
return dict(z=unit_rows(centers + noise), label=labels) # unit normA site fine-tunes for its own purposes; it does not constrain the result to be an isometry of the public model. We therefore model the per-site drift as a non-isometric linear map \(B_k = Q_k D_k\) — a random rotation \(Q_k\) composed with a diagonal stretch \(D_k = \mathrm{diag}(e^{\beta g})\), \(g\) standard normal. The parameter \(\beta\) is the non-isometry magnitude: \(\beta = 0\) recovers an exactly orthogonal drift (the special case where the geometry is merely rotated), and \(\beta > 0\) stretches the embedding directions anisotropically, as a freely fine-tuned model generically would. The site’s embeddings are unit-normalized after the map, so the drift acts on the sphere.
import inspect
print(inspect.getsource(sim.random_drift))def random_drift(p: int, beta: float, rng: np.random.Generator) -> np.ndarray:
"""A fine-tune, modeled as ``B = Q D``.
A random rotation composed with an anisotropic stretch
``D = diag(exp(beta * g))``. ``beta = 0`` gives an exactly
orthogonal drift -- the special case where the geometry is merely
rotated, and the one the protocol walk-through uses. ``beta > 0``
stretches directions unevenly, which is what a freely fine-tuned
model generically does.
"""
q, _ = np.linalg.qr(rng.normal(size=(p, p)))
if beta == 0:
return q
return q @ np.diag(np.exp(beta * rng.normal(size=p)))
def embed_private(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 @ B_k.T, then
# unit-normalized (the drift acts on the sphere; under a
# non-isometric B the normalization is a genuine nonlinearity).
return unit_rows(z_public @ B_k.T)The public anchor cohort is a small, publicly-available set of patient examples. Each site embeds the anchors twice — once with the public model (giving \(Z_{\text{pub}}\), identical at every site by construction) and once with its own private fine-tuned model (giving \(Z_{\text{priv},k}\)). The pair drives a post-hoc compatibility adapter \(A_k\) that maps the site’s private geometry back to the public protocol, fit by
\[ \min_{A}\ \lVert Z_{\text{priv},k}\,A - Z_{\text{pub}}\rVert_F^2 \;+\; \mu\,\lVert A^\top A - I\rVert_F^2 . \]
The penalty parameter \(\mu\) traces a one-parameter family. At \(\mu \to \infty\) the adapter is forced orthogonal and the fit is the classical orthogonal Procrustes problem, solved in closed form by the singular value decomposition — this is the rigid endpoint, exact only when the drift is itself isometric. At \(\mu = 0\) the adapter is an unconstrained least-squares map, computed with a small ridge on the normal equations, the best public-space fit but generically non-isometric. Intermediate \(\mu\) interpolates: a near-orthogonal map. We deploy \(A_k\) by applying it to the encrypted query, \(\langle A_k q,\, v\rangle\), which recovers the public-space cosine \(\langle q, u\rangle\) when the adapter inverts the drift (exactly, in the orthogonal limit).
print(inspect.getsource(sim.fit_adapter_mu))def fit_adapter_mu(
anchor_public: np.ndarray,
anchor_site: np.ndarray,
mu: float,
) -> np.ndarray:
"""The adapter family, indexed by a near-isometry penalty.
Minimizes ``||Z_site A - Z_public||^2 + mu * ||A^T A - I||^2``.
At ``mu = inf`` the constraint binds and this is the orthogonal
Procrustes solution of :func:`fit_adapter`; at ``mu = 0`` it is
unconstrained least squares, computed with a small ridge on the
normal equations; in between, near-orthogonal.
"""
p = anchor_site.shape[1]
if np.isinf(mu):
return fit_adapter(anchor_public, anchor_site)
gram = anchor_site.T @ anchor_site + 1e-6 * np.eye(p)
a_ls = np.linalg.solve(gram, anchor_site.T @ anchor_public)
if mu == 0:
return a_ls
from scipy.optimize import minimize
def fn(par):
a = par.reshape(p, p)
resid = anchor_site @ a - anchor_public
skew = a.T @ a - np.eye(p)
return float(np.sum(resid**2) + mu * np.sum(skew**2))
def gr(par):
a = par.reshape(p, p)
resid = anchor_site @ a - anchor_public
skew = a.T @ a - np.eye(p)
return (2 * anchor_site.T @ resid + 4 * mu * (a @ skew)).ravel()
out = minimize(
fn, a_ls.ravel(), jac=gr, method="L-BFGS-B", options=dict(maxiter=400)
)
return out.x.reshape(p, p)
A convex alternative is Gram (metric) learning: fit a positive-semidefinite \(M\) matching the public Gram matrix, \(\min_{M\succeq0}\lVert Z_{\text{priv}}MZ_{\text{priv}}^\top - Z_{\text{pub}}Z_{\text{pub}}^\top\rVert_F^2\). Its minimizer is \(M = A_{\text{LS}}A_{\text{LS}}^\top\), and we deploy its symmetric square root. The sweep below compares it with the \(\mu\) family.
print(inspect.getsource(sim.fit_gram))def fit_gram(anchor_public: np.ndarray, anchor_site: np.ndarray) -> np.ndarray:
"""The convex alternative: match Gram matrices, deploy the square root.
Included only so the comparison below is against something, rather
than against nothing. Its minimizer is analytic.
"""
p = anchor_site.shape[1]
a_ls = np.linalg.solve(
anchor_site.T @ anchor_site + 1e-6 * np.eye(p),
anchor_site.T @ anchor_public,
)
m = a_ls @ a_ls.T
vals, vecs = np.linalg.eigh(m)
return vecs @ (np.sqrt(np.maximum(vals, 0.0))[:, None] * vecs.T)
A moderate non-isometry magnitude drives the protocol walk-through; the sweeps later vary it:
beta_demo = 0.6
mu_demo = np.inf # orthogonal endpoint for the encrypted walk-throughThe site cohorts and anchor cohort are sampled from the same phenotype-mixture model (in real deployments the anchor cohort is publicly distributed and its phenotype distribution is set once; we approximate by independent sampling).
public_anchor = embed_public(n_anchor)
public_query = embed_public(1) # one query for the protocol walk-through
site_cohorts = [embed_public(n) for n in cohort_sizes]We later fit one adapter \(A_k\) per site, build both the raw-private and folded public-compatible databases, and report retrieval recall as the non-isometry \(\beta\) and the penalty \(\mu\) vary.
Threshold key generation
The three sites jointly construct a CKKS keypair so that the secret key is split \(n\)-of-\(n\) across them. make_threshold_master() walks that chain across the sites: each site generates its own share, keeps it, and passes on only a public key. The master is handed the joint public key and holds no secret material at all.
Because the joint key is built from the sites, the sites are constructed first.
from homomorphepy import fhe_context, make_worker, make_threshold_master, Ct, packed_codec
from homomorphepy.examples.similarity import CKKS_PARAMS
cc = fhe_context("CKKS", **CKKS_PARAMS)
sites = [make_worker(f"Site {k + 1}", None, lambda data, theta: None)
for k in range(n_sites)]
master = make_threshold_master("master", cc, 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 slot a key share could sit in.
pub = sites[0].params
pub<OpenFHEParams CKKS
public key b886b9581854f4531a1dabbb76e4a433
secret material: none>
A CKKS encrypted value holds a vector of numbers, one per slot, and arithmetic acts on all slots at once. Setting batch_size to \(p\) gives each encrypted value exactly \(p\) = 32 slots, so one embedding fills one encrypted value, a coordinate per slot.
Site \(k\)’s share lives in sites[k]._share and nowhere else. master.decrypt(ct, length) runs the \(n\)-of-\(n\) ceremony by asking each site for a partial decryption — no party ever holds a usable secret unilaterally, the aggregator included.
Everything below encrypts with pub. No party consults the aggregator again; each holds what it was given and works from that.
Joint rotation keys
The matrix–vector multiply step in the protocol is implemented as a diagonal-encoded matvec. For a \(p \times p\) matrix \(M\) and an encrypted vector \(q \in \mathbf{R}^p\) held one component per slot:
\[ M \cdot q \;=\; \sum_{i=0}^{p-1} d_i \odot \mathrm{rot}(q, i), \]
where \(d_i\) is the \(i\)-th diagonal of \(M\) (a \(p\)-vector, and never encrypted), \(\odot\) is slot-wise multiplication, and \(\mathrm{rot}(q, i)\) cyclically rotates the slots of \(q\) by \(i\). Each rotation requires a precomputed rotation key. In single-key CKKS these are generated from the secret key; in threshold CKKS they are generated by an \(n\)-of-\(n\) ceremony that mirrors the encryption-key ceremony.
The inner-product step adds a second use of rotation keys: the log-\(p\) rotation-and-add reduction that sums the slots of an encoded \(p\)-vector into slot 0. The same set of rotation indices serves both purposes; we generate keys for indices \(1, \ldots, p-1\) and rely on the subset that each step needs.
make_joint_rotation_keys() runs that ceremony at the sites: the lead site generates its rotation keys under its own share, each remaining site folds in its own share of the keys in turn, and the accumulated map is registered under the joint public key’s tag. Nothing collects the shares into one place.
from homomorphepy.actors import make_joint_rotation_keys
rotation_indices = range(1, p)
make_joint_rotation_keys(master, rotation_indices)After insertion, any value encrypted under the joint public key can be rotated by any index in rotation_indices via EvalRotate(ct, idx).
A round-trip smoke test confirms the joint rotation key works. We encrypt a known vector, rotate by 3 slots, decrypt via the \(n\)-of-\(n\) ceremony, and check that slot \(i\) now holds \(x_{(i+3) \bmod p}\):
x = np.arange(1, p + 1, dtype=float)
ct_x = pub.encrypt(x.tolist())
ct_rot = Ct(cc.cc.EvalRotate(ct_x.raw, 3), cc.cc)
recovered = np.asarray(master.decrypt(ct_rot, length=p), dtype=float)
expected = x[(np.arange(p) + 3) % p]
rot_err = float(np.max(np.abs(recovered - expected)))
assert rot_err < 1e-6
print(f"rotation round-trip max error: {rot_err:.2e}")rotation round-trip max error: 2.70e-13
Per-site adapter fit and database setup
For the encrypted walk-through we fit the adapter at the orthogonal endpoint (\(\mu = \infty\)), so the map applied to the query is exactly norm-preserving and the matrix–vector multiply below produces a unit-norm result — the regime in which the raw-private-database matvec (Design 2) is a valid cosine. The recall sweeps later vary \(\mu\) and the drift. Each site’s fine-tuned model is the non-isometric \(B_k\); the adapter \(A_k\) is fit post hoc on the public anchor cohort.
B = [sim.random_drift(p, beta_demo, rng) for _ in range(n_sites)]
# Each site's private database, embedded under that site's
# fine-tuned model and unit-normalized on the sphere.
db = [dict(z=embed_private(c["z"], B[k]), label=c["label"])
for k, c in enumerate(site_cohorts)]
# Each site fits its compatibility adapter on the anchor cohort.
A_hat = [sim.fit_adapter_mu(public_anchor["z"],
embed_private(public_anchor["z"], B[k]), mu_demo)
for k in range(n_sites)]
# 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 = [dict(z=unit_rows(db[k]["z"] @ A_hat[k]), label=db[k]["label"])
for k in range(n_sites)]
# Setup diagnostics the master would receive: anchor-reconstruction
# error and the adapter's departure from isometry.
print(f"Per-site adapter diagnostics (beta = {beta_demo}, mu = inf):")
for k in range(n_sites):
Zr = embed_private(public_anchor["z"], B[k])
recon = np.linalg.norm(Zr @ A_hat[k] - public_anchor["z"])
aniso = np.linalg.norm(A_hat[k].T @ A_hat[k] - np.eye(p))
print(f" site {k + 1}: anchor recon {recon:.2e}, ||A^T A - I||_F {aniso:.2e}")Per-site adapter diagnostics (beta = 0.6, mu = inf):
site 1: anchor recon 4.07e+00, ||A^T A - I||_F 9.75e-15
site 2: anchor recon 2.39e+00, ||A^T A - I||_F 9.91e-15
site 3: anchor recon 3.43e+00, ||A^T A - I||_F 9.01e-15
Diagonal-encoded matrix-vector multiply
The CKKS-friendly way to multiply an unencrypted \(p \times p\) matrix into an encrypted vector is the diagonal encoding. The matrix \(M\) is decomposed into its \(p\) generalized diagonals, each an unencrypted length-\(p\) vector:
\[ d_i[j] = M[\,j,\, (j + i) \bmod p\,], \qquad i = 0, 1, \ldots, p-1. \]
The matrix–vector product becomes
\[ M \cdot q \;=\; \sum_{i=0}^{p-1} d_i \,\odot\, \mathrm{rot}(q, i), \]
where \(\odot\) is slot-wise multiplication and \(\mathrm{rot}(q, i)\) is the \(i\)-step cyclic slot rotation. Each term multiplies an encrypted vector by an unencrypted one and adds the result to a running encrypted total. The cost is \(p\) rotations and \(p\) such multiplications per matrix-vector product, and it consumes a single level of the precision budget.
Multiplying by an unencrypted vector rather than an encrypted one is what keeps that cost down: the adapter \(A_k\) is the site’s own, so it never needs encrypting, and the operation is correspondingly cheaper than multiplying two encrypted quantities together.
print(inspect.getsource(sim.diagonals))def diagonals(m: np.ndarray) -> list[np.ndarray]:
"""Generalized diagonals of ``m``: ``d_i[j] = m[j, (j + i) % p]``.
The encoding that turns a matrix-vector product into rotations.
"""
p = m.shape[0]
idx = np.arange(p)
return [m[idx, (idx + i) % p].copy() for i in range(p)]
def encrypted_matvec(ct_q, M, cc, p):
diags = sim.diagonals(M)
codec = packed_codec(cc)
ct_acc = None
for i in range(p):
d_pt = codec.encode(diags[i].tolist())
ct_rot = ct_q if i == 0 else Ct(cc.cc.EvalRotate(ct_q.raw, i), cc.cc)
ct_term = Ct(cc.cc.EvalMult(ct_rot.raw, d_pt), cc.cc)
ct_acc = ct_term if ct_acc is None else Ct(cc.cc.EvalAdd(ct_acc.raw, ct_term.raw), cc.cc)
return ct_accA round-trip smoke test on a known query vector confirms the matvec recovers \(A_1 \cdot q\) to floating-point precision:
q_demo = public_query["z"][0]
ct_q = pub.encrypt(q_demo.tolist())
ct_Aq = encrypted_matvec(ct_q, A_hat[0], cc, p)
Aq_recovered = np.asarray(master.decrypt(ct_Aq, length=p), dtype=float)
Aq_expected = A_hat[0] @ q_demo
matvec_err = float(np.max(np.abs(Aq_recovered - Aq_expected)))
print(f"matvec max error (site 1): {matvec_err:.2e}")matvec max error (site 1): 1.09e-12
Inner product against the local database
After the matvec, the encrypted query has been mapped by the adapter into site \(k\)’s private geometry (unit-norm at the orthogonal endpoint used here). The cosine similarity against a private-database vector \(v\) (also unit-norm, and never encrypted since it never leaves the site) is just
\[ \cos(A_k q,\, v) \;=\; \langle A_k q,\, v \rangle \;=\; \sum_{j=1}^{p} (A_k q)_j \cdot v_j. \]
In CKKS this is a slot-wise multiply of the matvec output by the unencrypted \(v\), followed by a log-\(p\) rotate-and-add slot-summation reduction that places the summed inner product in slot 0.
def slot_sum_reduction(ct, p, cc):
# 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 // 2
while step >= 1:
ct = Ct(cc.cc.EvalAdd(ct.raw, cc.cc.EvalRotate(ct.raw, step)), cc.cc)
step //= 2
return ct
def encrypted_inner_product(ct_x, v_plain, cc, p):
pt_v = packed_codec(cc).encode(np.asarray(v_plain).tolist())
ct_prod = Ct(cc.cc.EvalMult(ct_x.raw, pt_v), cc.cc)
return slot_sum_reduction(ct_prod, p, cc)A smoke test against a single database vector confirms the inner product matches the unencrypted computation at slot 0:
v_test = db[0]["z"][0]
ct_score = encrypted_inner_product(ct_Aq, v_test, cc, p)
score_recovered = float(master.decrypt(ct_score, length=1))
score_expected = float(Aq_expected @ v_test)
ip_err = abs(score_recovered - score_expected)
print(f"inner-product error (site 1, patient 1): {ip_err:.2e}")inner-product error (site 1, patient 1): 3.02e-13
The same encrypted_inner_product serves the Design 1 deployment without any matvec: the encrypted public-model query is scored directly against the folded, public-compatible database vectors. At the orthogonal endpoint this matches the matvec route exactly; for a non-isometric adapter it is the route that stays a valid cosine.
ct_score_fold = encrypted_inner_product(ct_q, db_fold[0]["z"][0], cc, p)
fold_recovered = float(master.decrypt(ct_score_fold, length=1))
fold_expected = float(q_demo @ db_fold[0]["z"][0])
fold_err = abs(fold_recovered - fold_expected)
print(f"Design-1 inner-product error (site 1, patient 1): {fold_err:.2e}")Design-1 inner-product error (site 1, patient 1): 8.06e-14
What is in the slots we do not read
The score is read from slot 0, but the decrypted value has all 32 slots, and length=1 only limits how many are shown. Whoever decrypts can read the rest. If those slots held partial sums, they would reveal more than the score.
They do not. The rotations wrap around exactly \(p\) slots, so every slot ends up holding the full sum. The check below confirms that all 32 slots hold the released score.
all_slots = np.asarray(master.decrypt(ct_score, length=p), dtype=float)
print(f"slots holding the released score: "
f"{int(np.sum(np.abs(all_slots - score_recovered) < 1e-6))} of {p} "
f"(max deviation {np.max(np.abs(all_slots - score_recovered)):.2e})")slots holding the released score: 32 of 32 (max deviation 8.53e-13)
This holds only because the sum covers the whole batch. If several patients were packed into one encrypted value, each sum would cover only some of the slots, and the other slots would hold partial sums. Multiplying by an unencrypted vector that is 1 in slot 0 and 0 elsewhere clears them. That multiplication uses one level of multiplicative depth. With multiplicative_depth = 3, the matrix–vector multiply and the inner product use two levels, which leaves one for it.
Site function: full per-site protocol step
The site-side computation closes over the site’s adapter \(A_k\) and database \(D_k\), takes the encrypted query as input, and returns a list of encrypted scores plus the corresponding local indices in the clear. This is the Design 2 branch: the adapter is applied to the encrypted query (the matvec), and scoring runs against the site’s raw private database.
def make_similarity_site_fn(A_k, db_k, cc, p):
def site_fn(ct_q):
ct_Aq = encrypted_matvec(ct_q, A_k, cc, p)
ct_scores = [encrypted_inner_product(ct_Aq, v, cc, p) for v in db_k["z"]]
return dict(scores=ct_scores,
local_index=np.arange(len(ct_scores)),
label=db_k["label"])
return site_fn
site_fns = [make_similarity_site_fn(A_hat[k], db[k], cc, p) for k in range(n_sites)]A timed end-to-end run on the demo query through site 1 returns one encrypted score per local patient:
import time
t0 = time.perf_counter()
site1_out = site_fns[0](ct_q)
site1_elapsed = time.perf_counter() - t0
print(f"site 1 produced {len(site1_out['scores'])} encrypted scores in {site1_elapsed:.2f} s")site 1 produced 80 encrypted scores in 3.26 s
Master orchestration and threshold-decrypted top-k
The master broadcasts the encrypted query to every site, collects per-site encrypted scores plus the local patient indices and labels in the clear, and runs the \(n\)-of-\(n\) threshold decryption ceremony for each score so it can sort them once they are in the clear and return the top-k.
The decryption pattern is per-patient: each encrypted score goes through one threshold-decrypt round. master.decrypt() sends the encrypted score to each site in turn — MultipartyDecryptLead at site 1, MultipartyDecryptMain at each of the rest, every site applying its own share — and fuses the returned partials with MultipartyDecryptFusion, which needs only the public context. For our 240 total patients this runs in a few seconds; production deployments would pack many patients into the slots of a single encrypted value and amortize the ceremony.
def run_similarity_query(ct_q, site_fns, master, top_k):
# Fan out to every site.
rows = []
for k, fn in enumerate(site_fns):
out = fn(ct_q)
# 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.
for i, ct in enumerate(out["scores"]):
score = float(master.decrypt(ct, length=1))
rows.append(dict(site_id=k + 1, local_index=int(out["local_index"][i]),
label=int(out["label"][i]), score=score))
scored = pd.DataFrame(rows).sort_values("score", ascending=False)
return scored.head(top_k)
t0 = time.perf_counter()
top_result = run_similarity_query(ct_q, site_fns, master, top_k=top_k)
elapsed = time.perf_counter() - t0
print(f"Top-{top_k} retrieval over {n_sites} sites and {sum(cohort_sizes)} patients in {elapsed:.1f} s")
print(f"Query phenotype label: {int(public_query['label'][0])}")
print(top_result.to_string(index=False))Top-5 retrieval over 3 sites and 240 patients in 17.5 s
Query phenotype label: 3
site_id local_index label score
1 5 3 0.902989
1 64 3 0.901950
1 72 3 0.898401
1 13 3 0.895986
2 42 3 0.895100
Reference-truth comparison
The protocol output is meaningful only if the inner products computed under encryption agree with the same quantities computed in the clear, to within CKKS approximation error. The reference computation runs the same \(A_k \cdot q\) adapter application and inner-product reduction unencrypted on each site’s database:
def plaintext_top_k(q, site_data, A_list, top_k):
rows = []
for k, s_k in enumerate(site_data):
Aq = A_list[k] @ q
scores = s_k["z"] @ Aq
rows += [dict(site_id=k + 1, local_index=i, label=int(s_k["label"][i]),
score=float(scores[i]))
for i in range(len(scores))]
scored = pd.DataFrame(rows).sort_values("score", ascending=False)
return scored.head(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 = top_result.merge(plain_top, on=["site_id", "local_index"],
suffixes=("_enc", "_plain"))
score_err = float(np.max(np.abs(compare["score_enc"] - compare["score_plain"])))
print(f"Top-{top_k} encrypted vs cleartext score max error: {score_err:.2e}")
# Whether the encrypted-domain top-k contains the same
# (site_id, local_index) pairs as the cleartext reference.
enc_set = set(zip(top_result["site_id"], top_result["local_index"]))
plain_set = set(zip(plain_top["site_id"], plain_top["local_index"]))
set_match = len(enc_set & plain_set)
print(f"Top-{top_k} set match: {set_match} of {top_k}")Top-5 encrypted vs cleartext score max error: 5.61e-13
Top-5 set match: 5 of 5
The same run, in one call
run() runs the same protocol on the package’s built-in cohort of three small sites with an exactly orthogonal drift:
r = sim.run()
r.n_candidates, r.max_abs_score_error, r.rank_agreement(18, 1.0571266084724584e-12, True)
What the adapter buys: fidelity, posture, and a sweet spot
The encrypted mechanics work to floating-point precision, as the smoke tests and reference comparison above confirm (max error \(\sim 10^{-12}\)). So we now investigate three questions:
- Fidelity. When the drift is genuinely non-isometric, does a near-orthogonal or least-squares adapter recover retrieval that the rigid orthogonal Procrustes endpoint (\(\mu\to\infty\)) cannot?
- Posture. Where do the two deployments agree — Design 1 (fold the adapter into the database offline) and Design 2 (apply it to the encrypted query, the matvec) — and where must we prefer one? They coincide only when the adapter is orthogonal, so the per-vector norms it produces are all one.
- Sweet spot. Does an intermediate \(\mu\) ever beat both endpoints, and when?
The sweeps use a less-forgiving cluster configuration than the walk-through (unit-norm phenotype centers, so noise competes with separation), letting alignment quality rather than trivial separability drive recall.
The sweeps are recorded once by docs/_recorded/record_similarity_sweeps.py, which calls the three functions below; the tables and figure read the JSON it writes. Federated recall of a query population under one deployment:
print(inspect.getsource(sim.fed_recall))def fed_recall(query_pop, databases, adapters, design: int, top_k: int) -> float:
"""Mean recall@k of a query population across the federation.
``design=1`` folds the adapter into the database offline and scores
``<q, unit(z A)>``; ``design=2`` applies it to the query and scores
``<z, A q>``. They coincide exactly when ``A`` is orthogonal, and
part company otherwise -- which is the point of measuring both.
"""
hits = []
for qi in range(query_pop["z"].shape[0]):
q = query_pop["z"][qi]
labels, scores = [], []
for db, a in zip(databases, adapters, strict=True):
if design == 1:
sc = _unit_rows(db["z"] @ a) @ q
else:
sc = db["z"] @ (a @ q)
labels.append(db["label"])
scores.append(sc)
labels = np.concatenate(labels)
scores = np.concatenate(scores)
top = np.argsort(-scores)[:top_k]
hits.append(float(np.sum(labels[top] == query_pop["label"][qi]) / top_k))
return float(np.mean(hits))
mu_main = sim.mu_sweep(beta=0.6, n_anchor=100) # ample calibration
mu_scarce = sim.mu_sweep(beta=0.6, n_anchor=24) # anchor < p = 32
beta_tab = sim.beta_sweep(n_anchor=100) # Procrustes vs near-orthogonal vs LSimport json
from pathlib import Path
rec = json.loads(Path("_recorded/similarity_sweeps.json").read_text())
ample, scarce = rec["mu_ample"], rec["mu_scarce"]
beta_tab = pd.DataFrame(rec["beta_sweep"])
print(f"mu sweep (beta=0.6, anchor=100): ideal={ample['ideal']:.3f} "
f"unaligned={ample['unaligned']:.3f} Gram={ample['gram']:.3f}")
print(pd.DataFrame({"mu": ample["mu_labels"], "d1": ample["d1"],
"d2": ample["d2"], "aniso": ample["aniso"]}).round(3).to_string(index=False))
print("\nbeta sweep (anchor=100, common random numbers):")
print(beta_tab[["beta", "least_squares", "near_orthogonal", "procrustes", "ideal"]]
.round(3).to_string(index=False))
print("\nmu sweep at scarce anchor=24:")
print(pd.DataFrame({"mu": scarce["mu_labels"], "d1": scarce["d1"],
"d2": scarce["d2"]}).round(3).to_string(index=False))mu sweep (beta=0.6, anchor=100): ideal=0.482 unaligned=0.163 Gram=0.172
mu d1 d2 aniso
0 0.485 0.457 46.654
0.1 0.497 0.470 4.586
1 0.493 0.472 1.613
10 0.472 0.468 0.494
inf 0.463 0.463 0.000
beta sweep (anchor=100, common random numbers):
beta least_squares near_orthogonal procrustes ideal
0.0 0.475 0.475 0.475 0.475
0.3 0.478 0.490 0.487 0.482
0.6 0.485 0.493 0.463 0.482
1.0 0.480 0.480 0.433 0.482
mu sweep at scarce anchor=24:
mu d1 d2
0 0.435 0.418
0.1 0.443 0.442
1 0.415 0.425
10 0.417 0.417
inf 0.398 0.398
Three readings come out of the sweeps. Fidelity (β panel): at \(\beta = 0\) the drift is isometric and Procrustes, the near-orthogonal map, and least squares all match the no-drift ideal; as the drift bends, Procrustes falls away while the relaxed adapters track the ideal — the richer map earns its keep exactly when fine-tuning is non-isometric. Posture (μ panel): Design 1 and Design 2 coincide at \(\mu \to \infty\) (where the adapter is orthogonal and the per-vector norms are all one) and separate as \(\mu\) shrinks, so the raw-database matvec is a valid cosine only near the orthogonal endpoint; the departure-from- isometry panel quantifies the budget. Sweet spot (scarce- anchor panel): when the calibration cohort is smaller than the embedding dimension, an intermediate \(\mu\) regularizes the adapter and can beat both endpoints — a secondary effect that appears only under scarce calibration. The Gram-PSD reference sits near the unaligned floor: its rotation invariance loses the public-frame orientation, so it is the wrong tool for a public-query retrieval.
import matplotlib.pyplot as plt
x = np.arange(len(rec["mu_grid"]))
k = rec["config"]["top_k"]
fig, ax = plt.subplots(2, 2, figsize=(9, 7), constrained_layout=True)
# (1) recall vs mu: Design 1 vs Design 2, with Gram / ideal / floor
a = ax[0, 0]
a.plot(x, ample["d1"], "o-", color="black", label="Design 1 (fold)")
a.plot(x, ample["d2"], "o--", color="firebrick", mfc="none",
label="Design 2 (matvec)")
a.axhline(ample["gram"], ls="-.", lw=2, color="orange", label="Gram-PSD")
a.axhline(ample["ideal"], ls=":", color="darkgreen", label="ideal")
a.axhline(ample["unaligned"], ls=":", color="gray", label="unaligned")
a.set(xticks=x, xticklabels=rec["mu_grid"], ylim=(0, 1),
xlabel="μ (0 = LS -> inf = Procrustes)",
ylabel=f"recall@{k}", title=f"Recall vs μ (β = {ample['beta']})")
a.legend(loc="right", fontsize=7, frameon=False)
# (2) recall vs beta: Procrustes vs near-orthogonal vs LS
a = ax[0, 1]
a.plot(beta_tab["beta"], beta_tab["procrustes"], "o-", color="black",
label="Procrustes (μ = inf)")
a.plot(beta_tab["beta"], beta_tab["near_orthogonal"], "o--", color="blue",
mfc="none", label="near-orth (μ = 1)")
a.plot(beta_tab["beta"], beta_tab["least_squares"], "^:", color="firebrick",
mfc="none", label="LS (μ = 0)")
a.plot(beta_tab["beta"], beta_tab["ideal"], ":", color="darkgreen",
label="ideal")
a.set(ylim=(0, 1), xlabel="β (non-isometry)", ylabel=f"recall@{k}",
title="Procrustes vs relaxed adapters")
a.legend(loc="lower left", fontsize=7, frameon=False)
# (3) recall vs mu at scarce anchor (< p): the regularization sweet spot
a = ax[1, 0]
a.plot(x, scarce["d1"], "o-", color="black", label="Design 1 (fold)")
a.plot(x, scarce["d2"], "o--", color="firebrick", mfc="none",
label="Design 2 (matvec)")
a.set(xticks=x, xticklabels=rec["mu_grid"], ylim=(0, 1), xlabel="μ",
ylabel=f"recall@{k} (Design 1)",
title=f"Scarce anchor (n = {scarce['n_anchor']} < p = {rec['config']['p']})")
a.legend(loc="lower left", fontsize=7, frameon=False)
# (4) departure from isometry vs mu (the Design-2 / matvec budget)
a = ax[1, 1]
a.semilogy(x, np.asarray(ample["aniso"]) + 1e-12, "o-", color="black")
a.set(xticks=x, xticklabels=rec["mu_grid"], xlabel="μ",
ylabel=r"$\|A^\top A - I\|_F$", title="Departure from isometry")
plt.show()
Discussion
- Federated retrieval across heterogeneous fine-tuned models. Each of \(n\) sites holds its own privately fine-tuned variant of a public foundation model; the protocol returns the top-\(k\) patients across all sites whose private-model embeddings are most cosine-similar to a public-model query. The master sees every candidate’s score.
- Site-private compatibility adapter on the \(\mu\) axis. Each site fits an adapter \(A_k\) on a public anchor cohort along the family \(\min_A \lVert Z_{\text{priv}}A - Z_{\text{pub}}\rVert^2 + \mu\lVert A^\top A - I\rVert^2\) — orthogonal Procrustes at \(\mu\to\infty\), least squares at \(\mu = 0\). The adapter never leaves the site and never appears in the clear at the master.
- Two deployments, one axis. The adapter is either applied to the encrypted query (Design 2, the matvec) or folded into the database offline (Design 1, a plain encrypted inner product). The two coincide exactly at the orthogonal endpoint (where every per-vector norm is one) and part company as \(\mu\) relaxes; Design 1 then gives correct cosines while Design 2 does not. The matvec is thus the large-\(\mu\) branch.
- Threshold key generation, \(n\)-of-\(n\). The CKKS secret key is split across all sites. The master holds no usable secret material. Any subset short of the full \(n\) cannot decrypt anything along the way.
- Diagonal-encoded matvec under threshold CKKS. A \(p \times p\) matrix–vector multiply on an encrypted vector, implemented as \(p-1\) nontrivial rotations (the \(i = 0\) diagonal needs none), \(p\) multiplications by unencrypted vectors, and \(p-1\) encrypted additions, all within one level of the precision budget. The joint rotation keys for the cyclic-slot rotations come from an \(n\)-of-\(n\) ceremony that mirrors the encryption-key ceremony.
- Inner-product reduction via \(\log p\) rotation-and-add. Slot-wise multiply by the site’s own unencrypted database vector followed by a halve-and-fold reduction places the cosine similarity in slot 0.
- Encrypted scoring, cleartext ranking. Each site’s scores are encrypted and only the \(n\)-of-\(n\) ceremony can open them, so no single party — the aggregator included — learns a score on its own. Ranking, though, happens after decryption: the aggregator threshold-decrypts every candidate score and sorts, so what it ends up holding is each candidate’s site, local index, label, and score, not a top-\(k\) release. See the Limitations below, which is where encrypted top-\(k\) would change this. The encrypted protocol’s scores agree with the same computation run unencrypted to \(\sim 10^{-13}\).
- Fidelity under non-isometric drift. When fine-tuning is a genuine non-isometry, the rigid orthogonal endpoint leaves recall on the table while a relaxed (least-squares or near-orthogonal) adapter tracks the no-drift ideal. At the isometric special case the two coincide — that special case is the orthogonal protocol of the earlier design.
Limitations
- Encrypted top-\(k\) selection. We sort in the clear after threshold-decrypting the per-patient scores, so the ceremony runs once per candidate and the aggregator ends up holding, for every candidate at every site: the site it came from, the local index, the label, and the score. That is more than a top-\(k\) release, and it is the honest description of what this implementation discloses. Ranking under encryption cannot be avoided by ordering the steps differently — without encrypted comparison, nothing can know the global top-\(k\) without first learning enough scores to rank them. Encrypted argmax / top-\(k\) via polynomial sign approximation is feasible in CKKS but adds depth and complexity orthogonal to this page’s pedagogical aim.
- Normalization, and what the \(\mu\)-knob trades. The near-isometry of the adapter controls the per-vector norms \(\lVert A_k^\top v\rVert\). At the orthogonal end they are all one, which buys two things at once: the homomorphic inverse-square-root (the depth-dominating Newton iteration of Qu and Xu (2023) and Prantl et al. (2024)) is avoided, and every per-site score is a cosine in \([-1,1]\) on the same scale across sites, so the raw-database matvec (Design 2) is directly comparable. As \(\mu\) relaxes for better fidelity, those norms spread; the cure is to fold the adapter into the database offline and unit-normalize there (Design 1), which restores comparable cosines at the cost of materializing public-compatible embeddings and re-folding them on a public-model upgrade. So \(\mu\) trades fidelity against the raw-database-matvec posture, with offline folding as the release valve — not a single forced choice.
- Slot-tiling for production scale. The page runs at \(p = 32\) with one encrypted value per database vector. Production at \(p = 512\) would pack many database vectors into the slots of a single encrypted value, amortizing the matvec and inner-product cost across patients within a site.
- Real-model validation. The synthetic data above drifts by a parameterized non-isometric \(B_k = Q_k D_k\), a controllable stand-in for free fine-tuning. Real fine-tuned foundation models (BiomedCLIP fine-tuned per site, PubMedBERT, etc.) drift in ways no parameterized family captures exactly. Validating the alignment on a real fine-tuned model is future work.
- Adaptive-query leakage. Repeated adaptive queries by the querier (or a coalition with the querier) leak structural information about the cohort — the same Hyrum-style observation that applies to any retrieval system. We do not eliminate this leakage; we acknowledge it as the cost of any released-function output.
- Malicious-secure threshold protocol. The trust model here is honest-but-curious. A malicious-secure variant would require zero-knowledge proofs of correct partial decryption and verifiable computation on the encrypted scores; that is heavy machinery and out of scope.