shap-svg 0.2.1 → 0.2.2

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) => ({
@@ -1245,12 +1283,14 @@ var BEESWARM_MISSING_COLOR = "#777777";
1245
1283
  var BEESWARM_ROW_HEIGHT = 0.4;
1246
1284
  var NBINS = 100;
1247
1285
  var AXIS_HEIGHT3 = 52;
1248
- var COLOR_BAR = {
1249
- colormap: "red_blue",
1250
- tickLabels: ["Low", "High"],
1251
- label: "Feature value",
1252
- labelPad: 0
1253
- };
1286
+ function colorBarSpec(labels) {
1287
+ return {
1288
+ colormap: "red_blue",
1289
+ tickLabels: [labels.featureValueLow, labels.featureValueHigh],
1290
+ label: labels.featureValue,
1291
+ labelPad: 0
1292
+ };
1293
+ }
1254
1294
  var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1255
1295
  var TICK_LABEL_PT4 = 11;
1256
1296
  var TITLE_PT2 = 13;
@@ -1316,7 +1356,7 @@ function spreadPoints(xs, rowIndex, seed) {
1316
1356
  const scale = 0.9 * (BEESWARM_ROW_HEIGHT / (maxPositiveOffset + 1));
1317
1357
  return offsets.map((offset) => rowIndex + offset * scale);
1318
1358
  }
1319
- function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance") {
1359
+ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance", labels = shapLabels) {
1320
1360
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1321
1361
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1322
1362
  }
@@ -1331,7 +1371,8 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1331
1371
  importance,
1332
1372
  order,
1333
1373
  maxDisplay,
1334
- faithfulOtherRow
1374
+ faithfulOtherRow,
1375
+ labels
1335
1376
  ),
1336
1377
  rowSort,
1337
1378
  explanation.data
@@ -1383,14 +1424,14 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1383
1424
  });
1384
1425
  return { rows, collapsedCount: display.collapsedCount };
1385
1426
  }
1386
- function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1427
+ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
1387
1428
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1388
1429
  return {
1389
1430
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1390
1431
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
1391
1432
  xTitle: {
1392
- // _labels.py:5, labels["VALUE"].
1393
- text: "SHAP value (impact on model output)",
1433
+ // The default is _labels.py:5, labels["VALUE"].
1434
+ text: title,
1394
1435
  x: marginLeft + plotWidth / 2,
1395
1436
  y: plotBottom + AXIS_TITLE_DY,
1396
1437
  fontSize: TITLE_PT2
@@ -1398,14 +1439,17 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1398
1439
  };
1399
1440
  }
1400
1441
  function beeswarmLayout(valueRows, opts) {
1442
+ var _a;
1401
1443
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1444
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1445
+ const colorBar = colorBarSpec(labels);
1402
1446
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1403
1447
  const fit = opts.colorBar ? fitColorBar({
1404
1448
  plotWidth: width - marginLeft - marginRight,
1405
1449
  available: marginRight,
1406
1450
  gapRatio: COLOR_BAR_GAP_RATIO,
1407
1451
  minGap: 0,
1408
- extent: colorBarExtent(plotBottom - marginTop, COLOR_BAR)
1452
+ extent: colorBarExtent(plotBottom - marginTop, colorBar)
1409
1453
  }) : null;
1410
1454
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1411
1455
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
@@ -1436,10 +1480,10 @@ function beeswarmLayout(valueRows, opts) {
1436
1480
  rows,
1437
1481
  xDomain: [min, max],
1438
1482
  xZero: toX(0),
1439
- ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1483
+ ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, labels.shapValueAxis),
1440
1484
  plotWidth,
1441
1485
  plotBottom,
1442
- colorBar: fit ? colorBarLayout(COLOR_BAR, {
1486
+ colorBar: fit ? colorBarLayout(colorBar, {
1443
1487
  x: marginLeft + plotWidth + fit.gap,
1444
1488
  y1: marginTop,
1445
1489
  y2: plotBottom
@@ -1469,7 +1513,7 @@ function percentile2(values, fraction) {
1469
1513
  const weight = position - lower;
1470
1514
  return sorted[lower] + (sorted[upper] - sorted[lower]) * weight;
1471
1515
  }
1472
- function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance") {
1516
+ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance", labels = shapLabels) {
1473
1517
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1474
1518
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1475
1519
  }
@@ -1481,7 +1525,8 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1481
1525
  importance,
1482
1526
  featureOrder,
1483
1527
  maxDisplay,
1484
- faithfulOtherRow
1528
+ faithfulOtherRow,
1529
+ labels
1485
1530
  ),
1486
1531
  rowSort,
1487
1532
  explanation.data
@@ -1545,7 +1590,7 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1545
1590
  sampleLabelColumn: explanation.sampleLabelColumn
1546
1591
  };
1547
1592
  }
1548
- function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom) {
1593
+ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom, title) {
1549
1594
  const { ticks, step } = niceTicks(
1550
1595
  -0.5,
1551
1596
  sampleCount - 0.5,
@@ -1560,7 +1605,7 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1560
1605
  })),
