shap-svg 0.1.0 → 0.2.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/README.md CHANGED
@@ -26,23 +26,107 @@ affiliated with the SHAP authors.
26
26
 
27
27
  | Component | SHAP counterpart | Shows |
28
28
  | --- | --- | --- |
29
- | `ShapBar` | `shap.plots.bar` | mean(\|SHAP value\|) per feature across samples |
30
- | `ShapBeeswarm` | `shap.plots.beeswarm` | one dot per sample per feature, coloured by feature value |
31
- | `ShapHeatmap` | `shap.plots.heatmap` | samples × features coloured by SHAP value, with the f(x) line above |
32
- | `ShapWaterfall` | `shap.plots.waterfall` | how one sample's prediction is built from E[f(X)] to f(x) |
29
+ | `Plots.bar` | `shap.plots.bar` | mean(\|SHAP value\|) per feature across samples |
30
+ | `Plots.beeswarm` | `shap.plots.beeswarm` | one dot per sample per feature, coloured by feature value |
31
+ | `Plots.heatmap` | `shap.plots.heatmap` | samples × features coloured by SHAP value, with the f(x) line above |
32
+ | `Plots.waterfall` | `shap.plots.waterfall` | how one sample's prediction is built from E[f(X)] to f(x) |
33
33
 
34
34
  Every chart is a pure component: all state that changes what is drawn arrives through props, so the
35
35
  host application owns its own controls. The only internal state is hover highlighting.
36
36
 
37
37
  ## Usage
38
38
 
