diff --git a/CONTEXT.md b/CONTEXT.md index 5437040b..9d938dc1 100644 --- a/CONTEXT.md +++ b/CONTEXT.md @@ -122,7 +122,7 @@ _Avoid_: "order of differencing" for `d_max` — `d_max` is the maximum across t The number of independent long-run relationships among integrated series, from the Johansen procedure (`johansen_test`). Both sequential tests are reported — `rank_trace` and `rank_max_eigen` — and `rank` is `rank_trace` by documented convention. Decisions rest on **critical values, not p-values** (MacKinnon-Haug-Michelis 1996 tables, as vendored by statsmodels), which is why `alpha` is restricted to 0.10 / 0.05 / 0.01. Rank ≥ 1 means differencing every series discards the long-run relationship. A vector error-correction model (VECM) is **out of scope**; the recommended response is a VAR in levels (the Sims–Stock–Watson stance; the Minnesota prior already shrinks toward random walks). _Avoid_: "number of cointegrating vectors" in API surface (fine in prose); "cointegration test" without saying which statistic, since trace and max-eigen can disagree. **Convergence report**: -The VAR-aware verdict on whether a fitted posterior is usable, produced by `convergence_report()` (or the delegating `FittedVAR.convergence_report()` / `IdentifiedVAR.convergence_report()`). It reports R-hat and both effective sample sizes per *parameter block* with the worst coordinate named, the global divergence count, and the posterior distribution of the spectral radius, and carries machine-readable `DiagnosticMessage` codes for the two VAR-specific failure modes. Its `status` — `"passed"` / `"warnings"` / `"failed"` — reserves `"failed"` for sampler pathology; explosive draws warn but never fail. See `docs/adr/0008-convergence-report-blocks-and-thresholds.md`. +The VAR-aware verdict on whether a fitted posterior is usable, produced by `convergence_report()` (or the delegating `FittedVAR.convergence_report()` / `IdentifiedVAR.convergence_report()`). It reports R-hat and both effective sample sizes per *parameter block* with the worst coordinate named, the global divergence, energy (E-BFMI) and max-treedepth statistics, and the posterior distribution of the spectral radius, and carries machine-readable `DiagnosticMessage` codes for the two VAR-specific failure modes. Its `status` — `"passed"` / `"warnings"` / `"failed"` — reserves `"failed"` for sampler pathology; explosive draws, low E-BFMI and treedepth saturation warn but never fail. See `docs/adr/0008-convergence-report-blocks-and-thresholds.md`. _Avoid_: "diagnostics" as a synonym for this one object — the diagnostics family is wider, and the name `.diagnostics()` is reserved. **Parameter block**: @@ -191,4 +191,4 @@ _Avoid_: "unstable" for a whole posterior — stability is a per-draw property, - "Σ" now means the *scale* matrix under `StudentT` errors and the covariance under `Gaussian` errors. `sigma()` returns the same object either way; when the number has to be a variance, say so and use `innovation_covariance()`. - "Counterfactual" in the wider literature spans shock-path edits (Impulso's meaning), policy-rule replacement (Sims–Zha style; out of scope), and Lucas-robust constructions (McKay–Wolf; out of scope). When comparing with external work, say which one is meant. - "Companion" is overloaded: the *companion matrix* is the stacked first-order form of a VAR(p), while the ADPRR "calibrated companion" `q_cal` is the plausibility statistic's partner quantity. They share nothing. Always write "companion matrix" in full; never shorten it to "the companion". -- `StabilitySummary` (the convergence report's spectral-radius block) is distinct from the `StabilityResult` planned for the ecological-stability work: the former summarises one scalar per draw for a diagnostic verdict, the latter will carry the full complex eigenvalue spectrum and the reactivity/return-rate measures derived from it. Both read `companion_eigenvalues`; neither subsumes the other. +- `StabilitySummary` (the convergence report's spectral-radius block) is distinct from the `StabilityResult` planned for the ecological-stability work: the former summarises one scalar per draw for a diagnostic verdict, keeping a capped subset of the complex roots only so `plot()` has something to scatter, whereas the latter will carry the full spectrum for every draw and the reactivity/return-rate measures derived from it. Both read `companion_eigenvalues`; neither subsumes the other. diff --git a/docs/adr/0008-convergence-report-blocks-and-thresholds.md b/docs/adr/0008-convergence-report-blocks-and-thresholds.md index 5faa75a4..f89e85b3 100644 --- a/docs/adr/0008-convergence-report-blocks-and-thresholds.md +++ b/docs/adr/0008-convergence-report-blocks-and-thresholds.md @@ -23,9 +23,13 @@ The `identification` block exists but is normally empty: the structural shock ma | R-hat | 1.01 | 1.05 | Vehtari et al. (2021); classic Gelman–Rubin | | Effective sample size | 400 | 100 | 100 per chain at four chains | | Divergence rate | any divergence | 1% | Betancourt (2017) | +| E-BFMI | 0.3 | *never* | Betancourt (2016), arXiv:1604.00695 | +| Max-treedepth saturation rate | 1% | *never* | — | | Explosive draw fraction | 5% | *never* | — | -R-hat and ESS comparisons are strict, so a metric sitting exactly on a threshold passes; the two rate thresholds (divergence rate, explosive fraction) trigger at the boundary. Thresholds live in a frozen `ConvergenceThresholds` model rather than as module constants so a caller can tighten them for a specific study and the report echoes back what it used. +R-hat, ESS and E-BFMI comparisons are strict, so a metric sitting exactly on a threshold passes; the three rate thresholds (divergence rate, treedepth saturation rate, explosive fraction) trigger at the boundary. Thresholds live in a frozen `ConvergenceThresholds` model rather than as module constants so a caller can tighten them for a specific study and the report echoes back what it used. + +The treedepth threshold has no literature source, unlike the others. Stan and PyMC surface any saturation at all, but a handful of hits in a long run costs wall-clock time and nothing else, so reporting them would train users to skim past the message. One transition in a hundred is the point at which the sampler is spending real effort on trajectories it never gets to finish. Both backends record the flag under different names — `reached_max_treedepth` in PyMC, `maxdepth_reached` in nutpie — so this is the one statistic the report resolves through a name map rather than reading directly. ## Why explosive draws never fail @@ -33,10 +37,17 @@ Posterior mass on parameter draws whose companion matrix has spectral radius at The reason is that explosiveness is a property of the *model*, not of the sampler. Macroeconomic data in levels under a Minnesota prior centred on a random walk puts substantial mass near the unit circle by construction; that is the prior doing its job, and a fraction of draws crossing it is expected rather than pathological. Failing the report there would train users to ignore `"failed"`, which must keep meaning "these draws do not describe the posterior". Convergence and stability are different questions and are reported as such. +## Why E-BFMI and max-treedepth never fail either + +Both are statements about *efficiency*, not about wrongness. A saturated tree depth means NUTS stopped a trajectory before it turned back on itself, so the draws are more autocorrelated than they need to be; a low E-BFMI means momentum resampling is exploring the energy distribution slowly, so the tails are undersampled relative to the bulk. Neither says the retained draws come from the wrong distribution — unlike a divergence, which says the sampler could not follow the geometry at all, or an unmixed R-hat, which says the chains are not describing one distribution. + +Both are also remediable by changing the sampler or the parameterisation without touching the model, so failing on them would block work that is merely slower than it should be. Both metrics are on the report regardless, so a caller who wants them fatal reads `ebfmi` and `treedepth_saturation_rate` and decides. + ## Rejected alternatives - **A thin wrapper over `az.summary`.** Rejected: it produces one row per coordinate with no block structure, no stability, and no VAR-specific interpretation — exactly the output users already have and cannot act on. - **Living in `results.py` alongside the other result objects.** Rejected: `VARResultBase` contracts for `median`/`hdi`/`to_dataframe`/`plot` over a posterior-predictive DataArray, and a convergence report has no such array. Following `LagOrderResult`'s precedent would have forced a fake `plot` and a fake `median`. A dedicated `diagnostics.py` also gives the diagnostics family (issue #57's umbrella) somewhere to grow. - **Per-block divergence attribution.** Rejected: a divergence is a property of a trajectory through the whole parameter space. Splitting the count by block would invent an attribution the sampler never made. - **Warning through `warnings.warn`.** Rejected: the report object carries `status`, `messages`, and the per-block table, so the caller decides whether to print, raise, or ignore. A diagnostic that emits warnings cannot be used inside a loop over model specifications. -- **A `.plot()` method in v1.** Deferred: issue #57 owns diagnostic visuals, and the raw `(chain, draw)` radius array is exposed so a histogram or unit-circle scatter is a few lines away. +- **A `.plot()` method in v1.** Deferred, then added on `StabilitySummary` alone (issue #179): `report.stability.plot()` gives the spectral-radius histogram beside the companion-root scatter, matching how every other result object in the library is plotted. `ConvergenceReport` itself still has no `plot` — a block table of R-hat and ESS is a table, and issue #57 owns whatever diagnostic visuals go beyond stability. +- **Retaining every companion eigenvalue on `StabilitySummary`.** Rejected: the array grows as `draws × n_vars × n_lags` complex numbers — roughly 15 MB for 4000 draws of a 240×240 companion matrix — on an object whose every statistic is derived from the radii. The summary keeps a chain-pooled, deterministically strided subset of at most 200 draws, computed from the single eigendecomposition the radii already require, which is more points than the scatter panel can distinguish and a fixed ceiling regardless of posterior size. diff --git a/docs/reference/diagnostics.md b/docs/reference/diagnostics.md index 4d5a64b3..f1116d07 100644 --- a/docs/reference/diagnostics.md +++ b/docs/reference/diagnostics.md @@ -2,11 +2,13 @@ Convergence and dynamic-stability diagnostics for a fitted VAR posterior. `convergence_report` reports R-hat and effective sample size *per parameter -block* with the offending coordinate named, counts divergences globally, and -summarises the posterior distribution of the companion-matrix spectral -radius. Reach for it through `FittedVAR.convergence_report()` or +block* with the offending coordinate named; counts divergences, energy +pathologies (E-BFMI) and max-treedepth saturation globally; and summarises +the posterior distribution of the companion-matrix spectral radius. Reach +for it through `FittedVAR.convergence_report()` or `IdentifiedVAR.convergence_report()`; the free function is the entry point -for posteriors built by hand. +for posteriors built by hand. `report.stability.plot()` draws the radius +posterior beside the companion roots on the unit circle. ```{eval-rst} .. currentmodule:: impulso.diagnostics diff --git a/docs/reference/plotting.md b/docs/reference/plotting.md index ebe61488..d6eff984 100644 --- a/docs/reference/plotting.md +++ b/docs/reference/plotting.md @@ -17,4 +17,5 @@ plot_counterfactual plot_volatility plot_sv_forecast + plot_stability ``` diff --git a/src/impulso/diagnostics.py b/src/impulso/diagnostics.py index 244b03cd..c776733c 100644 --- a/src/impulso/diagnostics.py +++ b/src/impulso/diagnostics.py @@ -40,10 +40,11 @@ from pydantic import Field, model_validator from impulso._base import ImpulsoBaseModel, ImpulsoModel -from impulso._stability import spectral_radius +from impulso._stability import companion_eigenvalues if TYPE_CHECKING: import xarray as xr + from matplotlib.figure import Figure from impulso.protocols import VolatilityProcess @@ -82,10 +83,26 @@ # `v{i}_` prefix (see `impulso.sv.spec`), so the whole family maps by pattern. _SV_PREFIX = re.compile(r"^v\d+_") -# The only sampler statistic PyMC and nutpie agree on the name of. Every other -# stat (tree depth, acceptance rate, energy) is spelled differently by the two -# backends, so the report reads this one and nothing else. +# PyMC and nutpie agree on the name of exactly two sampler statistics, and the +# report reads only those two directly. _DIVERGING_KEY = "diverging" +_ENERGY_KEY = "energy" + +# Everything else is spelled differently by the two backends. Max-treedepth +# saturation is a boolean per transition called `reached_max_treedepth` by PyMC +# and `maxdepth_reached` by nutpie, so it needs an explicit name map; the first +# name present wins, and a posterior carrying neither simply has no treedepth +# statistics. +_MAX_TREEDEPTH_KEYS: tuple[str, ...] = ("reached_max_treedepth", "maxdepth_reached") + +# The eigenvalues behind the spectral radii are kept only so that +# `StabilitySummary.plot` has something to scatter, and the full set grows as +# `draws * n_vars * n_lags` complex numbers — roughly 15 MB for 4000 draws of a +# 240x240 companion matrix, which has no business living on a frozen result. +# The summary therefore retains a chain-pooled, deterministically strided +# subset of at most this many draws: more points than a scatter plot can +# distinguish, and a fixed ceiling regardless of posterior size. +_PLOT_EIGENVALUE_DRAWS = 200 def assign_blocks( @@ -143,11 +160,12 @@ def assign_blocks( class ConvergenceThresholds(ImpulsoModel): """Cut-offs separating a passing report from warnings and failures. - R-hat and ESS comparisons are strict, so a metric sitting exactly on a - threshold passes: `max_rhat == 1.01` does not warn. The rate thresholds - trigger at the boundary: a divergence rate of exactly - `divergence_fail_rate` fails, and an explosive fraction of exactly - `explosive_warn` escalates to a warning. + R-hat, ESS and E-BFMI comparisons are strict, so a metric sitting exactly + on a threshold passes: `max_rhat == 1.01` does not warn, and an E-BFMI of + exactly 0.3 does not warn. The rate thresholds trigger at the boundary: a + divergence rate of exactly `divergence_fail_rate` fails, a treedepth + saturation rate of exactly `treedepth_warn_rate` warns, and an explosive + fraction of exactly `explosive_warn` escalates to a warning. Attributes: rhat_warn: R-hat above this warns. Default 1.01, the rank-normalised @@ -159,6 +177,17 @@ class ConvergenceThresholds(ImpulsoModel): ess_fail: Effective sample size below this fails. Default 100. divergence_fail_rate: Divergence rate at or above which the report fails. Default 0.01; any divergence at all warns. + ebfmi_warn: Energy Bayesian fraction of missing information (E-BFMI) + below this warns. Default 0.3, the cut-off Betancourt (2016, + arXiv:1604.00695) proposes and ArviZ documents on `arviz.bfmi`. + E-BFMI never fails a report. + treedepth_warn_rate: Fraction of post-warmup transitions that + saturated the sampler's maximum tree depth, at or above which the + report warns. Default 0.01. Stan and PyMC surface any saturation + at all, but a handful of hits in a long run costs only wall-clock + time and says nothing about the draws, so Impulso sets the bar at + one transition in a hundred and stays silent below it. Treedepth + saturation never fails a report. explosive_warn: Fraction of explosive draws at or above which the explosive-draw message is raised from informational to a warning. Default 0.05. Explosive draws never fail a report. @@ -169,6 +198,8 @@ class ConvergenceThresholds(ImpulsoModel): ess_warn: float = 400.0 ess_fail: float = 100.0 divergence_fail_rate: float = 0.01 + ebfmi_warn: float = 0.3 + treedepth_warn_rate: float = 0.01 explosive_warn: float = 0.05 @@ -236,6 +267,11 @@ class StabilitySummary(ImpulsoBaseModel): Attributes: radius: Read-only spectral radii with shape `(chain, draw)` — after thinning, if `stability_draws` was used. + eigenvalues: Read-only complex companion-matrix roots with shape + `(draws, n_vars * n_lags)`, pooled over chains and strided down + to at most 200 draws. Every statistic on this object is derived + from `radius`; these are retained for `plot` alone, which is why + they are capped rather than kept in full. p_explosive: Fraction of draws with radius >= 1. max_radius: Largest radius over all draws. n_vars: Number of endogenous variables. @@ -247,6 +283,7 @@ class StabilitySummary(ImpulsoBaseModel): """ radius: np.ndarray = Field(repr=False) + eigenvalues: np.ndarray = Field(repr=False) p_explosive: float max_radius: float n_vars: int @@ -256,9 +293,10 @@ class StabilitySummary(ImpulsoBaseModel): @model_validator(mode="after") def _make_readonly(self) -> Self: - radius = np.asarray(self.radius).copy() - radius.flags.writeable = False - object.__setattr__(self, "radius", radius) + for field in ("radius", "eigenvalues"): + array = np.asarray(getattr(self, field)).copy() + array.flags.writeable = False + object.__setattr__(self, field, array) return self def median(self) -> float: @@ -299,6 +337,20 @@ def to_dataframe(self) -> pd.DataFrame: index=pd.Index(["stability"], name="quantity"), ) + def plot(self) -> Figure: + """Spectral-radius histogram beside the eigenvalue scatter. + + This is the single plotting entry point for stability; a report + reaches it as `report.stability.plot()`, matching every other result + object in the library. + + Returns: + Matplotlib Figure with two axes. + """ + from impulso.plotting import plot_stability + + return plot_stability(self) + class ConvergenceReport(ImpulsoBaseModel): """VAR-aware convergence and stability diagnostics for one posterior. @@ -314,6 +366,14 @@ class ConvergenceReport(ImpulsoBaseModel): n_transitions: Total post-warmup transitions, or None. divergence_rate: `divergences / n_transitions`, or None. sampler_stats_available: Whether divergence statistics were found. + ebfmi: Per-chain energy Bayesian fraction of missing information, in + chain order, or None when the posterior carries no `energy` + statistic (every conjugate and hand-built posterior). + treedepth_saturations: Number of post-warmup transitions that hit the + sampler's maximum tree depth, or None when neither backend's + treedepth statistic is present. + treedepth_saturation_rate: `treedepth_saturations / n_transitions`, + or None. n_chains: Number of chains in the posterior. n_draws: Number of post-warmup draws per chain. thresholds: Thresholds used to derive `status`. @@ -328,12 +388,20 @@ class ConvergenceReport(ImpulsoBaseModel): n_transitions: int | None divergence_rate: float | None sampler_stats_available: bool + ebfmi: list[float] | None + treedepth_saturations: int | None + treedepth_saturation_rate: float | None n_chains: int n_draws: int thresholds: ConvergenceThresholds messages: list[DiagnosticMessage] status: Literal["passed", "warnings", "failed"] + @property + def min_ebfmi(self) -> float | None: + """Worst per-chain E-BFMI, or None when energy was not recorded.""" + return None if not self.ebfmi else min(self.ebfmi) + @property def max_rhat(self) -> float | None: """Worst R-hat across all blocks, or None if undefined everywhere.""" @@ -369,10 +437,12 @@ def to_dataframe(self) -> pd.DataFrame: def summary(self) -> str: """Multi-line human-readable rendering of the whole report.""" divergences = "unavailable" if self.divergences is None else str(self.divergences) - header = [ - f"Convergence report: {self.status.upper()}", - f" {self.n_chains} chains x {self.n_draws} draws | divergences: {divergences}", - ] + sampler = [f"{self.n_chains} chains x {self.n_draws} draws", f"divergences: {divergences}"] + if self.min_ebfmi is not None: + sampler.append(f"min E-BFMI: {self.min_ebfmi:.2f}") + if self.treedepth_saturations is not None and self.treedepth_saturation_rate is not None: + sampler.append(f"max-treedepth hits: {self.treedepth_saturations} ({self.treedepth_saturation_rate:.2%})") + header = [f"Convergence report: {self.status.upper()}", " " + " | ".join(sampler)] table = self.to_dataframe()[["max_rhat", "min_ess_bulk", "min_ess_tail", "max_rhat_coord"]] lower, upper = self.stability.hdi() stability = ( @@ -526,6 +596,49 @@ def _divergences(idata: az.InferenceData) -> tuple[int | None, int | None, float return count, total, (count / total if total else 0.0), True +def _ebfmi(idata: az.InferenceData) -> list[float] | None: + """Per-chain E-BFMI, or None when the posterior carries no energy trace. + + The array is handed to `arviz.bfmi` rather than the whole + `InferenceData`, so only the post-warmup group is ever read — nutpie's + `warmup_sample_stats` also carries an `energy` variable, and adaptation + energy says nothing about the retained draws. + """ + if "sample_stats" not in idata.groups() or _ENERGY_KEY not in idata.sample_stats: + return None + energy = idata.sample_stats[_ENERGY_KEY] + if {"chain", "draw"} <= set(energy.dims): + energy = energy.transpose("chain", "draw") + values = np.asarray(energy.values, dtype=float) + # A constant energy trace has zero variance and divides by zero inside + # ArviZ. That is not a diagnosis, so it is reported as absent, not as a + # pathology, and numpy's notice is silenced so a report never warns. + with warnings.catch_warnings(): + warnings.simplefilter("ignore", RuntimeWarning) + bfmi = np.atleast_1d(np.asarray(az.bfmi(values), dtype=float)) + if not bool(np.all(np.isfinite(bfmi))): + return None + return [float(value) for value in bfmi] + + +def _treedepth(idata: az.InferenceData) -> tuple[int | None, int | None, float | None]: + """Max-treedepth saturation count, transition total and rate. + + Both backends record a boolean per transition; only the name differs + (see `_MAX_TREEDEPTH_KEYS`). As with divergences, the warmup group is + never read: saturation during adaptation is expected and harmless. + """ + if "sample_stats" not in idata.groups(): + return None, None, None + key = next((name for name in _MAX_TREEDEPTH_KEYS if name in idata.sample_stats), None) + if key is None: + return None, None, None + saturated = np.asarray(idata.sample_stats[key].values).astype(bool) + count = int(saturated.sum()) + total = int(saturated.size) + return count, total, (count / total if total else 0.0) + + # -------------------------------------------------------------------------- # Messages and status # -------------------------------------------------------------------------- @@ -666,6 +779,64 @@ def _divergence_messages( ] +def _ebfmi_messages( + ebfmi: list[float] | None, + thresholds: ConvergenceThresholds, +) -> list[DiagnosticMessage]: + """The low-energy-fraction finding. A warning, never a failure.""" + if not ebfmi: + return [] + worst = min(ebfmi) + if worst >= thresholds.ebfmi_warn: + return [] + chain = ebfmi.index(worst) + return [ + DiagnosticMessage( + code="low_ebfmi", + severity="warning", + message=( + f"E-BFMI falls to {worst:.2f} on chain {chain}, below the warning threshold " + f"of {thresholds.ebfmi_warn}. Momentum resampling is not matching the " + "marginal energy distribution, so the sampler explores the tails of the " + "posterior slowly and effective sample size there is worse than the bulk " + "figures suggest. In a VAR this usually means a funnel between the " + "shrinkage hyperparameters and the coefficients they govern: reparameterise " + "the hierarchy non-centred, tighten the hyperprior, or lengthen `tune`. " + "Longer chains alone rarely fix it." + ), + ) + ] + + +def _treedepth_messages( + saturations: int | None, + rate: float | None, + n_transitions: int | None, + thresholds: ConvergenceThresholds, +) -> list[DiagnosticMessage]: + """The max-treedepth finding. A warning, never a failure.""" + if not saturations or rate is None or rate < thresholds.treedepth_warn_rate: + return [] + total = n_transitions if n_transitions is not None else 0 + return [ + DiagnosticMessage( + code="treedepth_saturation", + severity="warning", + message=( + f"{saturations} of {total} transitions saturated the maximum tree depth " + f"({rate:.2%}, at or above the warning threshold of " + f"{thresholds.treedepth_warn_rate:.2%}). NUTS was cut off before its " + "trajectory turned back on itself, so those transitions moved less far than " + "they should have and the draws are more autocorrelated than the tuning " + "suggests. This costs efficiency, not correctness. Raise `max_treedepth`, or " + "better, fix the ill-conditioning that makes long trajectories necessary: " + "nutpie's low-rank mass matrix — " + f"{_NUTPIE_REMEDY} — usually removes the need for deep trees in a VAR." + ), + ) + ] + + def _chain_messages(n_chains: int) -> list[DiagnosticMessage]: """The single-chain note: R-hat is a between-chain statistic.""" if n_chains >= 2: @@ -754,9 +925,15 @@ def _stability_summary( B = B[:, ::stride] thinned_from = n_draws - radius = spectral_radius(B, n_lags) + # One eigendecomposition serves both outputs: the radii are the row-wise + # maximum modulus, and `plot` gets a strided subset of the same roots. + eigenvalues = companion_eigenvalues(B, n_lags) + radius = np.max(np.abs(eigenvalues), axis=-1) + pooled = eigenvalues.reshape(-1, eigenvalues.shape[-1]) + stride = max(1, -(-pooled.shape[0] // _PLOT_EIGENVALUE_DRAWS)) return StabilitySummary( radius=radius, + eigenvalues=pooled[::stride], p_explosive=float(np.mean(radius >= 1.0)), max_radius=float(np.max(radius)), n_vars=B.shape[-2], @@ -779,9 +956,10 @@ def convergence_report( """Build a VAR-aware convergence and stability report for a posterior. Reports R-hat and both effective sample sizes per parameter block, each - with the coordinate that attains the worst value; the global divergence - count; and the posterior distribution of the companion-matrix spectral - radius. See the module docstring for why a VAR needs its own report. + with the coordinate that attains the worst value; the global divergence, + energy and max-treedepth statistics; and the posterior distribution of + the companion-matrix spectral radius. See the module docstring for why a + VAR needs its own report. `"failed"` is reserved for sampler pathology — R-hat above `rhat_fail`, effective sample size below `ess_fail`, or a divergence @@ -789,6 +967,20 @@ def convergence_report( fail: mass near a unit root is a legitimate posterior statement about level data, not evidence that the sampler misbehaved. + Low E-BFMI and max-treedepth saturation also warn rather than fail, for + a different reason: both are statements about *efficiency*, not about + wrongness. A saturated tree depth means NUTS stopped a trajectory early, + so the draws are more autocorrelated than they need to be; a low E-BFMI + means momentum resampling is exploring the energy distribution slowly, + so the tails are undersampled. Neither says the retained draws come from + the wrong distribution — unlike a divergence, which says the sampler + could not follow the geometry at all, or an unmixed R-hat, which says + the chains are not describing one distribution. Both are also remediable + by changing the sampler or the parameterisation without changing the + model, so failing a report on them would block work that is merely + slower than it should be. The metrics are on the report either way, so a + caller who wants them to be fatal can read the fields and decide. + Cost is dominated by the eigendecomposition of one `(n * p, n * p)` companion matrix per draw. Pass `stability_draws` to compute the radii on a deterministically strided subset when the posterior is large. @@ -840,6 +1032,8 @@ def convergence_report( ] stability = _stability_summary(posterior, n_lags, hdi_prob, stability_draws) divergences, n_transitions, rate, stats_available = _divergences(idata) + ebfmi = _ebfmi(idata) + treedepth_saturations, treedepth_total, treedepth_rate = _treedepth(idata) max_rhat = _extreme([block.max_rhat for block in blocks], "max") rhat_coord = next((block.max_rhat_coord for block in blocks if block.max_rhat == max_rhat), None) @@ -854,6 +1048,8 @@ def convergence_report( *_ess_messages(min_bulk, bulk_coord, "bulk", thresholds), *_ess_messages(min_tail, tail_coord, "tail", thresholds), *_divergence_messages(divergences, n_transitions, rate, stats_available, thresholds), + *_ebfmi_messages(ebfmi, thresholds), + *_treedepth_messages(treedepth_saturations, treedepth_rate, treedepth_total, thresholds), *_chain_messages(n_chains), *_stability_messages(stability, thresholds), ] @@ -865,6 +1061,9 @@ def convergence_report( n_transitions=n_transitions, divergence_rate=rate, sampler_stats_available=stats_available, + ebfmi=ebfmi, + treedepth_saturations=treedepth_saturations, + treedepth_saturation_rate=treedepth_rate, n_chains=n_chains, n_draws=int(posterior.sizes["draw"]), thresholds=thresholds, diff --git a/src/impulso/plotting/__init__.py b/src/impulso/plotting/__init__.py index 94a84743..d3237611 100644 --- a/src/impulso/plotting/__init__.py +++ b/src/impulso/plotting/__init__.py @@ -7,6 +7,7 @@ from impulso.plotting._forecast import plot_forecast from impulso.plotting._historical_decomposition import plot_historical_decomposition from impulso.plotting._irf import plot_irf +from impulso.plotting._stability import plot_stability from impulso.plotting._structural_scenario import plot_structural_scenario from impulso.plotting._sv_forecast import plot_sv_forecast from impulso.plotting._sv_volatility import plot_volatility @@ -19,6 +20,7 @@ "plot_forecast", "plot_historical_decomposition", "plot_irf", + "plot_stability", "plot_structural_scenario", "plot_sv_forecast", "plot_volatility", diff --git a/src/impulso/plotting/_stability.py b/src/impulso/plotting/_stability.py new file mode 100644 index 00000000..7da44dff --- /dev/null +++ b/src/impulso/plotting/_stability.py @@ -0,0 +1,71 @@ +"""Companion-matrix stability plotting.""" + +from typing import TYPE_CHECKING + +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.figure import Figure + +if TYPE_CHECKING: + from impulso.diagnostics import StabilitySummary + +# Legend labels for the two reference marks, named so tests can find the +# artists by label rather than by position in `ax.lines`. +UNIT_CIRCLE_LABEL = "unit circle" +UNIT_ROOT_LABEL = "unit root" + + +def plot_stability( + summary: "StabilitySummary", + figsize: tuple[float, float] = (11, 4.5), + bins: int = 40, +) -> Figure: + """Plot the spectral-radius posterior beside the eigenvalue scatter. + + Two views of the same question. The left panel is the posterior of the + companion-matrix spectral radius with the unit root marked, answering + *how much mass is explosive*. The right panel scatters the individual + companion roots on the complex plane against the unit circle, answering + *which* roots drive it: a single real root creeping outward is a + near-unit-root level series, whereas a conjugate pair approaching the + circle is an oscillatory mode with a long-lived cycle. + + The scatter uses the eigenvalue subset the summary retains (at most 200 + pooled draws), not every draw, so the panel is bounded in cost and in + ink regardless of posterior size. + + Args: + summary: `StabilitySummary` from a convergence report. + figsize: Figure size. + bins: Histogram bin count for the spectral-radius panel. + + Returns: + Matplotlib Figure with two axes. + """ + radius = np.asarray(summary.radius, dtype=float).reshape(-1) + eigenvalues = np.asarray(summary.eigenvalues).reshape(-1) + + fig, (ax_hist, ax_scatter) = plt.subplots(1, 2, figsize=figsize) + fig.suptitle("Dynamic stability") + + ax_hist.hist(radius, bins=bins, color="C0", alpha=0.8) + ax_hist.axvline(1.0, color="C3", linestyle="--", linewidth=1.2, label=UNIT_ROOT_LABEL) + ax_hist.set_xlabel("Spectral radius") + ax_hist.set_ylabel("Draws") + ax_hist.set_title(f"Explosive draws: {summary.p_explosive:.1%}") + ax_hist.legend(fontsize=8) + ax_hist.grid(alpha=0.3) + + theta = np.linspace(0.0, 2.0 * np.pi, 361) + ax_scatter.plot(np.cos(theta), np.sin(theta), color="C3", linewidth=1.2, label=UNIT_CIRCLE_LABEL) + ax_scatter.scatter(eigenvalues.real, eigenvalues.imag, s=6, alpha=0.3, color="C0", label="eigenvalues") + ax_scatter.axhline(0.0, color="0.7", linewidth=0.6, zorder=0) + ax_scatter.axvline(0.0, color="0.7", linewidth=0.6, zorder=0) + ax_scatter.set_aspect("equal", adjustable="datalim") + ax_scatter.set_xlabel("Re") + ax_scatter.set_ylabel("Im") + ax_scatter.set_title(f"Companion roots ({summary.eigenvalues.shape[0]} draws)") + ax_scatter.legend(fontsize=8) + + fig.tight_layout() + return fig diff --git a/src/impulso/protocols.py b/src/impulso/protocols.py index a1b25b2a..193ec838 100644 --- a/src/impulso/protocols.py +++ b/src/impulso/protocols.py @@ -221,6 +221,12 @@ class VolatilityProcess(Protocol): adapter's own variables into the report's covariance or volatility block. Adapters that omit it lose nothing beyond block attribution: unrecognised variables land in the report's `other` block. + + The `v{i}_` posterior-variable prefix is reserved: `assign_blocks` + routes any variable matching `^v\\d+_` to the volatility block on sight, + so an adapter that registers an unrelated variable under that spelling + (say `v2_gdp`) will see it misattributed. Name custom variables anything + else and claim them through `posterior_var_names()`. """ name: str diff --git a/tests/conftest.py b/tests/conftest.py index 9941d771..e3683b27 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -195,11 +195,13 @@ def permanent_transitory_2v(): # --------------- Diagnostics posterior factory --------------- -# PyMC and nutpie spell almost every sampler statistic differently; only -# `diverging` is common. The factory can emit either shape so tests can prove -# the report reads that one key and ignores the rest. -_PYMC_EXTRA_STATS = ("reached_max_treedepth", "tree_depth", "acceptance_rate", "lp", "energy", "step_size") -_NUTPIE_EXTRA_STATS = ("maxdepth_reached", "depth", "mean_tree_accept", "logp", "energy", "step_size", "tuning") +# PyMC and nutpie spell almost every sampler statistic differently; `diverging` +# and `energy` are the only common ones, and max-treedepth saturation needs a +# name map. The factory can emit either shape so tests can prove the report +# reads the right key per backend and ignores the rest. +_MAX_TREEDEPTH_NAME = {False: "reached_max_treedepth", True: "maxdepth_reached"} +_PYMC_NOISE_STATS = ("tree_depth", "acceptance_rate", "lp", "step_size") +_NUTPIE_NOISE_STATS = ("depth", "mean_tree_accept", "logp", "step_size", "tuning") def _coefficient_draws(rng, shape, explosive_frac, bad_coord): @@ -233,18 +235,48 @@ def _cholesky_draws(rng, sigma_sd, n_chains, n_draws, n_vars): return L -def _sampler_stats(rng, n_chains, n_draws, divergences, nutpie_shaped): - """A `sample_stats` group in either backend's shape, plus nutpie's warmup group.""" +def _flag_draws(n_chains, n_draws, count): + """Boolean per-transition flags, `count` of them set at fixed positions.""" flat = np.zeros(n_chains * n_draws, dtype=bool) - flat[:divergences] = True - stats = {"diverging": (("chain", "draw"), flat.reshape(n_chains, n_draws))} - for name in _NUTPIE_EXTRA_STATS if nutpie_shaped else _PYMC_EXTRA_STATS: + flat[:count] = True + return flat.reshape(n_chains, n_draws) + + +def _energy_draws(rng, n_chains, n_draws, energy_rho): + """AR(1) energy trace whose E-BFMI is approximately `2 * (1 - energy_rho)`. + + `energy_rho=0.3` gives a healthy trace (E-BFMI near 1.4); pushing it + toward 1 makes successive energies nearly identical, which is exactly the + slow energy exploration a low E-BFMI reports. + """ + energy = np.empty((n_chains, n_draws)) + energy[:, 0] = rng.standard_normal(n_chains) + innovation = np.sqrt(1.0 - energy_rho**2) + for t in range(1, n_draws): + energy[:, t] = energy_rho * energy[:, t - 1] + innovation * rng.standard_normal(n_chains) + return energy + + +def _sampler_stats(rng, n_chains, n_draws, divergences, treedepth_hits, energy_rho, nutpie_shaped): + """A `sample_stats` group in either backend's shape, plus nutpie's warmup group.""" + stats = { + "diverging": (("chain", "draw"), _flag_draws(n_chains, n_draws, divergences)), + "energy": (("chain", "draw"), _energy_draws(rng, n_chains, n_draws, energy_rho)), + _MAX_TREEDEPTH_NAME[nutpie_shaped]: ( + ("chain", "draw"), + _flag_draws(n_chains, n_draws, treedepth_hits), + ), + } + for name in _NUTPIE_NOISE_STATS if nutpie_shaped else _PYMC_NOISE_STATS: stats[name] = (("chain", "draw"), rng.standard_normal((n_chains, n_draws))) groups = {"sample_stats": xr.Dataset(stats)} if nutpie_shaped: - # Warmup divergences belong to adaptation and must not be counted. + # Warmup divergences, saturations and energy belong to adaptation and + # must not be counted. groups["warmup_sample_stats"] = xr.Dataset({ - "diverging": (("chain", "draw"), np.ones((n_chains, n_draws), dtype=bool)) + "diverging": (("chain", "draw"), np.ones((n_chains, n_draws), dtype=bool)), + "maxdepth_reached": (("chain", "draw"), np.ones((n_chains, n_draws), dtype=bool)), + "energy": (("chain", "draw"), np.zeros((n_chains, n_draws))), }) return groups @@ -266,11 +298,16 @@ def make_var_posterior(): indicator is neither autocorrelated nor chain-dependent. * `divergences` — divergent-transition count, placed at fixed positions. `None` omits the `sample_stats` group entirely. + * `treedepth_hits` — number of transitions flagged as having saturated + the maximum tree depth, under whichever backend name is in force. + * `energy_rho` — AR(1) coefficient of the energy trace. The default 0.3 + is healthy; values near 1 drive E-BFMI below its warning threshold. * `extra_vars` — mapping of posterior variable name to trailing shape. * `coords` — attach `var`/`coeff` coords (as `VAR.fit` does) or leave the dims bare (as `ConjugateVAR` does). * `nutpie_shaped` — emit nutpie's sampler-stat names plus a - `warmup_sample_stats` group full of divergences that must be ignored. + `warmup_sample_stats` group full of divergences and saturations that + must be ignored. """ def _make( @@ -282,6 +319,8 @@ def _make( bad_coord=None, explosive_frac=0.0, divergences=0, + treedepth_hits=0, + energy_rho=0.3, extra_vars=None, coords=True, nutpie_shaped=False, @@ -312,7 +351,7 @@ def _make( groups = {"posterior": xr.Dataset(data_vars, coords=posterior_coords)} if divergences is not None: - groups |= _sampler_stats(rng, n_chains, n_draws, divergences, nutpie_shaped) + groups |= _sampler_stats(rng, n_chains, n_draws, divergences, treedepth_hits, energy_rho, nutpie_shaped) return az.InferenceData(**groups) diff --git a/tests/test_diagnostics.py b/tests/test_diagnostics.py index f3f1065f..8f669982 100644 --- a/tests/test_diagnostics.py +++ b/tests/test_diagnostics.py @@ -187,6 +187,150 @@ def test_warmup_divergences_are_ignored(self, make_var_posterior): assert report.status == "passed" +class TestEnergy: + def test_healthy_energy_is_reported_without_a_message(self, make_var_posterior): + report = convergence_report(make_var_posterior(), n_lags=1) + assert report.ebfmi is not None + assert len(report.ebfmi) == 4 + assert all(value > 0.3 for value in report.ebfmi) + assert report.min_ebfmi == min(report.ebfmi) + assert "low_ebfmi" not in codes(report) + + def test_slow_energy_exploration_warns(self, make_var_posterior): + report = convergence_report(make_var_posterior(energy_rho=0.97), n_lags=1) + assert report.min_ebfmi is not None + assert report.min_ebfmi < 0.3 + message = next(m for m in report.messages if m.code == "low_ebfmi") + assert message.severity == "warning" + assert report.status == "warnings" + + def test_message_names_the_worst_chain(self, make_var_posterior): + report = convergence_report(make_var_posterior(energy_rho=0.97), n_lags=1) + worst = report.ebfmi.index(min(report.ebfmi)) + message = next(m for m in report.messages if m.code == "low_ebfmi") + assert f"chain {worst}" in message.message + + def test_threshold_is_respected(self, make_var_posterior): + idata = make_var_posterior(energy_rho=0.97) + lenient = ConvergenceThresholds(ebfmi_warn=0.0) + assert "low_ebfmi" not in codes(convergence_report(idata, n_lags=1, thresholds=lenient)) + + def test_comparison_is_strict_at_the_threshold(self, make_var_posterior): + idata = make_var_posterior(energy_rho=0.97) + observed = convergence_report(idata, n_lags=1).min_ebfmi + assert observed is not None + on_threshold = ConvergenceThresholds(ebfmi_warn=observed) + just_above = ConvergenceThresholds(ebfmi_warn=observed + 1e-12) + assert "low_ebfmi" not in codes(convergence_report(idata, n_lags=1, thresholds=on_threshold)) + assert "low_ebfmi" in codes(convergence_report(idata, n_lags=1, thresholds=just_above)) + + def test_low_ebfmi_never_fails_the_report(self, make_var_posterior): + # An efficiency pathology, not evidence the draws are wrong. + report = convergence_report(make_var_posterior(energy_rho=0.99), n_lags=1) + assert report.status == "warnings" + assert not any(m.severity == "failure" for m in report.messages) + + def test_nutpie_shaped_energy_matches_pymc_shaped(self, make_var_posterior): + nutpie = convergence_report(make_var_posterior(nutpie_shaped=True), n_lags=1) + pymc = convergence_report(make_var_posterior(), n_lags=1) + assert nutpie.ebfmi == pytest.approx(pymc.ebfmi) + + def test_warmup_energy_is_ignored(self, make_var_posterior): + # The warmup group carries a constant energy trace, whose BFMI is + # undefined; reading it would erase the post-warmup diagnosis. + idata = make_var_posterior(nutpie_shaped=True) + assert "warmup_sample_stats" in idata.groups() + report = convergence_report(idata, n_lags=1) + assert report.min_ebfmi is not None + assert report.min_ebfmi > 0.3 + + +class TestMaxTreedepth: + def test_no_saturation_is_reported_as_zero(self, make_var_posterior): + report = convergence_report(make_var_posterior(), n_lags=1) + assert report.treedepth_saturations == 0 + assert report.treedepth_saturation_rate == 0.0 + assert "treedepth_saturation" not in codes(report) + + def test_counts_are_exact_from_pymc_shaped_stats(self, make_var_posterior): + report = convergence_report(make_var_posterior(treedepth_hits=40), n_lags=1) + assert report.treedepth_saturations == 40 + assert report.treedepth_saturation_rate == pytest.approx(40 / 800) + + def test_counts_are_exact_from_nutpie_shaped_stats(self, make_var_posterior): + report = convergence_report(make_var_posterior(treedepth_hits=40, nutpie_shaped=True), n_lags=1) + assert report.treedepth_saturations == 40 + assert report.treedepth_saturation_rate == pytest.approx(40 / 800) + + def test_backends_agree(self, make_var_posterior): + nutpie = convergence_report(make_var_posterior(treedepth_hits=13, nutpie_shaped=True), n_lags=1) + pymc = convergence_report(make_var_posterior(treedepth_hits=13), n_lags=1) + assert nutpie.treedepth_saturations == pymc.treedepth_saturations == 13 + assert codes(nutpie) == codes(pymc) + + def test_warmup_saturations_are_ignored(self, make_var_posterior): + idata = make_var_posterior(treedepth_hits=0, nutpie_shaped=True) + assert "warmup_sample_stats" in idata.groups() + assert convergence_report(idata, n_lags=1).treedepth_saturations == 0 + + def test_rate_below_the_threshold_is_silent(self, make_var_posterior): + report = convergence_report(make_var_posterior(treedepth_hits=7), n_lags=1) + assert report.treedepth_saturation_rate is not None + assert report.treedepth_saturation_rate < 0.01 + assert "treedepth_saturation" not in codes(report) + + def test_rate_at_the_threshold_warns(self, make_var_posterior): + report = convergence_report(make_var_posterior(treedepth_hits=8), n_lags=1) + assert report.treedepth_saturation_rate == pytest.approx(0.01) + message = next(m for m in report.messages if m.code == "treedepth_saturation") + assert message.severity == "warning" + assert report.status == "warnings" + + def test_saturation_never_fails_the_report(self, make_var_posterior): + # Deep trees cost wall-clock time and autocorrelation, not correctness. + report = convergence_report(make_var_posterior(treedepth_hits=800), n_lags=1) + assert report.status == "warnings" + assert not any(m.severity == "failure" for m in report.messages) + + def test_message_quotes_the_low_rank_mass_matrix_remedy(self, make_var_posterior): + report = convergence_report(make_var_posterior(treedepth_hits=80), n_lags=1) + message = next(m for m in report.messages if m.code == "treedepth_saturation") + assert "max_treedepth" in message.message + assert "low_rank_modified_mass_matrix" in message.message + + def test_custom_threshold_flips_the_message(self, make_var_posterior): + idata = make_var_posterior(treedepth_hits=4) + strict = ConvergenceThresholds(treedepth_warn_rate=0.001) + assert "treedepth_saturation" not in codes(convergence_report(idata, n_lags=1)) + assert "treedepth_saturation" in codes(convergence_report(idata, n_lags=1, thresholds=strict)) + + +class TestMissingEfficiencyStats: + def test_stats_present_but_energy_and_treedepth_absent(self, make_var_posterior): + idata = make_var_posterior() + stripped = InferenceData( + posterior=idata.posterior, + sample_stats=idata.sample_stats.drop_vars(["energy", "reached_max_treedepth"]), + ) + report = convergence_report(stripped, n_lags=1) + assert report.sampler_stats_available is True + assert report.divergences == 0 + assert report.ebfmi is None + assert report.min_ebfmi is None + assert report.treedepth_saturations is None + assert report.treedepth_saturation_rate is None + assert report.status == "passed" + + def test_constant_energy_is_reported_as_absent_not_as_a_pathology(self, make_var_posterior): + idata = make_var_posterior() + stats = idata.sample_stats.copy() + stats["energy"] = (("chain", "draw"), np.zeros(stats["energy"].shape)) + report = convergence_report(InferenceData(posterior=idata.posterior, sample_stats=stats), n_lags=1) + assert report.ebfmi is None + assert "low_ebfmi" not in codes(report) + assert report.status == "passed" + + class TestMissingSamplerStats: @pytest.fixture def report(self, make_var_posterior): @@ -198,6 +342,12 @@ def test_availability_flag_and_none_counts(self, report): assert report.n_transitions is None assert report.divergence_rate is None + def test_efficiency_metrics_are_none_without_stats(self, report): + assert report.ebfmi is None + assert report.min_ebfmi is None + assert report.treedepth_saturations is None + assert report.treedepth_saturation_rate is None + def test_message_is_informational_only(self, report): message = next(m for m in report.messages if m.code == "sampler_stats_missing") assert message.severity == "info" @@ -344,6 +494,8 @@ def test_defaults_match_the_documented_values(self): assert (thresholds.rhat_warn, thresholds.rhat_fail) == (1.01, 1.05) assert (thresholds.ess_warn, thresholds.ess_fail) == (400.0, 100.0) assert thresholds.divergence_fail_rate == 0.01 + assert thresholds.ebfmi_warn == 0.3 + assert thresholds.treedepth_warn_rate == 0.01 assert thresholds.explosive_warn == 0.05 @@ -374,6 +526,13 @@ def test_summary_mentions_blocks_divergences_and_headlines(self, make_var_poster def test_summary_reports_unavailable_stats(self, make_var_posterior): text = convergence_report(make_var_posterior(divergences=None), n_lags=1).summary() assert "divergences: unavailable" in text + assert "E-BFMI" not in text + assert "max-treedepth" not in text + + def test_summary_carries_the_efficiency_headlines(self, make_var_posterior): + text = convergence_report(make_var_posterior(treedepth_hits=8), n_lags=1).summary() + assert "min E-BFMI:" in text + assert "max-treedepth hits: 8 (1.00%)" in text def test_repr_is_one_line(self, make_var_posterior): text = repr(convergence_report(make_var_posterior(), n_lags=1)) @@ -419,6 +578,29 @@ def test_thinning_larger_than_the_posterior_is_a_no_op(self, make_var_posterior) assert stability.thinned_from is None assert stability.radius.shape == (4, 200) + def test_eigenvalues_are_read_only(self, make_var_posterior): + eigenvalues = convergence_report(make_var_posterior(), n_lags=1).stability.eigenvalues + with pytest.raises(ValueError, match="read-only"): + eigenvalues[0, 0] = 99.0 + + def test_eigenvalues_are_pooled_capped_and_strided(self, make_var_posterior): + # 4 x 200 = 800 pooled draws, capped at 200: a stride of 4. + stability = convergence_report(make_var_posterior(), n_lags=1).stability + assert stability.eigenvalues.shape == (200, 2) + assert np.iscomplexobj(stability.eigenvalues) + np.testing.assert_allclose( + np.max(np.abs(stability.eigenvalues), axis=-1), + stability.radius.reshape(-1)[::4], + ) + + def test_small_posteriors_keep_every_eigenvalue(self, make_var_posterior): + stability = convergence_report(make_var_posterior(n_chains=1, n_draws=50), n_lags=1).stability + assert stability.eigenvalues.shape == (50, 2) + + def test_eigenvalue_trailing_axis_is_the_companion_dimension(self, make_var_posterior): + stability = convergence_report(make_var_posterior(n_vars=3, n_lags=2), n_lags=2).stability + assert stability.eigenvalues.shape[-1] == 6 + def test_rejects_non_positive_thinning(self, make_var_posterior): with pytest.raises(ValueError, match="stability_draws must be positive"): convergence_report(make_var_posterior(), n_lags=1, stability_draws=0) diff --git a/tests/test_plotting.py b/tests/test_plotting.py index 559afb5b..eb2cc879 100644 --- a/tests/test_plotting.py +++ b/tests/test_plotting.py @@ -3,15 +3,26 @@ import arviz as az import matplotlib import numpy as np +import pytest import xarray as xr from matplotlib.figure import Figure matplotlib.use("Agg") -from impulso.plotting import plot_fevd, plot_forecast, plot_historical_decomposition, plot_irf +from impulso.plotting import plot_fevd, plot_forecast, plot_historical_decomposition, plot_irf, plot_stability +from impulso.plotting._stability import UNIT_CIRCLE_LABEL, UNIT_ROOT_LABEL from impulso.results import FEVDResult, ForecastResult, HistoricalDecompositionResult, IRFResult +@pytest.fixture(autouse=True) +def _close_figures(): + """Plot functions return an open pyplot figure; drop it after each test.""" + yield + import matplotlib.pyplot as plt + + plt.close("all") + + def _make_forecast_result(n_vars=2, steps=8) -> ForecastResult: rng = np.random.default_rng(42) names = [f"y{i + 1}" for i in range(n_vars)] @@ -157,6 +168,67 @@ def test_deviation_line_matches_median_of_sum(self): np.testing.assert_allclose(lines[0].get_ydata(), expected.sel(response=resp).values, atol=1e-12) +class TestPlotStability: + @staticmethod + def _summary(make_var_posterior, **kwargs): + from impulso.diagnostics import convergence_report + + return convergence_report(make_var_posterior(**kwargs), n_lags=1).stability + + @staticmethod + def _line(ax, label): + lines = [line for line in ax.get_lines() if line.get_label() == label] + assert len(lines) == 1 + return lines[0] + + def test_returns_two_panels(self, make_var_posterior): + fig = plot_stability(self._summary(make_var_posterior)) + assert isinstance(fig, Figure) + assert len(fig.axes) == 2 + assert fig._suptitle.get_text() == "Dynamic stability" + + def test_histogram_bin_count_is_the_argument(self, make_var_posterior): + summary = self._summary(make_var_posterior) + assert len(plot_stability(summary).axes[0].patches) == 40 + assert len(plot_stability(summary, bins=15).axes[0].patches) == 15 + + def test_histogram_covers_every_radius_draw(self, make_var_posterior): + summary = self._summary(make_var_posterior) + ax = plot_stability(summary).axes[0] + assert sum(patch.get_height() for patch in ax.patches) == summary.radius.size + + def test_unit_root_line_sits_at_one(self, make_var_posterior): + ax = plot_stability(self._summary(make_var_posterior)).axes[0] + np.testing.assert_array_equal(self._line(ax, UNIT_ROOT_LABEL).get_xdata(), [1.0, 1.0]) + + def test_unit_circle_has_modulus_one(self, make_var_posterior): + ax = plot_stability(self._summary(make_var_posterior)).axes[1] + line = self._line(ax, UNIT_CIRCLE_LABEL) + np.testing.assert_allclose(np.hypot(line.get_xdata(), line.get_ydata()), 1.0) + + def test_scatter_plots_every_retained_eigenvalue(self, make_var_posterior): + summary = self._summary(make_var_posterior) + ax = plot_stability(summary).axes[1] + assert len(ax.collections) == 1 + offsets = ax.collections[0].get_offsets() + assert offsets.shape == (summary.eigenvalues.size, 2) + np.testing.assert_allclose(np.sort(offsets[:, 0]), np.sort(summary.eigenvalues.real.reshape(-1))) + + def test_scatter_axes_are_equally_scaled(self, make_var_posterior): + ax = plot_stability(self._summary(make_var_posterior)).axes[1] + assert ax.get_aspect() == 1.0 + + def test_explosive_mass_reaches_the_histogram_title(self, make_var_posterior): + fig = plot_stability(self._summary(make_var_posterior, explosive_frac=0.15)) + assert "15.0%" in fig.axes[0].get_title() + + def test_summary_plot_method_delegates(self, make_var_posterior): + # The house entry point: `report.stability.plot()`. + fig = self._summary(make_var_posterior).plot() + assert isinstance(fig, Figure) + assert len(fig.axes) == 2 + + def test_plot_volatility_returns_figure(synthetic_sv_idata): import pandas as pd from matplotlib.figure import Figure diff --git a/tests/test_protocols.py b/tests/test_protocols.py index 8253a7a2..126d9253 100644 --- a/tests/test_protocols.py +++ b/tests/test_protocols.py @@ -97,6 +97,13 @@ def test_volatility_process_isinstance_discriminates(self): assert isinstance(_ConformingVolatility(), VolatilityProcess) assert not isinstance(_Empty(), VolatilityProcess) + def test_docstring_reserves_the_sv_variable_prefix(self): + # Adapter authors read the protocol, not ADR-0008; the reserved + # prefix has to be visible where a custom adapter is written. + doc = VolatilityProcess.__doc__ or "" + assert "v{i}_" in doc + assert "posterior_var_names()" in doc + def test_pymc_volatility_process_adds_builder(self): # Sub-protocol extends the query surface with build_pymc_latent. assert is_protocol(PyMCVolatilityProcess)