shap-svg 0.2.0 → 0.2.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/dist/index.cjs CHANGED
@@ -34,6 +34,9 @@ __export(index_exports, {
34
34
  beeswarmLayout: () => beeswarmLayout,
35
35
  beeswarmRows: () => beeswarmRows,
36
36
  collapseToDisplay: () => collapseToDisplay,
37
+ colorBarExtent: () => colorBarExtent,
38
+ colorBarLayout: () => colorBarLayout,
39
+ fitColorBar: () => fitColorBar,
37
40
  formatFeatureLabel: () => formatFeatureLabel,
38
41
  formatLevel: () => formatLevel,
39
42
  formatShapValue: () => formatShapValue,
@@ -46,6 +49,7 @@ __export(index_exports, {
46
49
  orderFeatures: () => orderFeatures,
47
50
  parseExplanation: () => parseExplanation,
48
51
  sampleColormap: () => sampleColormap,
52
+ scalarFormatterLabels: () => scalarFormatterLabels,
49
53
  sortDisplayRows: () => sortDisplayRows,
50
54
  waterfallLayout: () => waterfallLayout,
51
55
  waterfallRows: () => waterfallRows
@@ -1193,12 +1197,118 @@ function sampleColormap(name, t) {
1193
1197
  return `#${hexByte(red)}${hexByte(green)}${hexByte(blue)}`;
1194
1198
  }
1195
1199
 
1200
+ // src/core/colorBar.ts
1201
+ var ASPECT = 80;
1202
+ var TICK_PAD = 3.5;
1203
+ var TICK_LABEL_PT3 = 11;
1204
+ var LABEL_PT = 12;
1205
+ var LABEL_THICKNESS_EM = 1.08;
1206
+ var LABEL_BASELINE_EM = 0.84;
1207
+ var CHAR_EM = 0.6;
1208
+ var STEPS = 64;
1209
+ function textWidth(text, fontSize) {
1210
+ return text.length * fontSize * CHAR_EM;
1211
+ }
1212
+ function tickColumnWidth(spec) {
1213
+ return Math.max(...spec.tickLabels.map((label) => textWidth(label, TICK_LABEL_PT3)));
1214
+ }
1215
+ function colorBarExtent(height, spec) {
1216
+ return height / ASPECT + TICK_PAD + Math.max(0, tickColumnWidth(spec) + spec.labelPad) + LABEL_PT * LABEL_THICKNESS_EM;
1217
+ }
1218
+ function fitColorBar(opts) {
1219
+ const { plotWidth, available, gapRatio, minGap, extent } = opts;
1220
+ const gapFor = (width) => Math.max(gapRatio * width, minGap);
1221
+ if (gapFor(plotWidth) + extent <= available) {
1222
+ return { plotWidth, gap: gapFor(plotWidth) };
1223
+ }
1224
+ const byRatio = (available + plotWidth - extent) / (1 + gapRatio);
1225
+ const fitted = Math.max(
1226
+ 0,
1227
+ gapRatio * byRatio >= minGap ? byRatio : plotWidth - (minGap + extent - available)
1228
+ );
1229
+ return { plotWidth: fitted, gap: gapFor(fitted) };
1230
+ }
1231
+ function colorBarLayout(spec, position) {
1232
+ const { x, y1, y2 } = position;
1233
+ const width = (y2 - y1) / ASPECT;
1234
+ const step = (y2 - y1) / STEPS;
1235
+ const steps = Array.from({ length: STEPS }, (_, index) => ({
1236
+ y: y2 - (index + 1) * step,
1237
+ // Every band but the lowest reaches half a pixel into the one below, so
1238
+ // anti-aliasing cannot open a hairline seam between them.
1239
+ height: index === 0 ? step : step + 0.5,
1240
+ color: sampleColormap(spec.colormap, index / (STEPS - 1))
1241
+ }));
1242
+ const tickX = x + width + TICK_PAD;
1243
+ const labelLeft = tickX + tickColumnWidth(spec) + spec.labelPad;
1244
+ return {
1245
+ x,
1246
+ y1,
1247
+ y2,
1248
+ width,
1249
+ steps,
1250
+ tickX,
1251
+ tickFontSize: TICK_LABEL_PT3,
1252
+ ticks: [
1253
+ { y: y2, label: spec.tickLabels[0] },
1254
+ { y: y1, label: spec.tickLabels[1] }
1255
+ ],
1256
+ ...spec.offsetText ? { offsetText: { x: tickX, y: y1 - TICK_LABEL_PT3, text: spec.offsetText } } : {},
1257
+ label: {
1258
+ x: labelLeft + LABEL_PT * LABEL_BASELINE_EM,
1259
+ y: (y1 + y2) / 2,
1260
+ text: spec.label,
1261
+ fontSize: LABEL_PT
1262
+ },
1263
+ right: Math.max(
1264
+ tickX + tickColumnWidth(spec),
1265
+ labelLeft + LABEL_PT * LABEL_THICKNESS_EM
1266
+ )
1267
+ };
1268
+ }
1269
+ var MINUS3 = "\u2212";
1270
+ var POWER_LIMITS = [-5, 6];
1271
+ function roundTo(value, decimals) {
1272
+ const factor = 10 ** decimals;
1273
+ return Math.round(value * factor) / factor;
1274
+ }
1275
+ function scalarFormatterLabels(locs) {
1276
+ const largest = Math.max(0, ...locs.map(Math.abs));
1277
+ const magnitude = largest === 0 ? 0 : Math.floor(Math.log10(largest));
1278
+ const order = magnitude <= POWER_LIMITS[0] || magnitude >= POWER_LIMITS[1] ? magnitude : 0;
1279
+ const scaled = locs.map((value) => value / 10 ** order);
1280
+ let range = Math.max(...scaled) - Math.min(...scaled);
1281
+ if (range === 0) range = Math.max(...scaled.map(Math.abs));
1282
+ if (range === 0) range = 1;
1283
+ const rangeMagnitude = Math.floor(Math.log10(range));
1284
+ const threshold = 1e-3 * 10 ** rangeMagnitude;
1285
+ let decimals = Math.max(0, 3 - rangeMagnitude);
1286
+ while (decimals >= 0) {
1287
+ const error = Math.max(...scaled.map((value) => Math.abs(value - roundTo(value, decimals))));
1288
+ if (error < threshold) decimals -= 1;
1289
+ else break;
1290
+ }
1291
+ decimals += 1;
1292
+ const labels = scaled.map((value) => {
1293
+ const text = value.toFixed(decimals);
1294
+ return Number(text) === 0 ? text.replace("-", "") : text.replace("-", MINUS3);
1295
+ });
1296
+ return order === 0 ? { labels } : { labels, offsetText: `1e${String(order).replace("-", MINUS3)}` };
1297
+ }
1298
+
1196
1299
  // src/core/beeswarmLayout.ts
1197
1300
  var BEESWARM_MISSING_COLOR = "#777777";
1198
1301
  var BEESWARM_ROW_HEIGHT = 0.4;
1199
1302
  var NBINS = 100;
1200
- var AXIS_HEIGHT3 = 74;
1201
- var TICK_LABEL_PT3 = 11;
1303
+ var AXIS_HEIGHT3 = 52;
1304
+ var COLOR_BAR = {
1305
+ colormap: "red_blue",
1306
+ tickLabels: ["Low", "High"],
1307
+ label: "Feature value",
1308
+ labelPad: 0
1309
+ };
1310
+ var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1311
+ var TICK_LABEL_PT4 = 11;
1202
1312
  var TITLE_PT2 = 13;
1203
1313
  var X_MARGIN2 = 0.05;
1204
1314
  function percentile(values, percent) {
@@ -1330,7 +1440,7 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1330
1440
  return { rows, collapsedCount: display.collapsedCount };
1331
1441
  }
1332
1442
  function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1333
- const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT3));
1443
+ const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1334
1444
  return {
1335
1445
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1336
1446
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
@@ -1345,7 +1455,15 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1345
1455
  }
1346
1456
  function beeswarmLayout(valueRows, opts) {
1347
1457
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1348
- const plotWidth = width - marginLeft - marginRight;
1458
+ const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1459
+ const fit = opts.colorBar ? fitColorBar({
1460
+ plotWidth: width - marginLeft - marginRight,
1461
+ available: marginRight,
1462
+ gapRatio: COLOR_BAR_GAP_RATIO,
1463
+ minGap: 0,
1464
+ extent: colorBarExtent(plotBottom - marginTop, COLOR_BAR)
1465
+ }) : null;
1466
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1349
1467
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
1350
1468
  const dataMin = Math.min(0, ...values);
1351
1469
  const dataMax = Math.max(0, ...values);
@@ -1370,7 +1488,6 @@ function beeswarmLayout(valueRows, opts) {
1370
1488
  }))
1371
1489
  };
