shap-svg 0.2.0 → 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/react.cjs CHANGED
@@ -271,8 +271,34 @@ function formatFeatureLabel(name) {
271
271
  return name.replace(/_/g, " ");
272
272
  }
273
273
 
274
+ // src/core/labels.ts
275
+ var shapLabels = {
276
+ shapValue: "SHAP value",
277
+ shapValueAxis: "SHAP value (impact on model output)",
278
+ meanAbsShapValue: "mean(|SHAP value|)",
279
+ featureValue: "Feature value",
280
+ featureValueLow: "Low",
281
+ featureValueHigh: "High",
282
+ missingFeatureValue: "missing",
283
+ samples: "Instances",
284
+ sampleTotal: "\u03A3\u03C6",
285
+ sampleFallback: (sampleNumber) => `Sample ${sampleNumber}`,
286
+ baseValue: "E[f(X)]",
287
+ modelOutput: "f(x)",
288
+ otherFeatures: (count, style) => style === "sum" ? `Sum of ${count} other features` : `${count} other features`
289
+ };
290
+ function resolveLabels(labels) {
291
+ if (!labels) return shapLabels;
292
+ const resolved = { ...shapLabels };
293
+ for (const key of Object.keys(labels)) {
294
+ const value = labels[key];
295
+ if (value !== void 0) resolved[key] = value;
296
+ }
297
+ return resolved;
298
+ }
299
+
274
300
  // src/core/collapse.ts
