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/dist/index.cjs CHANGED
@@ -48,8 +48,10 @@ __export(index_exports, {
48
48
  heatmapRows: () => heatmapRows,
49
49
  orderFeatures: () => orderFeatures,
50
50
  parseExplanation: () => parseExplanation,
51
+ resolveLabels: () => resolveLabels,
51
52
  sampleColormap: () => sampleColormap,
52
53
  scalarFormatterLabels: () => scalarFormatterLabels,
54
+ shapLabels: () => shapLabels,
53
55
  sortDisplayRows: () => sortDisplayRows,
54
56
  waterfallLayout: () => waterfallLayout,
55
57
  waterfallRows: () => waterfallRows
@@ -319,8 +321,34 @@ function formatFeatureLabel(name) {
319
321
  return name.replace(/_/g, " ");
320
322
  }
321
323
 
324
+ // src/core/labels.ts
325
+ var shapLabels = {
326
+ shapValue: "SHAP value",
327
+ shapValueAxis: "SHAP value (impact on model output)",
328
+ meanAbsShapValue: "mean(|SHAP value|)",
329
+ featureValue: "Feature value",
330
+ featureValueLow: "Low",
331
+ featureValueHigh: "High",
332
+ missingFeatureValue: "missing",
333
+ samples: "Instances",
334
+ sampleTotal: "\u03A3\u03C6",
335
+ sampleFallback: (sampleNumber) => `Sample ${sampleNumber}`,
336
+ baseValue: "E[f(X)]",
337
+ modelOutput: "f(x)",
338
+ otherFeatures: (count, style) => style === "sum" ? `Sum of ${count} other features` : `${count} other features`
339
+ };
340
+ function resolveLabels(labels) {
341
+ if (!labels) return shapLabels;
342
+ const resolved = { ...shapLabels };
343
+ for (const key of Object.keys(labels)) {
344
+ const value = labels[key];
345
+ if (value !== void 0) resolved[key] = value;
346
+ }
347
+ return resolved;
348
+ }
349
+
322
350
  // src/core/collapse.ts
323
- function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow) {
351
+ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow, labels = shapLabels) {
324
352
  const p = order.length;
325
353
  if (maxDisplay >= p) {
326
354
  return {
@@ -343,7 +371,7 @@ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithful
343
371
  const collapsed = order.slice(realCount);
344
372
  const collapsedValue = collapsed.reduce((sum, index) => sum + importance[index], 0);
345
373
  rows.push({
346
- label: faithfulOtherRow ? `Sum of ${collapsed.length} other features` : `${collapsed.length} other features`,
374
+ label: labels.otherFeatures(collapsed.length, faithfulOtherRow ? "sum" : "count"),
347
375
  featureIndex: null,
348
376
  value: collapsedValue,
349
377
  isOtherRow: true
@@ -410,15 +438,15 @@ var AXIS_HEIGHT = 52;
410
438
  var TICK_LABEL_PT = 11;
411
439
  var TITLE_PT = 13;
412
440
  var X_MARGIN = 0.05;
413
- function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
441
+ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
414
442
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT));
415
443
  return {
416
444
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
417
445
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
418
446
  xTitle: {
419
- // _bar.py:143-150 builds this from the Explanation's transform history:
420
- // "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
421
- text: "mean(|SHAP value|)",
447
+ // _bar.py:143-150 builds the default from the Explanation's transform
448
+ // history: "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
449
+ text: title,
422
450
  x: marginLeft + plotWidth / 2,
423
451
  y: plotBottom + AXIS_TITLE_DY,
424
452
  fontSize: TITLE_PT
@@ -426,6 +454,7 @@ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
426
454
  };
427
455
  }
428
456
  function barLayout(rows, opts) {
457
+ var _a;
429
458
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
430
459
  const plotWidth = width - marginLeft - marginRight;
431
460
  const values = rows.rows.map((r) => r.value);
@@ -466,7 +495,15 @@ function barLayout(rows, opts) {
466
495
  xZero,
467
496
  plotWidth,
468
497
  plotBottom: marginTop + rows.rows.length * rowHeight,
469
- ...barXAxis(min, max, toX, marginLeft, plotWidth, marginTop + rows.rows.length * rowHeight),
498
+ ...barXAxis(
499
+ min,
500
+ max,
501
+ toX,
502
+ marginLeft,
503
+ plotWidth,
504
+ marginTop + rows.rows.length * rowHeight,
505
+ ((_a = opts.labels) != null ? _a : shapLabels).meanAbsShapValue
506
+ ),
470
507
  zeroLine: { x: xZero, y1: marginTop, y2: marginTop + rows.rows.length * rowHeight },
471
508
  height: marginTop + rows.rows.length * rowHeight + AXIS_HEIGHT
472
509
  };
@@ -500,7 +537,7 @@ function placeValueLabel(value, startX, endX, gutterX, decimals) {
500
537
  }
501
538
  return { ...base, x: startX + VALUE_LABEL_GAP, anchor: "start", inside: false };
502
539
  }
503
- function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
540
+ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow, labels = shapLabels) {
504
541
  if (!Number.isInteger(sampleIndex) || sampleIndex < 0 || sampleIndex >= explanation.nSamples) {
505
542
  throw new RangeError(
506
543
  `sampleIndex must identify a Sample from 0 to ${explanation.nSamples - 1}, received ${sampleIndex}`
@@ -538,7 +575,8 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
538
575
  if (hasOtherRow) {
539
576
  const value = collapsed.reduce((sum, featureIndex) => sum + values[featureIndex], 0);
540
577
  rows.push({
541
- label: `${collapsed.length} other features`,
578
+ // SHAP's waterfall counts the hidden Features whatever the collapse mode.
579
+ label: labels.otherFeatures(collapsed.length, "count"),
542
580
  featureIndex: null,
543
581
  isOtherRow: true,
544
582
  value,
@@ -556,7 +594,9 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
556
594
  };
557
595
  }
558
596
  function waterfallLayout(valueRows, opts) {
597
+ var _a;
559
598
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
599
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
560
600
  const plotWidth = width - marginLeft - marginRight;
561
601
  const coordinates = [valueRows.baseValue, valueRows.modelOutput];
562
602
  for (const row of valueRows.rows) coordinates.push(row.left, row.left + row.width);
@@ -623,7 +663,7 @@ function waterfallLayout(valueRows, opts) {
623
663
  // got, which is the whole reason SHAP hides the left spine.
624
664
  y1: plotBottom - rowHeight,
625
665
  y2: plotBottom,
626
- label: `E[f(X)] = ${formatLevel(valueRows.baseValue, opts.decimals)}`
666
+ label: `${labels.baseValue} = ${formatLevel(valueRows.baseValue, opts.decimals)}`
627
667
  },
628
668
  {
629
669
  kind: "output",
@@ -632,7 +672,7 @@ function waterfallLayout(valueRows, opts) {
632
672
  // axvline(fx, 0, 1) — the full height.
633
673
  y1: marginTop,
634
674
  y2: plotBottom,
635
- label: `f(x) = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
675
+ label: `${labels.modelOutput} = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
636
676
  }
637
677
  ],
638
678
  separators: valueRows.rows.map((_, index) => ({
@@ -1301,12 +1341,14 @@ var BEESWARM_MISSING_COLOR = "#777777";
1301
1341
  var BEESWARM_ROW_HEIGHT = 0.4;
1302
1342
  var NBINS = 100;
1303
1343
  var AXIS_HEIGHT3 = 52;
1304
- var COLOR_BAR = {
1305
- colormap: "red_blue",
1306
- tickLabels: ["Low", "High"],
1307
- label: "Feature value",
1308
- labelPad: 0
1309
- };
1344
+ function colorBarSpec(labels) {
1345
+ return {
1346
+ colormap: "red_blue",
1347
+ tickLabels: [labels.featureValueLow, labels.featureValueHigh],
1348
+ label: labels.featureValue,
1349
+ labelPad: 0
1350
+ };
1351
+ }
1310
1352
  var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1311
1353
  var TICK_LABEL_PT4 = 11;
1312
1354
  var TITLE_PT2 = 13;
@@ -1372,7 +1414,7 @@ function spreadPoints(xs, rowIndex, seed) {
1372
1414
  const scale = 0.9 * (BEESWARM_ROW_HEIGHT / (maxPositiveOffset + 1));
1373
1415
  return offsets.map((offset) => rowIndex + offset * scale);
1374
1416
  }
1375
- function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance") {
1417
+ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance", labels = shapLabels) {
1376
1418
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1377
1419
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1378
1420
  }
@@ -1387,7 +1429,8 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1387
1429
  importance,
1388
1430
  order,
1389
1431
  maxDisplay,
1390
- faithfulOtherRow
1432
+ faithfulOtherRow,
1433
+ labels
1391
1434
  ),
1392
1435
  rowSort,
1393
1436
  explanation.data
@@ -1439,14 +1482,14 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1439
1482
  });
1440
1483
  return { rows, collapsedCount: display.collapsedCount };
1441
1484
  }
1442
- function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1485
+ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
1443
1486
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1444
1487
  return {
1445
1488
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1446
1489
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
1447
1490
  xTitle: {
1448
- // _labels.py:5, labels["VALUE"].
1449
- text: "SHAP value (impact on model output)",
1491
+ // The default is _labels.py:5, labels["VALUE"].
1492
+ text: title,
1450
1493
  x: marginLeft + plotWidth / 2,
1451
1494
  y: plotBottom + AXIS_TITLE_DY,
1452
1495
  fontSize: TITLE_PT2
@@ -1454,14 +1497,17 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1454
1497
  };
1455
1498
  }
1456
1499
  function beeswarmLayout(valueRows, opts) {
1500
+ var _a;
1457
1501
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1502
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1503
+ const colorBar = colorBarSpec(labels);
1458
1504
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1459
1505
  const fit = opts.colorBar ? fitColorBar({
1460
1506
  plotWidth: width - marginLeft - marginRight,
1461
1507
  available: marginRight,
1462
1508
  gapRatio: COLOR_BAR_GAP_RATIO,
1463
1509
  minGap: 0,
1464
- extent: colorBarExtent(plotBottom - marginTop, COLOR_BAR)
1510
+ extent: colorBarExtent(plotBottom - marginTop, colorBar)
1465
1511
  }) : null;
1466
1512
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1467
1513
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
@@ -1492,10 +1538,10 @@ function beeswarmLayout(valueRows, opts) {
1492
1538
  rows,
1493
1539
  xDomain: [min, max],
1494
1540
  xZero: toX(0),
1495
- ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1541
+ ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, labels.shapValueAxis),
1496
1542
  plotWidth,
1497
1543
  plotBottom,
1498
- colorBar: fit ? colorBarLayout(COLOR_BAR, {
1544
+ colorBar: fit ? colorBarLayout(colorBar, {
1499
1545
  x: marginLeft + plotWidth + fit.gap,
1500
1546
  y1: marginTop,
1501
1547
  y2: plotBottom
@@ -1525,7 +1571,7 @@ function percentile2(values, fraction) {
1525
1571
  const weight = position - lower;
1526
1572
  return sorted[lower] + (sorted[upper] - sorted[lower]) * weight;
1527
1573
  }
1528
- function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance") {
1574
+ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance", labels = shapLabels) {
1529
1575
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1530
1576
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1531
1577
  }
@@ -1537,7 +1583,8 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1537
1583
  importance,
1538
1584
  featureOrder,
1539
1585
  maxDisplay,
1540
- faithfulOtherRow
1586
+ faithfulOtherRow,
1587
+ labels
1541
1588
  ),
1542
1589
  rowSort,
1543
1590
  explanation.data
@@ -1601,7 +1648,7 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1601
1648
  sampleLabelColumn: explanation.sampleLabelColumn
1602
1649
  };
1603
1650
  }
1604
- function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom) {
1651
+ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom, title) {
1605
1652
  const { ticks, step } = niceTicks(
1606
1653
  -0.5,
1607
1654
  sampleCount - 0.5,
@@ -1616,7 +1663,7 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1616
1663
  })),
1617
1664
  xSpine: null,
1618
1665
  xTitle: {
1619
- text: "Instances",
1666
+ text: title,
1620
1667
  x: marginLeft + plotWidth / 2,
1621
1668
  y: plotBottom + AXIS_TITLE_DY,
1622
1669
  fontSize: TICK_LABEL_PT5
@@ -1624,15 +1671,17 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1624
1671
  };
1625
1672
  }
1626
1673
  function heatmapLayout(valueRows, opts) {
1674
+ var _a;
1627
1675
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1676
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1628
1677
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1629
1678
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1630
1679
  const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1631
- const colorBarSpec = {
1680
+ const colorBarSpec2 = {
1632
1681
  colormap: "red_white_blue",
1633
1682
  tickLabels: [ticks.labels[0], ticks.labels[1]],
1634
1683
  ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1635
- label: "SHAP value (impact on model output)",
1684
+ label: labels.shapValueAxis,
1636
1685
  labelPad: -10
1637
1686
  };
1638
1687
  const colorBarTop = FX_TOP;
@@ -1641,7 +1690,7 @@ function heatmapLayout(valueRows, opts) {
1641
1690
  available: marginRight,
1642
1691
  gapRatio: COLOR_BAR_GAP_RATIO2,
1643
1692
  minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1644
- extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec)
1693
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec2)
1645
1694
  }) : null;
1646
1695
  const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1647
1696
  const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
@@ -1714,11 +1763,18 @@ function heatmapLayout(valueRows, opts) {
1714
1763
  x1: marginLeft - Y_TICK_LENGTH,
1715
1764
  x2: marginLeft
1716
1765
  })),
1717
- ...heatmapXAxis(valueRows.columns.length, marginLeft, cellWidth, plotWidth, plotBottom),
1766
+ ...heatmapXAxis(
1767
+ valueRows.columns.length,
1768
+ marginLeft,
1769
+ cellWidth,
1770
+ plotWidth,
1771
+ plotBottom,
1772
+ labels.samples
1773
+ ),
1718
1774
  sampleLabelColumn: valueRows.sampleLabelColumn,
1719
1775
  plotWidth,
1720
1776
  cellWidth,
1721
- colorBar: fit ? colorBarLayout(colorBarSpec, {
1777
+ colorBar: fit ? colorBarLayout(colorBarSpec2, {
1722
1778
  x: gridRight + fit.gap,
1723
1779
  y1: colorBarTop,
1724
1780
  y2: plotBottom
@@ -1756,8 +1812,10 @@ function heatmapLayout(valueRows, opts) {
1756
1812
  heatmapRows,
1757
1813
  orderFeatures,
1758
1814
  parseExplanation,
1815
+ resolveLabels,
1759
1816
  sampleColormap,
1760
1817
  scalarFormatterLabels,
1818
+ shapLabels,
1761
1819
  sortDisplayRows,
1762
1820
  waterfallLayout,
1763
1821
  waterfallRows