Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions src/adaptiverg_qec/rhat.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,9 @@ class DiagnosticState(StrEnum):
Ketten um +10 verschoben wieder als konvergiert gemeldet. Die Pruefung ist als
``not (ptp > tol)`` formuliert, damit ein NaN (Ueberlauf des Medians bei ~1e308)
fail-closed als entartet zaehlt.
Absoluter Boden ``16 * np.spacing(scale)``: im Subnormal-Bereich unterlaeuft
``eps * scale`` auf 0, die Rundung betraegt dort aber 1 ulp (5e-324). Fuer normale
Zahlen ist ``spacing(scale)`` ~ ``eps * scale`` und aendert nichts.
"""


Expand Down Expand Up @@ -152,8 +155,10 @@ class RhatResult:
"""False bei konstanten Draws oder entarteter folded-Komponente.

Dann ist rhat KEIN definiertes R-hat im Vehtari-Sinn: bei konstanten Draws ein
Konventionswert, bei DEGENERATE_FOLDED das bulk-R-hat ohne messbare Skalen-
komponente.
Konventionswert; bei DEGENERATE_FOLDED weiterhin ``max(bulk_rhat, folded_rhat)``,
wobei ``folded_rhat`` nicht aussagekraeftig ist (Konventionswert 1.0 oder auf
Rundungsrauschen gerechnet) -- rhat ist dann oft 1.0 und NICHT das bulk-R-hat
(Delta-Review: 234 von 400 Faellen).
"""

@property
Expand Down Expand Up @@ -412,7 +417,8 @@ def split_rhat(draws: np.ndarray, *, expected_constant: bool = False) -> RhatRes
folded = np.abs(chains - median)
folded_scale = max(float(np.max(folded)), abs(median))
folded_degenerate = not constant_draws and not (
float(np.ptp(folded)) > _FOLDED_DEGENERACY_RTOL * folded_scale
float(np.ptp(folded))
> max(_FOLDED_DEGENERACY_RTOL * folded_scale, 16.0 * float(np.spacing(folded_scale)))
)
if folded_degenerate:
diagnostic_state = DiagnosticState.DEGENERATE_FOLDED
Expand Down
20 changes: 16 additions & 4 deletions tests/test_rhat.py
Original file line number Diff line number Diff line change
Expand Up @@ -305,7 +305,7 @@ def test_unbalanced_two_point_chains_keep_a_defined_folded_rhat() -> None:

@pytest.mark.parametrize(
("low", "high"),
[(10.1, 10.3), (1e6 + 0.1, 1e6 + 0.3), (-1e3 - 0.25, -1e3 + 0.5), (1e308, 1.5e308)],
[(10.1, 10.3), (1e6 + 0.1, 1e6 + 0.3), (-1000.3, -1000.1), (1e308, 1.5e308)],
)
def test_folded_degeneracy_is_detected_at_any_location(low: float, high: float) -> None:
"""Delta-Review zu f0d36e8: die Toleranz war nur relativ zu max|theta - median|.
Expand All @@ -323,10 +323,22 @@ def test_folded_degeneracy_is_detected_at_any_location(low: float, high: float)


def test_folded_tolerance_does_not_swallow_real_scale_differences() -> None:
"""Obere Grenze der Toleranz: ein echter Unterschied von 1e-13 bei Skala 1 ist
rund 30 eps und muss als messbar gelten (vorher war ein 1000x lockerer Wert blind)."""
values = np.concatenate([np.full(2000, -1.0), np.full(1000, 1.0), np.full(1000, 1.0 + 1e-13)])
"""Obere Grenze der Toleranz: ein echter Unterschied von 1e-14 bei Skala 1 ist
rund 45 eps (knapp das 3-fache der Toleranz von 16 eps) und muss als messbar gelten.
Gemessen: eine 4x lockerere Toleranz macht diesen Fall entartet (Test rot), eine
2x lockerere nicht -- die Toleranz ist damit auf Faktor 2-4 festgenagelt. (Mit dem
frueheren 1e-13 = 450 eps blieb sogar eine 16x lockerere Toleranz unentdeckt.)"""
values = np.concatenate([np.full(2000, -1.0), np.full(1000, 1.0), np.full(1000, 1.0 + 1e-14)])
chains = np.random.default_rng(13).permutation(values).reshape(4, 1000)
r = rhat.split_rhat(chains)
assert r.diagnostic_state is rhat.DiagnosticState.OK
assert r.rhat_defined


def test_folded_degeneracy_is_detected_for_subnormal_values() -> None:
"""Delta-Review 162e3ac: im Subnormal-Bereich unterlaeuft eps * scale auf 0, waehrend
die Rundung der Faltung 1 ulp (5e-324) betraegt -- ohne absoluten Boden blieb eine
ausbalancierte Kette OK/converged."""
r = rhat.split_rhat(_balanced_two_point_chains(1e-315, 2e-315 + 5e-324, seed=1))
assert r.diagnostic_state is rhat.DiagnosticState.DEGENERATE_FOLDED
assert not r.converged
Loading