{
 "nbformat": 4,
 "nbformat_minor": 5,
 "metadata": {
  "kernelspec": {
   "name": "python3",
   "display_name": "Python 3",
   "language": "python"
  },
  "language_info": {
   "name": "python"
  },
  "colab": {
   "name": "membershipinferenceattack.ipynb"
  }
 },
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "# membership inference attack \u2014 Python demo\n\nNumerical companion to the entry [membership inference attack](https://dictionaryofml.org/terms/membershipinferenceattack.html) of the [Dictionary of Applied Machine Learning](https://dictionaryofml.org/): it recomputes what the entry states and prints one line per check.\n\nShows 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.\n\nRequires NumPy and Matplotlib only, and uses fixed seeds, so the printed numbers reproduce exactly. Generated from [`pythondemos/membershipinferenceattack.py`](https://dictionaryofml.org/terms/membershipinferenceattack.py); CC BY 4.0."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "# Notebook shim: the script resolves output paths relative to __file__,\n# which a notebook kernel does not define; everything lands in the\n# working directory instead.\nimport os\n__file__ = os.path.join(os.getcwd(), \"membershipinferenceattack.py\")\nos.makedirs(\"pythondemos\", exist_ok=True)"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "\"\"\"\nmembershipinferenceattack.py \u2014 numerical companion to the glossary entry\n'membership inference attack'.\n\nPurpose\n-------\nShows the signal the attack uses and what bounds it: a hypothesis that\nmemorizes its training set incurs no loss on its members and a visible\nloss on data points it never saw, so a loss threshold decides\nmembership; and the attack's accuracy rises and falls with the gap\nbetween training error and validation error, so a hypothesis that\ngeneralizes leaks little.  Self-contained (numpy and matplotlib only),\nfixed seed.\n\nSetup\n-----\nRecords of 60 patients form the training set: a feature vector of one\nmeasurement x, the 60 values of an evenly spaced grid on [0, 1], and a\nlabel y = sin(6x) plus noise.  Another 60 records with measurements\ndrawn at random from [0, 1] are not in the training set.  ERM with the\nloss (y - prediction)^2 fits hypotheses of flexibility d, the weighted\nsums of the first d cosine functions cos(j x pi), for d = 2, ..., 60;\nwith d = 60 the hypothesis passes through every training record.  The\nadversary queries the published hypothesis at a record's feature vector\nand declares the record a member when the loss of the prediction lies\nbelow a threshold, chosen so that half of the 120 records lie below it.\n\nBlocks\n------\n[B-memorize] With d = 60 the training error is zero while the average\n             loss on the 60 unseen records is well above it: every member\n             is reproduced exactly, every non-member is missed.\n[B-attack]   For d = 2, ..., 60 the attack's accuracy, the fraction of\n             the 120 records whose membership it decides correctly, grows\n             with the gap between validation error and training error:\n             near one half where the gap is near zero, above 0.9 where\n             the hypothesis memorizes.\n\nOutputs\n-------\nmembershipinferenceattack_losses.csv : loss of every record under the\n                                       d = 60 hypothesis, with a member\n                                       flag.\nmembershipinferenceattack_gap.csv    : d, training error, validation\n                                       error, their gap, attack accuracy.\nmembershipinferenceattack.png        : matplotlib preview (checking only).\n\"\"\"\n\nimport numpy as np\nimport matplotlib\n\nmatplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\n\nfrom pathlib import Path\n\nOUT_DIR = Path(__file__).parent\n\nreport = []\n\n\ndef check(name, ok):\n    report.append((name, bool(ok)))\n    print(f\"  [{'ok' if ok else 'FAIL'}] {name}\")\n\n\nrng = np.random.default_rng(0)\nm = 60\nx_in = np.linspace(0.0, 1.0, m)                    # members, on a grid\nx_out = np.sort(rng.uniform(0.0, 1.0, m))          # non-members\n\n\ndef labels(xs):\n    return np.sin(6.0 * xs) + 0.3 * rng.standard_normal(len(xs))\n\n\ny_in, y_out = labels(x_in), labels(x_out)\n\n\ndef design(xs, d):\n    return np.cos(np.outer(xs, np.arange(d)) * np.pi)\n\n\ndef erm(xs, ys, d):\n    return np.linalg.lstsq(design(xs, d), ys, rcond=None)[0]\n\n\ndef losses(xs, ys, w):\n    return (ys - design(xs, len(w)) @ w) ** 2\n\n\ndef attack_accuracy(w):\n    \"\"\"Member iff loss below the threshold that half of all 120 records lie below.\"\"\"\n    l_in, l_out = losses(x_in, y_in, w), losses(x_out, y_out, w)\n    tau = np.median(np.r_[l_in, l_out])\n    correct = np.sum(l_in < tau) + np.sum(l_out >= tau)\n    return float(correct / (2 * m)), l_in, l_out"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[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."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "w_full = erm(x_in, y_in, m)\nacc_full, l_in_full, l_out_full = attack_accuracy(w_full)\ntr_full, va_full = float(l_in_full.mean()), float(l_out_full.mean())\ncheck(f\"[B-memorize] d = {m}: training error {tr_full:.1e}, average loss \"\n      f\"on the unseen records {va_full:.2f}\",\n      tr_full < 1e-8 and va_full > 0.1)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[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."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "ds = list(range(2, m + 1, 2))\nrows = []\nfor d in ds:\n    w = erm(x_in, y_in, d)\n    acc, l_in, l_out = attack_accuracy(w)\n    tr, va = float(l_in.mean()), float(l_out.mean())\n    rows.append((d, tr, va, va - tr, acc))\ngaps = np.array([r[3] for r in rows]); accs = np.array([r[4] for r in rows])\nsmall = [r for r in rows if abs(r[3]) < 0.02]   # gap near zero\nagree = float(np.corrcoef(np.argsort(np.argsort(gaps)),\n                          np.argsort(np.argsort(accs)))[0, 1])\ncheck(f\"[B-attack]   attack accuracy rises with the gap (rank agreement \"\n      f\"{agree:.2f}); d = 2: gap {rows[0][3]:.3f}, accuracy \"\n      f\"{rows[0][4]:.2f}; d = {m}: gap {rows[-1][3]:.2f}, accuracy \"\n      f\"{rows[-1][4]:.2f}\",\n      agree > 0.8 and rows[-1][4] > 0.9\n      and all(abs(r[4] - 0.5) < 0.1 for r in small))\n\n# ---------------------------------------------------------------- CSV\nwith open(OUT_DIR / \"membershipinferenceattack_losses.csv\", \"w\") as fh:\n    fh.write(\"x,loss,member\\n\")\n    for a, b in zip(x_in, l_in_full):\n        fh.write(f\"{a:.4f},{b:.6f},1\\n\")\n    for a, b in zip(x_out, l_out_full):\n        fh.write(f\"{a:.4f},{b:.6f},0\\n\")\nwith open(OUT_DIR / \"membershipinferenceattack_gap.csv\", \"w\") as fh:\n    fh.write(\"d,trainerr,valerr,gap,accuracy\\n\")\n    for d, tr, va, gap, acc in rows:\n        fh.write(f\"{d},{tr:.5f},{va:.5f},{gap:.5f},{acc:.4f}\\n\")\n\n# -------------------------------------------------------------- preview\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(9.4, 3.9))\nax1.plot(x_in, l_in_full, \"ko\", ms=4, label=\"members (training set)\")\nax1.plot(x_out, l_out_full, \"ks\", mfc=\"none\", ms=4, label=\"non-members\")\nax1.axhline(np.median(np.r_[l_in_full, l_out_full]), color=\"black\", ls=\"--\",\n            label=\"threshold\")\nax1.set_yscale(\"symlog\", linthresh=1e-8)\nax1.set_xlabel(\"feature $x$\")\nax1.set_ylabel(\"loss of the prediction\")\nax1.set_title(f\"losses under the memorizing hypothesis ($d = {m}$)\")\nax1.legend(frameon=False, fontsize=8)\nax2.plot(ds, [r[3] for r in rows], \"ko-\", label=\"validation error minus training error\")\nax2.plot(ds, [r[4] for r in rows], \"ks--\", mfc=\"none\", label=\"attack accuracy\")\nax2.set_xlabel(\"flexibility $d$ of the hypothesis\")\nax2.set_ylabel(\"value\")\nax2.set_title(\"the gap bounds what the attack learns\")\nax2.legend(frameon=False, fontsize=8)\nfig.tight_layout()\nfig.savefig(OUT_DIR / \"membershipinferenceattack.png\", dpi=110)\n\nn_ok = sum(ok for _, ok in report)\nprint(f\"\\n{n_ok}/{len(report)} checks pass\")\nprint(f\"wrote {OUT_DIR / 'membershipinferenceattack_losses.csv'}, \"\n      f\"{OUT_DIR / 'membershipinferenceattack_gap.csv'}, \"\n      f\"{OUT_DIR / 'membershipinferenceattack.png'}\")\nif n_ok != len(report):\n    raise SystemExit(1)"
  }
 ]
}