275
- function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow) {
301
+ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow, labels = shapLabels) {
276
302
  const p = order.length;
277
303
  if (maxDisplay >= p) {
278
304
  return {
@@ -295,7 +321,7 @@ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithful
295
321
  const collapsed = order.slice(realCount);
296
322
  const collapsedValue = collapsed.reduce((sum, index) => sum + importance[index], 0);
297
323
  rows.push({
298
- label: faithfulOtherRow ? `Sum of ${collapsed.length} other features` : `${collapsed.length} other features`,
324
+ label: labels.otherFeatures(collapsed.length, faithfulOtherRow ? "sum" : "count"),
299
325
  featureIndex: null,
300
326
  value: collapsedValue,
301
327
  isOtherRow: true
@@ -364,15 +390,15 @@ var AXIS_HEIGHT = 52;
364
390
  var TICK_LABEL_PT = 11;
365
391
  var TITLE_PT = 13;
366
392
  var X_MARGIN = 0.05;
367
- function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
393
+ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
368
394
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT));
369
395
  return {
370
396
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
371
397
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
372
398
  xTitle: {
373
- // _bar.py:143-150 builds this from the Explanation's transform history:
374
- // "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
375
- text: "mean(|SHAP value|)",
399
+ // _bar.py:143-150 builds the default from the Explanation's transform
400
+ // history: "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
401
+ text: title,
376
402
  x: marginLeft + plotWidth / 2,
377
403
  y: plotBottom + AXIS_TITLE_DY,
378
404
  fontSize: TITLE_PT
@@ -380,6 +406,7 @@ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
380
406
  };
381
407
  }
382
408
  function barLayout(rows, opts) {
409
+ var _a;
383
410
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
384
411
  const plotWidth = width - marginLeft - marginRight;
385
412
  const values = rows.rows.map((r) => r.value);
@@ -420,7 +447,15 @@ function barLayout(rows, opts) {
420
447
  xZero,
421
448
  plotWidth,
422
449
  plotBottom: marginTop + rows.rows.length * rowHeight,
423
- ...barXAxis(min, max, toX, marginLeft, plotWidth, marginTop + rows.rows.length * rowHeight),
450
+ ...barXAxis(
451
+ min,
452
+ max,
453
+ toX,
454
+ marginLeft,
455
+ plotWidth,
456
+ marginTop + rows.rows.length * rowHeight,
457
+ ((_a = opts.labels) != null ? _a : shapLabels).meanAbsShapValue
458
+ ),
424
459
  zeroLine: { x: xZero, y1: marginTop, y2: marginTop + rows.rows.length * rowHeight },
425
460
  height: marginTop + rows.rows.length * rowHeight + AXIS_HEIGHT
426
461
  };
@@ -495,9 +530,11 @@ function ShapBar({
495
530
  classIndex = 1,
496
531
  width = 720,
497
532
  rowHeight = 26,
533
+ labels,
498
534
  onFeatureClick
499
535
  }) {
500
536
  const [hovered, setHovered] = (0, import_react.useState)(null);
537
+ const words = (0, import_react.useMemo)(() => resolveLabels(labels), [labels]);
501
538
  const layout = (0, import_react.useMemo)(() => {
502
539
  const raw = parseExplanation(explanation, { classIndex });
503
540
  const parsed = groupByGenus2 ? groupExplanationByGenus(raw) : raw;
@@ -508,23 +545,25 @@ function ShapBar({
508
545
  importance,
509
546
  order,
510
547
  maxDisplay,
511
- faithfulOtherRow
548
+ faithfulOtherRow,
549
+ words
512
550
  );
513
551
  return barLayout(rows, {
514
552
  width,
515
553
  rowHeight,
516
554
  marginLeft: 260,
517
555
  marginRight: 90,
518
- marginTop: 8
556
+ marginTop: 8,
557
+ labels: words
519
558
  });
520
- }, [groupByGenus2, explanation, maxDisplay, faithfulOtherRow, classIndex, width, rowHeight]);
559
+ }, [groupByGenus2, explanation, maxDisplay, faithfulOtherRow, classIndex, width, rowHeight, words]);
521
560
  return /* @__PURE__ */ (0, import_jsx_runtime2.jsxs)(
522
561
  "svg",
523
562
  {
524
563
  width,
525
564
  height: layout.height,
526
565
  role: "img",
527
- "aria-label": "Mean absolute SHAP value per feature",
566
+ "aria-label": `Mean absolute ${words.shapValue} per feature`,
528
567
  children: [
529
568
  /* @__PURE__ */ (0, import_jsx_runtime2.jsx)(
530
569
  "line",
@@ -1161,6 +1200,105 @@ function sampleColormap(name, t) {
1161
1200
  return `#${hexByte(red)}${hexByte(green)}${hexByte(blue)}`;
1162
1201
  }
1163
1202
 
1203
+ // src/core/colorBar.ts
1204
+ var ASPECT = 80;
1205
+ var TICK_PAD = 3.5;
1206
+ var TICK_LABEL_PT2 = 11;
1207
+ var LABEL_PT = 12;
1208
+ var LABEL_THICKNESS_EM = 1.08;
1209
+ var LABEL_BASELINE_EM = 0.84;
1210
+ var CHAR_EM = 0.6;
1211
+ var STEPS = 64;
1212
+ function textWidth(text, fontSize) {
1213
+ return text.length * fontSize * CHAR_EM;
1214
+ }
1215
+ function tickColumnWidth(spec) {
1216
+ return Math.max(...spec.tickLabels.map((label) => textWidth(label, TICK_LABEL_PT2)));
1217
+ }
1218
+ function colorBarExtent(height, spec) {
1219
+ return height / ASPECT + TICK_PAD + Math.max(0, tickColumnWidth(spec) + spec.labelPad) + LABEL_PT * LABEL_THICKNESS_EM;
1220
+ }
1221
+ function fitColorBar(opts) {
1222
+ const { plotWidth, available, gapRatio, minGap, extent } = opts;
1223
+ const gapFor = (width) => Math.max(gapRatio * width, minGap);
1224
+ if (gapFor(plotWidth) + extent <= available) {
1225
+ return { plotWidth, gap: gapFor(plotWidth) };
1226
+ }
1227
+ const byRatio = (available + plotWidth - extent) / (1 + gapRatio);
1228
+ const fitted = Math.max(
1229
+ 0,
1230
+ gapRatio * byRatio >= minGap ? byRatio : plotWidth - (minGap + extent - available)
1231
+ );
1232
+ return { plotWidth: fitted, gap: gapFor(fitted) };
1233
+ }
1234
+ function colorBarLayout(spec, position) {
1235
+ const { x, y1, y2 } = position;
1236
+ const width = (y2 - y1) / ASPECT;
1237
+ const step = (y2 - y1) / STEPS;
1238
+ const steps = Array.from({ length: STEPS }, (_, index) => ({
1239
+ y: y2 - (index + 1) * step,
1240
+ // Every band but the lowest reaches half a pixel into the one below, so
1241
+ // anti-aliasing cannot open a hairline seam between them.
1242
+ height: index === 0 ? step : step + 0.5,
1243
+ color: sampleColormap(spec.colormap, index / (STEPS - 1))
1244
+ }));
1245
+ const tickX = x + width + TICK_PAD;
1246
+ const labelLeft = tickX + tickColumnWidth(spec) + spec.labelPad;
1247
+ return {
1248
+ x,
1249
+ y1,
1250
+ y2,
1251
+ width,
1252
+ steps,
1253
+ tickX,
1254
+ tickFontSize: TICK_LABEL_PT2,
1255
+ ticks: [
1256
+ { y: y2, label: spec.tickLabels[0] },
1257
+ { y: y1, label: spec.tickLabels[1] }
1258
+ ],
1259
+ ...spec.offsetText ? { offsetText: { x: tickX, y: y1 - TICK_LABEL_PT2, text: spec.offsetText } } : {},
1260
+ label: {
1261
+ x: labelLeft + LABEL_PT * LABEL_BASELINE_EM,
1262
+ y: (y1 + y2) / 2,
1263
+ text: spec.label,
1264
+ fontSize: LABEL_PT
1265
+ },
1266
+ right: Math.max(
1267
+ tickX + tickColumnWidth(spec),
1268
+ labelLeft + LABEL_PT * LABEL_THICKNESS_EM
1269
+ )
1270
+ };
1271
+ }
1272
+ var MINUS3 = "\u2212";
1273
+ var POWER_LIMITS = [-5, 6];
1274
+ function roundTo(value, decimals) {
1275
+ const factor = 10 ** decimals;
1276
+ return Math.round(value * factor) / factor;
1277
+ }
1278
+ function scalarFormatterLabels(locs) {
1279
+ const largest = Math.max(0, ...locs.map(Math.abs));
1280
+ const magnitude = largest === 0 ? 0 : Math.floor(Math.log10(largest));
1281
+ const order = magnitude <= POWER_LIMITS[0] || magnitude >= POWER_LIMITS[1] ? magnitude : 0;
1282
+ const scaled = locs.map((value) => value / 10 ** order);
1283
+ let range = Math.max(...scaled) - Math.min(...scaled);
1284
+ if (range === 0) range = Math.max(...scaled.map(Math.abs));
1285
+ if (range === 0) range = 1;
1286
+ const rangeMagnitude = Math.floor(Math.log10(range));
1287
+ const threshold = 1e-3 * 10 ** rangeMagnitude;
1288
+ let decimals = Math.max(0, 3 - rangeMagnitude);
1289
+ while (decimals >= 0) {
1290
+ const error = Math.max(...scaled.map((value) => Math.abs(value - roundTo(value, decimals))));
1291
+ if (error < threshold) decimals -= 1;
1292
+ else break;
1293
+ }
1294
+ decimals += 1;
1295
+ const labels = scaled.map((value) => {
1296
+ const text = value.toFixed(decimals);
1297
+ return Number(text) === 0 ? text.replace("-", "") : text.replace("-", MINUS3);
1298
+ });
1299
+ return order === 0 ? { labels } : { labels, offsetText: `1e${String(order).replace("-", MINUS3)}` };
1300
+ }
1301
+
1164
1302
  // src/core/rowSort.ts
1165
1303
  function sortDisplayRows(display, sort, data) {
1166
1304
  if (sort === "importance") return display;
@@ -1184,8 +1322,17 @@ function sortDisplayRows(display, sort, data) {
1184
1322
  var BEESWARM_MISSING_COLOR = "#777777";
1185
1323
  var BEESWARM_ROW_HEIGHT = 0.4;
1186
1324
  var NBINS = 100;
1187
- var AXIS_HEIGHT2 = 74;
1188
- var TICK_LABEL_PT2 = 11;
1325
+ var AXIS_HEIGHT2 = 52;
1326
+ function colorBarSpec(labels) {
1327
+ return {
1328
+ colormap: "red_blue",
1329
+ tickLabels: [labels.featureValueLow, labels.featureValueHigh],
1330
+ label: labels.featureValue,
1331
+ labelPad: 0
1332
+ };
1333
+ }
1334
+ var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1335
+ var TICK_LABEL_PT3 = 11;
1189
1336
  var TITLE_PT2 = 13;
1190
1337
  var X_MARGIN2 = 0.05;
1191
1338
  function percentile(values, percent) {
@@ -1249,7 +1396,7 @@ function spreadPoints(xs, rowIndex, seed) {
1249
1396
  const scale = 0.9 * (BEESWARM_ROW_HEIGHT / (maxPositiveOffset + 1));
1250
1397
  return offsets.map((offset) => rowIndex + offset * scale);
1251
1398
  }
1252
- function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance") {
1399
+ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance", labels = shapLabels) {
1253
1400
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1254
1401
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1255
1402
  }
@@ -1264,7 +1411,8 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1264
1411
  importance,
1265
1412
  order,
1266
1413
  maxDisplay,
1267
- faithfulOtherRow
1414
+ faithfulOtherRow,
1415
+ labels
1268
1416
  ),
1269
1417
  rowSort,
1270
1418
  explanation.data
@@ -1316,14 +1464,14 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1316
1464
  });
1317
1465
  return { rows, collapsedCount: display.collapsedCount };
1318
1466
  }
1319
- function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1320
- const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT2));
1467
+ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
1468
+ const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT3));
1321
1469
  return {
1322
1470
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1323
1471
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
1324
1472
  xTitle: {
1325
- // _labels.py:5, labels["VALUE"].
1326
- text: "SHAP value (impact on model output)",
1473
+ // The default is _labels.py:5, labels["VALUE"].
1474
+ text: title,
1327
1475
  x: marginLeft + plotWidth / 2,
1328
1476
  y: plotBottom + AXIS_TITLE_DY,
1329
1477
  fontSize: TITLE_PT2
@@ -1331,8 +1479,19 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1331
1479
  };
1332
1480
  }
1333
1481
  function beeswarmLayout(valueRows, opts) {
1482
+ var _a;
1334
1483
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1335
- const plotWidth = width - marginLeft - marginRight;
1484
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1485
+ const colorBar = colorBarSpec(labels);
1486
+ const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1487
+ const fit = opts.colorBar ? fitColorBar({
1488
+ plotWidth: width - marginLeft - marginRight,
1489
+ available: marginRight,
1490
+ gapRatio: COLOR_BAR_GAP_RATIO,
1491
+ minGap: 0,
1492
+ extent: colorBarExtent(plotBottom - marginTop, colorBar)
1493
+ }) : null;
1494
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1336
1495
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
1337
1496
  const dataMin = Math.min(0, ...values);
1338
1497
  const dataMax = Math.max(0, ...values);
@@ -1357,20 +1516,76 @@ function beeswarmLayout(valueRows, opts) {
1357
1516
  }))
1358
1517
  };
