argonx 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.
argonx/__init__.py ADDED
@@ -0,0 +1,11 @@
1
+ """
2
+ argonx: Bayesian decision engine for robust A/B testing.
3
+
4
+ This package provides a comprehensive framework for Bayesian experimentation,
5
+ integrating primary metrics with guardrails, sequential stopping rules, and
6
+ hierarchical partial pooling.
7
+ """
8
+
9
+ from .experiment import Experiment
10
+
11
+ __all__ = ["Experiment"]
@@ -0,0 +1,99 @@
1
+ """
2
+ Decision rules and analysis engines for the Bayesian A/B testing framework.
3
+
4
+ This subpackage implements the full analytical pipeline that transforms raw posterior
5
+ samples into actionable experiment outcomes. It covers five distinct concerns:
6
+
7
+ - **Metrics** (`metrics.py`): Core Bayesian decision metrics — P(best), expected loss,
8
+ CVaR, ROPE analysis, and HDI-bounded lift. All metrics are derived from the same
9
+ posterior draws to ensure coherence.
10
+
11
+ - **Guardrails** (`guardrails.py`): Safety constraint evaluation. Computes P(degraded)
12
+ per metric per variant and surfaces conflicts when the primary metric clears its bar
13
+ but a secondary metric fails.
14
+
15
+ - **Joint** (`joint.py`): Joint probability of simultaneously satisfying both the primary
16
+ metric and all selected guardrails, with correlation diagnostics to expose when
17
+ metric co-movement helps or hurts the compound decision.
18
+
19
+ - **Composite** (`composite.py`): Weighted scoring of multiple metrics into a single
20
+ posterior distribution of business value, with asymmetric deterioration weights and
21
+ optional guardrail penalties.
22
+
23
+ - **Engine** (`engine.py`): Orchestration layer. Calls all of the above in dependency
24
+ order and assembles a structured `DecisionResult` with a plain-English recommendation.
25
+ """
26
+
27
+ from .composite import (
28
+ CompositeResult,
29
+ compute_composite_score,
30
+ )
31
+
32
+ from .engine import (
33
+ DecisionResult,
34
+ run_engine,
35
+ )
36
+
37
+ from .guardrails import (
38
+ ConflictResult,
39
+ GuardrailBundle,
40
+ GuardrailResult,
41
+ compute_all_guardrails,
42
+ compute_guardrail,
43
+ )
44
+
45
+ from .joint import (
46
+ JointResult,
47
+ compute_joint_probability,
48
+ )
49
+
50
+ from .metrics import (
51
+ CVaRResult,
52
+ LiftResult,
53
+ LossResult,
54
+ MetricsBundle,
55
+ PBestResult,
56
+ ROPEResult,
57
+ compute_all_metrics,
58
+ compute_cvar,
59
+ compute_expected_loss,
60
+ compute_lift_hdi,
61
+ compute_prob_best,
62
+ compute_rope,
63
+ )
64
+
65
+ __all__ = [
66
+ # Dataclasses — composite
67
+ "CompositeResult",
68
+ # Dataclasses — engine
69
+ "DecisionResult",
70
+ # Dataclasses — guardrails
71
+ "ConflictResult",
72
+ "GuardrailBundle",
73
+ "GuardrailResult",
74
+ # Dataclasses — joint
75
+ "JointResult",
76
+ # Dataclasses — metrics
77
+ "CVaRResult",
78
+ "LiftResult",
79
+ "LossResult",
80
+ "MetricsBundle",
81
+ "PBestResult",
82
+ "ROPEResult",
83
+ # Functions — composite
84
+ "compute_composite_score",
85
+ # Functions — engine
86
+ "run_engine",
87
+ # Functions — guardrails
88
+ "compute_all_guardrails",
89
+ "compute_guardrail",
90
+ # Functions — joint
91
+ "compute_joint_probability",
92
+ # Functions — metrics
93
+ "compute_all_metrics",
94
+ "compute_cvar",
95
+ "compute_expected_loss",
96
+ "compute_lift_hdi",
97
+ "compute_prob_best",
98
+ "compute_rope",
99
+ ]
@@ -0,0 +1,203 @@
1
+ from __future__ import annotations
2
+
3
+ import numpy as np
4
+ import warnings
5
+ from dataclasses import dataclass, field
6
+
7
+ from argonx.decision_rules.guardrails import GuardrailBundle
8
+
9
+
10
+ @dataclass
11
+ class CompositeResult:
12
+ """
13
+ Structured outcome of a composite business metric calculation across variants.
14
+
15
+ This encapsulates the final scores, the distributions representing Bayesian uncertainty,
16
+ and the relative contributions of individual metrics. It serves as the primary data
17
+ payload for rendering plotting results and making ultimate ship decisions.
18
+ """
19
+
20
+ score: dict[str, float]
21
+ score_distribution: dict[str, np.ndarray]
22
+ metric_contributions: dict[str, dict[str, float]]
23
+ prob_exceeds_threshold: dict[str, float]
24
+ gap_distribution: dict[str, np.ndarray]
25
+ gap_hdi: dict[str, tuple[float, float]]
26
+ best_variant: str
27
+ threshold: float
28
+ warnings: list[str] = field(default_factory=list)
29
+
30
+
31
+ def _compute_hdi(samples: np.ndarray, prob: float = 0.94) -> tuple[float, float]:
32
+ """
33
+ Calculate the Highest Density Interval (HDI) for a given set of posterior draws.
34
+
35
+ This function finds the shortest contiguous interval that contains a specified
36
+ proportion of the distribution, which is characteristic of the HDI.
37
+
38
+ Parameters
39
+ ----------
40
+ samples : np.ndarray
41
+ A 1-dimensional array of posterior draws.
42
+ prob : float, optional
43
+ The probability mass that should be contained within the interval, by default 0.94.
44
+
45
+ Returns
46
+ -------
47
+ tuple[float, float]
48
+ The lower and upper bounds of the computed credible interval.
49
+ """
50
+ sorted_samples = np.sort(samples)
51
+ n = len(samples)
52
+ interval_idx = int(np.floor(prob * n))
53
+ widths = sorted_samples[interval_idx:] - sorted_samples[: n - interval_idx]
54
+ min_idx = np.argmin(widths)
55
+ return (
56
+ float(sorted_samples[min_idx]),
57
+ float(sorted_samples[min_idx + interval_idx]),
58
+ )
59
+
60
+
61
+ def compute_composite_score(
62
+ primary_samples: np.ndarray,
63
+ guardrail_samples: dict[str, np.ndarray],
64
+ variant_names: list[str],
65
+ control: str,
66
+ weights: dict[str, float],
67
+ guardrail_bundle: GuardrailBundle,
68
+ deterioration_weights: dict[str, float] | None = None,
69
+ guardrail_penalty: float = 0.0,
70
+ threshold: float = 0.0,
71
+ ) -> CompositeResult:
72
+ """
73
+ Calculate a unified decision score integrating primary metrics and guardrails.
74
+
75
+ This function isolates the lift of each variant over control, applies directional
76
+ weights to model relative business impact, enforces penalties for failing categorical
77
+ guardrails, and aggregates them into a singular posterior distribution of scores.
78
+
79
+ Parameters
80
+ ----------
81
+ primary_samples : np.ndarray
82
+ Posterior samples for the primary target metric (draws x variants).
83
+ guardrail_samples : dict[str, np.ndarray]
84
+ A mapping of guardrail metric names to their corresponding posterior draws.
85
+ variant_names : list[str]
86
+ A sequential list of variant names corresponding to the columns in samples arrays.
87
+ control : str
88
+ The designated control variant used as the baseline for lift computation.
89
+ weights : dict[str, float]
90
+ Linear weights defining the reward (or penalty) for positive shifts per metric.
91
+ guardrail_bundle : GuardrailBundle
92
+ A evaluated bundle containing boolean pass/fail states for variants across guardrails.
93
+ deterioration_weights : dict[str, float] | None, optional
94
+ Asymmetric weights applied exclusively to negative degradation. Defaults to `weights`.
95
+ guardrail_penalty : float, optional
96
+ Fixed deduction applied to a variant's total score if it fails the bundle. Defaults to 0.0.
97
+ threshold : float, optional
98
+ The minimal composite score required to consider a variant viably improved. Defaults to 0.0.
99
+
100
+ Returns
101
+ -------
102
+ CompositeResult
103
+ The comprehensive scoring breakdown, including expected scores, underlying distributions,
104
+ and interval estimations.
105
+ """
106
+
107
+ collected: list[str] = []
108
+
109
+ if not weights:
110
+ raise ValueError("weights must be provided")
111
+
112
+ valid_metrics = {"primary"} | set(guardrail_samples.keys())
113
+ matched = [k for k in weights if k in valid_metrics]
114
+ if not matched:
115
+ raise ValueError(
116
+ f"No valid metric keys in weights. Expected one of {valid_metrics}"
117
+ )
118
+
119
+ if deterioration_weights is None:
120
+ deterioration_weights = weights.copy()
121
+ warnings.warn(
122
+ "deterioration_weights not provided, using symmetric weights",
123
+ UserWarning,
124
+ stacklevel=2,
125
+ )
126
+ collected.append("Using symmetric deterioration weights")
127
+
128
+ if guardrail_penalty == 0:
129
+ warnings.warn(
130
+ "guardrail_penalty is zero, guardrails have no effect",
131
+ UserWarning,
132
+ stacklevel=2,
133
+ )
134
+ collected.append("Guardrail penalty is zero")
135
+
136
+ c_idx = variant_names.index(control)
137
+ variants = [v for v in variant_names if v != control]
138
+
139
+ deltas = {}
140
+ control_vals = primary_samples[:, c_idx]
141
+ deltas["primary"] = primary_samples - control_vals[:, None]
142
+
143
+ for m, s in guardrail_samples.items():
144
+ deltas[m] = s - s[:, c_idx][:, None]
145
+
146
+ score_per_draw: dict[str, np.ndarray] = {}
147
+ contributions: dict[str, dict[str, float]] = {}
148
+
149
+ for i, v in enumerate(variant_names):
150
+ if v == control:
151
+ continue
152
+
153
+ total = np.zeros(primary_samples.shape[0])
154
+ contrib = {}
155
+
156
+ for m, d in deltas.items():
157
+ w = weights.get(m, 0.0)
158
+ dw = deterioration_weights.get(m, w)
159
+
160
+ val = d[:, i]
161
+
162
+ pos = np.maximum(val, 0)
163
+ neg = np.minimum(val, 0)
164
+
165
+ comp = w * pos + dw * neg
166
+ total += comp
167
+ contrib[m] = float(np.mean(comp))
168
+
169
+ variant_failed = not guardrail_bundle.variant_passed.get(v, True)
170
+ penalty = guardrail_penalty * (1 if variant_failed else 0)
171
+ total -= penalty
172
+
173
+ score_per_draw[v] = total
174
+ contributions[v] = contrib
175
+
176
+ score = {v: float(np.mean(score_per_draw[v])) for v in variants}
177
+
178
+ prob_exceeds = {v: float(np.mean(score_per_draw[v] > threshold)) for v in variants}
179
+
180
+ gap_distribution = {v: score_per_draw[v] - threshold for v in variants}
181
+
182
+ gap_hdi = {v: _compute_hdi(gap_distribution[v]) for v in variants}
183
+
184
+ best = max(score, key=lambda k: score[k])
185
+
186
+ if all(s < threshold for s in score.values()):
187
+ collected.append("No variant exceeds composite threshold")
188
+
189
+ for v, s in score.items():
190
+ if s < 0:
191
+ collected.append(f"{v} has negative composite score")
192
+
193
+ return CompositeResult(
194
+ score=score,
195
+ score_distribution=score_per_draw,
196
+ metric_contributions=contributions,
197
+ prob_exceeds_threshold=prob_exceeds,
198
+ gap_distribution=gap_distribution,
199
+ gap_hdi=gap_hdi,
200
+ best_variant=best,
201
+ threshold=threshold,
202
+ warnings=collected,
203
+ )
@@ -0,0 +1,335 @@
1
+ from __future__ import annotations
2
+
3
+ from dataclasses import dataclass, field
4
+
5
+ import numpy as np
6
+
7
+ from .metrics import MetricsBundle, compute_all_metrics
8
+ from .guardrails import GuardrailBundle, compute_all_guardrails
9
+ from .joint import JointResult, compute_joint_probability
10
+ from .composite import CompositeResult, compute_composite_score
11
+
12
+
13
+ @dataclass
14
+ class DecisionResult:
15
+ """Final decision output with structured interpretation."""
16
+
17
+ state: str
18
+ recommendation: str
19
+ best_variant: str
20
+
21
+ primary_strength: str
22
+ risk_level: str
23
+ practical_significance: str
24
+ guardrail_status: str
25
+ confidence: str
26
+
27
+ metrics: MetricsBundle
28
+ guardrails: GuardrailBundle
29
+ joint: JointResult | None
30
+ composite: CompositeResult | None
31
+
32
+ reasons: list[str] = field(default_factory=list)
33
+ notes: list[str] = field(default_factory=list)
34
+
35
+
36
+ def _evaluate_primary_strength(metrics: MetricsBundle, config: dict) -> str:
37
+ """Compute strength of primary signal."""
38
+ best = metrics.prob_best.best_variant
39
+ p_best = metrics.prob_best.probabilities[best]
40
+ loss = metrics.loss.expected_loss[best]
41
+ practical = metrics.rope.prob_practical.get(best, 0.0)
42
+
43
+ if (
44
+ p_best >= config["prob_best_strong"]
45
+ and loss <= config["expected_loss_max"]
46
+ and practical >= config["rope_practical_min"]
47
+ ):
48
+ return "strong"
49
+
50
+ if p_best >= config["prob_best_moderate"]:
51
+ return "moderate"
52
+
53
+ return "weak"
54
+
55
+
56
+ def _evaluate_risk(metrics: MetricsBundle, config: dict) -> str:
57
+ """Classify risk using expected loss and CVaR."""
58
+ best = metrics.prob_best.best_variant
59
+ el = metrics.loss.expected_loss[best]
60
+ cv = metrics.cvar.cvar[best]
61
+
62
+ if el <= config["expected_loss_max"] and (
63
+ el == 0 or cv / max(el, 1e-12) <= config["cvar_ratio_max"]
64
+ ):
65
+ return "low"
66
+
67
+ if el <= config["expected_loss_max"] * 5:
68
+ return "medium"
69
+
70
+ return "high"
71
+
72
+
73
+ def _evaluate_practical_significance(metrics: MetricsBundle, config: dict) -> str:
74
+ """Classify practical significance using ROPE."""
75
+ best = metrics.prob_best.best_variant
76
+ prob = metrics.rope.prob_practical.get(best, 0.0)
77
+
78
+ if prob >= config["rope_practical_min"]:
79
+ return "yes"
80
+
81
+ if prob >= 0.5:
82
+ return "uncertain"
83
+
84
+ return "no"
85
+
86
+
87
+ def _evaluate_guardrails(guardrails: GuardrailBundle) -> str:
88
+ """Summarize guardrail status."""
89
+ return "pass" if guardrails.all_passed else "fail"
90
+
91
+
92
+ def _evaluate_confidence(metrics: MetricsBundle, config: dict) -> str:
93
+ """Assess confidence using posterior certainty."""
94
+ best = metrics.prob_best.best_variant
95
+ p_best = metrics.prob_best.probabilities[best]
96
+
97
+ low = metrics.lift.hdi_low[best]
98
+ high = metrics.lift.hdi_high[best]
99
+ width = abs(high - low)
100
+
101
+ if p_best >= 0.95 and width < 0.1:
102
+ return "high"
103
+
104
+ if p_best >= 0.8:
105
+ return "medium"
106
+
107
+ return "low"
108
+
109
+
110
+ def _determine_state(
111
+ primary_strength: str,
112
+ risk_level: str,
113
+ practical_significance: str,
114
+ guardrails: GuardrailBundle,
115
+ ) -> str:
116
+ """Determine overall decision state."""
117
+ if guardrails.conflicts:
118
+ return "guardrail conflicts"
119
+
120
+ if (
121
+ primary_strength == "strong"
122
+ and risk_level == "low"
123
+ and practical_significance == "yes"
124
+ and guardrails.all_passed
125
+ ):
126
+ return "strong win"
127
+
128
+ if (
129
+ primary_strength == "moderate"
130
+ and guardrails.all_passed
131
+ and risk_level == "low"
132
+ and practical_significance == "yes"
133
+ ):
134
+ return "weak win"
135
+
136
+ if risk_level == "high":
137
+ return "high risk"
138
+
139
+ return "inconclusive"
140
+
141
+
142
+ def _map_recommendation(state: str) -> str:
143
+ """Map decision state to recommendation."""
144
+ return {
145
+ "strong win": "ship variant",
146
+ "weak win": "consider shipping",
147
+ "high risk": "do not ship",
148
+ "guardrail conflicts": "review required",
149
+ "inconclusive": "continue experiment",
150
+ }[state]
151
+
152
+
153
+ def _build_reasons(
154
+ primary_strength: str,
155
+ risk_level: str,
156
+ practical_significance: str,
157
+ guardrail_status: str,
158
+ ) -> list[str]:
159
+ """Generate human readable reasoning signals."""
160
+ reasons = [
161
+ f"Primary signal is {primary_strength}",
162
+ f"Risk level is {risk_level}",
163
+ f"Practical significance is {practical_significance}",
164
+ ]
165
+
166
+ if guardrail_status == "pass":
167
+ reasons.append("All guardrails passed")
168
+ else:
169
+ reasons.append("One or more guardrails failed")
170
+
171
+ return reasons
172
+
173
+
174
+ def _collect_notes(
175
+ metrics: MetricsBundle,
176
+ guardrails: GuardrailBundle,
177
+ joint: JointResult,
178
+ composite: CompositeResult,
179
+ ) -> list[str]:
180
+ """Collect warnings and edge case signals from all components."""
181
+ notes = []
182
+
183
+ notes.extend(metrics.warnings)
184
+ notes.extend(guardrails.warnings)
185
+
186
+ if joint is not None:
187
+ notes.extend(joint.warnings)
188
+
189
+ if composite is not None:
190
+ notes.extend(composite.warnings)
191
+
192
+ for conflict in guardrails.conflicts:
193
+ if conflict.severity == "high":
194
+ notes.append(
195
+ f"High severity guardrail violation on '{conflict.metric}' "
196
+ f"for variant '{conflict.variant}'"
197
+ )
198
+
199
+ for v in metrics.loss.expected_loss:
200
+ el = metrics.loss.expected_loss[v]
201
+ cv = metrics.cvar.cvar[v]
202
+ if el > 0 and cv / el > 5:
203
+ notes.append(f"Variant '{v}' has extreme tail risk")
204
+
205
+ return notes
206
+
207
+
208
+ def run_engine(
209
+ samples: np.ndarray,
210
+ variant_names: list[str],
211
+ control: str,
212
+ guardrail_samples: dict[str, np.ndarray],
213
+ config: dict,
214
+ ) -> DecisionResult:
215
+ """
216
+ Run full decision pipeline from posterior samples.
217
+
218
+ Processes primary Bayesian updates, checks safety guardrails, isolates joint constraint
219
+ probabilities, and coalesces performance metrics into a composite decision score.
220
+ Returns a structured recommendation detailing business risk, primary strength,
221
+ and significance.
222
+
223
+ Parameters
224
+ ----------
225
+ samples : np.ndarray
226
+ Array containing primary metric posterior draws.
227
+ variant_names : list[str]
228
+ Ordered collection of variants identifying the draw columns.
229
+ control : str
230
+ The primary baseline variant used for comparisons.
231
+ guardrail_samples : dict[str, np.ndarray]
232
+ Mapping of safety metric designations to their posterior distributions.
233
+ config : dict
234
+ A master configuration dictionary determining parameters like decision thresholds,
235
+ cvar bounds, rope regions, and guardrail limits.
236
+
237
+ Returns
238
+ -------
239
+ DecisionResult
240
+ A bundled diagnostic report mapping qualitative shipping recommendations to
241
+ underlying statistical confidence intervals and rule triggers.
242
+ """
243
+
244
+ metrics = compute_all_metrics(
245
+ samples=samples,
246
+ variant_names=variant_names,
247
+ control=control,
248
+ rope_bounds=config["rope_bounds"],
249
+ alpha=config.get("alpha", 0.95),
250
+ hdi_prob=config.get("hdi_prob", 0.95),
251
+ )
252
+
253
+ primary_strength = _evaluate_primary_strength(metrics, config)
254
+ risk_level = _evaluate_risk(metrics, config)
255
+ practical_significance = _evaluate_practical_significance(metrics, config)
256
+
257
+ primary_passed = (
258
+ primary_strength == "strong"
259
+ and risk_level != "high"
260
+ and practical_significance == "yes"
261
+ )
262
+
263
+ guardrails = compute_all_guardrails(
264
+ guardrail_samples=guardrail_samples,
265
+ variant_names=variant_names,
266
+ control=control,
267
+ thresholds=config.get("guardrail_thresholds", {}),
268
+ primary_passed=primary_passed,
269
+ lower_is_better=config.get("lower_is_better", {}),
270
+ )
271
+
272
+ guardrail_status = _evaluate_guardrails(guardrails)
273
+ confidence = _evaluate_confidence(metrics, config)
274
+
275
+ joint = None
276
+ if guardrail_samples:
277
+ joint = compute_joint_probability(
278
+ primary_samples=samples,
279
+ guardrail_samples=guardrail_samples,
280
+ variant_names=variant_names,
281
+ control=control,
282
+ primary_lower_is_better=config.get("primary_lower_is_better", False),
283
+ lower_is_better=config.get("lower_is_better", {}),
284
+ guardrail_thresholds=config.get("guardrail_thresholds", {}),
285
+ metrics_to_join=config.get("metrics_to_join", None),
286
+ )
287
+
288
+ composite = None
289
+ if "composite_weights" in config:
290
+ composite = compute_composite_score(
291
+ primary_samples=samples,
292
+ guardrail_samples=guardrail_samples,
293
+ variant_names=variant_names,
294
+ control=control,
295
+ weights=config["composite_weights"],
296
+ guardrail_bundle=guardrails,
297
+ deterioration_weights=config.get("deterioration_weights", None),
298
+ guardrail_penalty=config.get("guardrail_penalty", 0.0),
299
+ threshold=config.get("composite_threshold", 0.0),
300
+ )
301
+
302
+ state = _determine_state(
303
+ primary_strength,
304
+ risk_level,
305
+ practical_significance,
306
+ guardrails,
307
+ )
308
+
309
+ recommendation = _map_recommendation(state)
310
+
311
+ reasons = _build_reasons(
312
+ primary_strength,
313
+ risk_level,
314
+ practical_significance,
315
+ guardrail_status,
316
+ )
317
+
318
+ notes = _collect_notes(metrics, guardrails, joint, composite)
319
+
320
+ return DecisionResult(
321
+ state=state,
322
+ recommendation=recommendation,
323
+ best_variant=metrics.prob_best.best_variant,
324
+ primary_strength=primary_strength,
325
+ risk_level=risk_level,
326
+ practical_significance=practical_significance,
327
+ guardrail_status=guardrail_status,
328
+ confidence=confidence,
329
+ metrics=metrics,
330
+ guardrails=guardrails,
331
+ joint=joint,
332
+ composite=composite,
333
+ reasons=reasons,
334
+ notes=notes,
335
+ )