1372
1490
  });
1373
- const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1374
1491
  return {
1375
1492
  rows,
1376
1493
  xDomain: [min, max],
@@ -1378,6 +1495,11 @@ function beeswarmLayout(valueRows, opts) {
1378
1495
  ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1379
1496
  plotWidth,
1380
1497
  plotBottom,
1498
+ colorBar: fit ? colorBarLayout(COLOR_BAR, {
1499
+ x: marginLeft + plotWidth + fit.gap,
1500
+ y1: marginTop,
1501
+ y2: plotBottom
1502
+ }) : null,
1381
1503
  height: plotBottom + AXIS_HEIGHT3
1382
1504
  };
1383
1505
  }
@@ -1390,7 +1512,9 @@ var SIDE_BAR_GAP = 10;
1390
1512
  var SIDE_BAR_RIGHT_INSET = 40;
1391
1513
  var SIDE_BAR_HEIGHT_RATIO = 0.6;
1392
1514
  var AXIS_HEIGHT4 = 52;
1393
- var TICK_LABEL_PT4 = 10;
1515
+ var TICK_LABEL_PT5 = 10;
1516
+ var COLOR_BAR_GAP_RATIO2 = 0.1 / 0.89;
1517
+ var COLOR_BAR_CLEARANCE = 8;
1394
1518
  var Y_TICK_LENGTH = 5;
