uxplain 0.3.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.
- uxplain/__init__.py +82 -0
- uxplain/conformal/__init__.py +19 -0
- uxplain/conformal/_validation.py +34 -0
- uxplain/conformal/cqr_predictor.py +145 -0
- uxplain/conformal/crepes_classifier.py +186 -0
- uxplain/conformal/crepes_predictor.py +180 -0
- uxplain/datasets.py +526 -0
- uxplain/explainability/__init__.py +19 -0
- uxplain/explainability/fast_shap.py +223 -0
- uxplain/explainability/lime_explainer.py +170 -0
- uxplain/explainability/pdp_explainer.py +336 -0
- uxplain/explainability/shap_explainer.py +161 -0
- uxplain/plots.py +357 -0
- uxplain/plotting.py +746 -0
- uxplain/protocols.py +79 -0
- uxplain/uncertainty/__init__.py +23 -0
- uxplain/uncertainty/metrics.py +215 -0
- uxplain/uq_explainer.py +696 -0
- uxplain-0.3.0.dist-info/METADATA +134 -0
- uxplain-0.3.0.dist-info/RECORD +22 -0
- uxplain-0.3.0.dist-info/WHEEL +4 -0
- uxplain-0.3.0.dist-info/licenses/LICENSE +20 -0
uxplain/__init__.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
"""uxplain: explainability for conformal prediction uncertainty."""
|
|
2
|
+
|
|
3
|
+
__version__ = "0.3.0"
|
|
4
|
+
|
|
5
|
+
from uxplain.conformal.cqr_predictor import CQRConformalPredictor
|
|
6
|
+
from uxplain.conformal.crepes_classifier import (
|
|
7
|
+
ClassificationConformalMethod,
|
|
8
|
+
CrepesConformalClassifier,
|
|
9
|
+
)
|
|
10
|
+
from uxplain.conformal.crepes_predictor import (
|
|
11
|
+
ConformalMethod,
|
|
12
|
+
CrepesConformalPredictor,
|
|
13
|
+
)
|
|
14
|
+
from uxplain.explainability.lime_explainer import (
|
|
15
|
+
LIMEExplanation,
|
|
16
|
+
LimeUncertaintyExplainer,
|
|
17
|
+
)
|
|
18
|
+
from uxplain.explainability.pdp_explainer import (
|
|
19
|
+
PDPExplanation,
|
|
20
|
+
PDPUncertaintyExplainer,
|
|
21
|
+
)
|
|
22
|
+
from uxplain.explainability.shap_explainer import ShapUncertaintyExplainer
|
|
23
|
+
from uxplain.plotting import (
|
|
24
|
+
STYLE,
|
|
25
|
+
ice_curves,
|
|
26
|
+
lime_local,
|
|
27
|
+
pdp_curve,
|
|
28
|
+
pdp_importance,
|
|
29
|
+
pdp_interaction,
|
|
30
|
+
pdp_with_ice,
|
|
31
|
+
shap_bar,
|
|
32
|
+
shap_beeswarm,
|
|
33
|
+
shap_waterfall,
|
|
34
|
+
)
|
|
35
|
+
from uxplain.protocols import (
|
|
36
|
+
ConformalClassifierProtocol,
|
|
37
|
+
ConformalPredictorProtocol,
|
|
38
|
+
UncertaintyExplainerProtocol,
|
|
39
|
+
)
|
|
40
|
+
from uxplain.uncertainty.metrics import (
|
|
41
|
+
ClassificationMetric,
|
|
42
|
+
RegressionMetric,
|
|
43
|
+
UncertaintyMetric,
|
|
44
|
+
)
|
|
45
|
+
from uxplain.uq_explainer import (
|
|
46
|
+
ClassificationExplanationResult,
|
|
47
|
+
ExplanationResult,
|
|
48
|
+
UncertaintyExplanationPipeline,
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
__all__ = [
|
|
52
|
+
"STYLE",
|
|
53
|
+
"CQRConformalPredictor",
|
|
54
|
+
"ClassificationConformalMethod",
|
|
55
|
+
"ClassificationExplanationResult",
|
|
56
|
+
"ClassificationMetric",
|
|
57
|
+
"ConformalClassifierProtocol",
|
|
58
|
+
"ConformalMethod",
|
|
59
|
+
"ConformalPredictorProtocol",
|
|
60
|
+
"CrepesConformalClassifier",
|
|
61
|
+
"CrepesConformalPredictor",
|
|
62
|
+
"ExplanationResult",
|
|
63
|
+
"LIMEExplanation",
|
|
64
|
+
"LimeUncertaintyExplainer",
|
|
65
|
+
"PDPExplanation",
|
|
66
|
+
"PDPUncertaintyExplainer",
|
|
67
|
+
"RegressionMetric",
|
|
68
|
+
"ShapUncertaintyExplainer",
|
|
69
|
+
"UncertaintyExplainerProtocol",
|
|
70
|
+
"UncertaintyExplanationPipeline",
|
|
71
|
+
"UncertaintyMetric",
|
|
72
|
+
"__version__",
|
|
73
|
+
"ice_curves",
|
|
74
|
+
"lime_local",
|
|
75
|
+
"pdp_curve",
|
|
76
|
+
"pdp_importance",
|
|
77
|
+
"pdp_interaction",
|
|
78
|
+
"pdp_with_ice",
|
|
79
|
+
"shap_bar",
|
|
80
|
+
"shap_beeswarm",
|
|
81
|
+
"shap_waterfall",
|
|
82
|
+
]
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
from uxplain.conformal.cqr_predictor import (
|
|
2
|
+
CQRConformalPredictor,
|
|
3
|
+
)
|
|
4
|
+
from uxplain.conformal.crepes_classifier import (
|
|
5
|
+
ClassificationConformalMethod,
|
|
6
|
+
CrepesConformalClassifier,
|
|
7
|
+
)
|
|
8
|
+
from uxplain.conformal.crepes_predictor import (
|
|
9
|
+
ConformalMethod,
|
|
10
|
+
CrepesConformalPredictor,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"CQRConformalPredictor",
|
|
15
|
+
"ClassificationConformalMethod",
|
|
16
|
+
"ConformalMethod",
|
|
17
|
+
"CrepesConformalClassifier",
|
|
18
|
+
"CrepesConformalPredictor",
|
|
19
|
+
]
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
"""Shared validation for the single-output conformal backends."""
|
|
2
|
+
|
|
3
|
+
from numbers import Real
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def check_confidence(confidence):
|
|
9
|
+
if not isinstance(confidence, Real) or not 0 < confidence < 1:
|
|
10
|
+
raise ValueError(
|
|
11
|
+
f"confidence must be in (0, 1), got {confidence}. Pass the "
|
|
12
|
+
"coverage level as a fraction, e.g. 0.9 for 90%."
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def check_fit_data(X_train, y_train, X_calib, y_calib):
|
|
17
|
+
"""Prevent broadcasting of column targets and mismatched calibration data."""
|
|
18
|
+
arrays = []
|
|
19
|
+
for X, y, name in (
|
|
20
|
+
(X_train, y_train, "train"), (X_calib, y_calib, "calib")
|
|
21
|
+
):
|
|
22
|
+
X, y = np.asarray(X), np.asarray(y)
|
|
23
|
+
if X.ndim != 2 or min(X.shape) == 0:
|
|
24
|
+
raise ValueError(f"X_{name} must be a non-empty 2D array.")
|
|
25
|
+
if y.ndim == 2 and y.shape[1] == 1:
|
|
26
|
+
y = y[:, 0]
|
|
27
|
+
if y.ndim != 1:
|
|
28
|
+
raise ValueError(f"y_{name} must be one-dimensional (single output).")
|
|
29
|
+
if len(X) != len(y):
|
|
30
|
+
raise ValueError(f"X_{name} and y_{name} must have the same length.")
|
|
31
|
+
arrays.extend((X, y))
|
|
32
|
+
if arrays[0].shape[1] != arrays[2].shape[1]:
|
|
33
|
+
raise ValueError("X_train and X_calib must have the same number of features.")
|
|
34
|
+
return tuple(arrays)
|
|
@@ -0,0 +1,145 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Conformalized Quantile Regression predictor.
|
|
3
|
+
|
|
4
|
+
Implements CQR (Romano, Patterson & Candès, 2019).
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from ._validation import check_confidence, check_fit_data
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class CQRConformalPredictor:
|
|
15
|
+
"""
|
|
16
|
+
Conformalized Quantile Regression (CQR) predictor.
|
|
17
|
+
|
|
18
|
+
Fits two quantile regressors — one for the lower bound and one
|
|
19
|
+
for the upper bound — then calibrates them using nonconformity
|
|
20
|
+
scores on a held-out calibration set.
|
|
21
|
+
|
|
22
|
+
Parameters
|
|
23
|
+
----------
|
|
24
|
+
lower_model
|
|
25
|
+
Quantile regressor trained at quantile level ``alpha / 2``.
|
|
26
|
+
Must expose a sklearn-compatible ``fit(X, y)`` / ``predict(X)``
|
|
27
|
+
interface. Example::
|
|
28
|
+
|
|
29
|
+
GradientBoostingRegressor(loss="quantile", alpha=0.05)
|
|
30
|
+
|
|
31
|
+
upper_model
|
|
32
|
+
Quantile regressor trained at quantile level ``1 - alpha / 2``.
|
|
33
|
+
Example::
|
|
34
|
+
|
|
35
|
+
GradientBoostingRegressor(loss="quantile", alpha=0.95)
|
|
36
|
+
|
|
37
|
+
Notes
|
|
38
|
+
-----
|
|
39
|
+
The quantile levels baked into the models should match the
|
|
40
|
+
``confidence`` value used at predict time (e.g. models at 0.05 / 0.95
|
|
41
|
+
pair with ``confidence=0.90``). Changing ``confidence`` at predict
|
|
42
|
+
time still gives valid coverage via the calibration correction, but
|
|
43
|
+
intervals may be wider or narrower than optimal if the mismatch is
|
|
44
|
+
large.
|
|
45
|
+
|
|
46
|
+
CQR may return an empty prediction set (``lower > upper``), for example
|
|
47
|
+
after a negative calibration correction or quantile crossing. Endpoints
|
|
48
|
+
are preserved, so ``upper - lower`` is a signed span in that case, not
|
|
49
|
+
the nonnegative length of the empty set.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
def __init__(self, lower_model, upper_model):
|
|
53
|
+
self.lower_model = lower_model
|
|
54
|
+
self.upper_model = upper_model
|
|
55
|
+
|
|
56
|
+
self._scores: np.ndarray | None = None
|
|
57
|
+
self._n_calib: int | None = None
|
|
58
|
+
|
|
59
|
+
def fit(
|
|
60
|
+
self,
|
|
61
|
+
X_train: np.ndarray,
|
|
62
|
+
y_train: np.ndarray,
|
|
63
|
+
X_calib: np.ndarray,
|
|
64
|
+
y_calib: np.ndarray,
|
|
65
|
+
) -> None:
|
|
66
|
+
"""
|
|
67
|
+
Fit quantile models and calibrate nonconformity scores.
|
|
68
|
+
|
|
69
|
+
Parameters
|
|
70
|
+
----------
|
|
71
|
+
X_train, y_train
|
|
72
|
+
Training data.
|
|
73
|
+
X_calib, y_calib
|
|
74
|
+
Calibration data used to compute nonconformity scores.
|
|
75
|
+
"""
|
|
76
|
+
|
|
77
|
+
self._scores = None
|
|
78
|
+
self._n_calib = None
|
|
79
|
+
X_train, y_train, X_calib, y_calib = check_fit_data(
|
|
80
|
+
X_train, y_train, X_calib, y_calib,
|
|
81
|
+
)
|
|
82
|
+
y_calib = np.asarray(y_calib, dtype=float)
|
|
83
|
+
if not np.all(np.isfinite(y_calib)):
|
|
84
|
+
raise ValueError("y_calib must contain only finite values.")
|
|
85
|
+
|
|
86
|
+
self.lower_model.fit(X_train, y_train)
|
|
87
|
+
self.upper_model.fit(X_train, y_train)
|
|
88
|
+
|
|
89
|
+
q_low = self._predict_quantile(self.lower_model, X_calib)
|
|
90
|
+
q_high = self._predict_quantile(self.upper_model, X_calib)
|
|
91
|
+
|
|
92
|
+
# CQR nonconformity score: max(q_low - y, y - q_high)
|
|
93
|
+
self._scores = np.maximum(q_low - y_calib, y_calib - q_high)
|
|
94
|
+
self._n_calib = len(y_calib)
|
|
95
|
+
|
|
96
|
+
def predict(
|
|
97
|
+
self,
|
|
98
|
+
X: np.ndarray,
|
|
99
|
+
confidence: float = 0.9,
|
|
100
|
+
) -> tuple[np.ndarray, np.ndarray]:
|
|
101
|
+
"""
|
|
102
|
+
Generate calibrated prediction intervals.
|
|
103
|
+
|
|
104
|
+
Parameters
|
|
105
|
+
----------
|
|
106
|
+
X : np.ndarray
|
|
107
|
+
confidence : float
|
|
108
|
+
Desired marginal coverage level.
|
|
109
|
+
|
|
110
|
+
Returns
|
|
111
|
+
-------
|
|
112
|
+
lower : np.ndarray
|
|
113
|
+
upper : np.ndarray
|
|
114
|
+
"""
|
|
115
|
+
|
|
116
|
+
if self._scores is None:
|
|
117
|
+
raise RuntimeError("CQRConformalPredictor not fitted. Call fit() first.")
|
|
118
|
+
|
|
119
|
+
check_confidence(confidence)
|
|
120
|
+
X = np.asarray(X)
|
|
121
|
+
|
|
122
|
+
# Conformal quantile (Romano et al., 2019): the ceil((n + 1) * confidence)-th
|
|
123
|
+
# smallest score. When that rank exceeds n, no finite adjustment
|
|
124
|
+
# guarantees coverage, so the interval is unbounded. Do not subtract
|
|
125
|
+
# a tolerance: that can lower the rank below the requested coverage.
|
|
126
|
+
n = self._n_calib
|
|
127
|
+
k = int(np.ceil((n + 1) * confidence))
|
|
128
|
+
if k > n:
|
|
129
|
+
adjustment = np.inf
|
|
130
|
+
else:
|
|
131
|
+
adjustment = float(np.sort(self._scores)[k - 1])
|
|
132
|
+
|
|
133
|
+
lower = self._predict_quantile(self.lower_model, X) - adjustment
|
|
134
|
+
upper = self._predict_quantile(self.upper_model, X) + adjustment
|
|
135
|
+
|
|
136
|
+
return lower, upper
|
|
137
|
+
|
|
138
|
+
@staticmethod
|
|
139
|
+
def _predict_quantile(model, X):
|
|
140
|
+
values = np.asarray(model.predict(X), dtype=float)
|
|
141
|
+
if values.shape != (len(X),):
|
|
142
|
+
raise ValueError("Quantile models must predict one value per sample.")
|
|
143
|
+
if not np.all(np.isfinite(values)):
|
|
144
|
+
raise ValueError("Quantile model predictions must contain only finite values.")
|
|
145
|
+
return values
|
|
@@ -0,0 +1,186 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Conformal classification module.
|
|
3
|
+
|
|
4
|
+
Wrapper around crepes.WrapClassifier for prediction sets.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Literal
|
|
10
|
+
|
|
11
|
+
from crepes import WrapClassifier
|
|
12
|
+
import numpy as np
|
|
13
|
+
|
|
14
|
+
from ._validation import check_confidence, check_fit_data
|
|
15
|
+
|
|
16
|
+
ClassificationConformalMethod = Literal[
|
|
17
|
+
"standard",
|
|
18
|
+
"class_cond",
|
|
19
|
+
"mondrian",
|
|
20
|
+
]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class CrepesConformalClassifier:
|
|
24
|
+
"""
|
|
25
|
+
Wrapper for conformal classification using crepes.
|
|
26
|
+
|
|
27
|
+
Parameters
|
|
28
|
+
----------
|
|
29
|
+
model : object
|
|
30
|
+
Any sklearn-compatible classifier exposing
|
|
31
|
+
``fit``, ``predict``, and ``predict_proba``.
|
|
32
|
+
method : ClassificationConformalMethod
|
|
33
|
+
Conformal classification method:
|
|
34
|
+
|
|
35
|
+
- ``"standard"`` : basic (marginal) conformal classifier
|
|
36
|
+
- ``"class_cond"`` : class-conditional Mondrian (per-class coverage)
|
|
37
|
+
- ``"mondrian"`` : Mondrian categorizer based on predicted class
|
|
38
|
+
random_state : int, optional
|
|
39
|
+
Seed for the tie-breaking draws of smoothed p-values. Ignored when
|
|
40
|
+
``smoothing=False``.
|
|
41
|
+
smoothing : bool, default=False
|
|
42
|
+
Use randomized tie-breaking for prediction. Even with a fixed seed,
|
|
43
|
+
smoothed p-values can depend on row order and batch size, so built-in
|
|
44
|
+
explainers require ``smoothing=False``. Non-smoothed p-values are
|
|
45
|
+
deterministic for a fixed deterministic model and conservative under
|
|
46
|
+
the conformal assumptions; smoothing does not repair invalid sampling.
|
|
47
|
+
|
|
48
|
+
Notes
|
|
49
|
+
-----
|
|
50
|
+
Coverage requires calibration and future scores to be exchangeable after
|
|
51
|
+
training, within the selected groups for conditional methods. The fitted
|
|
52
|
+
model's classes must contain the full target label space. Unknown calibration
|
|
53
|
+
labels raise an error; future labels absent from ``classes_`` cannot be
|
|
54
|
+
covered. Do not select a split or tune the model on calibration outcomes.
|
|
55
|
+
|
|
56
|
+
Examples
|
|
57
|
+
--------
|
|
58
|
+
>>> from sklearn.ensemble import RandomForestClassifier
|
|
59
|
+
>>> cp = CrepesConformalClassifier(RandomForestClassifier(), method="class_cond")
|
|
60
|
+
>>> cp.fit(X_train, y_train, X_calib, y_calib)
|
|
61
|
+
>>> pred_set = cp.predict_set(X_test, confidence=0.9) # (n, n_classes) bool
|
|
62
|
+
>>> p_values = cp.predict_p(X_test) # (n, n_classes)
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
model,
|
|
68
|
+
method: ClassificationConformalMethod = "standard",
|
|
69
|
+
random_state: int | None = None,
|
|
70
|
+
smoothing: bool = False,
|
|
71
|
+
):
|
|
72
|
+
if method not in ("standard", "class_cond", "mondrian"):
|
|
73
|
+
raise ValueError(f"Unknown conformal classification method '{method}'.")
|
|
74
|
+
self.model = model
|
|
75
|
+
self.method = method
|
|
76
|
+
self.random_state = random_state
|
|
77
|
+
self.smoothing = smoothing
|
|
78
|
+
self.wrapper = None
|
|
79
|
+
self.classes_ = None
|
|
80
|
+
|
|
81
|
+
def fit(
|
|
82
|
+
self,
|
|
83
|
+
X_train: np.ndarray,
|
|
84
|
+
y_train: np.ndarray,
|
|
85
|
+
X_calib: np.ndarray,
|
|
86
|
+
y_calib: np.ndarray,
|
|
87
|
+
**kwargs,
|
|
88
|
+
) -> None:
|
|
89
|
+
"""
|
|
90
|
+
Fit conformal classifier.
|
|
91
|
+
|
|
92
|
+
Steps:
|
|
93
|
+
1. Fit underlying model
|
|
94
|
+
2. Optionally fit Mondrian categorizer
|
|
95
|
+
3. Calibrate conformal classifier
|
|
96
|
+
"""
|
|
97
|
+
|
|
98
|
+
self.wrapper = None
|
|
99
|
+
self.classes_ = None
|
|
100
|
+
X_train, y_train, X_calib, y_calib = check_fit_data(
|
|
101
|
+
X_train, y_train, X_calib, y_calib,
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
# Train
|
|
105
|
+
self.model.fit(X_train, y_train)
|
|
106
|
+
self.classes_ = np.asarray(self.model.classes_)
|
|
107
|
+
if not np.all(np.isin(y_calib, self.classes_)):
|
|
108
|
+
raise ValueError("y_calib contains classes absent from the fitted model.")
|
|
109
|
+
|
|
110
|
+
self.wrapper = WrapClassifier(self.model)
|
|
111
|
+
|
|
112
|
+
# Build calibration kwargs based on method
|
|
113
|
+
calibrate_kwargs = {}
|
|
114
|
+
|
|
115
|
+
if self.method == "class_cond":
|
|
116
|
+
calibrate_kwargs["class_cond"] = True
|
|
117
|
+
elif self.method == "mondrian":
|
|
118
|
+
calibrate_kwargs["mc"] = self.model.predict
|
|
119
|
+
|
|
120
|
+
# Calibrate
|
|
121
|
+
self.wrapper.calibrate(
|
|
122
|
+
X_calib,
|
|
123
|
+
y_calib,
|
|
124
|
+
**calibrate_kwargs,
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
def predict_set(
|
|
128
|
+
self,
|
|
129
|
+
X: np.ndarray,
|
|
130
|
+
confidence: float = 0.9,
|
|
131
|
+
) -> np.ndarray:
|
|
132
|
+
"""
|
|
133
|
+
Generate prediction sets.
|
|
134
|
+
|
|
135
|
+
Parameters
|
|
136
|
+
----------
|
|
137
|
+
X : np.ndarray
|
|
138
|
+
confidence : float
|
|
139
|
+
|
|
140
|
+
Returns
|
|
141
|
+
-------
|
|
142
|
+
prediction_set : np.ndarray of shape (n_samples, n_classes), dtype=bool
|
|
143
|
+
``True`` at position ``[i, k]`` means class ``k`` is included in
|
|
144
|
+
the prediction set for sample ``i``.
|
|
145
|
+
"""
|
|
146
|
+
|
|
147
|
+
if self.wrapper is None:
|
|
148
|
+
raise RuntimeError("CrepesConformalClassifier not fitted.")
|
|
149
|
+
|
|
150
|
+
check_confidence(confidence)
|
|
151
|
+
X = np.asarray(X)
|
|
152
|
+
# labels=False keeps the binary-array output; since crepes 0.9.1
|
|
153
|
+
# the default (labels=True) returns a list of lists of labels.
|
|
154
|
+
return self.wrapper.predict_set(
|
|
155
|
+
X, confidence=confidence, smoothing=self.smoothing,
|
|
156
|
+
seed=self.random_state, labels=False,
|
|
157
|
+
).astype(bool)
|
|
158
|
+
|
|
159
|
+
def predict_p(
|
|
160
|
+
self,
|
|
161
|
+
X: np.ndarray,
|
|
162
|
+
) -> np.ndarray:
|
|
163
|
+
"""
|
|
164
|
+
Compute conformal p-values for each class.
|
|
165
|
+
|
|
166
|
+
Returns
|
|
167
|
+
-------
|
|
168
|
+
p_values : np.ndarray of shape (n_samples, n_classes)
|
|
169
|
+
"""
|
|
170
|
+
|
|
171
|
+
if self.wrapper is None:
|
|
172
|
+
raise RuntimeError("CrepesConformalClassifier not fitted.")
|
|
173
|
+
|
|
174
|
+
X = np.asarray(X)
|
|
175
|
+
return self.wrapper.predict_p(
|
|
176
|
+
X, smoothing=self.smoothing, seed=self.random_state,
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
def predict_proba(
|
|
180
|
+
self,
|
|
181
|
+
X: np.ndarray,
|
|
182
|
+
) -> np.ndarray:
|
|
183
|
+
"""Underlying model probabilities."""
|
|
184
|
+
|
|
185
|
+
return self.model.predict_proba(np.asarray(X))
|
|
186
|
+
|
|
@@ -0,0 +1,180 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Conformal prediction module.
|
|
3
|
+
|
|
4
|
+
Wrapper around crepes.WrapRegressor
|
|
5
|
+
for interval prediction.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from typing import Literal
|
|
11
|
+
|
|
12
|
+
from crepes import WrapRegressor
|
|
13
|
+
from crepes.extras import (
|
|
14
|
+
DifficultyEstimator,
|
|
15
|
+
MondrianCategorizer,
|
|
16
|
+
)
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from ._validation import check_confidence, check_fit_data
|
|
20
|
+
|
|
21
|
+
ConformalMethod = Literal[
|
|
22
|
+
"standard",
|
|
23
|
+
"normalized",
|
|
24
|
+
"mondrian",
|
|
25
|
+
"normalized_mondrian",
|
|
26
|
+
]
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class _DeterministicMondrianCategorizer(MondrianCategorizer):
|
|
30
|
+
"""Fit and apply empirical bin thresholds without random tie jitter."""
|
|
31
|
+
|
|
32
|
+
def fit(self, X=None, f=None, de=None, no_bins=10):
|
|
33
|
+
self.f = f
|
|
34
|
+
self.de = de
|
|
35
|
+
scores = f(X) if f is not None else de.apply(X)
|
|
36
|
+
boundaries = np.unique(np.quantile(scores, np.linspace(0, 1, no_bins + 1)))
|
|
37
|
+
# Merge tied quantile boundaries. Constant scores form a single bin.
|
|
38
|
+
self.bin_thresholds = np.concatenate(([-np.inf], boundaries[1:-1], [np.inf]))
|
|
39
|
+
self.fitted = True
|
|
40
|
+
self.fitted_ = True
|
|
41
|
+
return self
|
|
42
|
+
|
|
43
|
+
def apply(self, X):
|
|
44
|
+
scores = self.f(X) if self.f is not None else self.de.apply(X)
|
|
45
|
+
# Equivalent to right-closed bins; tied scores always share a category.
|
|
46
|
+
return np.searchsorted(self.bin_thresholds[1:-1], scores, side="left")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class CrepesConformalPredictor:
|
|
50
|
+
"""
|
|
51
|
+
Wrapper for conformal regression using crepes.
|
|
52
|
+
|
|
53
|
+
Parameters
|
|
54
|
+
----------
|
|
55
|
+
model : object
|
|
56
|
+
Any sklearn-compatible regressor.
|
|
57
|
+
method : ConformalMethod
|
|
58
|
+
Conformal prediction method to use:
|
|
59
|
+
- "standard": basic conformal prediction
|
|
60
|
+
- "normalized": uses a DifficultyEstimator for
|
|
61
|
+
adaptive interval widths
|
|
62
|
+
- "mondrian": uses a MondrianCategorizer for
|
|
63
|
+
group-conditional coverage
|
|
64
|
+
- "normalized_mondrian": combines both
|
|
65
|
+
"""
|
|
66
|
+
|
|
67
|
+
def __init__(
|
|
68
|
+
self,
|
|
69
|
+
model,
|
|
70
|
+
method: ConformalMethod = "normalized",
|
|
71
|
+
):
|
|
72
|
+
if method not in ("standard", "normalized", "mondrian", "normalized_mondrian"):
|
|
73
|
+
raise ValueError(f"Unknown conformal regression method '{method}'.")
|
|
74
|
+
self.model = model
|
|
75
|
+
self.method = method
|
|
76
|
+
self.wrapper = None
|
|
77
|
+
self.difficulty_estimator = None
|
|
78
|
+
self.mondrian_categorizer = None
|
|
79
|
+
self._calibration_bins = None
|
|
80
|
+
|
|
81
|
+
def fit(
|
|
82
|
+
self,
|
|
83
|
+
X_train: np.ndarray,
|
|
84
|
+
y_train: np.ndarray,
|
|
85
|
+
X_calib: np.ndarray,
|
|
86
|
+
y_calib: np.ndarray,
|
|
87
|
+
**kwargs,
|
|
88
|
+
) -> None:
|
|
89
|
+
"""
|
|
90
|
+
Fit conformal predictor.
|
|
91
|
+
|
|
92
|
+
Steps:
|
|
93
|
+
1. Fit model
|
|
94
|
+
2. Fit difficulty estimator and/or Mondrian categorizer
|
|
95
|
+
3. Calibrate conformal predictor
|
|
96
|
+
"""
|
|
97
|
+
|
|
98
|
+
self.wrapper = None
|
|
99
|
+
self.difficulty_estimator = None
|
|
100
|
+
self.mondrian_categorizer = None
|
|
101
|
+
self._calibration_bins = None
|
|
102
|
+
X_train, y_train, X_calib, y_calib = check_fit_data(
|
|
103
|
+
X_train, y_train, X_calib, y_calib,
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
# Train
|
|
107
|
+
self.model.fit(
|
|
108
|
+
X_train,
|
|
109
|
+
y_train,
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
# Build calibration kwargs based on method
|
|
113
|
+
calibrate_kwargs = {}
|
|
114
|
+
|
|
115
|
+
if self.method in ("normalized", "normalized_mondrian"):
|
|
116
|
+
self.difficulty_estimator = DifficultyEstimator()
|
|
117
|
+
self.difficulty_estimator.fit(X_train, y=y_train)
|
|
118
|
+
calibrate_kwargs["de"] = self.difficulty_estimator
|
|
119
|
+
|
|
120
|
+
if self.method in ("mondrian", "normalized_mondrian"):
|
|
121
|
+
self.mondrian_categorizer = _DeterministicMondrianCategorizer()
|
|
122
|
+
mc_kwargs = {"de": self.difficulty_estimator} if self.method == "normalized_mondrian" else {"f": self.model.predict}
|
|
123
|
+
self.mondrian_categorizer.fit(X_train, **mc_kwargs)
|
|
124
|
+
calibrate_kwargs["mc"] = self.mondrian_categorizer
|
|
125
|
+
self._calibration_bins = np.unique(self.mondrian_categorizer.apply(X_calib))
|
|
126
|
+
|
|
127
|
+
self.wrapper = WrapRegressor(
|
|
128
|
+
self.model
|
|
129
|
+
)
|
|
130
|
+
|
|
131
|
+
# Calibrate
|
|
132
|
+
self.wrapper.calibrate(
|
|
133
|
+
X_calib,
|
|
134
|
+
y_calib,
|
|
135
|
+
**calibrate_kwargs,
|
|
136
|
+
)
|
|
137
|
+
|
|
138
|
+
def predict(
|
|
139
|
+
self,
|
|
140
|
+
X: np.ndarray,
|
|
141
|
+
confidence: float = 0.9,
|
|
142
|
+
) -> tuple[np.ndarray, np.ndarray]:
|
|
143
|
+
"""
|
|
144
|
+
Generate prediction intervals.
|
|
145
|
+
|
|
146
|
+
Parameters
|
|
147
|
+
----------
|
|
148
|
+
X : np.ndarray
|
|
149
|
+
|
|
150
|
+
confidence : float
|
|
151
|
+
|
|
152
|
+
Returns
|
|
153
|
+
-------
|
|
154
|
+
lower : np.ndarray
|
|
155
|
+
upper : np.ndarray
|
|
156
|
+
"""
|
|
157
|
+
|
|
158
|
+
if self.wrapper is None:
|
|
159
|
+
raise RuntimeError("CrepesConformalPredictor not fitted.")
|
|
160
|
+
|
|
161
|
+
check_confidence(confidence)
|
|
162
|
+
X = np.asarray(X)
|
|
163
|
+
|
|
164
|
+
intervals = self.wrapper.predict_int(
|
|
165
|
+
X,
|
|
166
|
+
confidence=confidence,
|
|
167
|
+
)
|
|
168
|
+
|
|
169
|
+
if self.mondrian_categorizer is not None:
|
|
170
|
+
bins = self.mondrian_categorizer.apply(X)
|
|
171
|
+
uncalibrated = ~np.isin(bins, self._calibration_bins)
|
|
172
|
+
# No finite group quantile exists without calibration observations.
|
|
173
|
+
# Some crepes versions otherwise leave these intervals at [0, 0].
|
|
174
|
+
intervals[uncalibrated, 0] = -np.inf
|
|
175
|
+
intervals[uncalibrated, 1] = np.inf
|
|
176
|
+
|
|
177
|
+
lower = intervals[:, 0]
|
|
178
|
+
upper = intervals[:, 1]
|
|
179
|
+
|
|
180
|
+
return lower, upper
|