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 CHANGED
@@ -172,6 +172,7 @@ Shared by all four charts:
172
172
  | `classIndex` | `1` | which output to draw for multi-output explanations |
173
173
  | `width` | `720` | SVG width in pixels |
174
174
  | `rowHeight` | per chart | pixels per feature row |
175
+ | `labels` | SHAP's wording | text the chart draws; see [Wording](#wording) |
175
176
  | `onFeatureClick` | — | called with the feature index, or `null` for the "other" row |
176
177
 
177
178
  Per chart:
@@ -179,11 +180,49 @@ Per chart:
179
180
  | Chart | Prop | Default | |
180
181
  | --- | --- | --- | --- |
181
182
  | `Plots.beeswarm`, `Plots.heatmap` | `rowSort` | `"importance"` | `"importance"`, `"name"` or `"featureValue"`; reorders the rows shown, never which rows are shown |
183
+ | `Plots.beeswarm`, `Plots.heatmap` | `colorBar` | `true` | SHAP's colour bar right of the plot — Low to High feature value on the beeswarm, the SHAP value range on the heatmap. The plot narrows when the right margin cannot hold it |
182
184
  | `Plots.beeswarm` | `seed`, `dotRadius` | `0`, `3` | jitter is seeded, so a chart is identical on every render |
183
185
  | `Plots.heatmap` | `onSampleClick` | — | called with the column's `sample_ids` entry |
184
186
  | `Plots.waterfall` | `sampleIndex` | `0` | which sample to explain |
185
187
  | `Plots.waterfall` | `decimals` | `2` | `2`, `3`, `4` or `"percent"`; display only |
186
188
 
189
+ ## Wording
190
+
191
+ Every piece of text a chart draws comes from `labels`, keyed by what it means rather than where it
192
+ appears — a word used in two places is changed once. Give only the keys you want to change; the rest
193
+ keep SHAP's wording, exported as `shapLabels`.
194
+
195
+ ```tsx
196
+ import { Plots } from "shap-svg/react";
197
+ import type { PlotLabels } from "shap-svg";
198
+
199
+ // A module constant: a new object on every render recomputes the layout on every render.
200
+ const researchLabels: Partial<PlotLabels> = {
201
+ shapValue: "Contribution",
202
+ shapValueAxis: "Contribution to predicted probability",
203
+ featureValue: "Relative abundance",
204
+ otherFeatures: (count) => `${count} other taxa`,
205
+ };
206
+
207
+ <Plots.beeswarm explanation={explanation} labels={researchLabels} />;
208
+ ```
209
+
210
+ | Key | Default | Drawn as |
211
+ | --- | --- | --- |
212
+ | `shapValue` | `SHAP value` | beeswarm tooltip |
213
+ | `shapValueAxis` | `SHAP value (impact on model output)` | beeswarm x axis, heatmap colour bar |
214
+ | `meanAbsShapValue` | `mean(\|SHAP value\|)` | bar x axis |
215
+ | `featureValue` | `Feature value` | beeswarm tooltip and colour bar |
216
+ | `featureValueLow`, `featureValueHigh` | `Low`, `High` | ends of the beeswarm colour bar |
217
+ | `missingFeatureValue` | `missing` | beeswarm tooltip, for a value the payload lacks |
218
+ | `samples` | `Instances` | heatmap x axis |
219
+ | `sampleTotal` | `Σφ` | heatmap tooltip |
220
+ | `sampleFallback(n)` | `Sample n` | a heatmap column with no `sample_labels` entry |
221
+ | `baseValue`, `modelOutput` | `E[f(X)]`, `f(x)` | the waterfall's two reference lines |
222
+ | `otherFeatures(count, style)` | `Sum of N other features` / `N other features` | the row for every feature not shown; `style` is `"sum"` in `faithfulOtherRow` mode on bar, beeswarm and heatmap |
223
+
224
+ Accessible names are built from the same words.
225
+
187
226
  ## Faithful to SHAP where it matters
188
227
 
189
228
  Ordering, the "other features" partition, colour maps, tick placement and the waterfall's geometry are