1395
1519
  function percentile2(values, fraction) {
1396
1520
  if (values.length === 0) return 0;
@@ -1481,7 +1605,7 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1481
1605
  const { ticks, step } = niceTicks(
1482
1606
  -0.5,
1483
1607
  sampleCount - 0.5,
1484
- tickSpace(plotWidth, TICK_LABEL_PT4),
1608
+ tickSpace(plotWidth, TICK_LABEL_PT5),
1485
1609
  { integer: true }
1486
1610
  );
1487
1611
  return {
@@ -1495,17 +1619,33 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1495
1619
  text: "Instances",
1496
1620
  x: marginLeft + plotWidth / 2,
1497
1621
  y: plotBottom + AXIS_TITLE_DY,
1498
- fontSize: TICK_LABEL_PT4
1622
+ fontSize: TICK_LABEL_PT5
1499
1623
  }
1500
1624
  };
1501
1625
  }
1502
1626
  function heatmapLayout(valueRows, opts) {
1503
1627
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1504
- const plotWidth = width - marginLeft - marginRight;
1505
- const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1506
- const gridRight = marginLeft + plotWidth;
1507
1628
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1508
1629
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1630
+ const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1631
+ const colorBarSpec = {
1632
+ colormap: "red_white_blue",
1633
+ tickLabels: [ticks.labels[0], ticks.labels[1]],
1634
+ ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1635
+ label: "SHAP value (impact on model output)",
1636
+ labelPad: -10
1637
+ };
1638
+ const colorBarTop = FX_TOP;
1639
+ const fit = opts.colorBar ? fitColorBar({
1640
+ plotWidth: width - marginLeft - marginRight,
1641
+ available: marginRight,
1642
+ gapRatio: COLOR_BAR_GAP_RATIO2,
1643
+ minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1644
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec)
1645
+ }) : null;
1646
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1647
+ const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1648
+ const gridRight = marginLeft + plotWidth;
1509
1649
  const columns = valueRows.columns.map((column, index) => ({
1510
1650
  ...column,
1511
1651
  x: marginLeft + index * cellWidth,
@@ -1578,6 +1718,11 @@ function heatmapLayout(valueRows, opts) {
1578
1718
  sampleLabelColumn: valueRows.sampleLabelColumn,
1579
1719
  plotWidth,
1580
1720
  cellWidth,
1721
+ colorBar: fit ? colorBarLayout(colorBarSpec, {
1722
+ x: gridRight + fit.gap,
1723
+ y1: colorBarTop,
1724
+ y2: plotBottom
1725
+ }) : null,
1581
1726
  height: plotBottom + AXIS_HEIGHT4
1582
1727
  };
1583
1728
  }
@@ -1597,6 +1742,9 @@ function heatmapLayout(valueRows, opts) {
1597
1742
  beeswarmLayout,
1598
1743
  beeswarmRows,
1599
1744
  collapseToDisplay,
1745
+ colorBarExtent,
1746
+ colorBarLayout,
1747
+ fitColorBar,
1600
1748
  formatFeatureLabel,
1601
1749
  formatLevel,
1602
1750
  formatShapValue,
@@ -1609,6 +1757,7 @@ function heatmapLayout(valueRows, opts) {
1609
1757
  orderFeatures,
1610
1758
  parseExplanation,
1611
1759
  sampleColormap,
1760
+ scalarFormatterLabels,
1612
1761
  sortDisplayRows,
1613
1762
  waterfallLayout,
1614
1763
  waterfallRows