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/README.md +39 -0
- package/dist/{chunk-OXFKP5I3.js → chunk-4X6PMJYT.js} +235 -34
- package/dist/chunk-4X6PMJYT.js.map +1 -0
- package/dist/index.cjs +240 -33
- package/dist/index.cjs.map +1 -1
- package/dist/index.d.cts +104 -7
- package/dist/index.d.ts +104 -7
- package/dist/index.js +13 -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 +369 -126
- package/dist/react.cjs.map +1 -1
- package/dist/react.d.cts +33 -5
- package/dist/react.d.ts +33 -5
- package/dist/react.js +143 -92
- package/dist/react.js.map +1 -1
- package/package.json +1 -1
- package/dist/chunk-OXFKP5I3.js.map +0 -1
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:
|
|
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
|
|
374
|
-
// "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
|
|
375
|
-
text:
|
|
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(
|
|
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":
|
|
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 =
|
|
1188
|
-
|
|
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,
|
|
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:
|
|
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
|
|
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/
|
|
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
|
-
|
|
1424
|
-
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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
|
-
|
|
1494
|
-
|
|
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,
|
|
1526
|
-
/* @__PURE__ */ (0,
|
|
1527
|
-
/* @__PURE__ */ (0,
|
|
1528
|
-
/* @__PURE__ */ (0,
|
|
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
|
|
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,
|
|
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:
|
|
1848
|
+
text: title,
|
|
1649
1849
|
x: marginLeft + plotWidth / 2,
|
|
1650
1850
|
y: plotBottom + AXIS_TITLE_DY,
|
|
1651
|
-
fontSize:
|
|
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
|
|
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(
|
|
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
|
|
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
|
|
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,
|
|
1774
|
-
layout.fxAxisMarks.map((mark) => /* @__PURE__ */ (0,
|
|
1775
|
-
/* @__PURE__ */ (0,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
1831
|
-
[layout.spines.left, layout.spines.right].map((spine, index) => /* @__PURE__ */ (0,
|
|
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,
|
|
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.
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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
|
|
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,
|
|
1931
|
-
/* @__PURE__ */ (0,
|
|
1932
|
-
/* @__PURE__ */ (0,
|
|
1933
|
-
/* @__PURE__ */ (0,
|
|
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
|
|
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
|
-
|
|
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,
|
|
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:
|
|
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:
|
|
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
|
|
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,
|
|
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,
|
|
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,
|
|
2179
|
-
/* @__PURE__ */ (0,
|
|
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,
|
|
2191
|
-
/* @__PURE__ */ (0,
|
|
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,
|
|
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,
|
|
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,
|
|
2229
|
-
/* @__PURE__ */ (0,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
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,
|
|
2536
|
+
/* @__PURE__ */ (0, import_jsx_runtime6.jsx)(
|
|
2294
2537
|
"text",
|
|
2295
2538
|
{
|
|
2296
2539
|
x: arrow.valueLabel.x,
|