{
 "nbformat": 4,
 "nbformat_minor": 5,
 "metadata": {
  "kernelspec": {
   "name": "python3",
   "display_name": "Python 3",
   "language": "python"
  },
  "language_info": {
   "name": "python"
  },
  "colab": {
   "name": "xaiterm.ipynb"
  }
 },
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "# explainable artificial intelligence (XAI) \u2014 Python demo\n\nNumerical companion to the entry [explainable artificial intelligence (XAI)](https://dictionaryofml.org/terms/xaiterm.html) of the [Dictionary of Applied Machine Learning](https://dictionaryofml.org/): it recomputes what the entry states and prints one line per check.\n\nBacks the entry's method claims with exact computations on the very hypothesis drawn in the entry's counterfactual figure, h(x) = 1.2 + 2.2 / (1 + exp(-1.8 (x - 3.5))) with decision threshold 2.6 and data point x0 = 2.2: LIME's local linear approximation, the counterfactual as the smallest prediction-altering change, and the additive (efficiency) property of SHAP. Self-contained (numpy only), fixed seed.\n\nRequires NumPy and Matplotlib only, and uses fixed seeds, so the printed numbers reproduce exactly. Generated from [`pythondemos/xaiterm.py`](https://dictionaryofml.org/terms/xaiterm.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(), \"xaiterm.py\")\nos.makedirs(\"pythondemos\", exist_ok=True)"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "\"\"\"\nxaiterm.py \u2014 numerical companion to the glossary entry 'explainable\nartificial intelligence (XAI)'.\n\nPurpose\n-------\nBacks the entry's method claims with exact computations on the very\nhypothesis drawn in the entry's counterfactual figure, h(x) = 1.2 +\n2.2 / (1 + exp(-1.8 (x - 3.5))) with decision threshold 2.6 and data point\nx0 = 2.2:  LIME's local linear approximation, the counterfactual as\nthe smallest prediction-altering change, and the additive (efficiency)\nproperty of SHAP.  Self-contained (numpy only), fixed seed.\n\nBlocks\n------\n[B-lime]  A proximity-weighted linear fit around x0 recovers the\n          tangent drawn in that figure: the fitted slope matches the\n          analytic derivative h'(x0) to 2%, and the fit approximates\n          h near x0 while erring at least 4x more far away.\n[B-cf]    The counterfactual x' is the SMALLEST change of x0\n          that alters the thresholded prediction: a grid search for\n          the nearest x with h(x) >= 2.6 agrees with the analytic\n          threshold crossing x' = 3.5 - ln(2.2/1.4 - 1)/1.8, and no\n          x closer to x0 crosses the threshold.  The threshold is the\n          one the entry's figure draws; a demo pinned to a stale value\n          would disagree with the picture it claims to compute.\n[B-faithful] Faithfulness has a price. Over the whole feature range the\n          best affine explanation still deviates from h by a wide gap,\n          and the deviation only falls towards zero as the explaining\n          function is allowed to grow until it is h itself. Measured as\n          the agreement rate |h - g| < 0.05 on a test grid, the local\n          surrogate agrees near x0 and disagrees away from it.\n[B-relevance] For an image, the explanation is one relevance score per\n          pixel. A linear classifier is fitted to 6x6 images of a seven\n          against shapes sharing its top bar; the contribution of pixel j\n          to the prediction is w_j x_j. The diagonal stroke carries the\n          largest scores and every unlit pixel scores exactly 0, as drawn\n          in the entry's Fig. 2. The shared top bar comes out small and\n          negative: it is not evidence for a seven, and the fit uses it to\n          bring the sum down to the label 1. Written to xaiterm.png.\n[B-morf]  The faithfulness of that map, tested rather than asserted. Pixels\n          are flipped in order of decreasing relevance (most relevant first,\n          MoRF) and, for comparison, in order of increasing relevance. The\n          MoRF order drives the score across the decision threshold after\n          far fewer flips, and does so for every one of 200 noisy images \u2014\n          which is what the entry means by a class activation map being\n          faithful for a prediction.\n[B-eerm]  Explainability can enter training instead. The user supplies\n          their own predictions for the training set, and a penalty\n          charges the part of the fitted predictions that those user\n          predictions do not already account for. As the penalty weight\n          grows, the fit becomes more predictable from the user's own\n          predictions while the average loss rises.\n[B-shap]  Exact Shapley values for a 3-feature model (computed by\n          enumerating all 8 coalitions, missing features replaced by\n          their baseline values) satisfy the efficiency property: the\n          contributions sum to f(x) - f(baseline) \u2014 SHAP decomposes\n          the prediction into additive feature contributions.  A\n          dummy feature that the model ignores receives contribution\n          exactly 0.\n\nOutputs\n-------\nxaiterm.png : preview (checking only) \u2014 the image, its relevance map, the\n              pixels MoRF flips to change the prediction, and the two\n              perturbation curves.\n\"\"\"\n\nimport itertools\nfrom math import factorial\n\nimport numpy as np\nimport matplotlib\n\nmatplotlib.use(\"Agg\")\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\n\nOUT_DIR = Path(__file__).parent\n\nreport = []                         # collects (check name, pass/fail) pairs\n\n\ndef check(name, ok):                # records and prints one verification\n    report.append((name, bool(ok)))\n    print(f\"  [{'ok' if ok else 'FAIL'}] {name}\")\n\n\n# the learned hypothesis of the entry's Fig. 1\ndef h(x):\n    return 1.2 + 2.2 / (1.0 + np.exp(-1.8 * (x - 3.5)))\n\n\ndef h_prime(x):                     # analytic derivative of h\n    s = 1.0 / (1.0 + np.exp(-1.8 * (x - 3.5)))\n    return 2.2 * 1.8 * s * (1.0 - s)\n\n\nx0 = 2.2                            # the data point explained in Fig. 1\ntau = 2.6                           # decision threshold drawn in the\n                                    # entry's counterfactual figure"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-lime]** A proximity-weighted linear fit around x0 recovers the tangent drawn in that figure: the fitted slope matches the analytic derivative h'(x0) to 2%, and the fit approximates h near x0 while erring at least 4x more far away."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "xs = np.linspace(0.3, 7.0, 400)\nw = np.exp(-((xs - x0) ** 2) / (2 * 0.15 ** 2))      # proximity weights at x0\nA = np.c_[np.ones_like(xs), xs]                      # linear design (1, x)\nsw = np.sqrt(w)\ncoef, *_ = np.linalg.lstsq(A * sw[:, None], h(xs) * sw, rcond=None)\ng = A @ coef                                          # LIME surrogate\nok_slope = abs(coef[1] - h_prime(x0)) < 0.02 * abs(h_prime(x0))\nnear = np.abs(xs - x0) < 0.3\nfar = np.abs(xs - 5.5) < 0.3\nerr_near = np.max(np.abs(h(xs) - g)[near])\nerr_far = np.max(np.abs(h(xs) - g)[far])\ncheck(\"[B-lime]  weighted linear fit recovers the tangent and is local\",\n      ok_slope and err_far > 4.0 * err_near)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-cf]** The counterfactual x' is the SMALLEST change of x0 that alters the thresholded prediction: a grid search for the nearest x with h(x) >= 2.6 agrees with the analytic threshold crossing x' = 3.5 - ln(2.2/1.4 - 1)/1.8, and no x closer to x0 crosses the threshold. The threshold is the one the entry's figure draws; a demo pinned to a stale value would disagree with the picture it claims to compute."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "x_cf = 3.5 - np.log(2.2 / (tau - 1.2) - 1.0) / 1.8   # analytic crossing\ncrossing = xs[h(xs) >= tau]                           # grid: threshold reached\nx_star = crossing[np.argmin(np.abs(crossing - x0))]  # nearest such point\nok_match = abs(x_star - x_cf) < 0.02                  # matches the analytic x'\n# minimality: no point strictly between x0 and x' crosses the threshold\nbetween = xs[(xs > x0) & (xs < x_cf - 0.01)]\nok_min = np.all(h(between) < tau) and h(x0) < tau\ncheck(\"[B-cf]    counterfactual = smallest prediction-altering change\",\n      ok_match and ok_min)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-shap]** Exact Shapley values for a 3-feature model (computed by enumerating all 8 coalitions, missing features replaced by their baseline values) satisfy the efficiency property: the contributions sum to f(x) - f(baseline) \u2014 SHAP decomposes the prediction into additive feature contributions. A dummy feature that the model ignores receives contribution exactly 0."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "def f(z):                           # 3-feature model; feature 2 is a dummy\n    return 2.0 * z[0] - 1.5 * z[1] + 0.8 * z[0] * z[1] + 0.0 * z[2]\n\n\nz = np.array([1.2, -0.7, 0.5])      # the data point to be explained\nbase = np.array([0.0, 0.0, 0.0])    # baseline (reference) feature values\n\n\ndef value(S):                       # coalition value: replace absent features\n    zz = base.copy()                # by their baseline values\n    for j in S:\n        zz[j] = z[j]\n    return f(zz)\n\nd = 3\nphi = np.zeros(d)                   # exact Shapley values by enumeration\nfor j in range(d):\n    others = [k for k in range(d) if k != j]\n    for m in range(len(others) + 1):\n        for S in itertools.combinations(others, m):\n            wgt = factorial(len(S)) * factorial(d - len(S) - 1) / factorial(d)\n            phi[j] += wgt * (value(S + (j,)) - value(S))\n\nok_eff = np.isclose(phi.sum(), f(z) - f(base))        # efficiency property\nok_dummy = np.isclose(phi[2], 0.0)                    # unused feature gets 0\ncheck(\"[B-shap]  Shapley contributions sum to f(x) - f(baseline); dummy \"\n      \"feature gets 0\", ok_eff and ok_dummy)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-faithful]** Faithfulness has a price. Over the whole feature range the best affine explanation still deviates from h by a wide gap, and the deviation only falls towards zero as the explaining function is allowed to grow until it is h itself. Measured as the agreement rate |h - g| < 0.05 on a test grid, the local surrogate agrees near x0 and disagrees away from it."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "# An explanation that agreed with h everywhere would be h. The best affine\n# explanation over the whole range is therefore stuck with a deviation, and\n# only a function allowed to approach h drives that deviation to zero.\ngrid = np.linspace(0.3, 7.0, 400)\ndev = []\nfor degree in (1, 2, 3, 5, 7):\n    c = np.polyfit(grid, h(grid), degree)\n    dev.append(float(np.max(np.abs(h(grid) - np.polyval(c, grid)))))\nprint(f\"    max deviation of the best fit of degree 1,2,3,5,7: \"\n      f\"{', '.join(f'{d:.3f}' for d in dev)}\")\ncheck(\"[B-faithful] the best affine explanation deviates from h somewhere\",\n      dev[0] > 0.2)\ncheck(\"[B-faithful] deviation falls only as the explanation grows\",\n      all(a > b for a, b in zip(dev, dev[1:])))\n# faithfulness as an agreement rate, the way the interpretableml entry states\n# it: how often the surrogate agrees with h to within a tolerance\ntol = 0.05\nagree_near = float(np.mean(np.abs(h(grid) - g)[near] < tol))\nagree_all = float(np.mean(np.abs(h(grid) - g) < tol))\nprint(f\"    agreement |h - g| < {tol}: {agree_near:.2f} near x0, \"\n      f\"{agree_all:.2f} over the whole range\")\ncheck(\"[B-faithful] the local surrogate is faithful near x0, not globally\",\n      agree_near > 0.9 and agree_all < 0.5)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-relevance]** For an image, the explanation is one relevance score per pixel. A linear classifier is fitted to 6x6 images of a seven against shapes sharing its top bar; the contribution of pixel j to the prediction is w_j x_j. The diagonal stroke carries the largest scores and every unlit pixel scores exactly 0, as drawn in the entry's Fig. 2. The shared top bar comes out small and negative: it is not evidence for a seven, and the fit uses it to bring the sum down to the label 1. Written to xaiterm.png."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "SEVEN = [(0, 5), (1, 5), (2, 5), (3, 5), (4, 5),\n         (4, 4), (3, 3), (3, 2), (2, 1), (2, 0)]     # the image of Fig. 2\n\n\ndef image(pixels):\n    img = np.zeros((6, 6))\n    for i, j in pixels:\n        img[5 - j, i] = 1.0\n    return img\n\n\nrng = np.random.default_rng(0)\nseven = image(SEVEN)\n# The negative shapes SHARE the top bar with the seven and differ in the\n# stroke below it. The top bar therefore separates nothing, and the diagonal\n# does \u2014 which is what makes it carry the larger relevance in Fig. 2.\nTOPBAR = [(0, 5), (1, 5), (2, 5), (3, 5), (4, 5)]\nothers = [image(TOPBAR + [(0, j) for j in range(5)]),      # top bar, left leg\n          image(TOPBAR)]                                   # top bar alone\nX, y = [], []\nfor _ in range(60):\n    X.append((seven + 0.05 * rng.standard_normal((6, 6))).ravel())\n    y.append(1.0)\n    other = others[rng.integers(len(others))]\n    X.append((other + 0.05 * rng.standard_normal((6, 6))).ravel())\n    y.append(-1.0)\nX, y = np.array(X), np.array(y)\nw_img, *_ = np.linalg.lstsq(X, y, rcond=None)        # the learned hypothesis\nrelevance = (w_img * seven.ravel()).reshape(6, 6)    # contribution of pixel j\nstroke = np.array([relevance[5 - j, i] for i, j in SEVEN[5:]])\ntopbar = np.array([relevance[5 - j, i] for i, j in SEVEN[:5]])\nunlit = relevance[seven == 0.0]\nprint(f\"    mean relevance: stroke {stroke.mean():.3f}, \"\n      f\"top bar {topbar.mean():.3f}, unlit pixels {np.abs(unlit).max():.3f}\")\ncheck(\"[B-relevance] every unlit pixel has relevance exactly 0\",\n      np.all(unlit == 0.0))\n# The stroke decides the prediction, and the top bar comes out SMALL AND\n# NEGATIVE: it is shared with the other shapes, so it is not evidence for a\n# seven, and the fit uses it to bring the sum down to the label 1.\ncheck(\"[B-relevance] the diagonal stroke dominates the shared top bar\",\n      stroke.mean() > 0.0\n      and abs(topbar.mean()) < 0.6 * stroke.mean())\ncheck(\"[B-relevance] the classifier separates the training images\",\n      np.all(np.sign(X @ w_img) == y))"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-morf]** The faithfulness of that map, tested rather than asserted. Pixels are flipped in order of decreasing relevance (most relevant first, MoRF) and, for comparison, in order of increasing relevance. The MoRF order drives the score across the decision threshold after far fewer flips, and does so for every one of 200 noisy images \u2014 which is what the entry means by a class activation map being faithful for a prediction."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "# Faithfulness of the map, as the entry states it: flipping the pixels it\n# scores highest must change the prediction more readily than flipping as\n# many of the pixels it scores low. \"Flip\" is the literal state flip of a\n# binary pixel, which is the perturbation the region-perturbation test was\n# first defined with.\nprint(\"[B-morf] flipping the highest-scoring pixels first changes the \"\n      \"prediction soonest\")\n\n\ndef flip_curve(img, order):\n    \"\"\"Score after flipping the first k pixels of `order`, k = 0, 1, 2, ...\"\"\"\n    x, out = img.ravel().copy(), [float(img.ravel() @ w_img)]\n    for j in order:\n        x[j] = 1.0 - x[j]                     # a binary pixel flips its state\n        out.append(float(x @ w_img))\n    return np.array(out)\n\n\ndef flips_to_change(curve):\n    \"\"\"How many flips until the prediction is no longer a seven.\"\"\"\n    below = np.where(curve <= 0.0)[0]\n    return int(below[0]) if len(below) else len(curve)\n\n\nmorf = np.argsort(-relevance.ravel())         # most relevant first\nlerf = np.argsort(relevance.ravel())          # least relevant first\ncurve_morf, curve_lerf = flip_curve(seven, morf), flip_curve(seven, lerf)\nk_morf, k_lerf = flips_to_change(curve_morf), flips_to_change(curve_lerf)\nprint(f\"    flips needed to change the prediction: {k_morf} guided by the \"\n      f\"map, {k_lerf} against it\")\ncheck(\"[B-morf]  the map-guided order changes the prediction sooner\",\n      k_morf < k_lerf)\n\n# the same comparison over many noisy images, so the claim is not one picture\nwins = 0\nfor _ in range(200):\n    img = np.clip(seven + 0.05 * rng.standard_normal((6, 6)), 0.0, 1.0)\n    rel = (w_img * img.ravel()).reshape(6, 6)\n    a = flips_to_change(flip_curve(img, np.argsort(-rel.ravel())))\n    b = flips_to_change(flip_curve(img, np.argsort(rel.ravel())))\n    wins += a < b\nprint(f\"    map-guided order wins on {wins}/200 noisy images\")\ncheck(\"[B-morf]  and it wins on every one of 200 noisy images\", wins == 200)\n\n# The map carries signs, so shading alone would not tell a reader which\n# pixels argue FOR a seven: the shade gives the size, a printed + or - the\n# direction, and neither channel is a colour.\nfig, ax = plt.subplots(2, 2, figsize=(7.6, 6.4))\n\nax[0, 0].imshow(seven, cmap=\"Greys\", vmin=0.0, vmax=1.0)\nax[0, 0].set_title(\"data point: the image\")\n\nim = ax[0, 1].imshow(np.abs(relevance), cmap=\"Greys\",\n                     vmin=0.0, vmax=np.abs(relevance).max())\nax[0, 1].set_title(\"explanation: relevance per pixel\")\nfor row in range(6):\n    for col in range(6):\n        val = relevance[row, col]\n        if val != 0.0:\n            ax[0, 1].text(col, row, \"+\" if val > 0 else \"-\",\n                          ha=\"center\", va=\"center\", fontsize=9,\n                          color=\"white\" if abs(val) > 0.5 * np.abs(relevance).max()\n                          else \"black\")\nfig.colorbar(im, ax=ax[0, 1], fraction=0.046, label=\"|relevance|\")\n\n# which pixels the map picks, and which of them the prediction turns on: the\n# flipped ones are ringed, and the order they were flipped in is printed, so\n# the reader sees they are the top of the relevance ranking and nothing else\nax[1, 0].imshow(seven, cmap=\"Greys\", vmin=0.0, vmax=1.0)\nfor rank, j_ in enumerate(morf[:k_morf]):\n    r, c = divmod(int(j_), 6)\n    lit = seven[r, c] > 0.5\n    ring = \"white\" if lit else \"black\"        # contrast against the pixel\n    ax[1, 0].plot(c, r, marker=\"o\", markersize=18, markerfacecolor=\"none\",\n                  markeredgecolor=ring, markeredgewidth=2.0)\n    ax[1, 0].text(c, r, str(rank + 1), ha=\"center\", va=\"center\", fontsize=8,\n                  color=ring)\nax[1, 0].set_title(f\"the {k_morf} flips that change the prediction\")\nax[1, 0].set_xlabel(\"pixel column\")\nax[1, 0].set_ylabel(\"pixel row\")\n\nax[1, 1].plot(np.arange(len(curve_morf)), curve_morf, \"k-\", marker=\"o\",\n              markersize=3, label=\"most relevant first\")\nax[1, 1].plot(np.arange(len(curve_lerf)), curve_lerf, \"k--\", marker=\"s\",\n              markersize=3, markerfacecolor=\"none\", label=\"least relevant first\")\nax[1, 1].axhline(0.0, color=\"k\", linewidth=0.8, linestyle=\":\")\nax[1, 1].annotate(\"prediction changes\", xy=(k_morf, 0.0),\n                  xytext=(k_morf + 4, -1.55), fontsize=8,\n                  arrowprops=dict(arrowstyle=\"->\", linewidth=0.8))\nax[1, 1].set_xlabel(\"number of pixels flipped\")\nax[1, 1].set_ylabel(\"score of the learned hypothesis\")\nax[1, 1].set_title(\"perturbation curves\")\nax[1, 1].legend(frameon=False, fontsize=8, loc=\"upper left\")\n\nfor a in (ax[0, 0], ax[0, 1], ax[1, 0]):\n    a.set_xlabel(\"pixel column\")\n    a.set_ylabel(\"pixel row\")\n    a.set_xticks(range(6))\n    a.set_yticks(range(6))\nfig.tight_layout()\nfig.savefig(OUT_DIR / \"xaiterm.png\", dpi=110)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "**[B-eerm]** Explainability can enter training instead. The user supplies their own predictions for the training set, and a penalty charges the part of the fitted predictions that those user predictions do not already account for. As the penalty weight grows, the fit becomes more predictable from the user's own predictions while the average loss rises."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "# Explainability built into training: the user supplies their own predictions\n# for the training set, and the penalty charges the part of the fitted\n# predictions that those user predictions do not already account for.\nn, d = 120, 3\nXe = rng.standard_normal((n, d))\nye = Xe @ np.array([1.5, -1.0, 0.7]) + 0.3 * rng.standard_normal(n)\nuser = Xe[:, 0]                       # this user reasons about feature 1 only\nU = np.c_[np.ones(n), user]           # what the user can already account for\nP = U @ np.linalg.pinv(U)             # the part of a fit that U accounts for\nM = np.eye(n) - P                     # and the part it does NOT account for\n\n\ndef fit(alpha):                       # EERM: average loss plus the penalty\n    A_ = Xe.T @ Xe + alpha * Xe.T @ M @ Xe\n    return np.linalg.solve(A_, Xe.T @ ye)\n\n\ndef explained(w_):                    # how far the fit follows the user\n    pred = Xe @ w_\n    resid = pred - P @ pred\n    return 1.0 - float(resid @ resid) / float(pred @ pred)\n\n\ndef avg_loss(w_):\n    r = Xe @ w_ - ye\n    return float(r @ r) / n\n\n\nrows = [(a, explained(fit(a)), avg_loss(fit(a))) for a in (0.0, 1.0, 10.0)]\nfor a, ex, te in rows:\n    print(f\"    alpha={a:<5} followed by the user summary: {ex:.3f}, \"\n          f\"average loss: {te:.3f}\")\ncheck(\"[B-eerm]  the penalty makes the fit follow the user summary\",\n      rows[0][1] < rows[1][1] < rows[2][1])\ncheck(\"[B-eerm]  and it costs average loss\",\n      rows[0][2] < rows[1][2] < rows[2][2])\n\nn_ok = sum(ok for _, ok in report)\nprint(f\"\\n{n_ok}/{len(report)} checks pass \"\n      f\"(x0 = {x0}, counterfactual x' = {x_cf:.3f}; \"\n      f\"phi = {np.round(phi, 3)}, sum = {phi.sum():.3f}, \"\n      f\"f(x) - f(base) = {f(z) - f(base):.3f})\")\nif n_ok != len(report):\n    raise SystemExit(1)"
  }
 ]
}