Threshold Cox with Gaussian Noise (Demonstration)

Introduction

This page is a demonstration. The previous page, the Cox threshold page, fits the Cox model under threshold FHE: the fit matches the cleartext fit, no single party can decrypt, and the aggregator sees only the joint log-likelihood \(\ell(\beta)\) at each optimizer query.

Here we ask what happens if each site also adds Gaussian noise of the kind an output differential privacy mechanism uses. The sensitivity \(\Delta = 1\) below is a placeholder, so the \(\varepsilon\) values are not a privacy guarantee. Adding the noise takes one extra rng.normal() call per site per query. With the optimizers used here, the fits are poor at any noise level that gives a small \(\varepsilon\). The page runs the fits and reports the numbers.

A brief output-DP primer

Given a query \(f: \mathcal{D} \to \mathbb{R}\) on a dataset \(D\), the Gaussian mechanism (Dwork and Roth 2014, sec. 3.5.3) releases \(\tilde f(D) = f(D) + \mathcal{N}(0, \sigma^2)\). If \(f\)’s sensitivity (the largest change \(f\) can undergo when one record is added or removed) is \(\Delta\), then for any \(\varepsilon \in (0, 1)\) and \(\delta > 0\), choosing \(\sigma \geq \Delta \sqrt{2\ln(1.25/\delta)}/\varepsilon\) makes a single release of \(\tilde f\) satisfy \((\varepsilon, \delta)\) differential privacy.

The optimizer issues many queries, so the per-query budget composes. We use zCDP composition (Bun and Steinke 2016): the Gaussian mechanism is \(\rho\)-zCDP with \(\rho = (\Delta/\sigma)^2/2\); \(T\) queries compose linearly to \(T \cdot \rho\) in zCDP units; and that converts to \((\varepsilon, \delta)\) via \(\varepsilon = \rho + 2\sqrt{\rho \log(1/\delta)}\).

The protocol

The setup is identical to the Cox threshold page, except each site adds an independent Gaussian noise term to its local contribution before encryption. Since variances of independent normals add, the decrypted total is \(\ell(\beta) + \mathcal{N}(0, \sigma^2)\) when each site adds \(\mathcal{N}(0, \sigma^2/N)\).

The Cox setup (same DLBCL data as the other Cox pages)

import numpy as np
from homomorphepy import load_dlbcl, site_order
from statsmodels.duration.hazard_regression import PHReg

df = load_dlbcl()
COVARIATES = ["GCB_sig", "LN_sig", "Prolif_sig", "BMP6", "MHC2_sig"]

def site_frame(name):
    m = (df["Subgroup"] == name).to_numpy()
    return {
        "time":   df["time"].to_numpy(dtype=float)[m],
        "status": df["status"].to_numpy(dtype=int)[m],
        "X":      np.column_stack([df[c].to_numpy(dtype=float)[m]
                                   for c in COVARIATES]),
    }

cox_data = {name: site_frame(name) for name in site_order()}

agg_model = PHReg(df["time"], df[COVARIATES], status=df["status"],
                  strata=df["Subgroup"].cat.codes, ties="efron")
agg_coef  = np.asarray(agg_model.fit().params, dtype=float)

def local_cox_nll(data, beta):
    b = np.asarray(beta, dtype=float).ravel()
    model = PHReg(data["time"], data["X"],
                  status=data["status"], ties="efron")
    return float(-model.loglike(b))

Threshold setup and DP-noised workers

The threshold setup is also the same as on the Cox threshold page. The only change is in each worker’s local function: it adds an independent \(\mathcal{N}(0, \sigma^2/N)\) draw to its local nLL before returning it. The noisy value is then encrypted as usual.

import inspect
from homomorphepy.examples import dp