@@ -261,8 +261,34 @@ function formatFeatureLabel(name) {
261
261
  return name.replace(/_/g, " ");
262
262
  }
263
263
 
264
+ // src/core/labels.ts
265
+ var shapLabels = {
266
+ shapValue: "SHAP value",
267
+ shapValueAxis: "SHAP value (impact on model output)",
268
+ meanAbsShapValue: "mean(|SHAP value|)",
269
+ featureValue: "Feature value",
270
+ featureValueLow: "Low",
271
+ featureValueHigh: "High",
272
+ missingFeatureValue: "missing",
273
+ samples: "Instances",
274
+ sampleTotal: "\u03A3\u03C6",
275
+ sampleFallback: (sampleNumber) => `Sample ${sampleNumber}`,
276
+ baseValue: "E[f(X)]",
277
+ modelOutput: "f(x)",
278
+ otherFeatures: (count, style) => style === "sum" ? `Sum of ${count} other features` : `${count} other features`
279
+ };
280
+ function resolveLabels(labels) {
281
+ if (!labels) return shapLabels;
282
+ const resolved = { ...shapLabels };
283
+ for (const key of Object.keys(labels)) {
284
+ const value = labels[key];
285
+ if (value !== void 0) resolved[key] = value;
286
+ }
287
+ return resolved;
288
+ }
289
+
264
290
  // src/core/collapse.ts
