shap-svg 0.2.2 → 0.2.4

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.js CHANGED
@@ -15,10 +15,11 @@ import {
15
15
  heatmapRows,
16
16
  orderFeatures,
17
17
  parseExplanation,
18
+ placeTooltip,
18
19
  resolveLabels,
19
20
  waterfallLayout,
20
21
  waterfallRows
21
- } from "./chunk-4X6PMJYT.js";
22
+ } from "./chunk-UVIVTJK3.js";
22
23
 
23
24
  // src/react/ShapBar.tsx
24
25
  import { useMemo, useState } from "react";
@@ -324,91 +325,124 @@ function ShapBeeswarm({
324
325
  const fallbackRow = layout.rows[0];
325
326
  const activeRow = hovered ? layout.rows[hovered.rowIndex] : fallbackRow;
326
327
  const activePoint = hovered ? activeRow == null ? void 0 : activeRow.points[hovered.pointIndex] : activeRow == null ? void 0 : activeRow.points[0];
327
- return /* @__PURE__ */ jsxs4("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP beeswarm", children: [
328
- /* @__PURE__ */ jsx4(
329
- "line",
330
- {
331
- x1: layout.xZero,
332
- x2: layout.xZero,
333
- y1: marginTop,
334
- y2: layout.plotBottom,
335
- stroke: "#777777",
336
- strokeWidth: 1
337
- }
338
- ),
339
- layout.rows.map((row, rowIndex) => /* @__PURE__ */ jsxs4(
340
- "g",
341
- {
342
- onMouseOver: (event) => handlePointHover(rowIndex, event),
343
- onMouseLeave: () => setHovered(null),
344
- onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
345
- style: { cursor: onFeatureClick ? "pointer" : "default" },
346
- children: [
347
- /* @__PURE__ */ jsx4(
348
- "rect",
349
- {
350
- x: 0,
351
- y: row.centerY - rowHeight / 2,
352
- width,
353
- height: rowHeight,
354
- fill: (hovered == null ? void 0 : hovered.rowIndex) === rowIndex ? "#00000008" : "transparent"
355
- }
356
- ),
357
- /* @__PURE__ */ jsx4(
358
- "text",
359
- {
360
- x: 250,
361
- y: row.centerY,
362
- textAnchor: "end",
363
- dominantBaseline: "middle",
364
- fontSize: 13,
365
- fill: "#333333",
366
- fontStyle: row.isOtherRow ? "normal" : "italic",
367
- children: row.label
368
- }
369
- ),
370
- row.points.map((point, pointIndex) => /* @__PURE__ */ jsx4(
371
- "circle",
372
- {
373
- "data-point-index": pointIndex,
374
- cx: point.x,
375
- cy: point.y,
376
- r: point.radius,
377
- fill: point.color,
378
- fillOpacity: 0.7
379
- },
380
- `point-${point.sampleIndex}`
381
- ))
382
- ]
383
- },
384
- `row-${rowIndex}`
385
- )),
386
- /* @__PURE__ */ jsx4(
387
- XAxis,
388
- {
389
- ticks: layout.xTicks,
390
- spine: layout.xSpine,
391
- title: layout.xTitle,
392
- plotBottom: layout.plotBottom,
393
- tickFontSize: 11
394
- }
395
- ),
396
- layout.colorBar && /* @__PURE__ */ jsx4(ColorBar, { bar: layout.colorBar }),
397
- activeRow && activePoint && /* @__PURE__ */ jsxs4(
398
- "g",
399
- {
400
- opacity: hovered ? 1 : 0,
401
- pointerEvents: "none",
402
- transform: `translate(${activePoint.x + 8} ${activePoint.y - 8})`,
403
- children: [
404
- /* @__PURE__ */ jsx4("rect", { x: 0, y: -16, width: 210, height: 54, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
405
- /* @__PURE__ */ jsx4("text", { x: 7, y: 0, fontSize: 11, fill: "#222222", children: activeRow.label }),
406
- /* @__PURE__ */ jsx4("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: `${words.shapValue}: ${formatShapValue(activePoint.valueX)}` }),
407
- /* @__PURE__ */ jsx4("text", { x: 7, y: 30, fontSize: 11, fill: "#222222", children: `${words.featureValue}: ${Number.isFinite(activePoint.featureValue) ? formatLevel(activePoint.featureValue) : words.missingFeatureValue}` })
408
- ]
409
- }
410
- )
411
- ] });
328
+ const tooltipLines = activeRow && activePoint ? [
329
+ activeRow.label,
330
+ `${words.shapValue}: ${formatShapValue(activePoint.valueX)}`,
331
+ // Unsigned: a feature value is an input, not a push in either direction.
332
+ `${words.featureValue}: ${Number.isFinite(activePoint.featureValue) ? formatLevel(activePoint.featureValue) : words.missingFeatureValue}`
333
+ ] : [];
334
+ const tooltip = activePoint ? placeTooltip({
335
+ anchorX: activePoint.x,
336
+ anchorY: activePoint.y,
337
+ lines: tooltipLines,
338
+ lineHeight: 15,
339
+ minWidth: 160,
340
+ chartWidth: width,
341
+ chartHeight: layout.height
342
+ }) : null;
343
+ return /* @__PURE__ */ jsxs4(
344
+ "svg",
345
+ {
346
+ width,
347
+ height: layout.height,
348
+ role: "img",
349
+ "aria-label": `${words.shapValue} of each feature, for every sample`,
350
+ children: [
351
+ /* @__PURE__ */ jsx4(
352
+ "line",
353
+ {
354
+ x1: layout.xZero,
355
+ x2: layout.xZero,
356
+ y1: marginTop,
357
+ y2: layout.plotBottom,
358
+ stroke: "#777777",
359
+ strokeWidth: 1
360
+ }
361
+ ),
362
+ layout.rows.map((row, rowIndex) => /* @__PURE__ */ jsxs4(
363
+ "g",
364
+ {
365
+ onMouseOver: (event) => handlePointHover(rowIndex, event),
366
+ onMouseLeave: () => setHovered(null),
367
+ onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
368
+ style: { cursor: onFeatureClick ? "pointer" : "default" },
369
+ children: [
370
+ /* @__PURE__ */ jsx4(
371
+ "rect",
372
+ {
373
+ x: 0,
374
+ y: row.centerY - rowHeight / 2,
375
+ width,
376
+ height: rowHeight,
377
+ fill: (hovered == null ? void 0 : hovered.rowIndex) === rowIndex ? "#00000008" : "transparent"
378
+ }
379
+ ),
380
+ /* @__PURE__ */ jsx4(
381
+ "text",
382
+ {
383
+ x: 250,
384
+ y: row.centerY,
385
+ textAnchor: "end",
386
+ dominantBaseline: "middle",
387
+ fontSize: 13,
388
+ fill: "#333333",
389
+ fontStyle: row.isOtherRow ? "normal" : "italic",
390
+ children: row.label
391
+ }
392
+ ),
393
+ row.points.map((point, pointIndex) => /* @__PURE__ */ jsx4(
394
+ "circle",
395
+ {
396
+ "data-point-index": pointIndex,
397
+ cx: point.x,
398
+ cy: point.y,
399
+ r: point.radius,
400
+ fill: point.color,
401
+ fillOpacity: 0.7
402
+ },
403
+ `point-${point.sampleIndex}`
404
+ ))
405
+ ]
406
+ },
407
+ `row-${rowIndex}`
408
+ )),
409
+ /* @__PURE__ */ jsx4(
410
+ XAxis,
411
+ {
412
+ ticks: layout.xTicks,
413
+ spine: layout.xSpine,
414
+ title: layout.xTitle,
415
+ plotBottom: layout.plotBottom,
416
+ tickFontSize: 11
417
+ }
418
+ ),
419
+ layout.colorBar && /* @__PURE__ */ jsx4(ColorBar, { bar: layout.colorBar }),
420
+ tooltip && /* @__PURE__ */ jsxs4(
421
+ "g",
422
+ {
423
+ opacity: hovered ? 1 : 0,
424
+ pointerEvents: "none",
425
+ transform: `translate(${tooltip.x} ${tooltip.y})`,
426
+ children: [
427
+ /* @__PURE__ */ jsx4(
428
+ "rect",
429
+ {
430
+ x: 0,
431
+ y: 0,
432
+ width: tooltip.width,
433
+ height: tooltip.height,
434
+ rx: 3,
435
+ fill: "#ffffff",
436
+ stroke: "#cccccc"
437
+ }
438
+ ),
439
+ tooltipLines.map((line, index) => /* @__PURE__ */ jsx4("text", { x: 7, y: 16 + index * 15, fontSize: 11, fill: "#222222", children: line }, `tooltip-${index}`))
440
+ ]
441
+ }
442
+ )
443
+ ]
444
+ }
445
+ );
412
446
  }
413
447
 
414
448
  // src/react/ShapHeatmap.tsx
@@ -452,170 +486,179 @@ function ShapHeatmap({
452
486
  };
453
487
  const tooltipWidth = activeColumn ? Math.max(190, nameOf(activeColumn).length * 6.5 + 16) : 190;
454
488
  const tooltipX = activeColumn ? Math.max(0, Math.min(activeColumn.centerX + 8, width - tooltipWidth - 8)) : 0;
455
- return /* @__PURE__ */ jsxs5("svg", { width, height: layout.height, role: "img", "aria-label": "Global SHAP heatmap", children: [
456
- layout.fxAxisMarks.map((mark) => /* @__PURE__ */ jsxs5("g", { children: [
457
- /* @__PURE__ */ jsx5(
458
- "line",
459
- {
460
- x1: layout.gridLeft - 4,
461
- x2: layout.gridRight,
462
- y1: mark.y,
463
- y2: mark.y,
464
- stroke: "#dddddd",
465
- strokeWidth: 1
466
- }
467
- ),
468
- /* @__PURE__ */ jsx5(
469
- "text",
470
- {
471
- x: layout.gridLeft - 8,
472
- y: mark.y,
473
- textAnchor: "end",
474
- dominantBaseline: "middle",
475
- fontSize: 10,
476
- fill: "#666666",
477
- children: mark.label
478
- }
479
- )
480
- ] }, `fx-axis-${mark.value}`)),
481
- /* @__PURE__ */ jsx5(
482
- "polyline",
483
- {
484
- points: layout.fxLine.map((point) => `${point.x},${point.y}`).join(" "),
485
- fill: "none",
486
- stroke: "#333333",
487
- strokeWidth: 1.5
488
- }
489
- ),
490
- /* @__PURE__ */ jsx5(
491
- "line",
492
- {
493
- x1: layout.gridLeft,
494
- x2: layout.gridRight,
495
- y1: layout.separatorY,
496
- y2: layout.separatorY,
497
- stroke: "#888888",
498
- strokeWidth: 1,
499
- strokeDasharray: "4 4"
500
- }
501
- ),
502
- /* @__PURE__ */ jsx5(
503
- XAxis,
504
- {
505
- ticks: layout.xTicks,
506
- spine: layout.xSpine,
507
- title: layout.xTitle,
508
- plotBottom: layout.plotBottom,
509
- tickFontSize: 10
510
- }
511
- ),
512
- /* @__PURE__ */ jsxs5("g", { "aria-hidden": "true", children: [
513
- [layout.spines.left, layout.spines.right].map((spine, index) => /* @__PURE__ */ jsx5(
514
- "line",
515
- {
516
- x1: spine.x,
517
- x2: spine.x,
518
- y1: spine.y1,
519
- y2: spine.y2,
520
- stroke: "#333333",
521
- strokeWidth: 1
522
- },
523
- `spine-${index}`
524
- )),
525
- layout.yTicks.map((tick, index) => /* @__PURE__ */ jsx5(
526
- "line",
527
- {
528
- x1: tick.x1,
529
- x2: tick.x2,
530
- y1: tick.y,
531
- y2: tick.y,
532
- stroke: "#333333",
533
- strokeWidth: 1
534
- },
535
- `ytick-${index}`
536
- ))
537
- ] }),
538
- layout.colorBar && /* @__PURE__ */ jsx5(ColorBar, { bar: layout.colorBar }),
539
- layout.rows.map((row, rowIndex) => /* @__PURE__ */ jsxs5(
540
- "g",
541
- {
542
- onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
543
- style: { cursor: onFeatureClick ? "pointer" : "default" },
544
- children: [
489
+ return /* @__PURE__ */ jsxs5(
490
+ "svg",
491
+ {
492
+ width,
493
+ height: layout.height,
494
+ role: "img",
495
+ "aria-label": `${words.shapValue} by feature and sample`,
496
+ children: [
497
+ layout.fxAxisMarks.map((mark) => /* @__PURE__ */ jsxs5("g", { children: [
498
+ /* @__PURE__ */ jsx5(
499
+ "line",
500
+ {
501
+ x1: layout.gridLeft - 4,
502
+ x2: layout.gridRight,
503
+ y1: mark.y,
504
+ y2: mark.y,
505
+ stroke: "#dddddd",
506
+ strokeWidth: 1
507
+ }
508
+ ),
545
509
  /* @__PURE__ */ jsx5(
546
510
  "text",
547
511
  {
548
- x: 250,
549
- y: row.centerY,
512
+ x: layout.gridLeft - 8,
513
+ y: mark.y,
550
514
  textAnchor: "end",
551
515
  dominantBaseline: "middle",
552
- fontSize: 13,
553
- fill: "#333333",
554
- fontStyle: row.isOtherRow ? "normal" : "italic",
555
- children: row.label
516
+ fontSize: 10,
517
+ fill: "#666666",
518
+ children: mark.label
556
519
  }
557
- ),
558
- row.cells.map((cell, columnIndex) => /* @__PURE__ */ jsx5(
559
- "rect",
520
+ )
521
+ ] }, `fx-axis-${mark.value}`)),
522
+ /* @__PURE__ */ jsx5(
523
+ "polyline",
524
+ {
525
+ points: layout.fxLine.map((point) => `${point.x},${point.y}`).join(" "),
526
+ fill: "none",
527
+ stroke: "#333333",
528
+ strokeWidth: 1.5
529
+ }
530
+ ),
531
+ /* @__PURE__ */ jsx5(
532
+ "line",
533
+ {
534
+ x1: layout.gridLeft,
535
+ x2: layout.gridRight,
536
+ y1: layout.separatorY,
537
+ y2: layout.separatorY,
538
+ stroke: "#888888",
539
+ strokeWidth: 1,
540
+ strokeDasharray: "4 4"
541
+ }
542
+ ),
543
+ /* @__PURE__ */ jsx5(
544
+ XAxis,
545
+ {
546
+ ticks: layout.xTicks,
547
+ spine: layout.xSpine,
548
+ title: layout.xTitle,
549
+ plotBottom: layout.plotBottom,
550
+ tickFontSize: 10
551
+ }
552
+ ),
553
+ /* @__PURE__ */ jsxs5("g", { "aria-hidden": "true", children: [
554
+ [layout.spines.left, layout.spines.right].map((spine, index) => /* @__PURE__ */ jsx5(
555
+ "line",
560
556
  {
561
- x: cell.x,
562
- y: cell.y,
563
- width: cell.width,
564
- height: cell.height,
565
- fill: cell.color
557
+ x1: spine.x,
558
+ x2: spine.x,
559
+ y1: spine.y1,
560
+ y2: spine.y2,
561
+ stroke: "#333333",
562
+ strokeWidth: 1
566
563
  },
567
- `cell-${columnIndex}`
564
+ `spine-${index}`
568
565
  )),
569
- /* @__PURE__ */ jsx5(
570
- "rect",
566
+ layout.yTicks.map((tick, index) => /* @__PURE__ */ jsx5(
567
+ "line",
571
568
  {
572
- x: row.sideBar.x,
573
- y: row.sideBar.y,
574
- width: row.sideBar.width,
575
- height: row.sideBar.height,
576
- fill: "#777777"
577
- }
578
- )
579
- ]
580
- },
581
- `row-${rowIndex}`
582
- )),
583
- activeColumn && /* @__PURE__ */ jsx5(
584
- "rect",
585
- {
586
- x: activeColumn.x,
587
- y: 8,
588
- width: activeColumn.width,
589
- height: layout.plotBottom - 8,
590
- fill: "#000000",
591
- fillOpacity: 0.1,
592
- pointerEvents: "none"
593
- }
594
- ),
595
- layout.columns.map((column, columnIndex) => /* @__PURE__ */ jsx5(
596
- "rect",
597
- {
598
- x: column.x,
599
- y: 8,
600
- width: column.width,
601
- height: layout.plotBottom - 8,
602
- fill: "transparent",
603
- "aria-label": `${nameOf(column)}, total ${words.shapValue} ${formatShapValue(column.total)}`,
604
- onMouseEnter: () => setHoveredColumn(columnIndex),
605
- onMouseLeave: () => setHoveredColumn(null),
606
- onClick: () => {
607
- if (column.sampleId !== void 0) onSampleClick == null ? void 0 : onSampleClick(column.sampleId);
608
- },
609
- style: { cursor: column.sampleId && onSampleClick ? "pointer" : "default" }
610
- },
611
- `column-hit-${column.sampleIndex}`
612
- )),
613
- activeColumn && /* @__PURE__ */ jsxs5("g", { pointerEvents: "none", transform: `translate(${tooltipX} 10)`, children: [
614
- /* @__PURE__ */ jsx5("rect", { x: 0, y: 0, width: tooltipWidth, height: 42, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
615
- /* @__PURE__ */ jsx5("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: nameOf(activeColumn) }),
616
- /* @__PURE__ */ jsx5("text", { x: 7, y: 31, fontSize: 11, fill: "#222222", children: `${words.sampleTotal}: ${formatShapValue(activeColumn.total)}` })
617
- ] })
618
- ] });
569
+ x1: tick.x1,
570
+ x2: tick.x2,
571
+ y1: tick.y,
572
+ y2: tick.y,
573
+ stroke: "#333333",
574
+ strokeWidth: 1
575
+ },
576
+ `ytick-${index}`
577
+ ))
578
+ ] }),
579
+ layout.colorBar && /* @__PURE__ */ jsx5(ColorBar, { bar: layout.colorBar }),
580
+ layout.rows.map((row, rowIndex) => /* @__PURE__ */ jsxs5(
581
+ "g",
582
+ {
583
+ onClick: () => onFeatureClick == null ? void 0 : onFeatureClick(row.featureIndex),
584
+ style: { cursor: onFeatureClick ? "pointer" : "default" },
585
+ children: [
586
+ /* @__PURE__ */ jsx5(
587
+ "text",
588
+ {
589
+ x: 250,
590
+ y: row.centerY,
591
+ textAnchor: "end",
592
+ dominantBaseline: "middle",
593
+ fontSize: 13,
594
+ fill: "#333333",
595
+ fontStyle: row.isOtherRow ? "normal" : "italic",
596
+ children: row.label
597
+ }
598
+ ),
599
+ row.cells.map((cell, columnIndex) => /* @__PURE__ */ jsx5(
600
+ "rect",
601
+ {
602
+ x: cell.x,
603
+ y: cell.y,
604
+ width: cell.width,
605
+ height: cell.height,
606
+ fill: cell.color
607
+ },
608
+ `cell-${columnIndex}`
609
+ )),
610
+ /* @__PURE__ */ jsx5(
611
+ "rect",
612
+ {
613
+ x: row.sideBar.x,
614
+ y: row.sideBar.y,
615
+ width: row.sideBar.width,
616
+ height: row.sideBar.height,
617
+ fill: "#777777"
618
+ }
619
+ )
620
+ ]
621
+ },
622
+ `row-${rowIndex}`
623
+ )),
624
+ activeColumn && /* @__PURE__ */ jsx5(
625
+ "rect",
626
+ {
627
+ x: activeColumn.x,
628
+ y: 8,
629
+ width: activeColumn.width,
630
+ height: layout.plotBottom - 8,
631
+ fill: "#000000",
632
+ fillOpacity: 0.1,
633
+ pointerEvents: "none"
634
+ }
635
+ ),
636
+ layout.columns.map((column, columnIndex) => /* @__PURE__ */ jsx5(
637
+ "rect",
638
+ {
639
+ x: column.x,
640
+ y: 8,
641
+ width: column.width,
642
+ height: layout.plotBottom - 8,
643
+ fill: "transparent",
644
+ "aria-label": `${nameOf(column)}, total ${words.shapValue} ${formatShapValue(column.total)}`,
645
+ onMouseEnter: () => setHoveredColumn(columnIndex),
646
+ onMouseLeave: () => setHoveredColumn(null),
647
+ onClick: () => {
648
+ if (column.sampleId !== void 0) onSampleClick == null ? void 0 : onSampleClick(column.sampleId);
649
+ },
650
+ style: { cursor: column.sampleId && onSampleClick ? "pointer" : "default" }
651
+ },
652
+ `column-hit-${column.sampleIndex}`
653
+ )),
654
+ activeColumn && /* @__PURE__ */ jsxs5("g", { pointerEvents: "none", transform: `translate(${tooltipX} 10)`, children: [
655
+ /* @__PURE__ */ jsx5("rect", { x: 0, y: 0, width: tooltipWidth, height: 42, rx: 3, fill: "#ffffff", stroke: "#cccccc" }),
656
+ /* @__PURE__ */ jsx5("text", { x: 7, y: 15, fontSize: 11, fill: "#222222", children: nameOf(activeColumn) }),
657
+ /* @__PURE__ */ jsx5("text", { x: 7, y: 31, fontSize: 11, fill: "#222222", children: `${words.sampleTotal}: ${formatShapValue(activeColumn.total)}` })
658
+ ] })
659
+ ]
660
+ }
661
+ );
619
662
  }
620
663
 
621
664
  // src/react/ShapWaterfall.tsx
@@ -637,11 +680,14 @@ function ShapWaterfall({
637
680
  const [hovered, setHovered] = useState4(null);
638
681
  const words = useMemo4(() => resolveLabels(labels), [labels]);
639
682
  const marginTop = 34;
640
- const layout = useMemo4(() => {
683
+ const { layout, sampleName } = useMemo4(() => {
684
+ var _a;
641
685
  const raw = parseExplanation(explanation, { classIndex });
642
686
  const parsed = groupByGenus ? groupExplanationByGenus(raw) : raw;
643
687
  const rows = waterfallRows(parsed, sampleIndex, maxDisplay, faithfulOtherRow, words);
644
- return waterfallLayout(rows, {
688
+ const sampleLabel = (_a = raw.sampleLabels) == null ? void 0 : _a[sampleIndex];
689
+ const sampleName2 = sampleLabel ? raw.sampleLabelColumn ? `${raw.sampleLabelColumn}: ${sampleLabel}` : sampleLabel : words.sampleFallback(sampleIndex + 1);
690
+ const layout2 = waterfallLayout(rows, {
645
691
  width,
646
692
  rowHeight,
647
693
  marginLeft: 260,
@@ -650,6 +696,7 @@ function ShapWaterfall({
650
696
  decimals,
651
697
  labels: words
652
698
  });
699
+ return { layout: layout2, sampleName: sampleName2 };
653
700
  }, [
654
701
  groupByGenus,
655
702
  explanation,
@@ -668,7 +715,7 @@ function ShapWaterfall({
668
715
  width,
669
716
  height: layout.height,
670
717
  role: "img",
671
- "aria-label": `Local SHAP waterfall for Sample ${sampleIndex}`,
718
+ "aria-label": `${words.shapValue} of each feature for ${sampleName}`,
672
719
  children: [
673
720
  layout.separators.map((separator, index) => /* @__PURE__ */ jsx6(
674
721
  "line",