Files
modeling_zhaocui/math/Q2/test_model.py
T

317 lines
16 KiB
Python

"""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()