1561
1606
  xSpine: null,
1562
1607
  xTitle: {
1563
- text: "Instances",
1608
+ text: title,
1564
1609
  x: marginLeft + plotWidth / 2,
1565
1610
  y: plotBottom + AXIS_TITLE_DY,
1566
1611
  fontSize: TICK_LABEL_PT5
@@ -1568,15 +1613,17 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1568
1613
  };
1569
1614
  }
1570
1615
  function heatmapLayout(valueRows, opts) {
1616
+ var _a;
1571
1617
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1618
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1572
1619
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1573
1620
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1574
1621
  const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1575
- const colorBarSpec = {
1622
+ const colorBarSpec2 = {
1576
1623
  colormap: "red_white_blue",
1577
1624
  tickLabels: [ticks.labels[0], ticks.labels[1]],
1578
1625
  ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1579
- label: "SHAP value (impact on model output)",
1626
+ label: labels.shapValueAxis,
1580
1627
  labelPad: -10
1581
1628
  };
1582
1629
  const colorBarTop = FX_TOP;
@@ -1585,7 +1632,7 @@ function heatmapLayout(valueRows, opts) {
1585
1632
  available: marginRight,
1586
1633
  gapRatio: COLOR_BAR_GAP_RATIO2,
1587
1634
  minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1588
- extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec)
1635
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec2)
1589
1636
  }) : null;
1590
1637
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1591
1638
  const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
@@ -1658,11 +1705,18 @@ function heatmapLayout(valueRows, opts) {
1658
1705
  x1: marginLeft - Y_TICK_LENGTH,
1659
1706
  x2: marginLeft
1660
1707
  })),
1661
- ...heatmapXAxis(valueRows.columns.length, marginLeft, cellWidth, plotWidth, plotBottom),
1708
+ ...heatmapXAxis(
1709
+ valueRows.columns.length,
1710
+ marginLeft,
1711
+ cellWidth,
1712
+ plotWidth,
1713
+ plotBottom,
1714
+ labels.samples
1715
+ ),
1662
1716
  sampleLabelColumn: valueRows.sampleLabelColumn,
1663
1717
  plotWidth,
1664
1718
  cellWidth,
1665
- colorBar: fit ? colorBarLayout(colorBarSpec, {
1719
+ colorBar: fit ? colorBarLayout(colorBarSpec2, {
1666
1720
  x: gridRight + fit.gap,
1667
1721
  y1: colorBarTop,
1668
1722
  y2: plotBottom
@@ -1685,6 +1739,8 @@ export {
1685
1739
  formatShapValue,
1686
1740
  formatLevel,
1687
1741
  formatFeatureLabel,
1742
+ shapLabels,
1743
+ resolveLabels,
1688
1744
  collapseToDisplay,
1689
1745
  TICK_LENGTH,
1690
1746
  TICK_LABEL_DY,
@@ -1708,4 +1764,4 @@ export {
1708
1764
  heatmapRows,
1709
1765
  heatmapLayout
1710
1766
  };
1711
- //# sourceMappingURL=chunk-ZWO4CN22.js.map
1767
+ //# sourceMappingURL=chunk-4X6PMJYT.js.map