shap-svg 0.2.1 → 0.2.3

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
@@ -172,6 +172,7 @@ Shared by all four charts:
172
172
  | `classIndex` | `1` | which output to draw for multi-output explanations |
173
173
  | `width` | `720` | SVG width in pixels |
174
174
  | `rowHeight` | per chart | pixels per feature row |
175
+ | `labels` | SHAP's wording | text the chart draws; see [Wording](#wording) |
175
176
  | `onFeatureClick` | — | called with the feature index, or `null` for the "other" row |
176
177
 
177
178
  Per chart:
@@ -185,6 +186,43 @@ Per chart:
185
186
  | `Plots.waterfall` | `sampleIndex` | `0` | which sample to explain |
186
187
  | `Plots.waterfall` | `decimals` | `2` | `2`, `3`, `4` or `"percent"`; display only |
187
188
 
189
+ ## Wording
190
+
191
+ Every piece of text a chart draws comes from `labels`, keyed by what it means rather than where it
192
+ appears — a word used in two places is changed once. Give only the keys you want to change; the rest
193
+ keep SHAP's wording, exported as `shapLabels`.
194
+
195
+ ```tsx
196
+ import { Plots } from "shap-svg/react";
197
+ import type { PlotLabels } from "shap-svg";
198
+
199
+ // A module constant: a new object on every render recomputes the layout on every render.
200
+ const researchLabels: Partial<PlotLabels> = {
201
+ shapValue: "Contribution",
202
+ shapValueAxis: "Contribution to predicted probability",
203
+ featureValue: "Relative abundance",
204
+ otherFeatures: (count) => `${count} other taxa`,
205
+ };
206
+
207
+ <Plots.beeswarm explanation={explanation} labels={researchLabels} />;
208
+ ```
209
+
210
+ | Key | Default | Drawn as |
211
+ | --- | --- | --- |
212
+ | `shapValue` | `SHAP value` | beeswarm tooltip |
213
+ | `shapValueAxis` | `SHAP value (impact on model output)` | beeswarm x axis, heatmap colour bar |
214
+ | `meanAbsShapValue` | `mean(\|SHAP value\|)` | bar x axis |
215
+ | `featureValue` | `Feature value` | beeswarm tooltip and colour bar |
216
+ | `featureValueLow`, `featureValueHigh` | `Low`, `High` | ends of the beeswarm colour bar |
217
+ | `missingFeatureValue` | `missing` | beeswarm tooltip, for a value the payload lacks |
218
+ | `samples` | `Instances` | heatmap x axis |
219
+ | `sampleTotal` | `Σφ` | heatmap tooltip |
220
+ | `sampleFallback(n)` | `Sample n` | a heatmap column with no `sample_labels` entry |
221
+ | `baseValue`, `modelOutput` | `E[f(X)]`, `f(x)` | the waterfall's two reference lines |
222
+ | `otherFeatures(count, style)` | `Sum of N other features` / `N other features` | the row for every feature not shown; `style` is `"sum"` in `faithfulOtherRow` mode on bar, beeswarm and heatmap |
223
+
224
+ Accessible names are built from the same words.
225
+
188
226
  ## Faithful to SHAP where it matters
189
227
 
190
228
  Ordering, the "other features" partition, colour maps, tick placement and the waterfall's geometry are
@@ -261,8 +261,34 @@ function formatFeatureLabel(name) {
261
261
  return name.replace(/_/g, " ");
262
262
  }
263
263
 