39
- ```tsx
40
- import { ShapBeeswarm, ShapWaterfall } from "shap-svg/react";
39
+ `shap-svg` draws SHAP values; it does not compute them. The values come from
40
+ [`shap`](https://github.com/shap/shap) in Python, travel through your server as JSON, and are handed to
41
+ a component in the browser:
42
+
43
+ ```text
44
+ Python: shap computes the values → server returns them as JSON → browser fetches → <Plots.waterfall explanation={…} />
45
+ ```
46
+
47
+ ### 1. Compute SHAP values in Python and serialise them
48
+
49
+ ```python
50
+ import numpy as np
51
+ import shap
52
+
53
+ explainer = shap.TreeExplainer(model) # any shap explainer works
54
+
55
+
56
+ def to_payload(explanation: shap.Explanation, sample_ids=None) -> dict:
57
+ """Serialise a shap.Explanation into the JSON shap-svg reads."""
58
+ payload = {
59
+ "contract_version": 1,
60
+ "values": np.asarray(explanation.values, dtype=float).tolist(),
61
+ "base_values": np.asarray(explanation.base_values, dtype=float).tolist(),
62
+ "data": np.asarray(explanation.data, dtype=float).tolist(),
63
+ "feature_names": [str(name) for name in explanation.feature_names],
64
+ }
65
+ if sample_ids is not None:
66
+ payload["sample_ids"] = [str(sample_id) for sample_id in sample_ids]
67
+ return payload
68
+ ```
69
+
70
+ - **Pass the explanation as it comes.** A scikit-learn binary classifier gives `values` shaped
71
+ `(samples, features, 2)` and `base_values` shaped `(samples, 2)`; `shap-svg` reads class 1 by default
72
+ (`classIndex` on every component). There is no need to pick a class in Python.
73
+ - **Send only finite numbers.** Fill or drop missing feature values first. Flask's `jsonify` writes
74
+ `NaN` unquoted, which is not valid JSON, and `parseExplanation` rejects it; serialising with
75
+ `json.dumps(payload, allow_nan=False)` makes the mistake fail on the server instead.
76
+ - **`sample_ids` are optional** but give each Sample a stable key: the heatmap hands it back from
77
+ `onSampleClick`. Add `sample_labels` for the names a person should read.
78
+
79
+ ### 2. Return the payload from your server
41
80
 
42
- <ShapBeeswarm explanation={explanation} maxDisplay={15} groupByGenus rowSort="name" />
43
- <ShapWaterfall explanation={explanation} sampleIndex={0} decimals="percent" />
81
+ ```python
82
+ from flask import Flask, jsonify
83
+
84
+ app = Flask(__name__)
85
+
86
+
87
+ @app.post("/api/explain")
88
+ def explain():
89
+ X = rows_to_explain() # a pandas DataFrame, e.g. built from the request body
90
+ explanation = explainer(X)
91
+ return jsonify(to_payload(explanation, sample_ids=X.index))
44
92
  ```
45
93
 
94
+ Any framework works — the contract is only the JSON above. Serve it gzipped if you can: the payload is
95
+ mostly repeated digits and compresses well.
96
+
97
+ ### 3. Fetch it in the browser and hand it to a component
98
+
99
+ ```tsx
100
+ import { useEffect, useState } from "react";
101
+ import type { Explanation } from "shap-svg";
102
+ import { Plots } from "shap-svg/react";
103
+
104
+ export function ExplanationView() {
105
+ const [explanation, setExplanation] = useState<Explanation>();
106
+
107
+ useEffect(() => {
108
+ fetch("/api/explain", { method: "POST" })
109
+ .then((response) => response.json())
110
+ .then(setExplanation);
111
+ }, []);
112
+
113
+ if (!explanation) return <p>Loading…</p>;
114
+
115
+ return (
116
+ <>
117
+ <Plots.beeswarm explanation={explanation} maxDisplay={15} groupByGenus rowSort="name" />
118
+ <Plots.waterfall explanation={explanation} sampleIndex={0} decimals="percent" />
119
+ </>
120
+ );
121
+ }
122
+ ```
123
+
124
+ The charts are named the way `shap` names them in Python: `shap.plots.bar` becomes `<Plots.bar />`.
125
+ `Plots` brings all four charts into your bundle, even if a page draws one — about 30 KB minified.
126
+
127
+ From here every control — how many features, grouping, sorting, precision — is a prop. Changing one
128
+ redraws from the payload already in memory; nothing goes back to the server.
129
+
46
130
  The framework-free core — parsing, ordering, collapsing, layout and colour — has no React import and
47
131
  can drive any renderer:
48
132
 
@@ -94,11 +178,12 @@ Per chart:
94
178
 
95
179
  | Chart | Prop | Default | |
96
180
  | --- | --- | --- | --- |
97
- | `ShapBeeswarm`, `ShapHeatmap` | `rowSort` | `"importance"` | `"importance"`, `"name"` or `"featureValue"`; reorders the rows shown, never which rows are shown |
98
- | `ShapBeeswarm` | `seed`, `dotRadius` | `0`, `3` | jitter is seeded, so a chart is identical on every render |
99
- | `ShapHeatmap` | `onSampleClick` | — | called with the column's `sample_ids` entry |
100
- | `ShapWaterfall` | `sampleIndex` | `0` | which sample to explain |
101
- | `ShapWaterfall` | `decimals` | `2` | `2`, `3`, `4` or `"percent"`; display only |
181
+ | `Plots.beeswarm`, `Plots.heatmap` | `rowSort` | `"importance"` | `"importance"`, `"name"` or `"featureValue"`; reorders the rows shown, never which rows are shown |
182
+ | `Plots.beeswarm`, `Plots.heatmap` | `colorBar` | `true` | SHAP's colour bar right of the plot — Low to High feature value on the beeswarm, the SHAP value range on the heatmap. The plot narrows when the right margin cannot hold it |
183
+ | `Plots.beeswarm` | `seed`, `dotRadius` | `0`, `3` | jitter is seeded, so a chart is identical on every render |
184
+ | `Plots.heatmap` | `onSampleClick` | — | called with the column's `sample_ids` entry |
185
+ | `Plots.waterfall` | `sampleIndex` | `0` | which sample to explain |
186
+ | `Plots.waterfall` | `decimals` | `2` | `2`, `3`, `4` or `"percent"`; display only |
102
187
 
103
188
  ## Faithful to SHAP where it matters
104
189
 
@@ -1141,12 +1141,118 @@ function sampleColormap(name, t) {
1141
1141
  return `#${hexByte(red)}${hexByte(green)}${hexByte(blue)}`;
1142
1142
  }
1143
1143
 
