"""Numerical checks for the Q2 state posterior, joint sampling, and decoder.""" from __future__ import annotations import unittest from types import SimpleNamespace import numpy as np import torch from scipy.special import betainc as scipy_betainc from crg import CRG, ReliabilityGRU, StructuredGaussianImputer from train import ( _decode_mixture, _calibrated_mixture_moments, _group_ids, _missing_rate_summary, _predictive_intervals, _trajectory_variance_components, continuous_mask, controlled_group_bootstrap, gate_diagnostic_rows, regularized_beta, smooth_group_risk, validate_attachment3_predictions, ) class StructuredGaussianTests(unittest.TestCase): def test_filter_nll_matches_dense_marginal_gaussian(self) -> None: torch.manual_seed(73) model = StructuredGaussianImputer((2, 2, 2)).double() xs = [torch.randn(1, 2, 2, dtype=torch.float64) for _ in range(3)] observed = torch.ones(1, 2, 3, dtype=torch.bool) got = model.observed_nll(xs, observed)[0] with torch.no_grad(): transition = model._transition() p0, q = model._covariances() emissions = model.emissions() emission = torch.cat(emissions, dim=0) noise = torch.block_diag(*[torch.diag(torch.nn.functional.softplus(raw) + 1e-4) for raw in model.r_raw]) offset = torch.cat(list(model.biases)) state_mean = torch.cat((model.mu0, transition @ model.mu0)) p01 = p0 @ transition.T p11 = transition @ p0 @ transition.T + q state_cov = torch.cat((torch.cat((p0, p01), dim=1), torch.cat((p01.T, p11), dim=1)), dim=0) observation_map = torch.block_diag(emission, emission) observation_cov = observation_map @ state_cov @ observation_map.T + torch.block_diag(noise, noise) observation_mean = torch.cat((offset + emission @ model.mu0, offset + emission @ (transition @ model.mu0))) values = torch.cat((torch.cat([xs[m][0, 0] for m in range(3)]), torch.cat([xs[m][0, 1] for m in range(3)]))) residual = values - observation_mean expected = 0.5 * ( residual @ torch.linalg.solve(observation_cov, residual) + torch.linalg.slogdet(observation_cov).logabsdet + len(values) * np.log(2.0 * np.pi) ) torch.testing.assert_close(got, expected, rtol=2e-4, atol=2e-4) def test_joint_trajectory_draws_retain_temporal_dependence(self) -> None: torch.manual_seed(19) model = StructuredGaussianImputer((2, 2, 2)) with torch.no_grad(): for emission in model.emission_raw: emission.zero_() model.emission_raw[1][0, 0] = 1.0 xs = [torch.zeros(1, 2, 2) for _ in range(3)] observed = torch.zeros(1, 2, 3, dtype=torch.bool) draws, _ = model.complete(xs, observed, 1600, joint_draws=True) temporal_correlation = float(np.corrcoef(draws[1][:, 0, 0, 0].cpu(), draws[1][:, 0, 1, 0].cpu())[0, 1]) self.assertGreater(temporal_correlation, 0.15) class LossAndMaskTests(unittest.TestCase): def test_beta_cdf_matches_scipy(self) -> None: a = torch.tensor([0.7, 2.0, 5.0]) b = torch.tensor([1.3, 3.0, 2.5]) x = torch.tensor([0.2, 0.8, 0.55]) actual = regularized_beta(x, a, b).detach().cpu().numpy() expected = scipy_betainc(a.numpy(), b.numpy(), x.numpy()) np.testing.assert_allclose(actual, expected, rtol=2e-5, atol=2e-6) def test_mask_is_contiguous_and_preserves_each_selected_source(self) -> None: original = np.ones((50, 3), dtype=bool) for mode in ("single", "sync", "partial", "async"): masked = continuous_mask(original, 0.5, mode, np.random.default_rng(101)) hidden = original & ~masked for modality in range(3): positions = np.flatnonzero(hidden[:, modality]) if len(positions): self.assertEqual(int(positions[-1] - positions[0] + 1), len(positions)) self.assertGreaterEqual(int(masked[:, modality].sum()), 10) def test_point_mask_keeps_rate_but_breaks_contiguous_span(self) -> None: original = np.ones((50, 3), dtype=bool) masked = continuous_mask(original, 0.3, "single", np.random.default_rng(887), modalities=(1,), kind="point") hidden = np.flatnonzero(original[:, 1] & ~masked[:, 1]) self.assertEqual(len(hidden), 15) self.assertGreaterEqual(int(masked[:, 1].sum()), 10) runs = np.split(hidden, np.flatnonzero(np.diff(hidden) > 1) + 1) self.assertGreater(len([run for run in runs if len(run)]), 1) def test_position_and_gap_structure_controls_hold_total_missing_fixed(self) -> None: original = np.ones((50, 3), dtype=bool) counts = [] for location in ("start", "middle", "end"): masked = continuous_mask( original, 0.3, "single", np.random.default_rng(22), modalities=(0,), location=location, ) hidden = np.flatnonzero(original[:, 0] & ~masked[:, 0]) counts.append(len(hidden)) if location == "start": self.assertEqual(int(hidden[0]), 0) elif location == "end": self.assertEqual(int(hidden[-1]), 49) else: self.assertLessEqual(abs(float(hidden.mean()) - 24.5), 1.0) self.assertEqual(counts, [15, 15, 15]) long = continuous_mask( original, 0.3, "single", np.random.default_rng(22), modalities=(0,), span_structure="long", ) short = continuous_mask( original, 0.3, "single", np.random.default_rng(22), modalities=(0,), span_structure="multi_short", ) long_hidden = np.flatnonzero(original[:, 0] & ~long[:, 0]) short_hidden = np.flatnonzero(original[:, 0] & ~short[:, 0]) self.assertEqual(len(long_hidden), len(short_hidden)) short_runs = np.split(short_hidden, np.flatnonzero(np.diff(short_hidden) > 1) + 1) self.assertGreaterEqual(len([run for run in short_runs if len(run)]), 2) def test_group_id_uses_any_newly_hidden_source(self) -> None: original = np.ones((2, 50, 3), dtype=bool) current = original.copy() current[0, 10:20, 1] = False current[1, 15:25, 2] = False groups = _group_ids(original, current) self.assertNotEqual(int(groups[0]), int(groups[1])) def test_missing_rates_follow_equal_modality_pdf_denominators(self) -> None: original = np.asarray([ [1, 1, 0], [1, 1, 0], [1, 0, 0], [1, 0, 0], ], dtype=bool) current = original.copy() current[0, 0] = False rates = _missing_rate_summary(original, current) np.testing.assert_allclose(rates["natural_by_modality"], [0.0, 0.5, 1.0]) self.assertAlmostEqual(rates["natural_global"], 0.5) np.testing.assert_allclose(rates["final_by_modality"], [0.25, 0.5, 1.0]) self.assertAlmostEqual(rates["final_global"], 7.0 / 12.0) np.testing.assert_allclose(rates["additional_by_modality"][:2], [0.25, 0.0]) self.assertTrue(np.isnan(rates["additional_by_modality"][2])) def test_smooth_group_risk_matches_prior_weighted_formula(self) -> None: losses = torch.tensor([1.0, 2.0, 4.0], dtype=torch.float64) group_ids = np.asarray([0, 0, 1]) lambda_group, tau = 0.2, 0.5 group_losses = torch.tensor([1.5, 4.0], dtype=torch.float64) priors = torch.tensor([2 / 3, 1 / 3], dtype=torch.float64) expected = ((1 - lambda_group) * (priors * group_losses).sum() + lambda_group * tau * torch.logsumexp(priors.log() + group_losses / tau, dim=0)) actual = smooth_group_risk(losses, group_ids, lambda_group, tau) torch.testing.assert_close(actual, expected) def test_controlled_group_bootstrap_is_paired_and_reports_aurc(self) -> None: split = SimpleNamespace( n=4, class_y=np.asarray([0, 0, 1, 2]), regression_y=np.asarray([-1.0, -0.5, 0.0, 1.0]), groups=np.asarray(["v1", "v1", "v2", "v3"]), mask=np.ones((4, 50, 3), dtype=bool), ) scenarios = ["0.0/none"] + [f"{rate:.1f}/{mode}" for mode in ("single", "sync", "partial", "async") for rate in (0.1, 0.3, 0.5, 0.7)] scenario_masks = {} for scenario in scenarios: mask = split.mask.copy() rate_name, pattern = scenario.split("/", 1) rate = float(rate_name) if rate > 0: modality = {"single": 0, "sync": 0, "partial": 1, "async": 2}[pattern] count = int(round(rate * 50)) mask[:, :count, modality] = False scenario_masks[scenario] = mask predictions = {} for model, shift in (("C0", 0.0), ("C1", 0.1)): predictions[model] = {} for index, scenario in enumerate(scenarios): predictions[model][scenario] = { "predicted_class": np.asarray([0, 1, 1, 2]), "predicted_score": split.regression_y + shift + index * 0.01, } rows = controlled_group_bootstrap(split, predictions, scenario_masks, repeats=20, seed=29) self.assertTrue(any(row["metric"] == "AURC_MAE" and row["model"] == "C1" for row in rows)) paired = next(row for row in rows if row["model"] == "C1" and row["scenario"] == "0.3/single" and row["metric"] == "mae") self.assertAlmostEqual(paired["delta_estimate"], 0.1) self.assertAlmostEqual(paired["delta_to_natural_mae"], 0.02) self.assertEqual(paired["replicates"], 20) def test_attachment3_submission_invariants(self) -> None: rows = [ {"case_id": "case-a", "predicted_class": 0, "predicted_sentiment": -0.2, "p_negative": 0.5, "p_neutral": 0.3, "p_positive": 0.2, "interval_90_lower": -1.0, "interval_90_upper": 0.5}, {"case_id": "case-b", "predicted_class": 1, "predicted_sentiment": 0.0, "p_negative": 0.2, "p_neutral": 0.6, "p_positive": 0.2, "interval_90_lower": -0.5, "interval_90_upper": 0.5}, ] validate_attachment3_predictions(["case-a", "case-b"], rows) rows[1]["predicted_sentiment"] = 1e-9 with self.assertRaisesRegex(ValueError, "polarity mismatch"): validate_attachment3_predictions(["case-a", "case-b"], rows) def test_gate_diagnostic_rows_keep_sample_position_and_modality(self) -> None: split = SimpleNamespace(ids=["v1$_$c1"], groups=np.asarray(["v1"]), mask=np.ones((1, 2, 3), dtype=bool)) scalar = np.zeros((1, 2, 3), dtype=np.float32) predictions = { "fusion_weights": np.full((1, 2, 3), 0.2, dtype=np.float32), "null_weights": np.full((1, 2), 0.4, dtype=np.float32), "time_pool_weights": np.full((1, 2), 0.5, dtype=np.float32), "reliability": np.ones((1, 2, 3), dtype=np.float32), "imputation_uncertainty": scalar, "gap": scalar, "span": scalar, "distance_before": scalar, "distance_after": scalar, } rows = gate_diagnostic_rows(split, predictions) self.assertEqual(len(rows), 6) self.assertEqual(rows[0]["sample_id"], "v1$_$c1") self.assertEqual(rows[-1]["modality"], "vision") def test_decoder_uses_neutral_priority_and_exact_zero(self) -> None: probabilities = np.asarray([[[1 / 3, 1 / 3, 1 / 3]], [[1 / 3, 1 / 3, 1 / 3]]], dtype=np.float32) beta = np.full((2, 1, 2, 2), 2.0, dtype=np.float32) _, classes, scores = _decode_mixture(probabilities, beta) self.assertEqual(int(classes[0]), 1) self.assertEqual(float(scores[0]), 0.0) def test_calibrated_signed_mixture_interval_and_variance_components(self) -> None: probabilities = np.asarray( [[[0.25, 0.5, 0.25]], [[0.4, 0.2, 0.4]]], dtype=np.float64, ) beta = np.full((2, 1, 2, 2), 2.0, dtype=np.float64) low, high = _predictive_intervals(probabilities, beta, temperature=1.5) self.assertLess(float(low[0]), 0.0) self.assertGreater(float(high[0]), 0.0) self.assertLess(float(low[0]), float(high[0])) total, within, between = _trajectory_variance_components(probabilities, beta) np.testing.assert_allclose(total, within + between, rtol=1e-6, atol=1e-7) mean_cold, variance_cold = _calibrated_mixture_moments(probabilities, beta, temperature=0.5) mean_warm, variance_warm = _calibrated_mixture_moments(probabilities, beta, temperature=2.0) self.assertTrue(np.isfinite(mean_cold).all() and np.isfinite(variance_cold).all()) self.assertGreater(abs(float(variance_cold[0] - variance_warm[0])), 1e-5) class RecurrentAndVariantTests(unittest.TestCase): def test_gru_reset_gate_is_applied_before_candidate_recurrent_map(self) -> None: model = ReliabilityGRU(input_dim=1, hidden=1) with torch.no_grad(): model.x_proj.weight.zero_() model.x_proj.bias.copy_(torch.tensor([10.0, 0.0, 1.0])) model.h_proj.weight.zero_() model.candidate_h.weight.fill_(2.0) x = torch.zeros(1, 2, 1) rho = torch.ones(1, 2) distance = torch.zeros(1, 2) actual = model._one_direction(x, rho, distance, reverse=False, reliability_update=False) z = torch.sigmoid(torch.tensor(10.0)) first = z * torch.tanh(torch.tensor(1.0)) reset = torch.sigmoid(torch.tensor(0.0)) candidate = torch.tanh(torch.tensor(1.0) + 2.0 * reset * first) expected = (1.0 - z) * first + z * candidate torch.testing.assert_close(actual[0, 1, 0], expected) def test_all_ablation_architectures_forward_and_backward(self) -> None: torch.manual_seed(9) options = { "C1": dict(use_imputer=False, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False), "C2": dict(use_imputer=True, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False), "C3": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=False, reliability_update=False, use_low_rank=False), "C4": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=False, use_low_rank=False), "C5": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=True, use_low_rank=False), "C6": dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=True, use_low_rank=True), } xs = [torch.randn(1, 4, width) for width in (3, 2, 2)] observed = torch.ones(1, 4, 3, dtype=torch.bool) observed[:, 1:3, 1] = False for name, flags in options.items(): with self.subTest(model=name): model = CRG(input_dims=(3, 2, 2), **flags) output = model(xs, observed, paths=2, joint_draws=flags["use_joint_draws"]) loss = output["class_logits"].sum() + output["beta_params"].sum() loss.backward() expected_paths = 2 if flags["use_imputer"] else 1 self.assertEqual(tuple(output["class_probs"].shape), (1, 3)) self.assertEqual(tuple(output["fusion_weights_by_path"].shape), (expected_paths, 1, 4, 3)) self.assertEqual(tuple(output["null_weights_by_path"].shape), (expected_paths, 1, 4)) self.assertEqual(tuple(output["time_pool_weights_by_path"].shape), (expected_paths, 1, 4)) torch.testing.assert_close( output["fusion_weights_by_path"].sum(dim=-1) + output["null_weights_by_path"], torch.ones((expected_paths, 1, 4)), ) torch.testing.assert_close( output["time_pool_weights_by_path"].sum(dim=-1), torch.ones((expected_paths, 1)), ) if __name__ == "__main__": unittest.main()