Dictionary of Applied Machine Learning · membership inference attack

membership inference attack — Python demo

Numerical companion to the entry membership inference attack: it recomputes what the entry states and prints one line per check

Shows the signal the attack uses and what bounds it: a hypothesis that memorizes its training set incurs no loss on its members and a visible loss on data points it never saw, so a loss threshold decides membership; and the attack's accuracy rises and falls with the gap between training error and validation error, so a hypothesis that generalizes leaks little. Self-contained (numpy and matplotlib only), fixed seed.

Run it without installing anything:
uv run https://dictionaryofml.org/terms/membershipinferenceattack.py
uv downloads this script and the pinned NumPy and Matplotlib it needs, then runs it; the script fetches any input file it uses. To keep the output files, download membershipinferenceattack.py into a folder and run uv run membershipinferenceattack.py there. With NumPy and Matplotlib already installed, python3 membershipinferenceattack.py, from any directory — it writes its output files into the current directory. Fixed seeds, so the printed numbers reproduce exactly. Download membershipinferenceattack.py · Notebook · Open in Colab

The script, block by block

One cell per block of the script: the code, and what that code printed when it last ran here

setup

"""
membershipinferenceattack.py — numerical companion to the glossary entry
'membership inference attack'.

Purpose
-------
Shows the signal the attack uses and what bounds it: a hypothesis that
memorizes its training set incurs no loss on its members and a visible
loss on data points it never saw, so a loss threshold decides
membership; and the attack's accuracy rises and falls with the gap
between training error and validation error, so a hypothesis that
generalizes leaks little.  Self-contained (numpy and matplotlib only),
fixed seed.

Setup
-----
Records of 60 patients form the training set: a feature vector of one
measurement x, the 60 values of an evenly spaced grid on [0, 1], and a
label y = sin(6x) plus noise.  Another 60 records with measurements
drawn at random from [0, 1] are not in the training set.  ERM with the
loss (y - prediction)^2 fits hypotheses of flexibility d, the weighted
sums of the first d cosine functions cos(j x pi), for d = 2, ..., 60;
with d = 60 the hypothesis passes through every training record.  The
adversary queries the published hypothesis at a record's feature vector
and declares the record a member when the loss of the prediction lies
below a threshold, chosen so that half of the 120 records lie below it.

Blocks
------
[B-memorize] With d = 60 the training error is zero while the average
             loss on the 60 unseen records is well above it: every member
             is reproduced exactly, every non-member is missed.
[B-attack]   For d = 2, ..., 60 the attack's accuracy, the fraction of
             the 120 records whose membership it decides correctly, grows
             with the gap between validation error and training error:
             near one half where the gap is near zero, above 0.9 where
             the hypothesis memorizes.

Outputs
-------
membershipinferenceattack_losses.csv : loss of every record under the
                                       d = 60 hypothesis, with a member
                                       flag.
membershipinferenceattack_gap.csv    : d, training error, validation
                                       error, their gap, attack accuracy.
membershipinferenceattack.png        : matplotlib preview (checking only).
"""

import numpy as np
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt

from pathlib import Path

OUT_DIR = Path(__file__).parent

report = []


def check(name, ok):
    report.append((name, bool(ok)))
    print(f"  [{'ok' if ok else 'FAIL'}] {name}")


rng = np.random.default_rng(0)
m = 60
x_in = np.linspace(0.0, 1.0, m)                    # members, on a grid
x_out = np.sort(rng.uniform(0.0, 1.0, m))          # non-members


def labels(xs):
    return np.sin(6.0 * xs) + 0.3 * rng.standard_normal(len(xs))


y_in, y_out = labels(x_in), labels(x_out)


def design(xs, d):
    return np.cos(np.outer(xs, np.arange(d)) * np.pi)


def erm(xs, ys, d):
    return np.linalg.lstsq(design(xs, d), ys, rcond=None)[0]


def losses(xs, ys, w):
    return (ys - design(xs, len(w)) @ w) ** 2


def attack_accuracy(w):
    """Member iff loss below the threshold that half of all 120 records lie below."""
    l_in, l_out = losses(x_in, y_in, w), losses(x_out, y_out, w)
    tau = np.median(np.r_[l_in, l_out])
    correct = np.sum(l_in < tau) + np.sum(l_out >= tau)
    return float(correct / (2 * m)), l_in, l_out

B-memorize

With d = 60 the training error is zero while the average loss on the 60 unseen records is well above it: every member is reproduced exactly, every non-member is missed.