265
- function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow) {
291
+ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithfulOtherRow, labels = shapLabels) {
266
292
  const p = order.length;
267
293
  if (maxDisplay >= p) {
268
294
  return {
@@ -285,7 +311,7 @@ function collapseToDisplay(featureNames, importance, order, maxDisplay, faithful
285
311
  const collapsed = order.slice(realCount);
286
312
  const collapsedValue = collapsed.reduce((sum, index) => sum + importance[index], 0);
287
313
  rows.push({
288
- label: faithfulOtherRow ? `Sum of ${collapsed.length} other features` : `${collapsed.length} other features`,
314
+ label: labels.otherFeatures(collapsed.length, faithfulOtherRow ? "sum" : "count"),
289
315
  featureIndex: null,
290
316
  value: collapsedValue,
291
317
  isOtherRow: true
@@ -354,15 +380,15 @@ var AXIS_HEIGHT = 52;
354
380
  var TICK_LABEL_PT = 11;
355
381
  var TITLE_PT = 13;
356
382
  var X_MARGIN = 0.05;
357
- function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
383
+ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
358
384
  const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT));
359
385
  return {
360
386
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
361
387
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
362
388
  xTitle: {
363
- // _bar.py:143-150 builds this from the Explanation's transform history:
364
- // "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
365
- text: "mean(|SHAP value|)",
389
+ // _bar.py:143-150 builds the default from the Explanation's transform
390
+ // history: "SHAP value" -> "|SHAP value|" -> "mean(|SHAP value|)".
391
+ text: title,
366
392
  x: marginLeft + plotWidth / 2,
367
393
  y: plotBottom + AXIS_TITLE_DY,
368
394
  fontSize: TITLE_PT
@@ -370,6 +396,7 @@ function barXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
370
396
  };
371
397
  }
372
398
  function barLayout(rows, opts) {
399
+ var _a;
373
400
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
374
401
  const plotWidth = width - marginLeft - marginRight;
375
402
  const values = rows.rows.map((r) => r.value);
@@ -410,7 +437,15 @@ function barLayout(rows, opts) {
410
437
  xZero,
411
438
  plotWidth,
412
439
  plotBottom: marginTop + rows.rows.length * rowHeight,
413
- ...barXAxis(min, max, toX, marginLeft, plotWidth, marginTop + rows.rows.length * rowHeight),
440
+ ...barXAxis(
441
+ min,
442
+ max,
443
+ toX,
444
+ marginLeft,
445
+ plotWidth,
446
+ marginTop + rows.rows.length * rowHeight,
447
+ ((_a = opts.labels) != null ? _a : shapLabels).meanAbsShapValue
448
+ ),
414
449
  zeroLine: { x: xZero, y1: marginTop, y2: marginTop + rows.rows.length * rowHeight },
415
450
  height: marginTop + rows.rows.length * rowHeight + AXIS_HEIGHT
416
451
  };
@@ -444,7 +479,7 @@ function placeValueLabel(value, startX, endX, gutterX, decimals) {
444
479
  }
445
480
  return { ...base, x: startX + VALUE_LABEL_GAP, anchor: "start", inside: false };
446
481
  }
447
- function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
482
+ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow, labels = shapLabels) {
448
483
  if (!Number.isInteger(sampleIndex) || sampleIndex < 0 || sampleIndex >= explanation.nSamples) {
449
484
  throw new RangeError(
450
485
  `sampleIndex must identify a Sample from 0 to ${explanation.nSamples - 1}, received ${sampleIndex}`
@@ -482,7 +517,8 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
482
517
  if (hasOtherRow) {
483
518
  const value = collapsed.reduce((sum, featureIndex) => sum + values[featureIndex], 0);
484
519
  rows.push({
485
- label: `${collapsed.length} other features`,
520
+ // SHAP's waterfall counts the hidden Features whatever the collapse mode.
521
+ label: labels.otherFeatures(collapsed.length, "count"),
486
522
  featureIndex: null,
487
523
  isOtherRow: true,
488
524
  value,
@@ -500,7 +536,9 @@ function waterfallRows(explanation, sampleIndex, maxDisplay, faithfulOtherRow) {
500
536
  };
501
537
  }
502
538
  function waterfallLayout(valueRows, opts) {
539
+ var _a;
503
540
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
541
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
504
542
  const plotWidth = width - marginLeft - marginRight;
505
543
  const coordinates = [valueRows.baseValue, valueRows.modelOutput];
506
544
  for (const row of valueRows.rows) coordinates.push(row.left, row.left + row.width);
@@ -567,7 +605,7 @@ function waterfallLayout(valueRows, opts) {
567
605
  // got, which is the whole reason SHAP hides the left spine.
568
606
  y1: plotBottom - rowHeight,
569
607
  y2: plotBottom,
570
- label: `E[f(X)] = ${formatLevel(valueRows.baseValue, opts.decimals)}`
608
+ label: `${labels.baseValue} = ${formatLevel(valueRows.baseValue, opts.decimals)}`
571
609
  },
572
610
  {
573
611
  kind: "output",
@@ -576,7 +614,7 @@ function waterfallLayout(valueRows, opts) {
576
614
  // axvline(fx, 0, 1) — the full height.
577
615
  y1: marginTop,
578
616
  y2: plotBottom,
579
- label: `f(x) = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
617
+ label: `${labels.modelOutput} = ${formatLevel(valueRows.modelOutput, opts.decimals)}`
580
618
  }
581
619
  ],
582
620
  separators: valueRows.rows.map((_, index) => ({
@@ -1141,12 +1179,120 @@ function sampleColormap(name, t) {
1141
1179
  return `#${hexByte(red)}${hexByte(green)}${hexByte(blue)}`;
1142
1180
  }
1143
1181
 
1182
+ // src/core/colorBar.ts
1183
+ var ASPECT = 80;
1184
+ var TICK_PAD = 3.5;
1185
+ var TICK_LABEL_PT3 = 11;
1186
+ var LABEL_PT = 12;
1187
+ var LABEL_THICKNESS_EM = 1.08;
1188
+ var LABEL_BASELINE_EM = 0.84;
1189
+ var CHAR_EM = 0.6;
1190
+ var STEPS = 64;
1191
+ function textWidth(text, fontSize) {
1192
+ return text.length * fontSize * CHAR_EM;
1193
+ }
1194
+ function tickColumnWidth(spec) {
1195
+ return Math.max(...spec.tickLabels.map((label) => textWidth(label, TICK_LABEL_PT3)));
1196
+ }
1197
+ function colorBarExtent(height, spec) {
1198
+ return height / ASPECT + TICK_PAD + Math.max(0, tickColumnWidth(spec) + spec.labelPad) + LABEL_PT * LABEL_THICKNESS_EM;
1199
+ }
1200
+ function fitColorBar(opts) {
1201
+ const { plotWidth, available, gapRatio, minGap, extent } = opts;
1202
+ const gapFor = (width) => Math.max(gapRatio * width, minGap);
1203
+ if (gapFor(plotWidth) + extent <= available) {
1204
+ return { plotWidth, gap: gapFor(plotWidth) };
1205
+ }
1206
+ const byRatio = (available + plotWidth - extent) / (1 + gapRatio);
1207
+ const fitted = Math.max(
1208
+ 0,
1209
+ gapRatio * byRatio >= minGap ? byRatio : plotWidth - (minGap + extent - available)
1210
+ );
1211
+ return { plotWidth: fitted, gap: gapFor(fitted) };
1212
+ }
1213
+ function colorBarLayout(spec, position) {
1214
+ const { x, y1, y2 } = position;
1215
+ const width = (y2 - y1) / ASPECT;
1216
+ const step = (y2 - y1) / STEPS;
1217
+ const steps = Array.from({ length: STEPS }, (_, index) => ({
1218
+ y: y2 - (index + 1) * step,
1219
+ // Every band but the lowest reaches half a pixel into the one below, so
1220
+ // anti-aliasing cannot open a hairline seam between them.
1221
+ height: index === 0 ? step : step + 0.5,
1222
+ color: sampleColormap(spec.colormap, index / (STEPS - 1))
1223
+ }));
1224
+ const tickX = x + width + TICK_PAD;
1225
+ const labelLeft = tickX + tickColumnWidth(spec) + spec.labelPad;
1226
+ return {
1227
+ x,
1228
+ y1,
1229
+ y2,
1230
+ width,
1231
+ steps,
1232
+ tickX,
1233
+ tickFontSize: TICK_LABEL_PT3,
1234
+ ticks: [
1235
+ { y: y2, label: spec.tickLabels[0] },
1236
+ { y: y1, label: spec.tickLabels[1] }
1237
+ ],
1238
+ ...spec.offsetText ? { offsetText: { x: tickX, y: y1 - TICK_LABEL_PT3, text: spec.offsetText } } : {},
1239
+ label: {
1240
+ x: labelLeft + LABEL_PT * LABEL_BASELINE_EM,
1241
+ y: (y1 + y2) / 2,
1242
+ text: spec.label,
1243
+ fontSize: LABEL_PT
1244
+ },
1245
+ right: Math.max(
1246
+ tickX + tickColumnWidth(spec),
1247
+ labelLeft + LABEL_PT * LABEL_THICKNESS_EM
1248
+ )
1249
+ };
1250
+ }
1251
+ var MINUS3 = "\u2212";
1252
+ var POWER_LIMITS = [-5, 6];
1253
+ function roundTo(value, decimals) {
1254
+ const factor = 10 ** decimals;
1255
+ return Math.round(value * factor) / factor;
1256
+ }
1257
+ function scalarFormatterLabels(locs) {
1258
+ const largest = Math.max(0, ...locs.map(Math.abs));
1259
+ const magnitude = largest === 0 ? 0 : Math.floor(Math.log10(largest));
1260
+ const order = magnitude <= POWER_LIMITS[0] || magnitude >= POWER_LIMITS[1] ? magnitude : 0;
1261
+ const scaled = locs.map((value) => value / 10 ** order);
1262
+ let range = Math.max(...scaled) - Math.min(...scaled);
1263
+ if (range === 0) range = Math.max(...scaled.map(Math.abs));
1264
+ if (range === 0) range = 1;
1265
+ const rangeMagnitude = Math.floor(Math.log10(range));
1266
+ const threshold = 1e-3 * 10 ** rangeMagnitude;
1267
+ let decimals = Math.max(0, 3 - rangeMagnitude);
1268
+ while (decimals >= 0) {
1269
+ const error = Math.max(...scaled.map((value) => Math.abs(value - roundTo(value, decimals))));
1270
+ if (error < threshold) decimals -= 1;
1271
+ else break;
1272
+ }
1273
+ decimals += 1;
1274
+ const labels = scaled.map((value) => {
1275
+ const text = value.toFixed(decimals);
1276
+ return Number(text) === 0 ? text.replace("-", "") : text.replace("-", MINUS3);
1277
+ });
1278
+ return order === 0 ? { labels } : { labels, offsetText: `1e${String(order).replace("-", MINUS3)}` };
1279
+ }
1280
+
1144
1281
  // src/core/beeswarmLayout.ts
1145
1282
  var BEESWARM_MISSING_COLOR = "#777777";
1146
1283
  var BEESWARM_ROW_HEIGHT = 0.4;
1147
1284
  var NBINS = 100;
1148
- var AXIS_HEIGHT3 = 74;
1149
- var TICK_LABEL_PT3 = 11;
1285
+ var AXIS_HEIGHT3 = 52;
1286
+ function colorBarSpec(labels) {
1287
+ return {
1288
+ colormap: "red_blue",
1289
+ tickLabels: [labels.featureValueLow, labels.featureValueHigh],
1290
+ label: labels.featureValue,
1291
+ labelPad: 0
1292
+ };
1293
+ }
1294
+ var COLOR_BAR_GAP_RATIO = 0.05 / 0.8;
1295
+ var TICK_LABEL_PT4 = 11;
1150
1296
  var TITLE_PT2 = 13;
1151
1297
  var X_MARGIN2 = 0.05;
1152
1298
  function percentile(values, percent) {
@@ -1210,7 +1356,7 @@ function spreadPoints(xs, rowIndex, seed) {
1210
1356
  const scale = 0.9 * (BEESWARM_ROW_HEIGHT / (maxPositiveOffset + 1));
1211
1357
  return offsets.map((offset) => rowIndex + offset * scale);
1212
1358
  }
1213
- function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance") {
1359
+ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSort = "importance", labels = shapLabels) {
1214
1360
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1215
1361
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1216
1362
  }
@@ -1225,7 +1371,8 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1225
1371
  importance,
1226
1372
  order,
1227
1373
  maxDisplay,
1228
- faithfulOtherRow
1374
+ faithfulOtherRow,
1375
+ labels
1229
1376
  ),
1230
1377
  rowSort,
1231
1378
  explanation.data
@@ -1277,14 +1424,14 @@ function beeswarmRows(explanation, maxDisplay, faithfulOtherRow, seed = 0, rowSo
1277
1424
  });
1278
1425
  return { rows, collapsedCount: display.collapsedCount };
1279
1426
  }
1280
- function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1281
- const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT3));
1427
+ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, title) {
1428
+ const { ticks, step } = niceTicks(min, max, tickSpace(plotWidth, TICK_LABEL_PT4));
1282
1429
  return {
1283
1430
  xTicks: ticks.map((value) => ({ value, x: toX(value), label: tickLabel(value, step) })),
1284
1431
  xSpine: { x1: marginLeft, x2: marginLeft + plotWidth, y: plotBottom },
1285
1432
  xTitle: {
1286
- // _labels.py:5, labels["VALUE"].
1287
- text: "SHAP value (impact on model output)",
1433
+ // The default is _labels.py:5, labels["VALUE"].
1434
+ text: title,
1288
1435
  x: marginLeft + plotWidth / 2,
1289
1436
  y: plotBottom + AXIS_TITLE_DY,
1290
1437
  fontSize: TITLE_PT2
@@ -1292,8 +1439,19 @@ function beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom) {
1292
1439
  };
1293
1440
  }
1294
1441
  function beeswarmLayout(valueRows, opts) {
1442
+ var _a;
1295
1443
  const { width, rowHeight, marginLeft, marginRight, marginTop, dotRadius } = opts;
1296
- const plotWidth = width - marginLeft - marginRight;
1444
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1445
+ const colorBar = colorBarSpec(labels);
1446
+ const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1447
+ const fit = opts.colorBar ? fitColorBar({
1448
+ plotWidth: width - marginLeft - marginRight,
1449
+ available: marginRight,
1450
+ gapRatio: COLOR_BAR_GAP_RATIO,
1451
+ minGap: 0,
1452
+ extent: colorBarExtent(plotBottom - marginTop, colorBar)
1453
+ }) : null;
1454
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1297
1455
  const values = valueRows.rows.flatMap((row) => row.points.map((point) => point.x));
1298
1456
  const dataMin = Math.min(0, ...values);
1299
1457
  const dataMax = Math.max(0, ...values);
@@ -1318,14 +1476,18 @@ function beeswarmLayout(valueRows, opts) {
1318
1476
  }))
1319
1477
  };
1320
1478
  });
1321
- const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1322
1479
  return {
1323
1480
  rows,
1324
1481
  xDomain: [min, max],
1325
1482
  xZero: toX(0),
1326
- ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom),
1483
+ ...beeswarmXAxis(min, max, toX, marginLeft, plotWidth, plotBottom, labels.shapValueAxis),
1327
1484
  plotWidth,
1328
1485
  plotBottom,
1486
+ colorBar: fit ? colorBarLayout(colorBar, {
1487
+ x: marginLeft + plotWidth + fit.gap,
1488
+ y1: marginTop,
1489
+ y2: plotBottom
1490
+ }) : null,
1329
1491
  height: plotBottom + AXIS_HEIGHT3
1330
1492
  };