print(inspect.getsource(dp.fit_at_sigma))
def fit_at_sigma(
    sigma: float,
    method: str = "Nelder-Mead",
    seed: int = 1,
    delta: float = DEFAULT_DELTA,
) -> DPFit:
    """Fit the threshold-Cox model with per-site Gaussian noise.

    Each site adds ``N(0, sigma^2/N)`` to its local negative
    log-likelihood before encryption, so the decrypted aggregate
    carries ``N(0, sigma^2)``.
    """
    sites = _cox_sites()
    names = list(sites)
    n_sites = len(names)
    rng = np.random.default_rng(seed)
    scale = sigma / math.sqrt(n_sites)

    def noisy_local(data, beta):
        value = local_cox_nll(data, beta)
        if value is None or math.isnan(value):
            return math.nan
        return value + (rng.normal(0.0, scale) if sigma > 0 else 0.0)

    ctx = fhe_context("CKKS", **CKKS_PARAMS)
    workers = [make_worker(n, sites[n], noisy_local) for n in names]
    master = make_threshold_master("Aggregator", ctx, workers)

    calls = {"n": 0}

    def objective(beta):
        calls["n"] += 1
        v = master.aggregate(np.asarray(beta, dtype=float))
        return 1e12 if (v is None or math.isnan(v)) else float(v)

    # The finite-difference step is LOAD-BEARING here, unlike in
    # examples/cox.py where sweeping it changed nothing. Without DP
    # noise the encrypted objective is accurate to ~1e-13 and any step
    # works. Under DP the gradient noise is the function noise divided
    # by h, so scipy's default (~1.5e-8) amplifies sigma by ~7e7 and
    # destroys the gradient even at sigma = 1e-4 -- the optimizer then
    # returns x0 unchanged and still reports success. At h = 1e-3 the
    # amplification is ~707x, which is survivable; the constant is
    # load-bearing rather than decorative here.
    x0 = np.zeros(len(COVARIATES))
    fun = objective
    if method == "BFGS":
        options = {"gtol": 1e-4, "finite_diff_rel_step": FINITE_DIFF_STEP}
    elif method == "L-BFGS-B":
        options = {"ftol": 1e-9, "finite_diff_rel_step": FINITE_DIFF_STEP}
    elif method == "Nelder-Mead":
        # Starting simplex: NM_SIMPLEX_STEP added to one coordinate per
        # vertex. Stop when the spread of objective values over the
        # simplex is at most NM_RELTOL * (|f(x0)| + NM_RELTOL), or after
        # NM_MAXFEV evaluations. f(x0) is evaluated once and reused as
        # the first simplex vertex, so it costs one query, not two.
        f0 = objective(x0)
        first = {"pending": True}

        def fun(beta):
            if first["pending"] and np.array_equal(beta, x0):
                first["pending"] = False
                return f0
            return objective(beta)

        options = {
            "initial_simplex": np.vstack([x0, x0 + NM_SIMPLEX_STEP * np.eye(len(x0))]),
            "xatol": np.inf,
            "fatol": NM_RELTOL * (abs(f0) + NM_RELTOL),
            "maxfev": NM_MAXFEV,
        }
    else:
        raise ValueError(f"unsupported method {method!r}")
    fit = minimize(fun, x0=x0, method=method, options=options)
    beta_hat = np.asarray(fit.x, dtype=float)

    # Centralized cleartext fit: what DP is degrading away from.
    from statsmodels.duration.hazard_regression import PHReg

    from homomorphepy.fixtures import load_dlbcl

    df = load_dlbcl()
    ref = PHReg(
        df["time"].to_numpy(dtype=float),
        np.column_stack([df[c].to_numpy(dtype=float) for c in COVARIATES]),
        status=df["status"].to_numpy(dtype=int),
        strata=df["Subgroup"].cat.codes.to_numpy(),
        ties="efron",
    ).fit()
    beta_ref = np.asarray(ref.params, dtype=float)

    return DPFit(
        sigma=sigma,
        method=method,
        coefficients=dict(zip(COVARIATES, beta_hat.tolist(), strict=True)),
        centralized=dict(zip(COVARIATES, beta_ref.tolist(), strict=True)),
        max_abs_diff=float(np.max(np.abs(beta_hat - beta_ref))),
        # How many coefficients still land on the correct side of zero:
        # the qualitative conclusion, which survives longer than the
        # point estimates do.
        sign_agreement=int(np.sum(np.sign(beta_hat) == np.sign(beta_ref))),
        n_queries=calls["n"],
        converged=bool(fit.success),
        epsilon=budget(calls["n"], sigma, delta)["epsilon"],
        context=ctx,
    )

Each fit is a full optimizer run through the encrypted channel, so the sweep is recorded once rather than run on every render:

OMP_NUM_THREADS=2 uv run python docs/_recorded/record_dp_sweep.py

The recorded sweep took 1.6 minutes.

Mechanical correctness: \(\sigma = 0\) reproduces the threshold fit

When the noise is zero the protocol reduces to the lossless threshold protocol. The fitted coefficients match the centralized fit to threshold-CKKS precision.

Threshold-DP protocol at \(\sigma = 0\) (BFGS) vs the centralized cleartext fit
Coefficient \(\hat\beta\), centralized \(\hat\beta\), protocol at \(\sigma = 0\) \(\lvert \text{difference} \rvert\)
GCB_sig -0.2638716 -0.2638715 \(1.167 \times 10^{-7}\)
LN_sig -0.2543592 -0.2543592 \(2.349 \times 10^{-8}\)
Prolif_sig 0.3031258 0.3031262 \(3.845 \times 10^{-7}\)
BMP6 0.3036375 0.3036377 \(1.353 \times 10^{-7}\)
MHC2_sig -0.3191467 -0.3191467 \(7.417 \times 10^{-9}\)

The maximum absolute coefficient difference at \(\sigma = 0\) is \(3.84 \times 10^{-7}\). So any deviation from the centralized fit below comes from the noise.

Increasing noise: where does it break?

We sweep \(\sigma\) across five orders of magnitude. Each fit is one full BFGS run over the threshold-DP encrypted nLL. BFGS estimates the gradient by finite differences of these nLL values.