1144
+ // src/core/colorBar.ts
1145
+ var ASPECT = 80;
1146
+ var TICK_PAD = 3.5;
1147
+ var TICK_LABEL_PT3 = 11;
1148
+ var LABEL_PT = 12;
1149
+ var LABEL_THICKNESS_EM = 1.08;
1150
+ var LABEL_BASELINE_EM = 0.84;
1151
+ var CHAR_EM = 0.6;
1152
+ var STEPS = 64;
1153
+ function textWidth(text, fontSize) {
1154
+ return text.length * fontSize * CHAR_EM;
1155
+ }
1156
+ function tickColumnWidth(spec) {
1157
+ return Math.max(...spec.tickLabels.map((label) => textWidth(label, TICK_LABEL_PT3)));
1158
+ }
1159
+ function colorBarExtent(height, spec) {
1160
+ return height / ASPECT + TICK_PAD + Math.max(0, tickColumnWidth(spec) + spec.labelPad) + LABEL_PT * LABEL_THICKNESS_EM;
1161
+ }
1162
+ function fitColorBar(opts) {
1163
+ const { plotWidth, available, gapRatio, minGap, extent } = opts;
1164
+ const gapFor = (width) => Math.max(gapRatio * width, minGap);
1165
+ if (gapFor(plotWidth) + extent <= available) {
1166
+ return { plotWidth, gap: gapFor(plotWidth) };
1167
+ }
1168
+ const byRatio = (available + plotWidth - extent) / (1 + gapRatio);
1169
+ const fitted = Math.max(
1170
+ 0,
1171
+ gapRatio * byRatio >= minGap ? byRatio : plotWidth - (minGap + extent - available)
1172
+ );
1173
+ return { plotWidth: fitted, gap: gapFor(fitted) };
1174
+ }
1175
+ function colorBarLayout(spec, position) {
1176
+ const { x, y1, y2 } = position;
1177
+ const width = (y2 - y1) / ASPECT;
1178
+ const step = (y2 - y1) / STEPS;
1179
+ const steps = Array.from({ length: STEPS }, (_, index) => ({
1180
+ y: y2 - (index + 1) * step,
1181
+ // Every band but the lowest reaches half a pixel into the one below, so
1182
+ // anti-aliasing cannot open a hairline seam between them.
1183
+ height: index === 0 ? step : step + 0.5,
1184
+ color: sampleColormap(spec.colormap, index / (STEPS - 1))
1185
+ }));
1186
+ const tickX = x + width + TICK_PAD;
1187
+ const labelLeft = tickX + tickColumnWidth(spec) + spec.labelPad;
1188
+ return {
1189
+ x,
1190
+ y1,
1191
+ y2,
1192
+ width,
1193
+ steps,
1194
+ tickX,
1195
+ tickFontSize: TICK_LABEL_PT3,
1196
+ ticks: [
1197
+ { y: y2, label: spec.tickLabels[0] },
1198
+ { y: y1, label: spec.tickLabels[1] }
1199
+ ],
1200
+ ...spec.offsetText ? { offsetText: { x: tickX, y: y1 - TICK_LABEL_PT3, text: spec.offsetText } } : {},
1201
+ label: {
1202
+ x: labelLeft + LABEL_PT * LABEL_BASELINE_EM,
1203
+ y: (y1 + y2) / 2,
1204
+ text: spec.label,
1205
+ fontSize: LABEL_PT
1206
+ },
1207
+ right: Math.max(
1208
+ tickX + tickColumnWidth(spec),
1209
+ labelLeft + LABEL_PT * LABEL_THICKNESS_EM
1210
+ )
1211
+ };
1212
+ }
1213
+ var MINUS3 = "\u2212";
1214
+ var POWER_LIMITS = [-5, 6];
1215
+ function roundTo(value, decimals) {
1216
+ const factor = 10 ** decimals;
1217
+ return Math.round(value * factor) / factor;
1218
+ }
1219
+ function scalarFormatterLabels(locs) {
1220
+ const largest = Math.max(0, ...locs.map(Math.abs));
1221
+ const magnitude = largest === 0 ? 0 : Math.floor(Math.log10(largest));
1222
+ const order = magnitude <= POWER_LIMITS[0] || magnitude >= POWER_LIMITS[1] ? magnitude : 0;
1223
+ const scaled = locs.map((value) => value / 10 ** order);
1224
+ let range = Math.max(...scaled) - Math.min(...scaled);
1225
+ if (range === 0) range = Math.max(...scaled.map(Math.abs));
1226
+ if (range === 0) range = 1;
1227
+ const rangeMagnitude = Math.floor(Math.log10(range));
1228
+ const threshold = 1e-3 * 10 ** rangeMagnitude;
1229
+ let decimals = Math.max(0, 3 - rangeMagnitude);
1230
+ while (decimals >= 0) {
1231
+ const error = Math.max(...scaled.map((value) => Math.abs(value - roundTo(value, decimals))));
1232
+ if (error < threshold) decimals -= 1;
1233
+ else break;
1234
+ }
1235
+ decimals += 1;
1236
+ const labels = scaled.map((value) => {
1237
+ const text = value.toFixed(decimals);
1238
+ return Number(text) === 0 ? text.replace("-", "") : text.replace("-", MINUS3);
1239
+ });
1240
+ return order === 0 ? { labels } : { labels, offsetText: `1e${String(order).replace("-", MINUS3)}` };
1241
+ }
1242
+
1144
1243
  // src/core/beeswarmLayout.ts
