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 +38 -0
- package/dist/{chunk-ZWO4CN22.js → chunk-4X6PMJYT.js} +91 -35
- package/dist/chunk-4X6PMJYT.js.map +1 -0
- package/dist/index.cjs +92 -34
- package/dist/index.cjs.map +1 -1
- package/dist/index.d.cts +15 -7
- package/dist/index.d.ts +15 -7
- package/dist/index.js +5 -1
- package/dist/{format-wEav_cPe.d.cts → labels-7wofev8M.d.cts} +44 -1
- package/dist/{format-wEav_cPe.d.ts → labels-7wofev8M.d.ts} +44 -1
- package/dist/react.cjs +120 -51
- package/dist/react.cjs.map +1 -1
- package/dist/react.d.cts +29 -5
- package/dist/react.d.ts +29 -5
- package/dist/react.js +35 -18
- package/dist/react.js.map +1 -1
- package/package.json +1 -1
- package/dist/chunk-ZWO4CN22.js.map +0 -1
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:
|
|
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
|
|
420
|
-
// "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
|
|
421
|
-
text:
|
|
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(
|
|
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
|
-
|
|
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:
|
|
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:
|
|
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
|
-
|
|
1305
|
-
|
|
1306
|
-
|
|
1307
|
-
|
|
1308
|
-
|
|
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:
|
|
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,
|
|
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(
|
|
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:
|
|
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
|
|
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:
|
|
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,
|
|
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(
|
|
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(
|
|
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
|