264
+ // src/core/labels.ts
265
+ var shapLabels = {
266
+ shapValue: "SHAP value",
267
+ shapValueAxis: "SHAP value (impact on model output)",
268
+ meanAbsShapValue: "mean(|SHAP value|)",
269
+ featureValue: "Feature value",
270
+ featureValueLow: "Low",
271
+ featureValueHigh: "High",
272
+ missingFeatureValue: "missing",
273
+ samples: "Instances",
274
+ sampleTotal: "\u03A3\u03C6",
275
+ sampleFallback: (sampleNumber) => `Sample ${sampleNumber}`,
276
+ baseValue: "E[f(X)]",
277
+ modelOutput: "f(x)",
278
+ otherFeatures: (count, style) => style === "sum" ? `Sum of ${count} other features` : `${count} other features`
279
+ };
280
+ function resolveLabels(labels) {
281
+ if (!labels) return shapLabels;
282
+ const resolved = { ...shapLabels };
283
+ for (const key of Object.keys(labels)) {
284
+ const value = labels[key];
285
+ if (value !== void 0) resolved[key] = value;
286
+ }
287
+ return resolved;
288
+ }
289
+
264
290
  // src/core/collapse.ts
265
- function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow) {
291
+ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow, labels = shapLabels) {
266
292
  const p = order.length;
267
293
  if (maxDisplay >= p) {
268
294
  return {
@@ -285,7 +311,7 @@ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithful
285
311
  const collapsed = order.slice(realCount);
286
312
  const collapsedValue = collapsed.reduce((sum, index) => sum + importance[index], 0);
287
313
  rows.push({
288
- label: faithfulOtherRow ? `Sum of ${collapsed.length} other features` : `${collapsed.length} other features`,
314
+ label: labels.otherFeatures(collapsed.length, faithfulOtherRow ? "sum" : "count"),
289
315
  featureIndex: null,
290
316
  value: collapsedValue,
291
317
  isOtherRow: true
@@ -354,15 +380,15 @@ var AXIS_HEIGHT = 52;
354
380
  var TICK_LABEL_PT = 11;
355
381
  var TITLE_PT = 13;
356
382
  var X_MARGIN = 0.05;
357
- function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
383
+ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
358
384
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT));
359
385
  return {
360
386
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
361
387
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
362
388
  xTitle: {
363
- // _bar.py:143-150 builds this from the Explanation's transform history:
364
- // "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
365
- text: "mean(|SHAP value|)",
389
+ // _bar.py:143-150 builds the default from the Explanation's transform
390
+ // history: "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
391
+ text: title,
366
392
  x: marginLeft + plotWidth / 2,
367
393
  y: plotBottom + AXIS_TITLE_DY,
368
394
  fontSize: TITLE_PT
@@ -370,6 +396,7 @@ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
370
396
  };
371
397
  }
372
398
  function barLayout(rows, opts) {
399
+ var _a;
373
400
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
374
401
  const plotWidth = width - marginLeft - marginRight;
375
402
  const values = rows.rows.map((r) => r.value);
@@ -410,7 +437,15 @@ function barLayout(rows, opts) {
410
437
  xZero,
411
438
  plotWidth,
412
439
  plotBottom: marginTop + rows.rows.length * rowHeight,
413
- ...barXAxis(min, max, toX, marginLeft, plotWidth, marginTop + rows.rows.length * rowHeight),
440
+ ...barXAxis(
441
+ min,
442
+ max,
443
+ toX,
444
+ marginLeft,
445
+ plotWidth,
446
+ marginTop + rows.rows.length * rowHeight,
447
+ ((_a = opts.labels) != null ? _a : shapLabels).meanAbsShapValue
448
+ ),
414
449
  zeroLine: { x: xZero, y1: marginTop, y2: marginTop + rows.rows.length * rowHeight },
415
450
  height: marginTop + rows.rows.length * rowHeight + AXIS_HEIGHT
416
451
  };
@@ -444,7 +479,7 @@ function placeValueLabel(value, startX, endX, gutterX, decimals) {
444
479
  }
445
480
  return { ...base, x: startX + VALUE_LABEL_GAP, anchor: "start", inside: false };
446
481
  }
447
- function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
482
+ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow, labels = shapLabels) {
448
483
  if (!Number.isInteger(sampleIndex) || sampleIndex < 0 || sampleIndex >= explanation.nSamples) {
449
484
  throw new RangeError(
450
485
  `sampleIndex must identify a Sample from 0 to ${explanation.nSamples - 1}, received ${sampleIndex}`
@@ -482,7 +517,8 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
482
517
  if (hasOtherRow) {
483
518
  const value = collapsed.reduce((sum, featureIndex) => sum + values[featureIndex], 0);
484
519
  rows.push({
485
- label: `${collapsed.length} other features`,
520
+ // SHAP's waterfall counts the hidden Features whatever the collapse mode.
521
+ label: labels.otherFeatures(collapsed.length, "count"),
486
522
  featureIndex: null,
487
523
  isOtherRow: true,
488
524
  value,
@@ -500,7 +536,9 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
500
536
  };
501
537
  }
502
538
  function waterfallLayout(valueRows, opts) {
539
+ var _a;
503
540
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
541
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
504
542
  const plotWidth = width - marginLeft - marginRight;
505
543
  const coordinates = [valueRows.baseValue, valueRows.modelOutput];
506
544
  for (const row of valueRows.rows) coordinates.push(row.left, row.left + row.width);
@@ -567,7 +605,7 @@ function waterfallLayout(valueRows, opts) {
567
605
  // got, which is the whole reason SHAP hides the left spine.
568
606
  y1: plotBottom - rowHeight,
569
607
  y2: plotBottom,
570
- label: `E[f(X)] = ${formatLevel(valueRows.baseValue, opts.decimals)}`
608
+ label: `${labels.baseValue} = ${formatLevel(valueRows.baseValue, opts.decimals)}`
571
609
  },
572
610
  {
573
611
  kind: "output",
@@ -576,7 +614,7 @@ function waterfallLayout(valueRows, opts) {
576
614
  // axvline(fx, 0, 1) — the full height.
577
615
  y1: marginTop,
578
616
  y2: plotBottom,
579
- label: `f(x) = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
617
+ label: `${labels.modelOutput} = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
580
618
  }
581
619
  ],
582
620
  separators: valueRows.rows.map((_, index) => ({
@@ -1240,17 +1278,35 @@ function scalarFormatterLabels(locs) {
1240
1278
  return order === 0 ? { labels } : { labels, offsetText: `1e${String(order).replace("-", MINUS3)}` };
1241
1279
  }
1242
1280
 
1281
+ // src/core/tooltip.ts
1282
+ var CHAR_PX = 6.5;
1283
+ var PADDING_X = 7;
1284
+ var OFFSET = 8;
1285
+ var RAISE = 24;
1286
+ function placeTooltip(opts) {
1287
+ const { anchorX, anchorY, lines, lineHeight, minWidth, chartWidth, chartHeight } = opts;
1288
+ const longest = lines.reduce((max, line) => Math.max(max, line.length), 0);
1289
+ const width = Math.max(minWidth, longest * CHAR_PX + 2 * PADDING_X);
1290
+ const height = lines.length * lineHeight + lineHeight * 0.6;
1291
+ const right = anchorX + OFFSET;
1292
+ const x = right + width <= chartWidth ? right : Math.max(0, anchorX - OFFSET - width);
1293
+ const y = Math.max(0, Math.min(anchorY - RAISE, chartHeight - height));
1294
+ return { x, y, width, height };
1295
+ }
1296
+
1243
1297
  // src/core/beeswarmLayout.ts
1244
1298
  var BEESWARM_MISSING_COLOR = "#777777";
1245
1299
  var BEESWARM_ROW_HEIGHT = 0.4;
1246
1300
  var NBINS = 100;
1247
1301
  var AXIS_HEIGHT3 = 52;
1248
- var COLOR_BAR = {
1249
- colormap: "red_blue",
1250
- tickLabels: ["Low", "High"],
1251
- label: "Feature value",
1252
- labelPad: 0
1253
- };
1302
+ function colorBarSpec(labels) {
1303
+ return {
1304
+ colormap: "red_blue",
1305
+ tickLabels: [labels.featureValueLow, labels.featureValueHigh],
1306
+ label: labels.featureValue,
1307
+ labelPad: 0
1308
+ };
1309
+ }
1254
1310
  var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1255
1311
  var TICK_LABEL_PT4 = 11;
1256
1312
  var TITLE_PT2 = 13;
@@ -1316,7 +1372,7 @@ function spreadPoints(xs, rowIndex, seed) {
1316
1372
  const scale = 0.9 * (BEESWARM_ROW_HEIGHT / (maxPositiveOffset + 1));
1317
1373
  return offsets.map((offset) => rowIndex + offset * scale);
1318
1374
  }
1319
- function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance") {
1375
+ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance", labels = shapLabels) {
1320
1376
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1321
1377
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1322
1378
  }
@@ -1331,7 +1387,8 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1331
1387
  importance,
1332
1388
  order,
1333
1389
  maxDisplay,
1334
- faithfulOtherRow
1390
+ faithfulOtherRow,
1391
+ labels
1335
1392
  ),
1336
1393
  rowSort,
1337
1394
  explanation.data
@@ -1383,14 +1440,14 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1383
1440
  });
1384
1441
  return { rows, collapsedCount: display.collapsedCount };
1385
1442
  }
1386
- function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1443
+ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
1387
1444
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1388
1445
  return {
1389
1446
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1390
1447
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
1391
1448
  xTitle: {
1392
- // _labels.py:5, labels["VALUE"].
1393
- text: "SHAP value (impact on model output)",
1449
+ // The default is _labels.py:5, labels["VALUE"].
1450
+ text: title,
1394
1451
  x: marginLeft + plotWidth / 2,
1395
1452
  y: plotBottom + AXIS_TITLE_DY,
1396
1453
  fontSize: TITLE_PT2
@@ -1398,14 +1455,17 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1398
1455
  };
1399
1456
  }
1400
1457
  function beeswarmLayout(valueRows, opts) {
1458
+ var _a;
1401
1459
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1460
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1461
+ const colorBar = colorBarSpec(labels);
1402
1462
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1403
1463
  const fit = opts.colorBar ? fitColorBar({
1404
1464
  plotWidth: width - marginLeft - marginRight,
1405
1465
  available: marginRight,
1406
1466
  gapRatio: COLOR_BAR_GAP_RATIO,
1407
1467
  minGap: 0,
1408
- extent: colorBarExtent(plotBottom - marginTop, COLOR_BAR)
1468
+ extent: colorBarExtent(plotBottom - marginTop, colorBar)
1409
1469
  }) : null;
1410
1470
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1411
1471
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
@@ -1436,10 +1496,10 @@ function beeswarmLayout(valueRows, opts) {
1436
1496
  rows,
1437
1497
  xDomain: [min, max],
1438
1498
  xZero: toX(0),
1439
- ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1499
+ ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, labels.shapValueAxis),
1440
1500
  plotWidth,
1441
1501
  plotBottom,
1442
- colorBar: fit ? colorBarLayout(COLOR_BAR, {
1502
+ colorBar: fit ? colorBarLayout(colorBar, {
1443
1503
  x: marginLeft + plotWidth + fit.gap,
1444
1504
  y1: marginTop,
1445
1505
  y2: plotBottom
@@ -1469,7 +1529,7 @@ function percentile2(values, fraction) {
1469
1529
  const weight = position - lower;
1470
1530
  return sorted[lower] + (sorted[upper] - sorted[lower]) * weight;
1471
1531
  }
1472
- function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance") {
1532
+ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance", labels = shapLabels) {
1473
1533
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1474
1534
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1475
1535
  }
@@ -1481,7 +1541,8 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1481
1541
  importance,
1482
1542
  featureOrder,
1483
1543
  maxDisplay,
1484
- faithfulOtherRow
1544
+ faithfulOtherRow,
1545
+ labels
1485
1546
  ),
1486
1547
  rowSort,
1487
1548
  explanation.data
@@ -1545,7 +1606,7 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1545
1606
  sampleLabelColumn: explanation.sampleLabelColumn
1546
1607
  };
1547
1608
  }
1548
- function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom) {
1609
+ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom, title) {
1549
1610
  const { ticks, step } = niceTicks(
1550
1611
  -0.5,
1551
1612
  sampleCount - 0.5,
@@ -1560,7 +1621,7 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1560
1621
  })),
1561
1622
  xSpine: null,
1562
1623
  xTitle: {
1563
- text: "Instances",
1624
+ text: title,
1564
1625
  x: marginLeft + plotWidth / 2,
1565
1626
  y: plotBottom + AXIS_TITLE_DY,
1566
1627
  fontSize: TICK_LABEL_PT5
@@ -1568,15 +1629,17 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1568
1629
  };
1569
1630
  }
1570
1631
  function heatmapLayout(valueRows, opts) {
1632
+ var _a;
1571
1633
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1634
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1572
1635
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1573
1636
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1574
1637
  const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1575
- const colorBarSpec = {
1638
+ const colorBarSpec2 = {
1576
1639
  colormap: "red_white_blue",
1577
1640
  tickLabels: [ticks.labels[0], ticks.labels[1]],
1578
1641
  ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1579
- label: "SHAP value (impact on model output)",
1642
+ label: labels.shapValueAxis,
1580
1643
  labelPad: -10
1581
1644
  };
1582
1645
  const colorBarTop = FX_TOP;
@@ -1585,7 +1648,7 @@ function heatmapLayout(valueRows, opts) {
1585
1648
  available: marginRight,
1586
1649
  gapRatio: COLOR_BAR_GAP_RATIO2,
1587
1650
  minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1588
- extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec)
1651
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec2)
1589
1652
  }) : null;
1590
1653
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1591
1654
  const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
@@ -1658,11 +1721,18 @@ function heatmapLayout(valueRows, opts) {
1658
1721
  x1: marginLeft - Y_TICK_LENGTH,
1659
1722
  x2: marginLeft
1660
1723
  })),