1145
1244
  var BEESWARM_MISSING_COLOR = "#777777";
1146
1245
  var BEESWARM_ROW_HEIGHT = 0.4;
1147
1246
  var NBINS = 100;
1148
- var AXIS_HEIGHT3 = 74;
1149
- var TICK_LABEL_PT3 = 11;
1247
+ var AXIS_HEIGHT3 = 52;
1248
+ var COLOR_BAR = {
1249
+ colormap: "red_blue",
1250
+ tickLabels: ["Low", "High"],
1251
+ label: "Feature value",
1252
+ labelPad: 0
1253
+ };
1254
+ var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1255
+ var TICK_LABEL_PT4 = 11;
1150
1256
  var TITLE_PT2 = 13;
1151
1257
  var X_MARGIN2 = 0.05;
1152
1258
  function percentile(values, percent) {
@@ -1278,7 +1384,7 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1278
1384
  return { rows, collapsedCount: display.collapsedCount };
1279
1385
  }
1280
1386
  function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1281
- const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT3));
1387
+ const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1282
1388
  return {
1283
1389
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1284
1390
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
@@ -1293,7 +1399,15 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1293
1399
  }
1294
1400
  function beeswarmLayout(valueRows, opts) {
1295
1401
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1296
- const plotWidth = width - marginLeft - marginRight;
1402
+ const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1403
+ const fit = opts.colorBar ? fitColorBar({
1404
+ plotWidth: width - marginLeft - marginRight,
1405
+ available: marginRight,
1406
+ gapRatio: COLOR_BAR_GAP_RATIO,
1407
+ minGap: 0,
1408
+ extent: colorBarExtent(plotBottom - marginTop, COLOR_BAR)
1409
+ }) : null;
1410
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1297
1411
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
1298
1412
  const dataMin = Math.min(0, ...values);
1299
1413
  const dataMax = Math.max(0, ...values);
@@ -1318,7 +1432,6 @@ function beeswarmLayout(valueRows, opts) {
1318
1432
  }))
1319
1433
  };
1320
1434
  });
1321
- const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1322
1435
  return {
1323
1436
  rows,
1324
1437
  xDomain: [min, max],
@@ -1326,6 +1439,11 @@ function beeswarmLayout(valueRows, opts) {
1326
1439
  ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1327
1440
  plotWidth,
1328
1441
  plotBottom,
1442
+ colorBar: fit ? colorBarLayout(COLOR_BAR, {
1443
+ x: marginLeft + plotWidth + fit.gap,
1444
+ y1: marginTop,
1445
+ y2: plotBottom
1446
+ }) : null,
1329
1447
  height: plotBottom + AXIS_HEIGHT3
1330
1448
  };
1331
1449
  }
