shaply 1.0.0__py3-none-any.whl
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.
- shaply/__init__.py +110 -0
- shaply/colors.py +96 -0
- shaply/config.py +241 -0
- shaply/enums.py +69 -0
- shaply/explanation.py +279 -0
- shaply/interaction.py +98 -0
- shaply/plots/__init__.py +52 -0
- shaply/plots/_common/__init__.py +8 -0
- shaply/plots/_common/cluster.py +122 -0
- shaply/plots/_common/layout.py +42 -0
- shaply/plots/_common/ordering.py +130 -0
- shaply/plots/_common/stats.py +42 -0
- shaply/plots/advanced/__init__.py +32 -0
- shaply/plots/advanced/beeswarm_ranges.py +197 -0
- shaply/plots/advanced/error_analysis.py +151 -0
- shaply/plots/advanced/explanation_archetypes.py +116 -0
- shaply/plots/advanced/feature_clustering.py +115 -0
- shaply/plots/advanced/importance_by_cohort.py +146 -0
- shaply/plots/advanced/importance_ci.py +118 -0
- shaply/plots/advanced/interaction_heatmap.py +93 -0
- shaply/plots/advanced/monotonicity.py +108 -0
- shaply/plots/advanced/response_curve.py +200 -0
- shaply/plots/advanced/shap_surface.py +148 -0
- shaply/plots/usual/__init__.py +21 -0
- shaply/plots/usual/bar.py +111 -0
- shaply/plots/usual/beeswarm.py +214 -0
- shaply/plots/usual/decision.py +140 -0
- shaply/plots/usual/force.py +164 -0
- shaply/plots/usual/heatmap.py +104 -0
- shaply/plots/usual/scatter.py +122 -0
- shaply/plots/usual/waterfall.py +129 -0
- shaply/py.typed +0 -0
- shaply-1.0.0.dist-info/METADATA +211 -0
- shaply-1.0.0.dist-info/RECORD +35 -0
- shaply-1.0.0.dist-info/WHEEL +4 -0
shaply/__init__.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""shaply - usual SHAP explainability figures rendered as Plotly charts.
|
|
2
|
+
|
|
3
|
+
The public API mirrors the familiar ``shap.plots`` entry points but returns
|
|
4
|
+
:class:`plotly.graph_objects.Figure` objects instead of matplotlib axes::
|
|
5
|
+
|
|
6
|
+
import shaply
|
|
7
|
+
|
|
8
|
+
fig = shaply.beeswarm(shap_values) # a shap.Explanation, ndarray or DataFrame
|
|
9
|
+
fig.show()
|
|
10
|
+
|
|
11
|
+
Every plotting function accepts a ``shap.Explanation``-like object, a numpy
|
|
12
|
+
array of SHAP values, or a :class:`pandas.DataFrame`, and an optional typed
|
|
13
|
+
config object from :mod:`shaply.config`.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
19
|
+
|
|
20
|
+
from shaply.config import (
|
|
21
|
+
BarConfig,
|
|
22
|
+
BeeswarmConfig,
|
|
23
|
+
BeeswarmRangesConfig,
|
|
24
|
+
DecisionConfig,
|
|
25
|
+
ErrorAnalysisConfig,
|
|
26
|
+
ExplanationArchetypesConfig,
|
|
27
|
+
FeatureClusteringConfig,
|
|
28
|
+
ForceConfig,
|
|
29
|
+
HeatmapConfig,
|
|
30
|
+
ImportanceByCohortConfig,
|
|
31
|
+
ImportanceCIConfig,
|
|
32
|
+
InteractionHeatmapConfig,
|
|
33
|
+
MonotonicityConfig,
|
|
34
|
+
ResponseCurveConfig,
|
|
35
|
+
ScatterConfig,
|
|
36
|
+
ShapSurfaceConfig,
|
|
37
|
+
WaterfallConfig,
|
|
38
|
+
)
|
|
39
|
+
from shaply.enums import ColorScale, FeatureOrdering, PlotType
|
|
40
|
+
from shaply.explanation import Explanation, to_explanation
|
|
41
|
+
from shaply.interaction import InteractionValues, to_interaction_values
|
|
42
|
+
from shaply.plots import (
|
|
43
|
+
bar,
|
|
44
|
+
beeswarm,
|
|
45
|
+
beeswarm_ranges,
|
|
46
|
+
decision,
|
|
47
|
+
error_analysis,
|
|
48
|
+
explanation_archetypes,
|
|
49
|
+
feature_clustering,
|
|
50
|
+
force,
|
|
51
|
+
heatmap,
|
|
52
|
+
importance_by_cohort,
|
|
53
|
+
importance_ci,
|
|
54
|
+
interaction_heatmap,
|
|
55
|
+
monotonicity_check,
|
|
56
|
+
response_curve,
|
|
57
|
+
scatter,
|
|
58
|
+
shap_surface,
|
|
59
|
+
waterfall,
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
try:
|
|
63
|
+
__version__ = version("shaply")
|
|
64
|
+
except PackageNotFoundError: # pragma: no cover - only during local dev without install
|
|
65
|
+
__version__ = "0.0.0"
|
|
66
|
+
|
|
67
|
+
__all__ = [
|
|
68
|
+
"BarConfig",
|
|
69
|
+
"BeeswarmConfig",
|
|
70
|
+
"BeeswarmRangesConfig",
|
|
71
|
+
"ColorScale",
|
|
72
|
+
"DecisionConfig",
|
|
73
|
+
"ErrorAnalysisConfig",
|
|
74
|
+
"Explanation",
|
|
75
|
+
"ExplanationArchetypesConfig",
|
|
76
|
+
"FeatureClusteringConfig",
|
|
77
|
+
"FeatureOrdering",
|
|
78
|
+
"ForceConfig",
|
|
79
|
+
"HeatmapConfig",
|
|
80
|
+
"ImportanceByCohortConfig",
|
|
81
|
+
"ImportanceCIConfig",
|
|
82
|
+
"InteractionHeatmapConfig",
|
|
83
|
+
"InteractionValues",
|
|
84
|
+
"MonotonicityConfig",
|
|
85
|
+
"PlotType",
|
|
86
|
+
"ResponseCurveConfig",
|
|
87
|
+
"ScatterConfig",
|
|
88
|
+
"ShapSurfaceConfig",
|
|
89
|
+
"WaterfallConfig",
|
|
90
|
+
"__version__",
|
|
91
|
+
"bar",
|
|
92
|
+
"beeswarm",
|
|
93
|
+
"beeswarm_ranges",
|
|
94
|
+
"decision",
|
|
95
|
+
"error_analysis",
|
|
96
|
+
"explanation_archetypes",
|
|
97
|
+
"feature_clustering",
|
|
98
|
+
"force",
|
|
99
|
+
"heatmap",
|
|
100
|
+
"importance_by_cohort",
|
|
101
|
+
"importance_ci",
|
|
102
|
+
"interaction_heatmap",
|
|
103
|
+
"monotonicity_check",
|
|
104
|
+
"response_curve",
|
|
105
|
+
"scatter",
|
|
106
|
+
"shap_surface",
|
|
107
|
+
"to_explanation",
|
|
108
|
+
"to_interaction_values",
|
|
109
|
+
"waterfall",
|
|
110
|
+
]
|
shaply/colors.py
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Color utilities reproducing SHAP's visual language for Plotly figures.
|
|
2
|
+
|
|
3
|
+
SHAP relies on a red/blue diverging scheme where blue encodes low feature
|
|
4
|
+
values and red encodes high ones. The exact anchor colors used here match the
|
|
5
|
+
defaults shipped by the ``shap`` package so figures feel familiar to its users.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from typing import Final
|
|
11
|
+
|
|
12
|
+
from shaply.enums import ColorScale
|
|
13
|
+
|
|
14
|
+
#: Canonical SHAP accent colors.
|
|
15
|
+
SHAP_RED: Final = "#ff0d57"
|
|
16
|
+
SHAP_BLUE: Final = "#008bff"
|
|
17
|
+
SHAP_GRAY: Final = "#777777"
|
|
18
|
+
|
|
19
|
+
#: Plotly colorscale definitions as ``(position, css_color)`` stops.
|
|
20
|
+
_ColorStop = tuple[float, str]
|
|
21
|
+
Colorscale = list[_ColorStop]
|
|
22
|
+
|
|
23
|
+
_RED_BLUE: Final[Colorscale] = [
|
|
24
|
+
(0.0, SHAP_BLUE),
|
|
25
|
+
(0.5, "#c6b7d6"),
|
|
26
|
+
(1.0, SHAP_RED),
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
_COOLWARM: Final[Colorscale] = [
|
|
30
|
+
(0.0, "#3b4cc0"),
|
|
31
|
+
(0.5, "#dddddd"),
|
|
32
|
+
(1.0, "#b40426"),
|
|
33
|
+
]
|
|
34
|
+
|
|
35
|
+
#: Sequential white->red scale for non-negative magnitudes (e.g. |interaction|).
|
|
36
|
+
_REDS: Final[Colorscale] = [
|
|
37
|
+
(0.0, "#fff5f0"),
|
|
38
|
+
(0.5, "#fca082"),
|
|
39
|
+
(1.0, SHAP_RED),
|
|
40
|
+
]
|
|
41
|
+
|
|
42
|
+
_SCALE_REGISTRY: Final[dict[ColorScale, Colorscale | str]] = {
|
|
43
|
+
ColorScale.RED_BLUE: _RED_BLUE,
|
|
44
|
+
ColorScale.COOLWARM: _COOLWARM,
|
|
45
|
+
ColorScale.REDS: _REDS,
|
|
46
|
+
ColorScale.VIRIDIS: "Viridis",
|
|
47
|
+
ColorScale.PLASMA: "Plasma",
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def resolve_colorscale(scale: ColorScale) -> Colorscale | str:
|
|
52
|
+
"""Return the Plotly-compatible colorscale for a :class:`ColorScale`.
|
|
53
|
+
|
|
54
|
+
Parameters
|
|
55
|
+
----------
|
|
56
|
+
scale
|
|
57
|
+
The named color scale to resolve.
|
|
58
|
+
|
|
59
|
+
Returns
|
|
60
|
+
-------
|
|
61
|
+
list of tuple or str
|
|
62
|
+
Either an explicit list of ``(position, color)`` stops or the name of a
|
|
63
|
+
built-in Plotly colorscale.
|
|
64
|
+
"""
|
|
65
|
+
return _SCALE_REGISTRY[scale]
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def sample_scale(scale: ColorScale, positions: list[float]) -> list[str]:
|
|
69
|
+
"""Sample a color scale at ``positions`` in ``[0, 1]``.
|
|
70
|
+
|
|
71
|
+
Parameters
|
|
72
|
+
----------
|
|
73
|
+
scale
|
|
74
|
+
The named color scale to sample.
|
|
75
|
+
positions
|
|
76
|
+
Fractions in ``[0, 1]`` at which to read the scale.
|
|
77
|
+
|
|
78
|
+
Returns
|
|
79
|
+
-------
|
|
80
|
+
list of str
|
|
81
|
+
One CSS ``rgb(...)`` color per requested position.
|
|
82
|
+
"""
|
|
83
|
+
from plotly.colors import sample_colorscale
|
|
84
|
+
|
|
85
|
+
resolved = resolve_colorscale(scale)
|
|
86
|
+
colorscale = resolved if isinstance(resolved, str) else [list(stop) for stop in resolved]
|
|
87
|
+
return list(sample_colorscale(colorscale, positions, colortype="rgb"))
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def signed_color(value: float) -> str:
|
|
91
|
+
"""Return the SHAP accent color matching the sign of ``value``.
|
|
92
|
+
|
|
93
|
+
Positive contributions map to red, negative ones to blue, matching the
|
|
94
|
+
waterfall and force conventions of ``shap``.
|
|
95
|
+
"""
|
|
96
|
+
return SHAP_RED if value >= 0 else SHAP_BLUE
|
shaply/config.py
ADDED
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""Public, validated plot configuration models (pydantic v2).
|
|
2
|
+
|
|
3
|
+
These are the external data models users can pass to tune figures. They are
|
|
4
|
+
deliberately kept separate from the internal :mod:`shaply.explanation`
|
|
5
|
+
dataclasses: pydantic validates and normalizes user input, dataclasses carry
|
|
6
|
+
already-trusted internal state.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from pydantic import BaseModel, ConfigDict, Field
|
|
12
|
+
|
|
13
|
+
from shaply.enums import ColorScale, FeatureOrdering
|
|
14
|
+
|
|
15
|
+
_DEFAULT_TEMPLATE = "plotly_white"
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class BasePlotConfig(BaseModel):
|
|
19
|
+
"""Options shared by every ``shaply`` figure."""
|
|
20
|
+
|
|
21
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
22
|
+
|
|
23
|
+
title: str | None = Field(default=None, description="Figure title.")
|
|
24
|
+
width: int | None = Field(default=None, gt=0, description="Figure width in pixels.")
|
|
25
|
+
height: int | None = Field(default=None, gt=0, description="Figure height in pixels.")
|
|
26
|
+
template: str = Field(default=_DEFAULT_TEMPLATE, description="Plotly layout template.")
|
|
27
|
+
show_grid: bool = Field(default=True, description="Whether to draw axis gridlines.")
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class BarConfig(BasePlotConfig):
|
|
31
|
+
"""Configuration for the global feature-importance bar plot."""
|
|
32
|
+
|
|
33
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
34
|
+
ordering: FeatureOrdering = Field(
|
|
35
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
36
|
+
description="How to order features along the axis.",
|
|
37
|
+
)
|
|
38
|
+
show_values: bool = Field(default=True, description="Annotate each bar with its numeric value.")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class BeeswarmConfig(BasePlotConfig):
|
|
42
|
+
"""Configuration for the beeswarm (summary) plot."""
|
|
43
|
+
|
|
44
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
45
|
+
ordering: FeatureOrdering = Field(
|
|
46
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
47
|
+
description="How to order features along the axis.",
|
|
48
|
+
)
|
|
49
|
+
color_scale: ColorScale = Field(
|
|
50
|
+
default=ColorScale.RED_BLUE,
|
|
51
|
+
description="Color scale encoding feature values.",
|
|
52
|
+
)
|
|
53
|
+
point_size: float = Field(default=5.0, gt=0, description="Marker size in pixels.")
|
|
54
|
+
jitter: float = Field(
|
|
55
|
+
default=0.35, ge=0.0, le=1.0, description="Vertical spread of overlapping points."
|
|
56
|
+
)
|
|
57
|
+
opacity: float = Field(default=0.8, gt=0.0, le=1.0, description="Marker opacity.")
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class BeeswarmRangesConfig(BeeswarmConfig):
|
|
61
|
+
"""Configuration for the combined beeswarm + feature-value-ranges figure.
|
|
62
|
+
|
|
63
|
+
Extends :class:`BeeswarmConfig` with the options controlling the right-hand
|
|
64
|
+
panel that shows the real distribution (violin + box) of each feature's
|
|
65
|
+
values.
|
|
66
|
+
"""
|
|
67
|
+
|
|
68
|
+
impact_panel_ratio: float = Field(
|
|
69
|
+
default=0.62,
|
|
70
|
+
gt=0.2,
|
|
71
|
+
lt=0.9,
|
|
72
|
+
description="Fraction of width given to the left (SHAP impact) panel.",
|
|
73
|
+
)
|
|
74
|
+
show_value_labels: bool = Field(
|
|
75
|
+
default=True,
|
|
76
|
+
description="Annotate each feature row with its real min and max value.",
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
class WaterfallConfig(BasePlotConfig):
|
|
81
|
+
"""Configuration for the single-prediction waterfall plot."""
|
|
82
|
+
|
|
83
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
84
|
+
show_values: bool = Field(default=True, description="Annotate each step with its contribution.")
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class ScatterConfig(BasePlotConfig):
|
|
88
|
+
"""Configuration for the dependence / scatter plot."""
|
|
89
|
+
|
|
90
|
+
color_scale: ColorScale = Field(
|
|
91
|
+
default=ColorScale.RED_BLUE,
|
|
92
|
+
description="Color scale used when coloring by an interaction feature.",
|
|
93
|
+
)
|
|
94
|
+
point_size: float = Field(default=6.0, gt=0, description="Marker size in pixels.")
|
|
95
|
+
opacity: float = Field(default=0.8, gt=0.0, le=1.0, description="Marker opacity.")
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
class HeatmapConfig(BasePlotConfig):
|
|
99
|
+
"""Configuration for the instances-by-features heatmap."""
|
|
100
|
+
|
|
101
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
102
|
+
ordering: FeatureOrdering = Field(
|
|
103
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
104
|
+
description="How to order features along the axis.",
|
|
105
|
+
)
|
|
106
|
+
color_scale: ColorScale = Field(
|
|
107
|
+
default=ColorScale.RED_BLUE,
|
|
108
|
+
description="Color scale encoding SHAP values.",
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
class ResponseCurveConfig(BasePlotConfig):
|
|
113
|
+
"""Configuration for the response-curve plot (smoothed effect + thresholds)."""
|
|
114
|
+
|
|
115
|
+
n_bins: int = Field(default=20, gt=2, description="Number of bins used to smooth.")
|
|
116
|
+
quantile_bins: bool = Field(
|
|
117
|
+
default=True,
|
|
118
|
+
description="Use equal-count quantile bins (stable) instead of equal-width.",
|
|
119
|
+
)
|
|
120
|
+
band: bool = Field(default=True, description="Draw a +/-1 std spread band.")
|
|
121
|
+
show_points: bool = Field(default=True, description="Overlay the raw scatter points.")
|
|
122
|
+
point_opacity: float = Field(default=0.25, gt=0.0, le=1.0, description="Raw point opacity.")
|
|
123
|
+
show_thresholds: bool = Field(
|
|
124
|
+
default=True, description="Mark the values where the mean effect crosses zero."
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
class InteractionHeatmapConfig(BasePlotConfig):
|
|
129
|
+
"""Configuration for the pairwise SHAP-interaction heatmap."""
|
|
130
|
+
|
|
131
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
132
|
+
color_scale: ColorScale = Field(
|
|
133
|
+
default=ColorScale.REDS,
|
|
134
|
+
description="Sequential color scale encoding interaction magnitude.",
|
|
135
|
+
)
|
|
136
|
+
show_diagonal: bool = Field(
|
|
137
|
+
default=False,
|
|
138
|
+
description="Keep the diagonal (main effects); hidden by default to reveal interactions.",
|
|
139
|
+
)
|
|
140
|
+
show_values: bool = Field(default=False, description="Annotate each cell with its value.")
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
class ErrorAnalysisConfig(BasePlotConfig):
|
|
144
|
+
"""Configuration for the error-analysis plot (SHAP drivers of errors)."""
|
|
145
|
+
|
|
146
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
147
|
+
show_values: bool = Field(default=True, description="Annotate each bar with its value.")
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
class ShapSurfaceConfig(BasePlotConfig):
|
|
151
|
+
"""Configuration for the 2D SHAP interaction surface."""
|
|
152
|
+
|
|
153
|
+
n_bins: int = Field(default=20, gt=2, description="Grid resolution per axis.")
|
|
154
|
+
color_scale: ColorScale = Field(
|
|
155
|
+
default=ColorScale.RED_BLUE,
|
|
156
|
+
description="Diverging color scale encoding the mean SHAP value.",
|
|
157
|
+
)
|
|
158
|
+
contour: bool = Field(default=False, description="Render smooth contours instead of a heatmap.")
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
class ImportanceByCohortConfig(BasePlotConfig):
|
|
162
|
+
"""Configuration for the cohort-wise importance plot."""
|
|
163
|
+
|
|
164
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
165
|
+
ordering: FeatureOrdering = Field(
|
|
166
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
167
|
+
description="How to order features (by overall importance across cohorts).",
|
|
168
|
+
)
|
|
169
|
+
n_cohorts: int = Field(
|
|
170
|
+
default=3, gt=1, description="Number of quantile cohorts when splitting by a feature."
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
class FeatureClusteringConfig(BasePlotConfig):
|
|
175
|
+
"""Configuration for the SHAP-similarity feature-clustering heatmap."""
|
|
176
|
+
|
|
177
|
+
color_scale: ColorScale = Field(
|
|
178
|
+
default=ColorScale.RED_BLUE,
|
|
179
|
+
description="Diverging color scale encoding SHAP correlation.",
|
|
180
|
+
)
|
|
181
|
+
show_values: bool = Field(default=False, description="Annotate each cell with its value.")
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
class ExplanationArchetypesConfig(BasePlotConfig):
|
|
185
|
+
"""Configuration for the SHAP-profile archetypes plot."""
|
|
186
|
+
|
|
187
|
+
n_clusters: int = Field(default=4, gt=1, description="Number of archetypes to extract.")
|
|
188
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
189
|
+
ordering: FeatureOrdering = Field(
|
|
190
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
191
|
+
description="How to order features along the axis.",
|
|
192
|
+
)
|
|
193
|
+
color_scale: ColorScale = Field(
|
|
194
|
+
default=ColorScale.RED_BLUE,
|
|
195
|
+
description="Diverging color scale encoding each archetype's mean SHAP.",
|
|
196
|
+
)
|
|
197
|
+
random_state: int = Field(default=0, description="Seed for the k-means clustering.")
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
class ImportanceCIConfig(BasePlotConfig):
|
|
201
|
+
"""Configuration for the bootstrap importance plot with confidence intervals."""
|
|
202
|
+
|
|
203
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
204
|
+
ordering: FeatureOrdering = Field(
|
|
205
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
206
|
+
description="How to order features along the axis.",
|
|
207
|
+
)
|
|
208
|
+
n_boot: int = Field(default=1000, gt=10, description="Number of bootstrap resamples.")
|
|
209
|
+
ci: float = Field(default=0.95, gt=0.0, lt=1.0, description="Confidence level.")
|
|
210
|
+
random_state: int = Field(default=0, description="Seed for the bootstrap resampling.")
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
class MonotonicityConfig(BasePlotConfig):
|
|
214
|
+
"""Configuration for the monotonicity-check plot."""
|
|
215
|
+
|
|
216
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
class ForceConfig(BasePlotConfig):
|
|
220
|
+
"""Configuration for the single-prediction force plot."""
|
|
221
|
+
|
|
222
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
223
|
+
show_values: bool = Field(
|
|
224
|
+
default=True, description="Annotate each segment with its contribution."
|
|
225
|
+
)
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
class DecisionConfig(BasePlotConfig):
|
|
229
|
+
"""Configuration for the decision plot."""
|
|
230
|
+
|
|
231
|
+
max_display: int = Field(default=10, gt=0, description="Max features to display.")
|
|
232
|
+
ordering: FeatureOrdering = Field(
|
|
233
|
+
default=FeatureOrdering.IMPORTANCE,
|
|
234
|
+
description="How to order features along the axis.",
|
|
235
|
+
)
|
|
236
|
+
color_scale: ColorScale = Field(
|
|
237
|
+
default=ColorScale.RED_BLUE,
|
|
238
|
+
description="Color scale encoding each instance's predicted output.",
|
|
239
|
+
)
|
|
240
|
+
line_width: float = Field(default=1.5, gt=0, description="Line width in pixels.")
|
|
241
|
+
opacity: float = Field(default=0.8, gt=0.0, le=1.0, description="Line opacity.")
|
shaply/enums.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
"""Enumerations used across ``shaply`` to keep public options explicit and typed."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from enum import IntEnum, StrEnum
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class PlotType(StrEnum):
|
|
9
|
+
"""Kind of SHAP figure that ``shaply`` can render."""
|
|
10
|
+
|
|
11
|
+
BAR = "bar"
|
|
12
|
+
BEESWARM = "beeswarm"
|
|
13
|
+
BEESWARM_RANGES = "beeswarm_ranges"
|
|
14
|
+
WATERFALL = "waterfall"
|
|
15
|
+
SCATTER = "scatter"
|
|
16
|
+
HEATMAP = "heatmap"
|
|
17
|
+
FORCE = "force"
|
|
18
|
+
DECISION = "decision"
|
|
19
|
+
RESPONSE_CURVE = "response_curve"
|
|
20
|
+
INTERACTION_HEATMAP = "interaction_heatmap"
|
|
21
|
+
ERROR_ANALYSIS = "error_analysis"
|
|
22
|
+
SHAP_SURFACE = "shap_surface"
|
|
23
|
+
IMPORTANCE_BY_COHORT = "importance_by_cohort"
|
|
24
|
+
FEATURE_CLUSTERING = "feature_clustering"
|
|
25
|
+
EXPLANATION_ARCHETYPES = "explanation_archetypes"
|
|
26
|
+
IMPORTANCE_CI = "importance_ci"
|
|
27
|
+
MONOTONICITY = "monotonicity"
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class FeatureOrdering(StrEnum):
|
|
31
|
+
"""Strategy used to order features along the categorical axis of a plot.
|
|
32
|
+
|
|
33
|
+
Attributes
|
|
34
|
+
----------
|
|
35
|
+
IMPORTANCE
|
|
36
|
+
Order by mean absolute SHAP value (most important first).
|
|
37
|
+
MAX_ABSOLUTE
|
|
38
|
+
Order by the single largest absolute SHAP value across samples.
|
|
39
|
+
ORIGINAL
|
|
40
|
+
Keep the order of ``feature_names`` as provided.
|
|
41
|
+
ALPHABETICAL
|
|
42
|
+
Order features alphabetically by name.
|
|
43
|
+
"""
|
|
44
|
+
|
|
45
|
+
IMPORTANCE = "importance"
|
|
46
|
+
MAX_ABSOLUTE = "max_absolute"
|
|
47
|
+
ORIGINAL = "original"
|
|
48
|
+
ALPHABETICAL = "alphabetical"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class ColorScale(StrEnum):
|
|
52
|
+
"""Named color scales available for continuous encodings.
|
|
53
|
+
|
|
54
|
+
``RED_BLUE`` reproduces the canonical SHAP diverging scheme
|
|
55
|
+
(blue for low feature values, red for high).
|
|
56
|
+
"""
|
|
57
|
+
|
|
58
|
+
RED_BLUE = "red_blue"
|
|
59
|
+
VIRIDIS = "viridis"
|
|
60
|
+
PLASMA = "plasma"
|
|
61
|
+
COOLWARM = "coolwarm"
|
|
62
|
+
REDS = "reds"
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class SortDirection(IntEnum):
|
|
66
|
+
"""Direction used when sorting numeric quantities."""
|
|
67
|
+
|
|
68
|
+
ASCENDING = 1
|
|
69
|
+
DESCENDING = -1
|