1331
1493
  }
@@ -1338,7 +1500,9 @@ var SIDE_BAR_GAP = 10;
1338
1500
  var SIDE_BAR_RIGHT_INSET = 40;
1339
1501
  var SIDE_BAR_HEIGHT_RATIO = 0.6;
1340
1502
  var AXIS_HEIGHT4 = 52;
1341
- var TICK_LABEL_PT4 = 10;
1503
+ var TICK_LABEL_PT5 = 10;
1504
+ var COLOR_BAR_GAP_RATIO2 = 0.1 / 0.89;
1505
+ var COLOR_BAR_CLEARANCE = 8;
1342
1506
  var Y_TICK_LENGTH = 5;
1343
1507
  function percentile2(values, fraction) {
1344
1508
  if (values.length === 0) return 0;
@@ -1349,7 +1513,7 @@ function percentile2(values, fraction) {
1349
1513
  const weight = position - lower;
1350
1514
  return sorted[lower] + (sorted[upper] - sorted[lower]) * weight;
1351
1515
  }
1352
- function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance") {
1516
+ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "importance", labels = shapLabels) {
1353
1517
  if (!Number.isInteger(maxDisplay) || maxDisplay <= 0) {
1354
1518
  throw new RangeError(`maxDisplay must be a positive integer, received ${maxDisplay}`);
1355
1519
  }
@@ -1361,7 +1525,8 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1361
1525
  importance,
1362
1526
  featureOrder,
1363
1527
  maxDisplay,
1364
- faithfulOtherRow
1528
+ faithfulOtherRow,
1529
+ labels
1365
1530
  ),