BFGS over the threshold-DP nLL at 6 values of \(\sigma\)
\(\sigma\) Queries GCB_sig LN_sig Prolif_sig BMP6 MHC2_sig \(\max_j \lvert \hat\beta_j - \hat\beta_j^{\text{centralized}} \rvert\)
\(10^{-5}\) 166 0.000000 0.000000 0.000000 0.000001 0.000000 0.319147
\(10^{-4}\) 144 0.000001 0.000001 0.000000 0.000002 0.000001 0.319147
\(10^{-3}\) 199 0.000012 0.000009 0.000005 0.000024 0.000010 0.319156
\(10^{-2}\) 199 0.000012 0.000009 0.000005 0.000025 0.000010 0.319156
\(10^{-1}\) 199 0.000012 0.000009 0.000005 0.000025 0.000010 0.319156
\(1\) 199 0.000012 0.000009 0.000005 0.000024 0.000010 0.319156

BFGS does not move from the zero starting point at any \(\sigma > 0\); the rows show the starting values.

Nelder–Mead at the same noise scales

The same sweep with Nelder–Mead. The starting simplex adds 0.1 to one coordinate per vertex. The search stops when the spread of objective values over the simplex is at most \(1 \times 10^{-7}\) relative to the starting value, or after 500 evaluations.

Nelder–Mead over the threshold-DP nLL at 6 values of \(\sigma\)
\(\sigma\) Queries GCB_sig LN_sig Prolif_sig BMP6 MHC2_sig \(\max_j \lvert \hat\beta_j - \hat\beta_j^{\text{centralized}} \rvert\)
\(10^{-5}\) 147 -0.264116 -0.253984 0.303746 0.304050 -0.319394 0.000620
\(10^{-4}\) 500 -0.264277 -0.253462 0.302743 0.302900 -0.319406 0.000897
\(10^{-3}\) 500 -0.261806 -0.255025 0.315837 0.301430 -0.319236 0.012712
\(10^{-2}\) 500 -0.260086 -0.243787 0.318555 0.309170 -0.314670 0.015429
\(10^{-1}\) 500 -0.254677 -0.232421 0.331196 0.328366 -0.325162 0.028070
\(1\) 500 -0.117900 -0.276456 0.331867 0.339219 -0.301667 0.145972

Fidelity decays as \(\sigma\) grows.

Privacy budget

For the fits at the first three values of \(\sigma\), with sensitivity \(\Delta = 1\) (placeholder) and target \(\delta = 10^{-5}\), zCDP composition gives:

def zcdp_to_eps(rho, delta=1e-5):
    return rho + 2 * np.sqrt(rho * np.log(1 / delta))

first_three = sorted(s for s in rows["sigma"].unique() if s > 0)[:3]
budget = pd.concat([
    rows[(rows["method"] == m) & rows["sigma"].isin(first_three)][["method", "sigma", "n_queries"]]
    for m in ("BFGS", "Nelder-Mead")])
budget["rho_per_query"] = (1 / budget["sigma"]) ** 2 / 2
budget["rho_total"]     = budget["n_queries"] * budget["rho_per_query"]
budget["epsilon_at_delta_1e_minus_5"] = zcdp_to_eps(budget["rho_total"])
zCDP composition; sensitivity \(\Delta = 1\), target \(\delta = 10^{-5}\)
Optimizer \(\sigma\) Queries \(k\) \(\rho\) per query \(\rho_{\text{total}} = k\rho\) \(\varepsilon\) at \(\delta = 10^{-5}\)
BFGS \(10^{-5}\) 166 \(5 \times 10^{9}\) \(8.3 \times 10^{11}\) \(8.3 \times 10^{11}\)
BFGS \(10^{-4}\) 144 \(5 \times 10^{7}\) \(7.2 \times 10^{9}\) \(7.201 \times 10^{9}\)
BFGS \(10^{-3}\) 199 \(5 \times 10^{5}\) \(9.95 \times 10^{7}\) \(9.957 \times 10^{7}\)
Nelder–Mead \(10^{-5}\) 147 \(5 \times 10^{9}\) \(7.35 \times 10^{11}\) \(7.35 \times 10^{11}\)
Nelder–Mead \(10^{-4}\) 500 \(5 \times 10^{7}\) \(2.5 \times 10^{10}\) \(2.5 \times 10^{10}\)
Nelder–Mead \(10^{-3}\) 500 \(5 \times 10^{5}\) \(2.5 \times 10^{8}\) \(2.501 \times 10^{8}\)

The \(\varepsilon\) column shows that at noise levels where the fits are still close, \(\varepsilon\) is large. Whether that is acceptable depends on the application.

One could also explore tighter sensitivity bound \(\Delta\) or allow far fewer queries, etc. We don’t do that here.

References

Bun, Mark, and Thomas Steinke. 2016. “Concentrated Differential Privacy: Simplifications, Extensions, and Lower Bounds.” Theory of Cryptography Conference (TCC), 635–58. https://doi.org/10.1007/978-3-662-53641-4_24.
Dwork, Cynthia, and Aaron Roth. 2014. The Algorithmic Foundations of Differential Privacy. Foundations and Trends in Theoretical Computer Science 9(3–4). Now Publishers. https://doi.org/10.1561/0400000042.