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.
Files changed (40) hide show
  1. conforme/__init__.py +81 -0
  2. conforme/backtest/__init__.py +6 -0
  3. conforme/backtest/forecasts.py +142 -0
  4. conforme/backtest/replay.py +96 -0
  5. conforme/conformal/__init__.py +42 -0
  6. conforme/conformal/calibrators/__init__.py +25 -0
  7. conforme/conformal/calibrators/aci.py +42 -0
  8. conforme/conformal/calibrators/base.py +97 -0
  9. conforme/conformal/calibrators/ranks.py +92 -0
  10. conforme/conformal/calibrators/risk.py +60 -0
  11. conforme/conformal/calibrators/split.py +120 -0
  12. conforme/conformal/calibrators/tracker.py +33 -0
  13. conforme/conformal/losses.py +31 -0
  14. conforme/conformal/scores.py +44 -0
  15. conforme/conformal/targets.py +47 -0
  16. conforme/data/__init__.py +6 -0
  17. conforme/data/hierarchy.py +103 -0
  18. conforme/data/panel.py +56 -0
  19. conforme/decision/__init__.py +6 -0
  20. conforme/decision/lost_sales.py +43 -0
  21. conforme/decision/policy.py +27 -0
  22. conforme/forecast/__init__.py +16 -0
  23. conforme/forecast/models/__init__.py +10 -0
  24. conforme/forecast/models/base.py +67 -0
  25. conforme/forecast/models/frames.py +83 -0
  26. conforme/forecast/models/mlforecast.py +100 -0
  27. conforme/forecast/models/naive.py +24 -0
  28. conforme/forecast/models/neuralforecast.py +77 -0
  29. conforme/forecast/models/statsforecast.py +32 -0
  30. conforme/forecast/reconcile.py +70 -0
  31. conforme/metrics.py +60 -0
  32. conforme/online/__init__.py +6 -0
  33. conforme/online/ledger.py +119 -0
  34. conforme/online/state.py +33 -0
  35. conforme/online/step.py +70 -0
  36. conforme/py.typed +0 -0
  37. conforme-0.1.0.dist-info/METADATA +216 -0
  38. conforme-0.1.0.dist-info/RECORD +40 -0
  39. conforme-0.1.0.dist-info/WHEEL +4 -0
  40. 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