@@ -1338,7 +1456,9 @@ var SIDE_BAR_GAP = 10;
1338
1456
  var SIDE_BAR_RIGHT_INSET = 40;
1339
1457
  var SIDE_BAR_HEIGHT_RATIO = 0.6;
1340
1458
  var AXIS_HEIGHT4 = 52;
1341
- var TICK_LABEL_PT4 = 10;
1459
+ var TICK_LABEL_PT5 = 10;
1460
+ var COLOR_BAR_GAP_RATIO2 = 0.1 / 0.89;
1461
+ var COLOR_BAR_CLEARANCE = 8;
1342
1462
  var Y_TICK_LENGTH = 5;
1343
1463
  function percentile2(values, fraction) {
1344
1464
  if (values.length === 0) return 0;
@@ -1429,7 +1549,7 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1429
1549
  const { ticks, step } = niceTicks(
1430
1550
  -0.5,
1431
1551
  sampleCount - 0.5,
1432
- tickSpace(plotWidth, TICK_LABEL_PT4),
1552
+ tickSpace(plotWidth, TICK_LABEL_PT5),
1433
1553
  { integer: true }
1434
1554
  );
1435
1555
  return {
@@ -1443,17 +1563,33 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1443
1563
  text: "Instances",
1444
1564
  x: marginLeft + plotWidth / 2,
1445
1565
  y: plotBottom + AXIS_TITLE_DY,
1446
- fontSize: TICK_LABEL_PT4
1566
+ fontSize: TICK_LABEL_PT5
1447
1567
  }
1448
1568
  };
1449
1569
  }
1450
1570
  function heatmapLayout(valueRows, opts) {
1451
1571
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1452
- const plotWidth = width - marginLeft - marginRight;
1453
- const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1454
- const gridRight = marginLeft + plotWidth;
1455
1572
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1456
1573
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1574
+ const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1575
+ const colorBarSpec = {
1576
+ colormap: "red_white_blue",
1577
+ tickLabels: [ticks.labels[0], ticks.labels[1]],
1578
+ ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1579
+ label: "SHAP value (impact on model output)",
1580
+ labelPad: -10
1581
+ };
1582
+ const colorBarTop = FX_TOP;
1583
+ const fit = opts.colorBar ? fitColorBar({
1584
+ plotWidth: width - marginLeft - marginRight,
1585
+ available: marginRight,
1586
+ gapRatio: COLOR_BAR_GAP_RATIO2,
1587
+ minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1588
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec)
1589
+ }) : null;
1590
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1591
+ const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1592
+ const gridRight = marginLeft + plotWidth;
1457
1593
  const columns = valueRows.columns.map((column, index) => ({
1458
1594
  ...column,
1459
1595
  x: marginLeft + index * cellWidth,
@@ -1526,6 +1662,11 @@ function heatmapLayout(valueRows, opts) {
1526
1662
  sampleLabelColumn: valueRows.sampleLabelColumn,
1527
1663
  plotWidth,
1528
1664
  cellWidth,
1665
+ colorBar: fit ? colorBarLayout(colorBarSpec, {
1666
+ x: gridRight + fit.gap,
1667
+ y1: colorBarTop,
1668
+ y2: plotBottom
1669
+ }) : null,
1529
1670
  height: plotBottom + AXIS_HEIGHT4
1530
1671
  };
1531
1672
  }
@@ -1556,6 +1697,10 @@ export {
1556
1697
  waterfallRows,
1557
1698
  waterfallLayout,
1558
1699
  sampleColormap,
1700
+ colorBarExtent,
1701
+ fitColorBar,
1702
+ colorBarLayout,
1703
+ scalarFormatterLabels,
1559
1704
  BEESWARM_MISSING_COLOR,
1560
1705
  BEESWARM_ROW_HEIGHT,
1561
1706
  beeswarmRows,
@@ -1563,4 +1708,4 @@ export {
1563
1708
  heatmapRows,
1564
1709
  heatmapLayout
1565
1710
  };
1566
- //# sourceMappingURL=chunk-OXFKP5I3.js.map
1711
+ //# sourceMappingURL=chunk-ZWO4CN22.js.map