w_full = erm(x_in, y_in, m)
acc_full, l_in_full, l_out_full = attack_accuracy(w_full)
tr_full, va_full = float(l_in_full.mean()), float(l_out_full.mean())
check(f"[B-memorize] d = {m}: training error {tr_full:.1e}, average loss "
      f"on the unseen records {va_full:.2f}",
      tr_full < 1e-8 and va_full > 0.1)
  [ok] [B-memorize] d = 60: training error 1.2e-29, average loss on the unseen records 0.22

B-attack

For d = 2, ..., 60 the attack's accuracy, the fraction of the 120 records whose membership it decides correctly, grows with the gap between validation error and training error: near one half where the gap is near zero, above 0.9 where the hypothesis memorizes.

ds = list(range(2, m + 1, 2))
rows = []
for d in ds:
    w = erm(x_in, y_in, d)
    acc, l_in, l_out = attack_accuracy(w)
    tr, va = float(l_in.mean()), float(l_out.mean())
    rows.append((d, tr, va, va - tr, acc))
gaps = np.array([r[3] for r in rows]); accs = np.array([r[4] for r in rows])
small = [r for r in rows if abs(r[3]) < 0.02]   # gap near zero
agree = float(np.corrcoef(np.argsort(np.argsort(gaps)),
                          np.argsort(np.argsort(accs)))[0, 1])
check(f"[B-attack]   attack accuracy rises with the gap (rank agreement "
      f"{agree:.2f}); d = 2: gap {rows[0][3]:.3f}, accuracy "
      f"{rows[0][4]:.2f}; d = {m}: gap {rows[-1][3]:.2f}, accuracy "
      f"{rows[-1][4]:.2f}",
      agree > 0.8 and rows[-1][4] > 0.9
      and all(abs(r[4] - 0.5) < 0.1 for r in small))

# ---------------------------------------------------------------- CSV
with open(OUT_DIR / "membershipinferenceattack_losses.csv", "w") as fh:
    fh.write("x,loss,member\n")
    for a, b in zip(x_in, l_in_full):
        fh.write(f"{a:.4f},{b:.6f},1\n")
    for a, b in zip(x_out, l_out_full):
        fh.write(f"{a:.4f},{b:.6f},0\n")
with open(OUT_DIR / "membershipinferenceattack_gap.csv", "w") as fh:
    fh.write("d,trainerr,valerr,gap,accuracy\n")
    for d, tr, va, gap, acc in rows:
        fh.write(f"{d},{tr:.5f},{va:.5f},{gap:.5f},{acc:.4f}\n")

# -------------------------------------------------------------- preview
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(9.4, 3.9))
ax1.plot(x_in, l_in_full, "ko", ms=4, label="members (training set)")
ax1.plot(x_out, l_out_full, "ks", mfc="none", ms=4, label="non-members")
ax1.axhline(np.median(np.r_[l_in_full, l_out_full]), color="black", ls="--",
            label="threshold")
ax1.set_yscale("symlog", linthresh=1e-8)
ax1.set_xlabel("feature $x$")
ax1.set_ylabel("loss of the prediction")
ax1.set_title(f"losses under the memorizing hypothesis ($d = {m}$)")
ax1.legend(frameon=False, fontsize=8)
ax2.plot(ds, [r[3] for r in rows], "ko-", label="validation error minus training error")
ax2.plot(ds, [r[4] for r in rows], "ks--", mfc="none", label="attack accuracy")
ax2.set_xlabel("flexibility $d$ of the hypothesis")
ax2.set_ylabel("value")
ax2.set_title("the gap bounds what the attack learns")
ax2.legend(frameon=False, fontsize=8)
fig.tight_layout()
fig.savefig(OUT_DIR / "membershipinferenceattack.png", dpi=110)

n_ok = sum(ok for _, ok in report)
print(f"\n{n_ok}/{len(report)} checks pass")
print(f"wrote {OUT_DIR / 'membershipinferenceattack_losses.csv'}, "
      f"{OUT_DIR / 'membershipinferenceattack_gap.csv'}, "
      f"{OUT_DIR / 'membershipinferenceattack.png'}")
if n_ok != len(report):
    raise SystemExit(1)
  [ok] [B-attack]   attack accuracy rises with the gap (rank agreement 0.97); d = 2: gap -0.004, accuracy 0.52; d = 60: gap 0.22, accuracy 1.00

2/2 checks pass
wrote /Users/junga1/dictionaryappliedml/pythondemos/membershipinferenceattack_losses.csv, /Users/junga1/dictionaryappliedml/pythondemos/membershipinferenceattack_gap.csv, /Users/junga1/dictionaryappliedml/pythondemos/membershipinferenceattack.png
Preview figure produced by membershipinferenceattack.py
The preview figure the block B-attack writes when the script runs