1359
1518
  });
1360
- const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1361
1519
  return {
1362
1520
  rows,
1363
1521
  xDomain: [min, max],
1364
1522
  xZero: toX(0),
1365
- ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1523
+ ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, labels.shapValueAxis),
1366
1524
  plotWidth,
1367
1525
  plotBottom,
1526
+ colorBar: fit ? colorBarLayout(colorBar, {
1527
+ x: marginLeft + plotWidth + fit.gap,
1528
+ y1: marginTop,
1529
+ y2: plotBottom
1530
+ }) : null,
1368
1531
  height: plotBottom + AXIS_HEIGHT2
1369
1532
  };
1370
1533
  }
1371
1534
 
1372
- // src/react/ShapBeeswarm.tsx
1535
+ // src/react/ColorBar.tsx
1373
1536
  var import_jsx_runtime3 = require("react/jsx-runtime");
1537
+ function ColorBar({ bar }) {
1538
+ return /* @__PURE__ */ (0, import_jsx_runtime3.jsxs)("g", { role: "img", "aria-label": `${bar.label.text}: ${bar.ticks[0].label} to ${bar.ticks[1].label}`, children: [
1539
+ bar.steps.map((step, index) => /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1540
+ "rect",
1541
+ {
1542
+ x: bar.x,
1543
+ y: step.y,
1544
+ width: bar.width,
1545
+ height: step.height,
1546
+ fill: step.color
1547
+ },
1548
+ `colorbar-step-${index}`
1549
+ )),
1550
+ bar.ticks.map((tick, index) => /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1551
+ "text",
1552
+ {
1553
+ x: bar.tickX,
1554
+ y: tick.y,
1555
+ dominantBaseline: "middle",
1556
+ fontSize: bar.tickFontSize,
1557
+ fill: "#333333",
1558
+ children: tick.label
1559
+ },
1560
+ `colorbar-tick-${index}`
1561
+ )),
1562
+ bar.offsetText && /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1563
+ "text",
1564
+ {
1565
+ x: bar.offsetText.x,
1566
+ y: bar.offsetText.y,
1567
+ fontSize: bar.tickFontSize,
1568
+ fill: "#333333",
1569
+ children: bar.offsetText.text
1570
+ }
1571
+ ),
1572
+ /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1573
+ "text",
1574
+ {
1575
+ x: bar.label.x,
1576
+ y: bar.label.y,
1577
+ transform: `rotate(-90 ${bar.label.x} ${bar.label.y})`,
1578
+ textAnchor: "middle",
1579
+ fontSize: bar.label.fontSize,
1580
+ fill: "#333333",
1581
+ children: bar.label.text
1582
+ }
1583
+ )
1584
+ ] });
1585
+ }
1586
+
1587
+ // src/react/ShapBeeswarm.tsx
1588
+ var import_jsx_runtime4 = require("react/jsx-runtime");
1374
1589
  function ShapBeeswarm({
1375
1590
  explanation,
1376
1591
  maxDisplay = 10,
@@ -1382,25 +1597,32 @@ function ShapBeeswarm({
1382
1597
  rowHeight = 28,
1383
1598
  seed = 0,
1384
1599
  dotRadius = 3,
1600
+ colorBar = true,
1601
+ labels,
1385
1602
  onFeatureClick
1386
1603
  }) {
1387
1604
  const [hovered, setHovered] = (0, import_react2.useState)(null);
1605
+ const words = (0, import_react2.useMemo)(() => resolveLabels(labels), [labels]);
1388
1606
  const marginTop = 8;
1389
1607
  const layout = (0, import_react2.useMemo)(() => {
1390
1608
  const raw = parseExplanation(explanation, { classIndex });
1391
1609
  const parsed = groupByGenus2 ? groupExplanationByGenus(raw) : raw;
1392
- const rows = beeswarmRows(parsed, maxDisplay, faithfulOtherRow, seed, rowSort);
1610
+ const rows = beeswarmRows(parsed, maxDisplay, faithfulOtherRow, seed, rowSort, words);
1393
1611
  return beeswarmLayout(rows, {
1394
1612
  width,
1395
1613
  rowHeight,
1396
1614
  marginLeft: 260,
1397
1615
  marginRight: 90,
1398
1616
  marginTop,
1399
- dotRadius
1617
+ dotRadius,
1618
+ colorBar,
1619
+ labels: words
1400
1620
  });
1401
1621
  }, [
1402
1622
  groupByGenus2,
1403
1623
  rowSort,
1624
+ colorBar,
1625
+ words,
1404
1626
  explanation,
1405
1627
  maxDisplay,
1406
1628
  faithfulOtherRow,
@@ -1420,9 +1642,8 @@ function ShapBeeswarm({
1420
1642
  const fallbackRow = layout.rows[0];
1421
1643
  const activeRow = hovered ? layout.rows[hovered.rowIndex] : fallbackRow;
1422
1644
  const activePoint = hovered ? activeRow == null ? void 0 : activeRow.points[hovered.pointIndex] : activeRow == null ? void 0 : activeRow.points[0];
1423
- const legendSteps = 32;
1424
- return /* @__PURE__ */ (0, import_jsx_runtime3.jsxs)("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP beeswarm", children: [
1425
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1645
+ return /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP beeswarm", children: [
1646
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
1426
1647
  "line",
1427
1648
  {
1428
1649
  x1: layout.xZero,
@@ -1433,7 +1654,7 @@ function ShapBeeswarm({
1433
1654
  strokeWidth: 1
1434
1655
  }
1435
1656
  ),
1436
- layout.rows.map((row, rowIndex) => /* @__PURE__ */ (0, import_jsx_runtime3.jsxs)(
1657
+ layout.rows.map((row, rowIndex) => /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)(
1437
1658
  "g",
1438
1659
  {
1439
1660
  onMouseOver: (event) => handlePointHover(rowIndex, event),
@@ -1441,7 +1662,7 @@ function ShapBeeswarm({
1441
1662
  onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
1442
1663
  style: { cursor: onFeatureClick ? "pointer" : "default" },
1443
1664
  children: [
1444
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1665
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
1445
1666
  "rect",
1446
1667
  {
1447
1668
  x: 0,
@@ -1451,7 +1672,7 @@ function ShapBeeswarm({
1451
1672
  fill: (hovered == null ? void 0 : hovered.rowIndex) === rowIndex ? "#00000008" : "transparent"
1452
1673
  }
1453
1674
  ),
1454
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1675
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
1455
1676
  "text",
1456
1677
  {
1457
1678
  x: 250,
@@ -1464,7 +1685,7 @@ function ShapBeeswarm({
1464
1685
  children: row.label
1465
1686
  }
1466
1687
  ),
1467
- row.points.map((point, pointIndex) => /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1688
+ row.points.map((point, pointIndex) => /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
1468
1689
  "circle",
1469
1690
  {
1470
1691
  "data-point-index": pointIndex,
@@ -1480,7 +1701,7 @@ function ShapBeeswarm({
1480
1701
  },
1481
1702
  `row-${rowIndex}`
1482
1703
  )),
1483
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1704
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
1484
1705
  XAxis,
1485
1706
  {
1486
1707
  ticks: layout.xTicks,
@@ -1490,42 +1711,18 @@ function ShapBeeswarm({
1490
1711
  tickFontSize: 11
1491
1712
  }
1492
1713
  ),
1493
- activeRow && /* @__PURE__ */ (0, import_jsx_runtime3.jsxs)("g", { "aria-label": "Feature value colour scale", children: [
1494
- Array.from({ length: legendSteps }, (_, index) => /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1495
- "rect",
1496
- {
1497
- x: width - 160 + index * 4,
1498
- y: layout.height - 22,
1499
- width: 4,
1500
- height: 7,
1501
- fill: sampleColormap("red_blue", index / (legendSteps - 1))
1502
- },
1503
- `legend-${index}`
1504
- )),
1505
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)("text", { x: width - 160, y: layout.height - 3, fontSize: 10, fill: "#555555", children: formatShapValue(activeRow.vmin) }),
1506
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)(
1507
- "text",
1508
- {
1509
- x: width - 32,
1510
- y: layout.height - 3,
1511
- textAnchor: "end",
1512
- fontSize: 10,
1513
- fill: "#555555",
1514
- children: formatShapValue(activeRow.vmax)
1515
- }
1516
- )
1517
- ] }),
1518
- activeRow && activePoint && /* @__PURE__ */ (0, import_jsx_runtime3.jsxs)(
1714
+ layout.colorBar && /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(ColorBar, { bar: layout.colorBar }),
1715
+ activeRow && activePoint && /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)(
1519
1716
  "g",
1520
1717
  {
1521
1718
  opacity: hovered ? 1 : 0,
1522
1719
  pointerEvents: "none",
1523
1720
  transform: `translate(${activePoint.x + 8} ${activePoint.y - 8})`,
1524
1721
  children: [
1525
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)("rect", { x: 0, y: -16, width: 210, height: 54, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
1526
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)("text", { x: 7, y: 0, fontSize: 11, fill: "#222222", children: activeRow.label }),
1527
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: `SHAP value: ${formatShapValue(activePoint.valueX)}` }),
1528
- /* @__PURE__ */ (0, import_jsx_runtime3.jsx)("text", { x: 7, y: 30, fontSize: 11, fill: "#222222", children: `Feature value: ${Number.isFinite(activePoint.featureValue) ? formatShapValue(activePoint.featureValue) : "missing"}` })
1722
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("rect", { x: 0, y: -16, width: 210, height: 54, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
1723
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("text", { x: 7, y: 0, fontSize: 11, fill: "#222222", children: activeRow.label }),
1724
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: `${words.shapValue}: ${formatShapValue(activePoint.valueX)}` }),
1725
+ /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("text", { x: 7, y: 30, fontSize: 11, fill: "#222222", children: `${words.featureValue}: ${Number.isFinite(activePoint.featureValue) ? formatLevel(activePoint.featureValue) : words.missingFeatureValue}` })
1529
1726
  ]
1530
1727
  }
1531
1728
  )
@@ -1543,7 +1740,9 @@ var SIDE_BAR_GAP = 10;
1543
1740
  var SIDE_BAR_RIGHT_INSET = 40;
1544
1741
  var SIDE_BAR_HEIGHT_RATIO = 0.6;
1545
1742
  var AXIS_HEIGHT3 = 52;
1546
- var TICK_LABEL_PT3 = 10;
1743
+ var TICK_LABEL_PT4 = 10;
1744
+ var COLOR_BAR_GAP_RATIO2 = 0.1 / 0.89;
1745
+ var COLOR_BAR_CLEARANCE = 8;
1547
1746
  var Y_TICK_LENGTH = 5;
1548
1747
  function percentile2(values, fraction) {
1549
1748
  if (values.length === 0) return 0;
@@ -1554,7 +1753,7 @@ function percentile2(values, fraction) {
1554
1753
  const weight = position - lower;
1555
1754
  return sorted[lower] + (sorted[upper] - sorted[lower]) * weight;
1556
1755
  }
1557
- function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance") {
1756
+ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance", labels = shapLabels) {
1558
1757
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1559
1758
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1560
1759
  }
@@ -1566,7 +1765,8 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1566
1765
  importance,
1567
1766
  featureOrder,
1568
1767
  maxDisplay,
1569
- faithfulOtherRow
1768
+ faithfulOtherRow,
1769
+ labels
1570
1770
  ),
1571
1771
  rowSort,
1572
1772
  explanation.data
@@ -1630,11 +1830,11 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1630
1830
  sampleLabelColumn: explanation.sampleLabelColumn
1631
1831
  };
1632
1832
  }
1633
- function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom) {
1833
+ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom, title) {
1634
1834
  const { ticks, step } = niceTicks(
1635
1835
  -0.5,
1636
1836
  sampleCount - 0.5,
1637
- tickSpace(plotWidth, TICK_LABEL_PT3),
1837
+ tickSpace(plotWidth, TICK_LABEL_PT4),
1638
1838
  { integer: true }
1639
1839
  );
1640
1840
  return {
@@ -1645,20 +1845,38 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1645
1845
  })),
1646
1846
  xSpine: null,
1647
1847
  xTitle: {
1648
- text: "Instances",
1848
+ text: title,
1649
1849
  x: marginLeft + plotWidth / 2,
1650
1850
  y: plotBottom + AXIS_TITLE_DY,
1651
- fontSize: TICK_LABEL_PT3
1851
+ fontSize: TICK_LABEL_PT4
1652
1852
  }
1653
1853
  };
1654
1854
  }
1655
1855
  function heatmapLayout(valueRows, opts) {
1856
+ var _a;
1656
1857
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1657
- const plotWidth = width - marginLeft - marginRight;
1658
- const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1659
- const gridRight = marginLeft + plotWidth;
1858
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1660
1859
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1661
1860
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1861
+ const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1862
+ const colorBarSpec2 = {
1863
+ colormap: "red_white_blue",
1864
+ tickLabels: [ticks.labels[0], ticks.labels[1]],
1865
+ ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1866
+ label: labels.shapValueAxis,
1867
+ labelPad: -10
1868
+ };
1869
+ const colorBarTop = FX_TOP;
1870
+ const fit = opts.colorBar ? fitColorBar({
1871
+ plotWidth: width - marginLeft - marginRight,
1872
+ available: marginRight,
1873
+ gapRatio: COLOR_BAR_GAP_RATIO2,
1874
+ minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1875
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec2)
1876
+ }) : null;
1877
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1878
+ const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1879
+ const gridRight = marginLeft + plotWidth;
1662
1880
  const columns = valueRows.columns.map((column, index) => ({
1663
1881
  ...column,
1664
1882
  x: marginLeft + index * cellWidth,
@@ -1727,16 +1945,28 @@ function heatmapLayout(valueRows, opts) {
1727
1945
  x1: marginLeft - Y_TICK_LENGTH,
1728
1946
  x2: marginLeft
1729
1947
  })),
1730
- ...heatmapXAxis(valueRows.columns.length, marginLeft, cellWidth, plotWidth, plotBottom),
1948
+ ...heatmapXAxis(
1949
+ valueRows.columns.length,
1950
+ marginLeft,
1951
+ cellWidth,
1952
+ plotWidth,
1953
+ plotBottom,
1954
+ labels.samples
1955
+ ),
1731
1956
  sampleLabelColumn: valueRows.sampleLabelColumn,
1732
1957
  plotWidth,
1733
1958
  cellWidth,
1959
+ colorBar: fit ? colorBarLayout(colorBarSpec2, {
1960
+ x: gridRight + fit.gap,
1961
+ y1: colorBarTop,
1962
+ y2: plotBottom
1963
+ }) : null,
1734
1964
  height: plotBottom + AXIS_HEIGHT3
1735
1965
  };
1736
1966
  }
1737
1967
 
1738
1968
  // src/react/ShapHeatmap.tsx
1739
- var import_jsx_runtime4 = require("react/jsx-runtime");
1969
+ var import_jsx_runtime5 = require("react/jsx-runtime");
1740
1970
  function ShapHeatmap({
1741
1971
  explanation,
1742
1972
  maxDisplay = 10,
@@ -1746,33 +1976,38 @@ function ShapHeatmap({
1746
1976
  classIndex = 1,
1747
1977
  width = 720,
1748
1978
  rowHeight = 26,
1979
+ colorBar = true,
1980
+ labels,
1749
1981
  onFeatureClick,
1750
1982
  onSampleClick
1751
1983
  }) {
1752
1984
  const [hoveredColumn, setHoveredColumn] = (0, import_react3.useState)(null);
1985
+ const words = (0, import_react3.useMemo)(() => resolveLabels(labels), [labels]);
1753
1986
  const marginTop = 72;
1754
1987
  const layout = (0, import_react3.useMemo)(() => {
1755
1988
  const raw = parseExplanation(explanation, { classIndex });
1756
1989
  const parsed = groupByGenus2 ? groupExplanationByGenus(raw) : raw;
1757
- const rows = heatmapRows(parsed, maxDisplay, faithfulOtherRow, rowSort);
1990
+ const rows = heatmapRows(parsed, maxDisplay, faithfulOtherRow, rowSort, words);
1758
1991
  return heatmapLayout(rows, {
1759
1992
  width,
1760
1993
  rowHeight,
1761
1994
  marginLeft: 260,
1762
1995
  marginRight: 100,
1763
- marginTop
1996
+ marginTop,
1997
+ colorBar,
1998
+ labels: words
1764
1999
  });
1765
- }, [groupByGenus2, rowSort, explanation, maxDisplay, faithfulOtherRow, classIndex, width, rowHeight]);
2000
+ }, [groupByGenus2, rowSort, colorBar, words, explanation, maxDisplay, faithfulOtherRow, classIndex, width, rowHeight]);
1766
2001
  const activeColumn = hoveredColumn === null ? void 0 : layout.columns[hoveredColumn];
1767
2002
  const nameOf = (column) => {
1768
- if (!column.sampleLabel) return `Sample ${column.sampleIndex + 1}`;
2003
+ if (!column.sampleLabel) return words.sampleFallback(column.sampleIndex + 1);
1769
2004
  return layout.sampleLabelColumn ? `${layout.sampleLabelColumn}: ${column.sampleLabel}` : column.sampleLabel;
1770
2005
  };
1771
2006
  const tooltipWidth = activeColumn ? Math.max(190, nameOf(activeColumn).length * 6.5 + 16) : 190;
1772
2007
  const tooltipX = activeColumn ? Math.max(0, Math.min(activeColumn.centerX + 8, width - tooltipWidth - 8)) : 0;
1773
- return /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP heatmap", children: [
1774
- layout.fxAxisMarks.map((mark) => /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)("g", { children: [
1775
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2008
+ return /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP heatmap", children: [
2009
+ layout.fxAxisMarks.map((mark) => /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { children: [
2010
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1776
2011
  "line",
1777
2012
  {
1778
2013
  x1: layout.gridLeft - 4,
@@ -1783,7 +2018,7 @@ function ShapHeatmap({
1783
2018
  strokeWidth: 1
1784
2019
  }
1785
2020
  ),
1786
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2021
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1787
2022
  "text",
1788
2023
  {
1789
2024
  x: layout.gridLeft - 8,
@@ -1796,7 +2031,7 @@ function ShapHeatmap({
1796
2031
  }
1797
2032
  )
1798
2033
  ] }, `fx-axis-${mark.value}`)),
1799
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2034
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1800
2035
  "polyline",
1801
2036
  {
1802
2037
  points: layout.fxLine.map((point) => `${point.x},${point.y}`).join(" "),
@@ -1805,7 +2040,7 @@ function ShapHeatmap({
1805
2040
  strokeWidth: 1.5
1806
2041
  }
1807
2042
  ),
1808
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2043
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1809
2044
  "line",
1810
2045
  {
1811
2046
  x1: layout.gridLeft,
@@ -1817,7 +2052,7 @@ function ShapHeatmap({
1817
2052
  strokeDasharray: "4 4"
1818
2053
  }
1819
2054
  ),
1820
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2055
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1821
2056
  XAxis,
1822
2057
  {
1823
2058
  ticks: layout.xTicks,
@@ -1827,8 +2062,8 @@ function ShapHeatmap({
1827
2062
  tickFontSize: 10
1828
2063
  }
1829
2064
  ),
1830
- /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)("g", { "aria-hidden": "true", children: [
1831
- [layout.spines.left, layout.spines.right].map((spine, index) => /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2065
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { "aria-hidden": "true", children: [
2066
+ [layout.spines.left, layout.spines.right].map((spine, index) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1832
2067
  "line",
1833
2068
  {
1834
2069
  x1: spine.x,
@@ -1840,7 +2075,7 @@ function ShapHeatmap({
1840
2075
  },
1841
2076
  `spine-${index}`
1842
2077
  )),
1843
- layout.yTicks.map((tick, index) => /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2078
+ layout.yTicks.map((tick, index) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1844
2079
  "line",
1845
2080
  {
1846
2081
  x1: tick.x1,
@@ -1853,13 +2088,14 @@ function ShapHeatmap({
1853
2088
  `ytick-${index}`
1854
2089
  ))
1855
2090
  ] }),
1856
- layout.rows.map((row, rowIndex) => /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)(
2091
+ layout.colorBar && /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(ColorBar, { bar: layout.colorBar }),
2092
+ layout.rows.map((row, rowIndex) => /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)(
1857
2093
  "g",
1858
2094
  {
1859
2095
  onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
1860
2096
  style: { cursor: onFeatureClick ? "pointer" : "default" },
1861
2097
  children: [
1862
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2098
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1863
2099
  "text",
1864
2100
  {
1865
2101
  x: 250,
@@ -1872,7 +2108,7 @@ function ShapHeatmap({
1872
2108
  children: row.label
1873
2109
  }
1874
2110
  ),
1875
- row.cells.map((cell, columnIndex) => /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2111
+ row.cells.map((cell, columnIndex) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1876
2112
  "rect",
1877
2113
  {
1878
2114
  x: cell.x,
@@ -1883,7 +2119,7 @@ function ShapHeatmap({
1883
2119
  },
1884
2120
  `cell-${columnIndex}`
1885
2121
  )),
1886
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2122
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1887
2123
  "rect",
1888
2124
  {
1889
2125
  x: row.sideBar.x,
@@ -1897,7 +2133,7 @@ function ShapHeatmap({
1897
2133
  },
1898
2134
  `row-${rowIndex}`
1899
2135
  )),
1900
- activeColumn && /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2136
+ activeColumn && /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1901
2137
  "rect",
1902
2138
  {
1903
2139
  x: activeColumn.x,
@@ -1909,7 +2145,7 @@ function ShapHeatmap({
1909
2145
  pointerEvents: "none"
1910
2146
  }
1911
2147
  ),
1912
- layout.columns.map((column, columnIndex) => /* @__PURE__ */ (0, import_jsx_runtime4.jsx)(
2148
+ layout.columns.map((column, columnIndex) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
1913
2149
  "rect",
1914
2150
  {
1915
2151
  x: column.x,
@@ -1917,7 +2153,7 @@ function ShapHeatmap({
1917
2153
  width: column.width,
1918
2154
  height: layout.plotBottom - 8,
1919
2155
  fill: "transparent",
1920
- "aria-label": `${nameOf(column)}, total SHAP value ${formatShapValue(column.total)}`,
2156
+ "aria-label": `${nameOf(column)}, total ${words.shapValue} ${formatShapValue(column.total)}`,
1921
2157
  onMouseEnter: () => setHoveredColumn(columnIndex),
1922
2158
  onMouseLeave: () => setHoveredColumn(null),
1923
2159
  onClick: () => {
@@ -1927,10 +2163,10 @@ function ShapHeatmap({
1927
2163
  },
1928
2164
  `column-hit-${column.sampleIndex}`
1929
2165
  )),
1930
- activeColumn && /* @__PURE__ */ (0, import_jsx_runtime4.jsxs)("g", { pointerEvents: "none", transform: `translate(${tooltipX} 10)`, children: [
1931
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("rect", { x: 0, y: 0, width: tooltipWidth, height: 42, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
1932
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: nameOf(activeColumn) }),
1933
- /* @__PURE__ */ (0, import_jsx_runtime4.jsx)("text", { x: 7, y: 31, fontSize: 11, fill: "#222222", children: `\u03A3\u03C6: ${formatShapValue(activeColumn.total)}` })
2166
+ activeColumn && /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { pointerEvents: "none", transform: `translate(${tooltipX} 10)`, children: [
2167
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)("rect", { x: 0, y: 0, width: tooltipWidth, height: 42, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
2168
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: nameOf(activeColumn) }),
2169
+ /* @__PURE__ */ (0, import_jsx_runtime5.jsx)("text", { x: 7, y: 31, fontSize: 11, fill: "#222222", children: `${words.sampleTotal}: ${formatShapValue(activeColumn.total)}` })
1934
2170
  ] })
1935
2171
  ] });
1936
2172
  }
@@ -1944,7 +2180,7 @@ var BAR_THICKNESS_RATIO2 = 0.8;
1944
2180
  var AXIS_HEIGHT4 = 52;
1945
2181
  var WATERFALL_TICK_LABEL_DY = 18;
1946
2182
  var WATERFALL_BASE_LABEL_DY = 36;
1947
- var TICK_LABEL_PT4 = 13;
2183
+ var TICK_LABEL_PT5 = 13;
1948
2184
  var colorFor = (value) => value < 0 ? NEGATIVE_COLOR : POSITIVE_COLOR;
1949
2185
  var VALUE_LABEL_GAP = 6;
1950
2186
  var VALUE_LABEL_FONT_SIZE = 12;
@@ -1966,7 +2202,7 @@ function placeValueLabel(value, startX, endX, gutterX, decimals) {
1966
2202
  }
1967
2203
  return { ...base, x: startX + VALUE_LABEL_GAP, anchor: "start", inside: false };
1968
2204
  }
1969
- function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
2205
+ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow, labels = shapLabels) {
1970
2206
  if (!Number.isInteger(sampleIndex) || sampleIndex < 0 || sampleIndex >= explanation.nSamples) {
1971
2207
  throw new RangeError(
1972
2208
  `sampleIndex must identify a Sample from 0 to ${explanation.nSamples - 1}, received ${sampleIndex}`
@@ -2004,7 +2240,8 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
2004
2240
  if (hasOtherRow) {
2005
2241
  const value = collapsed.reduce((sum, featureIndex) => sum + values[featureIndex], 0);
2006
2242
  rows.push({
2007
- label: `${collapsed.length} other features`,
2243
+ // SHAP's waterfall counts the hidden Features whatever the collapse mode.
2244
+ label: labels.otherFeatures(collapsed.length, "count"),
2008
2245
  featureIndex: null,
2009
2246
  isOtherRow: true,
2010
2247
  value,
@@ -2022,7 +2259,9 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
2022
2259
  };
2023
2260
  }
2024
2261
  function waterfallLayout(valueRows, opts) {
2262
+ var _a;
2025
2263
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
2264
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
2026
2265
  const plotWidth = width - marginLeft - marginRight;
2027
2266
  const coordinates = [valueRows.baseValue, valueRows.modelOutput];
2028
2267
  for (const row of valueRows.rows) coordinates.push(row.left, row.left + row.width);
@@ -2059,7 +2298,7 @@ function waterfallLayout(valueRows, opts) {
2059
2298
  };
2060
2299
  });
2061
2300
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
2062
- const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
2301
+ const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT5));
2063
2302
  const individualCount = valueRows.rows.filter((row) => !row.isOtherRow).length;
2064
2303
  const hasOtherRow = individualCount < valueRows.rows.length;
2065
2304
  const connectors = [];
@@ -2089,7 +2328,7 @@ function waterfallLayout(valueRows, opts) {
2089
2328
  // got, which is the whole reason SHAP hides the left spine.
2090
2329
  y1: plotBottom - rowHeight,
2091
2330
  y2: plotBottom,
2092
- label: `E[f(X)] = ${formatLevel(valueRows.baseValue, opts.decimals)}`
2331
+ label: `${labels.baseValue} = ${formatLevel(valueRows.baseValue, opts.decimals)}`
2093
2332
  },
2094
2333
  {
2095
2334
  kind: "output",
@@ -2098,7 +2337,7 @@ function waterfallLayout(valueRows, opts) {
2098
2337
  // axvline(fx, 0, 1) — the full height.
2099
2338
  y1: marginTop,
2100
2339
  y2: plotBottom,
2101
- label: `f(x) = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
2340
+ label: `${labels.modelOutput} = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
2102
2341
  }
2103
2342
  ],
2104
2343
  separators: valueRows.rows.map((_, index) => ({
@@ -2116,7 +2355,7 @@ function waterfallLayout(valueRows, opts) {
2116
2355
  }
2117
2356
 
2118
2357
  // src/react/ShapWaterfall.tsx
2119
- var import_jsx_runtime5 = require("react/jsx-runtime");
2358
+ var import_jsx_runtime6 = require("react/jsx-runtime");
2120
2359
  function ShapWaterfall({
2121
2360
  explanation,
2122
2361
  sampleIndex = 0,
@@ -2127,21 +2366,24 @@ function ShapWaterfall({
2127
2366
  width = 720,
2128
2367
  rowHeight = 30,
2129
2368
  decimals = 2,
2369
+ labels,
2130
2370
  onFeatureClick
2131
2371
  }) {
2132
2372
  const [hovered, setHovered] = (0, import_react4.useState)(null);
2373
+ const words = (0, import_react4.useMemo)(() => resolveLabels(labels), [labels]);
2133
2374
  const marginTop = 34;
2134
2375
  const layout = (0, import_react4.useMemo)(() => {
2135
2376
  const raw = parseExplanation(explanation, { classIndex });
2136
2377
  const parsed = groupByGenus2 ? groupExplanationByGenus(raw) : raw;
2137
- const rows = waterfallRows(parsed, sampleIndex, maxDisplay, faithfulOtherRow);
2378
+ const rows = waterfallRows(parsed, sampleIndex, maxDisplay, faithfulOtherRow, words);
2138
2379
  return waterfallLayout(rows, {
2139
2380
  width,
2140
2381
  rowHeight,
2141
2382
  marginLeft: 260,
2142
2383
  marginRight: 110,
2143
2384
  marginTop,
2144
- decimals
2385
+ decimals,
2386
+ labels: words
2145
2387
  });
2146
2388
  }, [
2147
2389
  groupByGenus2,
@@ -2152,9 +2394,10 @@ function ShapWaterfall({
2152
2394
  classIndex,
2153
2395
  width,
2154
2396
  rowHeight,
2155
- decimals
2397
+ decimals,
2398
+ words
2156
2399
  ]);
2157
- return /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)(
2400
+ return /* @__PURE__ */ (0, import_jsx_runtime6.jsxs)(
2158
2401
  "svg",
2159
2402
  {
2160
2403
  width,
@@ -2162,7 +2405,7 @@ function ShapWaterfall({
2162
2405
  role: "img",
2163
2406
  "aria-label": `Local SHAP waterfall for Sample ${sampleIndex}`,
2164
2407
  children: [
2165
- layout.separators.map((separator, index) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2408
+ layout.separators.map((separator, index) => /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2166
2409
  "line",
2167
2410
  {
2168
2411
  x1: separator.x1,
@@ -2175,8 +2418,8 @@ function ShapWaterfall({
2175
2418
  },
2176
2419
  `separator-${index}`
2177
2420
  )),
2178
- /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { "aria-hidden": "true", children: [
2179
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2421
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsxs)("g", { "aria-hidden": "true", children: [
2422
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2180
2423
  "line",
2181
2424
  {
2182
2425
  x1: layout.plotLeft,
@@ -2187,8 +2430,8 @@ function ShapWaterfall({
2187
2430
  strokeWidth: 1
2188
2431
  }
2189
2432
  ),
2190
- layout.xTicks.map((tick) => /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { children: [
2191
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2433
+ layout.xTicks.map((tick) => /* @__PURE__ */ (0, import_jsx_runtime6.jsxs)("g", { children: [
2434
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2192
2435
  "line",
2193
2436
  {
2194
2437
  x1: tick.x,
@@ -2199,7 +2442,7 @@ function ShapWaterfall({
2199
2442
  strokeWidth: 1
2200
2443
  }
2201
2444
  ),
2202
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2445
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2203
2446
  "text",
2204
2447
  {
2205
2448
  x: tick.x,
@@ -2212,7 +2455,7 @@ function ShapWaterfall({
2212
2455
  )
2213
2456
  ] }, `tick-${tick.value}`))
2214
2457
  ] }),
2215
- layout.connectors.map((connector, index) => /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2458
+ layout.connectors.map((connector, index) => /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2216
2459
  "line",
2217
2460
  {
2218
2461
  x1: connector.x,
@@ -2225,8 +2468,8 @@ function ShapWaterfall({
2225
2468
  },
2226
2469
  `connector-${index}`
2227
2470
  )),
2228
- layout.axisMarks.map((mark) => /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)("g", { children: [
2229
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2471
+ layout.axisMarks.map((mark) => /* @__PURE__ */ (0, import_jsx_runtime6.jsxs)("g", { children: [
2472
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2230
2473
  "line",
2231
2474
  {
2232
2475
  x1: mark.x,
@@ -2238,7 +2481,7 @@ function ShapWaterfall({
2238
2481
  strokeDasharray: "4 4"
2239
2482
  }
2240
2483
  ),
2241
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2484
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2242
2485
  "text",
2243
2486
  {
2244
2487
  x: mark.x,
@@ -2250,7 +2493,7 @@ function ShapWaterfall({
2250
2493
  }
2251
2494
  )
2252
2495
  ] }, mark.kind)),
2253
- layout.arrows.map((arrow, index) => /* @__PURE__ */ (0, import_jsx_runtime5.jsxs)(
2496
+ layout.arrows.map((arrow, index) => /* @__PURE__ */ (0, import_jsx_runtime6.jsxs)(
2254
2497
  "g",
2255
2498
  {
2256
2499
  onMouseEnter: () => setHovered(index),
@@ -2258,7 +2501,7 @@ function ShapWaterfall({
2258
2501
  onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(arrow.featureIndex),
2259
2502
  style: { cursor: onFeatureClick ? "pointer" : "default" },
2260
2503
  children: [
2261
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2504
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2262
2505
  "rect",
2263
2506
  {
2264
2507
  x: 0,
@@ -2268,7 +2511,7 @@ function ShapWaterfall({
2268
2511
  fill: hovered === index ? "#00000008" : "transparent"
2269
2512
  }
2270
2513
  ),
2271
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2514
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2272
2515
  "text",
2273
2516
  {
2274
2517
  x: 250,
@@ -2281,7 +2524,7 @@ function ShapWaterfall({
2281
2524
  children: arrow.label
2282
2525
  }
2283
2526
  ),
2284
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2527
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2285
2528
  "polygon",
2286
2529
  {
2287
2530
  points: arrow.points.map((point) => `${point.x},${point.y}`).join(" "),
@@ -2290,7 +2533,7 @@ function ShapWaterfall({
2290
2533
  strokeWidth: 1
2291
2534
  }
2292
2535
  ),
2293
- /* @__PURE__ */ (0, import_jsx_runtime5.jsx)(
2536
+ /* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
2294
2537
  "text",
2295
2538
  {
2296
2539
  x: arrow.valueLabel.x,