1366
1531
  rowSort,
1367
1532
  explanation.data
@@ -1425,11 +1590,11 @@ function heatmapRows(explanation, maxDisplay, faithfulOtherRow, rowSort = "impor
1425
1590
  sampleLabelColumn: explanation.sampleLabelColumn
1426
1591
  };
1427
1592
  }
1428
- function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom) {
1593
+ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom, title) {
1429
1594
  const { ticks, step } = niceTicks(
1430
1595
  -0.5,
1431
1596
  sampleCount - 0.5,
1432
- tickSpace(plotWidth, TICK_LABEL_PT4),
1597
+ tickSpace(plotWidth, TICK_LABEL_PT5),
1433
1598
  { integer: true }
1434
1599
  );
1435
1600
  return {
@@ -1440,20 +1605,38 @@ function heatmapXAxis(sampleCount, marginLeft, cellWidth, plotWidth, plotBottom)
1440
1605
  })),
1441
1606
  xSpine: null,
1442
1607
  xTitle: {
1443
- text: "Instances",
1608
+ text: title,
1444
1609
  x: marginLeft + plotWidth / 2,
1445
1610
  y: plotBottom + AXIS_TITLE_DY,
1446
- fontSize: TICK_LABEL_PT4
1611
+ fontSize: TICK_LABEL_PT5
1447
1612
  }
