conforme 0.1.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.
- conforme/__init__.py +81 -0
- conforme/backtest/__init__.py +6 -0
- conforme/backtest/forecasts.py +142 -0
- conforme/backtest/replay.py +96 -0
- conforme/conformal/__init__.py +42 -0
- conforme/conformal/calibrators/__init__.py +25 -0
- conforme/conformal/calibrators/aci.py +42 -0
- conforme/conformal/calibrators/base.py +97 -0
- conforme/conformal/calibrators/ranks.py +92 -0
- conforme/conformal/calibrators/risk.py +60 -0
- conforme/conformal/calibrators/split.py +120 -0
- conforme/conformal/calibrators/tracker.py +33 -0
- conforme/conformal/losses.py +31 -0
- conforme/conformal/scores.py +44 -0
- conforme/conformal/targets.py +47 -0
- conforme/data/__init__.py +6 -0
- conforme/data/hierarchy.py +103 -0
- conforme/data/panel.py +56 -0
- conforme/decision/__init__.py +6 -0
- conforme/decision/lost_sales.py +43 -0
- conforme/decision/policy.py +27 -0
- conforme/forecast/__init__.py +16 -0
- conforme/forecast/models/__init__.py +10 -0
- conforme/forecast/models/base.py +67 -0
- conforme/forecast/models/frames.py +83 -0
- conforme/forecast/models/mlforecast.py +100 -0
- conforme/forecast/models/naive.py +24 -0
- conforme/forecast/models/neuralforecast.py +77 -0
- conforme/forecast/models/statsforecast.py +32 -0
- conforme/forecast/reconcile.py +70 -0
- conforme/metrics.py +60 -0
- conforme/online/__init__.py +6 -0
- conforme/online/ledger.py +119 -0
- conforme/online/state.py +33 -0
- conforme/online/step.py +70 -0
- conforme/py.typed +0 -0
- conforme-0.1.0.dist-info/METADATA +216 -0
- conforme-0.1.0.dist-info/RECORD +40 -0
- conforme-0.1.0.dist-info/WHEEL +4 -0
- conforme-0.1.0.dist-info/licenses/LICENSE +21 -0
conforme/__init__.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
"""Conforme: point forecasts, hierarchy reconciliation, conformal bounds, and orders.
|
|
2
|
+
|
|
3
|
+
Vendor model adapters live in `conforme.forecast.models` and need their extra installed.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from conforme.backtest import Forecasts, Replay, replay, rolling_forecasts
|
|
7
|
+
from conforme.conformal import (
|
|
8
|
+
ACI,
|
|
9
|
+
Absolute,
|
|
10
|
+
Calibrator,
|
|
11
|
+
Feedback,
|
|
12
|
+
LeadTime,
|
|
13
|
+
Level,
|
|
14
|
+
Loss,
|
|
15
|
+
Miss,
|
|
16
|
+
QuantileCalibrator,
|
|
17
|
+
QuantileTracker,
|
|
18
|
+
RiskControl,
|
|
19
|
+
Score,
|
|
20
|
+
Signed,
|
|
21
|
+
SplitQuantile,
|
|
22
|
+
State,
|
|
23
|
+
Step,
|
|
24
|
+
Target,
|
|
25
|
+
)
|
|
26
|
+
from conforme.data import Hierarchy, Panel
|
|
27
|
+
from conforme.decision import Settlement, critical_ratio, order_up_to, settle
|
|
28
|
+
from conforme.forecast import (
|
|
29
|
+
BottomUp,
|
|
30
|
+
Covariate,
|
|
31
|
+
Fitted,
|
|
32
|
+
Forecaster,
|
|
33
|
+
Identity,
|
|
34
|
+
Reconciler,
|
|
35
|
+
SeasonalNaive,
|
|
36
|
+
Window,
|
|
37
|
+
WlsStruct,
|
|
38
|
+
)
|
|
39
|
+
from conforme.online import Issue, initial_state, step
|
|
40
|
+
|
|
41
|
+
__all__ = [
|
|
42
|
+
"ACI",
|
|
43
|
+
"Absolute",
|
|
44
|
+
"BottomUp",
|
|
45
|
+
"Calibrator",
|
|
46
|
+
"Covariate",
|
|
47
|
+
"Feedback",
|
|
48
|
+
"Fitted",
|
|
49
|
+
"Forecaster",
|
|
50
|
+
"Forecasts",
|
|
51
|
+
"Hierarchy",
|
|
52
|
+
"Identity",
|
|
53
|
+
"Issue",
|
|
54
|
+
"LeadTime",
|
|
55
|
+
"Level",
|
|
56
|
+
"Loss",
|
|
57
|
+
"Miss",
|
|
58
|
+
"QuantileCalibrator",
|
|
59
|
+
"Panel",
|
|
60
|
+
"QuantileTracker",
|
|
61
|
+
"RiskControl",
|
|
62
|
+
"Reconciler",
|
|
63
|
+
"Replay",
|
|
64
|
+
"Score",
|
|
65
|
+
"SeasonalNaive",
|
|
66
|
+
"Settlement",
|
|
67
|
+
"Signed",
|
|
68
|
+
"SplitQuantile",
|
|
69
|
+
"State",
|
|
70
|
+
"Step",
|
|
71
|
+
"Target",
|
|
72
|
+
"Window",
|
|
73
|
+
"WlsStruct",
|
|
74
|
+
"critical_ratio",
|
|
75
|
+
"initial_state",
|
|
76
|
+
"order_up_to",
|
|
77
|
+
"replay",
|
|
78
|
+
"rolling_forecasts",
|
|
79
|
+
"settle",
|
|
80
|
+
"step",
|
|
81
|
+
]
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
"""Backtests: rolling-origin forecasts, and online calibration replayed over them."""
|
|
2
|
+
|
|
3
|
+
from conforme.backtest.forecasts import Forecasts, rolling_forecasts
|
|
4
|
+
from conforme.backtest.replay import Replay, replay
|
|
5
|
+
|
|
6
|
+
__all__ = ["Forecasts", "Replay", "replay", "rolling_forecasts"]
|
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
"""Rolling-origin point forecasts: fit, predict, and reconcile at each origin.
|
|
2
|
+
|
|
3
|
+
`rolling_forecasts` keeps no state between calls and stores nothing. Each window is a
|
|
4
|
+
read-only slice that ends at its origin, so a model sees only what was known there.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from collections.abc import Mapping
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
import pandas as pd
|
|
12
|
+
|
|
13
|
+
from conforme.data.hierarchy import Hierarchy
|
|
14
|
+
from conforme.data.panel import Panel
|
|
15
|
+
from conforme.forecast.models.base import Covariate, Fitted, Forecaster, Window
|
|
16
|
+
from conforme.forecast.reconcile import Reconciler
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class Forecasts:
|
|
21
|
+
"""Reconciled points `[O, N, H]` for increasing origins `[O]`."""
|
|
22
|
+
|
|
23
|
+
origins: np.ndarray
|
|
24
|
+
points: np.ndarray
|
|
25
|
+
|
|
26
|
+
def residuals(self, actuals: np.ndarray) -> np.ndarray:
|
|
27
|
+
"""Return `actual - point` `[O, N, H]` from node actuals `[N, T]`.
|
|
28
|
+
|
|
29
|
+
A target after the last period is NaN, which means not yet known.
|
|
30
|
+
"""
|
|
31
|
+
n_periods = actuals.shape[1]
|
|
32
|
+
targets = self.origins[:, None] + np.arange(1, self.points.shape[2] + 1)
|
|
33
|
+
values = actuals[:, np.minimum(targets, n_periods - 1)].transpose(1, 0, 2)
|
|
34
|
+
resid = values.astype(np.float32) - self.points
|
|
35
|
+
return np.where((targets < n_periods)[:, None, :], resid, np.float32(np.nan))
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def rolling_forecasts(
|
|
39
|
+
panel: Panel,
|
|
40
|
+
hierarchy: Hierarchy,
|
|
41
|
+
model: Forecaster,
|
|
42
|
+
reconciler: Reconciler,
|
|
43
|
+
origins: np.ndarray,
|
|
44
|
+
horizon: int,
|
|
45
|
+
refit_every: int | None = None,
|
|
46
|
+
covariates: Mapping[str, Covariate] | None = None,
|
|
47
|
+
) -> Forecasts:
|
|
48
|
+
"""Fit, predict, and reconcile at each origin, the index of the last observed period.
|
|
49
|
+
|
|
50
|
+
The model forecasts the first S nodes, with S from the reconciler shape `(N, S)`. It
|
|
51
|
+
is fitted at the first origin and again every `refit_every` origins. None fits once.
|
|
52
|
+
"""
|
|
53
|
+
origins = _check_schedule(origins, len(panel.periods), horizon, refit_every)
|
|
54
|
+
covariates = dict(covariates or {})
|
|
55
|
+
n_nodes, n_series = reconciler.shape
|
|
56
|
+
if n_nodes != len(hierarchy.nodes) or n_series not in (hierarchy.n_bottom, n_nodes):
|
|
57
|
+
raise ValueError(f"reconciler shape {reconciler.shape} does not fit the hierarchy")
|
|
58
|
+
last = int(origins[-1])
|
|
59
|
+
for name, covariate in covariates.items():
|
|
60
|
+
_check_covariate(name, covariate, hierarchy.n_bottom, last, horizon)
|
|
61
|
+
|
|
62
|
+
history = _series_rows(panel.values, hierarchy, n_series, "sum")
|
|
63
|
+
inputs = {
|
|
64
|
+
name: _series_rows(covariate.values, hierarchy, n_series, covariate.aggregate)
|
|
65
|
+
for name, covariate in covariates.items()
|
|
66
|
+
}
|
|
67
|
+
reach = {name: _reach(covariate, horizon) for name, covariate in covariates.items()}
|
|
68
|
+
periods = pd.date_range(panel.periods[0], periods=last + 1 + horizon, freq=panel.freq)
|
|
69
|
+
|
|
70
|
+
points = np.empty((len(origins), n_nodes, horizon), dtype=np.float32)
|
|
71
|
+
fitted: Fitted | None = None
|
|
72
|
+
for index, origin in enumerate(origins):
|
|
73
|
+
seen = int(origin) + 1
|
|
74
|
+
window = Window(
|
|
75
|
+
y=history[:, :seen],
|
|
76
|
+
x={name: _visible(values, seen, reach[name]) for name, values in inputs.items()},
|
|
77
|
+
periods=periods[: seen + horizon],
|
|
78
|
+
horizon=horizon,
|
|
79
|
+
)
|
|
80
|
+
if fitted is None or (refit_every is not None and index % refit_every == 0):
|
|
81
|
+
fitted = model.fit(window)
|
|
82
|
+
base = np.asarray(fitted.predict(window))
|
|
83
|
+
if base.shape != (n_series, horizon):
|
|
84
|
+
raise ValueError(f"model returned shape {base.shape}, expected {(n_series, horizon)}")
|
|
85
|
+
if not np.isfinite(base).all():
|
|
86
|
+
raise ValueError(f"model returned nonfinite points at origin {origin}")
|
|
87
|
+
points[index] = reconciler(base)
|
|
88
|
+
return Forecasts(origins=origins, points=points)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def _check_schedule(
|
|
92
|
+
origins: np.ndarray, n_periods: int, horizon: int, refit_every: int | None
|
|
93
|
+
) -> np.ndarray:
|
|
94
|
+
origins = np.asarray(origins)
|
|
95
|
+
if origins.ndim != 1 or len(origins) == 0 or not np.issubdtype(origins.dtype, np.integer):
|
|
96
|
+
raise ValueError("origins must be a non-empty vector of period indexes")
|
|
97
|
+
if (np.diff(origins) <= 0).any():
|
|
98
|
+
raise ValueError("origins must be strictly increasing")
|
|
99
|
+
if origins[0] < 0 or origins[-1] >= n_periods:
|
|
100
|
+
raise ValueError(f"origins must be in [0, {n_periods - 1}]")
|
|
101
|
+
if horizon < 1:
|
|
102
|
+
raise ValueError(f"horizon must be at least 1, got {horizon}")
|
|
103
|
+
if refit_every is not None and refit_every < 1:
|
|
104
|
+
raise ValueError(f"refit_every must be at least 1, got {refit_every}")
|
|
105
|
+
return origins.astype(np.int64)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _check_covariate(
|
|
109
|
+
name: str, covariate: Covariate, n_bottom: int, last: int, horizon: int
|
|
110
|
+
) -> None:
|
|
111
|
+
rows, length = covariate.values.shape
|
|
112
|
+
if rows not in (1, n_bottom):
|
|
113
|
+
raise ValueError(f"covariate {name!r} must have one row or one row per bottom series")
|
|
114
|
+
needed = last + 1 + (horizon if covariate.known_ahead else 0)
|
|
115
|
+
if length != 1 and length < needed:
|
|
116
|
+
raise ValueError(f"covariate {name!r} has {length} periods, needs {needed}")
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def _series_rows(
|
|
120
|
+
values: np.ndarray, hierarchy: Hierarchy, n_series: int, aggregate: str
|
|
121
|
+
) -> np.ndarray:
|
|
122
|
+
"""Read-only rows for the first `n_series` nodes. A single row is shared by all."""
|
|
123
|
+
if len(values) > 1 and n_series > hierarchy.n_bottom:
|
|
124
|
+
nodes = hierarchy.summing[:n_series]
|
|
125
|
+
values = nodes @ values
|
|
126
|
+
if aggregate == "mean":
|
|
127
|
+
values = values / nodes.sum(axis=1)[:, None]
|
|
128
|
+
values = values.astype(np.float32)
|
|
129
|
+
view = values.view()
|
|
130
|
+
view.flags.writeable = False
|
|
131
|
+
return view
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def _reach(covariate: Covariate, horizon: int) -> int | None:
|
|
135
|
+
"""Periods visible after the origin. None for a static covariate."""
|
|
136
|
+
if covariate.values.shape[1] == 1:
|
|
137
|
+
return None
|
|
138
|
+
return horizon if covariate.known_ahead else 0
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _visible(values: np.ndarray, seen: int, reach: int | None) -> np.ndarray:
|
|
142
|
+
return values if reach is None else values[:, : seen + reach]
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Replay: online calibration over the origins of a backtest, kept for assessment."""
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from conforme.backtest.forecasts import Forecasts
|
|
8
|
+
from conforme.conformal.calibrators.base import Calibrator, State
|
|
9
|
+
from conforme.conformal.scores import Score
|
|
10
|
+
from conforme.conformal.targets import Target, columns
|
|
11
|
+
from conforme.online.step import initial_state, step
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@dataclass(frozen=True)
|
|
15
|
+
class Replay:
|
|
16
|
+
"""Calibration over origins `[O]`. Arrays are `[O, N, C]`.
|
|
17
|
+
|
|
18
|
+
`target` and `score` are NaN when the target is not known at the end of the data.
|
|
19
|
+
`state` continues the run with `conforme.online.step`.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
origins: np.ndarray
|
|
23
|
+
point: np.ndarray
|
|
24
|
+
threshold: np.ndarray
|
|
25
|
+
lower: np.ndarray
|
|
26
|
+
upper: np.ndarray
|
|
27
|
+
target: np.ndarray
|
|
28
|
+
score: np.ndarray
|
|
29
|
+
censored: np.ndarray
|
|
30
|
+
state: State
|
|
31
|
+
|
|
32
|
+
@property
|
|
33
|
+
def covered(self) -> np.ndarray:
|
|
34
|
+
"""1.0 when the score is within its threshold, NaN when not known."""
|
|
35
|
+
known = np.isfinite(self.score)
|
|
36
|
+
return np.where(known, (self.score <= self.threshold).astype(np.float32), np.nan)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def replay(
|
|
40
|
+
forecasts: Forecasts,
|
|
41
|
+
actuals: np.ndarray,
|
|
42
|
+
*,
|
|
43
|
+
target: Target,
|
|
44
|
+
score: Score,
|
|
45
|
+
calibrator: Calibrator,
|
|
46
|
+
censored: np.ndarray | None = None,
|
|
47
|
+
) -> Replay:
|
|
48
|
+
"""Run `conforme.online.step` at each origin of `forecasts` on node actuals `[N, T]`."""
|
|
49
|
+
origins = forecasts.origins
|
|
50
|
+
n_origins, n_nodes, horizon = forecasts.points.shape
|
|
51
|
+
if censored is None:
|
|
52
|
+
censored = np.zeros(actuals.shape, dtype=bool)
|
|
53
|
+
state = initial_state(target, calibrator, n_nodes, horizon)
|
|
54
|
+
issued = []
|
|
55
|
+
previous = -1
|
|
56
|
+
for index, origin in enumerate(origins):
|
|
57
|
+
seen = slice(previous + 1, int(origin) + 1)
|
|
58
|
+
state, out = step(
|
|
59
|
+
state,
|
|
60
|
+
int(origin),
|
|
61
|
+
forecasts.points[index],
|
|
62
|
+
actuals[:, seen],
|
|
63
|
+
target=target,
|
|
64
|
+
score=score,
|
|
65
|
+
calibrator=calibrator,
|
|
66
|
+
censored=censored[:, seen],
|
|
67
|
+
)
|
|
68
|
+
issued.append(out)
|
|
69
|
+
previous = int(origin)
|
|
70
|
+
cover = target.cover(horizon)
|
|
71
|
+
targets, target_censored = _targets(actuals, censored, origins, horizon, cover)
|
|
72
|
+
point = np.stack([out.point for out in issued])
|
|
73
|
+
return Replay(
|
|
74
|
+
origins=origins,
|
|
75
|
+
point=point,
|
|
76
|
+
threshold=np.stack([out.threshold for out in issued]),
|
|
77
|
+
lower=np.stack([out.lower for out in issued]),
|
|
78
|
+
upper=np.stack([out.upper for out in issued]),
|
|
79
|
+
target=targets,
|
|
80
|
+
score=score.score(targets, point),
|
|
81
|
+
censored=target_censored,
|
|
82
|
+
state=state,
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _targets(
|
|
87
|
+
actuals: np.ndarray, censored: np.ndarray, origins: np.ndarray, horizon: int, cover: np.ndarray
|
|
88
|
+
) -> tuple[np.ndarray, np.ndarray]:
|
|
89
|
+
"""Targets and censoring per column `[O, N, C]`, NaN after the last period."""
|
|
90
|
+
n_periods = actuals.shape[1]
|
|
91
|
+
steps = origins[:, None] + np.arange(1, horizon + 1)
|
|
92
|
+
known = steps < n_periods
|
|
93
|
+
clipped = np.minimum(steps, n_periods - 1)
|
|
94
|
+
values = np.where(known[:, None, :], actuals[:, clipped].transpose(1, 0, 2), np.nan)
|
|
95
|
+
flags = censored[:, clipped].transpose(1, 0, 2) & known[:, None, :]
|
|
96
|
+
return columns(values.astype(np.float32), cover), columns(flags, cover)
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""Conformal calibration: targets, scores, and calibration methods. No time, no origins.
|
|
2
|
+
|
|
3
|
+
A `Target` says which quantity each column bounds, a `Score` which bound each threshold
|
|
4
|
+
issues, a `Loss` what a bound cost, and a `Calibrator` turns the known targets into a
|
|
5
|
+
threshold per node and column.
|
|
6
|
+
`conforme.online` runs them origin by origin.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from conforme.conformal.calibrators import (
|
|
10
|
+
ACI,
|
|
11
|
+
Calibrator,
|
|
12
|
+
Feedback,
|
|
13
|
+
Level,
|
|
14
|
+
QuantileCalibrator,
|
|
15
|
+
QuantileTracker,
|
|
16
|
+
RiskControl,
|
|
17
|
+
SplitQuantile,
|
|
18
|
+
State,
|
|
19
|
+
)
|
|
20
|
+
from conforme.conformal.losses import Loss, Miss
|
|
21
|
+
from conforme.conformal.scores import Absolute, Score, Signed
|
|
22
|
+
from conforme.conformal.targets import LeadTime, Step, Target
|
|
23
|
+
|
|
24
|
+
__all__ = [
|
|
25
|
+
"ACI",
|
|
26
|
+
"Absolute",
|
|
27
|
+
"Calibrator",
|
|
28
|
+
"Feedback",
|
|
29
|
+
"LeadTime",
|
|
30
|
+
"Level",
|
|
31
|
+
"Loss",
|
|
32
|
+
"Miss",
|
|
33
|
+
"QuantileCalibrator",
|
|
34
|
+
"QuantileTracker",
|
|
35
|
+
"RiskControl",
|
|
36
|
+
"Score",
|
|
37
|
+
"Signed",
|
|
38
|
+
"SplitQuantile",
|
|
39
|
+
"State",
|
|
40
|
+
"Step",
|
|
41
|
+
"Target",
|
|
42
|
+
]
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
"""Calibration methods. A method is one module that implements `base.Calibrator`."""
|
|
2
|
+
|
|
3
|
+
from conforme.conformal.calibrators.aci import ACI
|
|
4
|
+
from conforme.conformal.calibrators.base import (
|
|
5
|
+
Calibrator,
|
|
6
|
+
Feedback,
|
|
7
|
+
Level,
|
|
8
|
+
QuantileCalibrator,
|
|
9
|
+
State,
|
|
10
|
+
)
|
|
11
|
+
from conforme.conformal.calibrators.risk import RiskControl
|
|
12
|
+
from conforme.conformal.calibrators.split import SplitQuantile
|
|
13
|
+
from conforme.conformal.calibrators.tracker import QuantileTracker
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"ACI",
|
|
17
|
+
"Calibrator",
|
|
18
|
+
"Feedback",
|
|
19
|
+
"Level",
|
|
20
|
+
"QuantileCalibrator",
|
|
21
|
+
"QuantileTracker",
|
|
22
|
+
"RiskControl",
|
|
23
|
+
"SplitQuantile",
|
|
24
|
+
"State",
|
|
25
|
+
]
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""Adaptive conformal inference: a working level that reacts to misses."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from conforme.conformal.calibrators.base import Calibrator, Feedback, QuantileCalibrator, State
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ACI(Calibrator):
|
|
9
|
+
"""Adaptive conformal inference (Gibbs and Candès 2021), per node and column.
|
|
10
|
+
|
|
11
|
+
After each known score, the working level moves by `gamma * (miss - (1 - level))`,
|
|
12
|
+
where a miss is a score above the threshold that was issued for it. The base
|
|
13
|
+
calibrator gives the threshold at the working level. The target level is the level
|
|
14
|
+
of the base. A working level at or above one gives an infinite threshold, at or
|
|
15
|
+
below zero an empty interval.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
def __init__(self, base: QuantileCalibrator, gamma: float = 0.005) -> None:
|
|
19
|
+
if gamma <= 0:
|
|
20
|
+
raise ValueError(f"gamma must be positive, got {gamma}")
|
|
21
|
+
self.base = base
|
|
22
|
+
self.level = base.level
|
|
23
|
+
self.gamma = gamma
|
|
24
|
+
|
|
25
|
+
def initial_state(self, n_nodes: int, n_columns: int) -> State:
|
|
26
|
+
return {
|
|
27
|
+
"base": self.base.initial_state(n_nodes, n_columns),
|
|
28
|
+
"level": np.full((n_nodes, n_columns), np.nan),
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
def update(self, state: State, feedback: Feedback) -> State:
|
|
32
|
+
working = np.where(np.isnan(state["level"]), self.level, state["level"])
|
|
33
|
+
misses, known = feedback.misses(working.shape[1])
|
|
34
|
+
working = working + self.gamma * (misses - (1 - self.level) * known)
|
|
35
|
+
return {"base": self.base.update(state["base"], feedback), "level": working}
|
|
36
|
+
|
|
37
|
+
def threshold(self, state: State) -> np.ndarray:
|
|
38
|
+
working = np.where(np.isnan(state["level"]), self.level, state["level"])
|
|
39
|
+
inside = np.clip(working, 1e-6, 1 - 1e-6)
|
|
40
|
+
out = self.base.threshold_at(state["base"], inside)
|
|
41
|
+
out = np.where(working >= 1, np.inf, out)
|
|
42
|
+
return np.where(working <= 0, -np.inf, out).astype(np.float32)
|
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""The calibrator contract: what a calibration method receives, keeps, and returns.
|
|
2
|
+
|
|
3
|
+
A calibrator turns the targets known so far into a threshold per node and column. A
|
|
4
|
+
threshold indexes a nested family of bounds, the `Score`: a quantile calibrator ranks
|
|
5
|
+
the scores, and a risk calibrator picks the smallest threshold whose bounds have a low
|
|
6
|
+
enough loss. Its
|
|
7
|
+
state is a nested dict of numpy arrays, and its methods are functions of that state.
|
|
8
|
+
So a product can save the state after each origin and continue later, and a backtest
|
|
9
|
+
is the same loop as production. A call owns the state it receives: it can write into
|
|
10
|
+
those arrays, and the caller continues only with the returned state.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from abc import ABC, abstractmethod
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
from typing import Any
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from conforme.conformal.scores import Score
|
|
20
|
+
|
|
21
|
+
State = dict[str, Any]
|
|
22
|
+
"""Nested dict of numpy arrays. `conforme.online.state` flattens it for storage."""
|
|
23
|
+
|
|
24
|
+
Level = float | np.ndarray
|
|
25
|
+
"""Target level: a scalar, or `[N, C]` per node and column. `[N, 1]` is one per node."""
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass(frozen=True)
|
|
29
|
+
class Feedback:
|
|
30
|
+
"""Targets that became known at one origin. Each row `[K]` is one target column of
|
|
31
|
+
one issuing origin, for all nodes.
|
|
32
|
+
|
|
33
|
+
Rows are in origin order. `point` and `target` are the column sums `[K, N]`, NaN
|
|
34
|
+
where the actual is missing. `issued` is the threshold that was issued for each
|
|
35
|
+
row, `censored` marks targets that are lower bounds, and `score` is the family of
|
|
36
|
+
bounds the thresholds index, so a calibrator can find the bounds of any threshold.
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
origin: np.ndarray
|
|
40
|
+
column: np.ndarray
|
|
41
|
+
point: np.ndarray
|
|
42
|
+
target: np.ndarray
|
|
43
|
+
issued: np.ndarray
|
|
44
|
+
censored: np.ndarray
|
|
45
|
+
score: Score
|
|
46
|
+
|
|
47
|
+
@property
|
|
48
|
+
def scores(self) -> np.ndarray:
|
|
49
|
+
"""Nonconformity scores `[K, N]` float32: the smallest threshold that holds the target."""
|
|
50
|
+
return self.score.score(self.target, self.point).astype(np.float32)
|
|
51
|
+
|
|
52
|
+
def misses(self, n_columns: int) -> tuple[np.ndarray, np.ndarray]:
|
|
53
|
+
"""Counts `[N, C]` of misses (score above the issued threshold) and known scores."""
|
|
54
|
+
scores = self.scores
|
|
55
|
+
known = np.isfinite(scores)
|
|
56
|
+
miss = (scores > self.issued) & known
|
|
57
|
+
columns = np.eye(n_columns)[self.column] # [K, C]
|
|
58
|
+
return miss.T @ columns, known.T @ columns
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class Calibrator(ABC):
|
|
62
|
+
"""A calibration method. It owns its target level, set at construction."""
|
|
63
|
+
|
|
64
|
+
@abstractmethod
|
|
65
|
+
def initial_state(self, n_nodes: int, n_columns: int) -> State:
|
|
66
|
+
"""Return the empty state."""
|
|
67
|
+
|
|
68
|
+
@abstractmethod
|
|
69
|
+
def update(self, state: State, feedback: Feedback) -> State:
|
|
70
|
+
"""Return the state after the feedback. It can reuse the input arrays."""
|
|
71
|
+
|
|
72
|
+
@abstractmethod
|
|
73
|
+
def threshold(self, state: State) -> np.ndarray:
|
|
74
|
+
"""Return thresholds `[N, C]` at the target level. inf means not ready."""
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
class QuantileCalibrator(Calibrator):
|
|
78
|
+
"""A calibrator that can give its threshold at any level, not only its own.
|
|
79
|
+
|
|
80
|
+
Wrappers such as `ACI` move the level and ask the base for the threshold there.
|
|
81
|
+
"""
|
|
82
|
+
|
|
83
|
+
level: Level
|
|
84
|
+
|
|
85
|
+
@abstractmethod
|
|
86
|
+
def threshold_at(self, state: State, level: Level) -> np.ndarray:
|
|
87
|
+
"""Return thresholds `[N, C]` at `level` instead of the target level."""
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def check_level(level: Level) -> Level:
|
|
91
|
+
"""A scalar or `[N, C]` level strictly between zero and one."""
|
|
92
|
+
values = np.asarray(level, dtype=np.float64)
|
|
93
|
+
if values.ndim not in (0, 2):
|
|
94
|
+
raise ValueError(f"level must be a scalar or [N, C], got shape {values.shape}")
|
|
95
|
+
if not np.isfinite(values).all() or ((values <= 0) | (values >= 1)).any():
|
|
96
|
+
raise ValueError("level must be strictly between zero and one")
|
|
97
|
+
return float(values) if values.ndim == 0 else values
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
"""Finite-sample score quantiles, with rolling retention and pooling.
|
|
2
|
+
|
|
3
|
+
These functions hold the rank rules of the library. Calibrators call them on the
|
|
4
|
+
scores they keep.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def score_quantile(scores: np.ndarray, level: float | np.ndarray) -> np.ndarray:
|
|
11
|
+
"""Finite-sample quantile of finite scores along axis 0, without interpolation.
|
|
12
|
+
|
|
13
|
+
A vector is one pool; a matrix is independent pools in columns. `level` is a
|
|
14
|
+
scalar or one value per output. The rank is ceil((n + 1) * level); insufficient
|
|
15
|
+
samples give inf.
|
|
16
|
+
"""
|
|
17
|
+
scores = np.asarray(scores)
|
|
18
|
+
level = np.asarray(level)
|
|
19
|
+
if scores.ndim not in (1, 2) or level.ndim > 1:
|
|
20
|
+
raise ValueError("scores must be a vector or matrix and level a scalar or vector")
|
|
21
|
+
if not np.isfinite(level).all() or ((level <= 0) | (level >= 1)).any():
|
|
22
|
+
raise ValueError("level must be strictly between zero and one")
|
|
23
|
+
if not np.isfinite(scores).all():
|
|
24
|
+
raise ValueError("scores must be finite; select only resolved observations")
|
|
25
|
+
shape = scores.shape[1:] if scores.ndim == 2 else np.shape(level)
|
|
26
|
+
ranks = np.ceil((len(scores) + 1) * np.broadcast_to(level, shape)).astype(np.int64)
|
|
27
|
+
out = np.full(shape, np.inf, dtype=np.result_type(scores.dtype, np.float32))
|
|
28
|
+
ready = ranks <= len(scores)
|
|
29
|
+
if ready.any():
|
|
30
|
+
ordered = np.sort(scores, axis=0)
|
|
31
|
+
if scores.ndim == 1:
|
|
32
|
+
out[ready] = ordered[ranks[ready] - 1]
|
|
33
|
+
else:
|
|
34
|
+
columns = np.nonzero(ready)[0]
|
|
35
|
+
out[ready] = ordered[ranks[ready] - 1, columns]
|
|
36
|
+
return out
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def retained_quantile(
|
|
40
|
+
scores: np.ndarray,
|
|
41
|
+
level: float | np.ndarray,
|
|
42
|
+
window: int | None = None,
|
|
43
|
+
groups: np.ndarray | None = None,
|
|
44
|
+
) -> np.ndarray:
|
|
45
|
+
"""Quantile per node `[N]` from scores `[K, N]` in origin order. NaN is no score.
|
|
46
|
+
|
|
47
|
+
`window` keeps the last K resolved rows per pool, before pooling. A row is resolved
|
|
48
|
+
for a pool when it has at least one finite score there, so holes do not consume
|
|
49
|
+
rows. `groups` has one integer per node: equal labels pool their scores, and every
|
|
50
|
+
member gets the pool quantile. `level` is a scalar or one value per node.
|
|
51
|
+
"""
|
|
52
|
+
if window is not None:
|
|
53
|
+
if window < 1:
|
|
54
|
+
raise ValueError(f"window must be at least 1, got {window}")
|
|
55
|
+
# Fast path: when the last K rows are complete, they are the retained rows.
|
|
56
|
+
if np.isfinite(scores[-window:]).all():
|
|
57
|
+
scores = scores[-window:]
|
|
58
|
+
if groups is None:
|
|
59
|
+
if window is not None and len(scores) > window:
|
|
60
|
+
finite = np.isfinite(scores)
|
|
61
|
+
retained = np.cumsum(finite[::-1], axis=0)[::-1] <= window
|
|
62
|
+
scores = np.where(finite & retained, scores, np.nan)
|
|
63
|
+
scores = scores[np.isfinite(scores).any(axis=1)]
|
|
64
|
+
return _finite_quantile(scores, level)
|
|
65
|
+
levels = np.broadcast_to(level, (scores.shape[1],))
|
|
66
|
+
out = np.full(scores.shape[1], np.inf, dtype=np.float32)
|
|
67
|
+
order = np.argsort(groups)
|
|
68
|
+
labels = groups[order]
|
|
69
|
+
for columns in np.split(order, np.flatnonzero(labels[1:] != labels[:-1]) + 1):
|
|
70
|
+
pool = scores[:, columns]
|
|
71
|
+
pool = pool[np.isfinite(pool).any(axis=1)]
|
|
72
|
+
if window is not None:
|
|
73
|
+
pool = pool[-window:]
|
|
74
|
+
out[columns] = score_quantile(pool[np.isfinite(pool)], levels[columns])
|
|
75
|
+
return out
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def _finite_quantile(scores: np.ndarray, level: float | np.ndarray) -> np.ndarray:
|
|
79
|
+
"""Independent column ranks from finite counts; NaN cells supply no score."""
|
|
80
|
+
if np.isfinite(scores).all():
|
|
81
|
+
return score_quantile(scores, level).astype(np.float32)
|
|
82
|
+
counts = np.isfinite(scores).sum(axis=0)
|
|
83
|
+
levels = np.broadcast_to(level, counts.shape)
|
|
84
|
+
# Match score_quantile's level dtype arithmetic, including float32 levels.
|
|
85
|
+
ranks = np.ceil((counts + 1).astype(levels.dtype) * levels).astype(np.int64)
|
|
86
|
+
out = np.full(counts.shape, np.inf, dtype=np.float32)
|
|
87
|
+
ready = (ranks <= counts) & (counts > 0)
|
|
88
|
+
if ready.any():
|
|
89
|
+
ordered = np.sort(np.where(np.isfinite(scores), scores, np.nan), axis=0)
|
|
90
|
+
columns = np.flatnonzero(ready)
|
|
91
|
+
out[columns] = ordered[ranks[columns] - 1, columns]
|
|
92
|
+
return out
|