1661
- ...heatmapXAxis(valueRows.columns.length, marginLeft, cellWidth, plotWidth, plotBottom),
1724
+ ...heatmapXAxis(
1725
+ valueRows.columns.length,
1726
+ marginLeft,
1727
+ cellWidth,
1728
+ plotWidth,
1729
+ plotBottom,
1730
+ labels.samples
1731
+ ),
1662
1732
  sampleLabelColumn: valueRows.sampleLabelColumn,
1663
1733
  plotWidth,
1664
1734
  cellWidth,
1665
- colorBar: fit ? colorBarLayout(colorBarSpec, {
1735
+ colorBar: fit ? colorBarLayout(colorBarSpec2, {
1666
1736
  x: gridRight + fit.gap,
1667
1737
  y1: colorBarTop,
1668
1738
  y2: plotBottom
@@ -1685,6 +1755,8 @@ export {
1685
1755
  formatShapValue,
1686
1756
  formatLevel,
1687
1757
  formatFeatureLabel,
1758
+ shapLabels,
1759
+ resolveLabels,
1688
1760
  collapseToDisplay,
1689
1761
  TICK_LENGTH,
1690
1762
  TICK_LABEL_DY,
@@ -1701,6 +1773,7 @@ export {
1701
1773
  fitColorBar,
1702
1774
  colorBarLayout,
1703
1775
  scalarFormatterLabels,
1776
+ placeTooltip,
1704
1777
  BEESWARM_MISSING_COLOR,
1705
1778
  BEESWARM_ROW_HEIGHT,
1706
1779
  beeswarmRows,
@@ -1708,4 +1781,4 @@ export {
1708
1781
  heatmapRows,
1709
1782
  heatmapLayout
1710
1783
  };
1711
- //# sourceMappingURL=chunk-ZWO4CN22.js.map
1784
+ //# sourceMappingURL=chunk-UVIVTJK3.js.map