def audit_backend_parity(
rstan_fit: Any,
cmdstanr_fit: Any,
variables: Sequence[str] | str | None = None,
mcse_multiplier: float = 3,
absolute_tolerance: float = 0,
relative_sd_tolerance: float = 0.10,
) -> BackendParityAudit:
"""Compare posterior summaries relative to Monte Carlo uncertainty."""
if mcse_multiplier < 0 or absolute_tolerance < 0 or relative_sd_tolerance < 0:
raise GP3BayesError("Backend parity tolerances must be non-negative.")
left = _draw_summary(rstan_fit, variables)
right = _draw_summary(cmdstanr_fit, variables)
lv = list(dict.fromkeys(left["variable"].astype(str)))
rv = list(dict.fromkeys(right["variable"].astype(str)))
common = [v for v in lv if v in set(rv)]
if not common:
raise GP3BayesError("No common posterior variables were available for comparison.")
left = left.set_index("variable").loc[common]
right = right.set_index("variable").loc[common]
combined = np.sqrt(
left["mcse_mean"].to_numpy(float) ** 2 + right["mcse_mean"].to_numpy(float) ** 2
)
mean_diff = left["mean"].to_numpy(float) - right["mean"].to_numpy(float)
allowed = np.maximum(float(absolute_tolerance), float(mcse_multiplier) * combined)
mean_ok = np.isfinite(combined) & (np.abs(mean_diff) <= allowed)
lsd = left["sd"].to_numpy(float)
rsd = right["sd"].to_numpy(float)
scale = np.maximum.reduce([np.abs(lsd), np.abs(rsd), np.full_like(lsd, np.finfo(float).eps)])
rel_sd = np.abs(lsd - rsd) / scale
sd_ok = np.isfinite(rel_sd) & (rel_sd <= float(relative_sd_tolerance))
table = pd.DataFrame(
{
"variable": common,
"left_mean": left["mean"].to_numpy(float),
"right_mean": right["mean"].to_numpy(float),
"mean_difference": mean_diff,
"combined_mcse": combined,
"allowed_mean_difference": allowed,
"mean_within_mcse": mean_ok,
"left_sd": lsd,
"right_sd": rsd,
"relative_sd_difference": rel_sd,
"sd_within_tolerance": sd_ok,
}
)
table["status"] = np.where(mean_ok & sd_ok, "pass", "review")
missing_left = tuple(v for v in rv if v not in set(lv))
missing_right = tuple(v for v in lv if v not in set(rv))
status = (
"pass"
if (table["status"] == "pass").all() and not missing_left and not missing_right
else "review"
)
return BackendParityAudit(
"0.2",
status,
table,
missing_left,
missing_right,
{
"mcse_multiplier": float(mcse_multiplier),
"absolute_tolerance": float(absolute_tolerance),
"relative_sd_tolerance": float(relative_sd_tolerance),
},
)