1448
1613
  };
1449
1614
  }
1450
1615
  function heatmapLayout(valueRows, opts) {
1616
+ var _a;
1451
1617
  const { width, rowHeight, marginLeft, marginRight, marginTop } = opts;
1452
- const plotWidth = width - marginLeft - marginRight;
1453
- const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1454
- const gridRight = marginLeft + plotWidth;
1618
+ const labels = (_a = opts.labels) != null ? _a : shapLabels;
1455
1619
  const plotBottom = marginTop + valueRows.rows.length * rowHeight;
1456
1620
  const sideBarWidth = Math.max(0, marginRight - SIDE_BAR_RIGHT_INSET - SIDE_BAR_GAP);
1621
+ const ticks = scalarFormatterLabels([valueRows.vmin, valueRows.vmax]);
1622
+ const colorBarSpec2 = {
1623
+ colormap: "red_white_blue",
1624
+ tickLabels: [ticks.labels[0], ticks.labels[1]],
1625
+ ...ticks.offsetText ? { offsetText: ticks.offsetText } : {},
1626
+ label: labels.shapValueAxis,
1627
+ labelPad: -10
1628
+ };
1629
+ const colorBarTop = FX_TOP;
1630
+ const fit = opts.colorBar ? fitColorBar({
1631
+ plotWidth: width - marginLeft - marginRight,
1632
+ available: marginRight,
1633
+ gapRatio: COLOR_BAR_GAP_RATIO2,
1634
+ minGap: SIDE_BAR_GAP + sideBarWidth + COLOR_BAR_CLEARANCE,
1635
+ extent: colorBarExtent(plotBottom - colorBarTop, colorBarSpec2)
1636
+ }) : null;
1637
+ const plotWidth = fit ? fit.plotWidth : width - marginLeft - marginRight;
1638
+ const cellWidth = valueRows.columns.length === 0 ? 0 : plotWidth / valueRows.columns.length;
1639
+ const gridRight = marginLeft + plotWidth;
1457
1640
  const columns = valueRows.columns.map((column, index) => ({
1458
1641
  ...column,
1459
1642
  x: marginLeft + index * cellWidth,
@@ -1522,10 +1705,22 @@ function heatmapLayout(valueRows, opts) {
1522
1705
  x1: marginLeft - Y_TICK_LENGTH,
1523
1706
  x2: marginLeft
1524
1707
  })),
1525
- ...heatmapXAxis(valueRows.columns.length, marginLeft, cellWidth, plotWidth, plotBottom),
1708
+ ...heatmapXAxis(
1709
+ valueRows.columns.length,
1710
+ marginLeft,
1711
+ cellWidth,
1712
+ plotWidth,
1713
+ plotBottom,
1714
+ labels.samples
1715
+ ),
1526
1716
  sampleLabelColumn: valueRows.sampleLabelColumn,
1527
1717
  plotWidth,
1528
1718
  cellWidth,
1719
+ colorBar: fit ? colorBarLayout(colorBarSpec2, {
1720
+ x: gridRight + fit.gap,
1721
+ y1: colorBarTop,
1722
+ y2: plotBottom
1723
+ }) : null,
1529
1724
  height: plotBottom + AXIS_HEIGHT4
1530
1725
  };
1531
1726
  }
@@ -1544,6 +1739,8 @@ export {
1544
1739
  formatShapValue,
1545
1740
  formatLevel,
1546
1741
  formatFeatureLabel,
1742
+ shapLabels,
1743
+ resolveLabels,
1547
1744
  collapseToDisplay,
1548
1745
  TICK_LENGTH,
1549
1746
  TICK_LABEL_DY,
@@ -1556,6 +1753,10 @@ export {
1556
1753
  waterfallRows,
1557
1754
  waterfallLayout,
1558
1755
  sampleColormap,
1756
+ colorBarExtent,
1757
+ fitColorBar,
1758
+ colorBarLayout,
1759
+ scalarFormatterLabels,
1559
1760
  BEESWARM_MISSING_COLOR,
1560
1761
  BEESWARM_ROW_HEIGHT,
1561
1762
  beeswarmRows,
@@ -1563,4 +1764,4 @@ export {
1563
1764
  heatmapRows,
1564
1765
  heatmapLayout
1565
1766
  };
1566
- //# sourceMappingURL=chunk-OXFKP5I3.js.map
1767
+ //# sourceMappingURL=chunk-4X6PMJYT.js.map