Files
modeling_zhaocui/math/Q2/train.py
T

2062 lines
111 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Train/evaluate the Q2 model under the V2 official split and ablation design."""
from __future__ import annotations
import argparse
import copy
import csv
from contextlib import contextmanager
import hashlib
import json
import math
import random
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from scipy.optimize import minimize_scalar
from scipy.special import betainc, betaincinv
from scipy.stats import pearsonr
from sklearn.linear_model import LogisticRegression, Ridge
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, mean_squared_error, recall_score
from sklearn.model_selection import GroupShuffleSplit
from torch.nn import functional as F
from transformers import AutoModel
from crg import CRG, INPUT_DIMS, MODALITIES, StructuredGaussianImputer
from data import (
ATTACHMENT3_ALIGNED,
ROOT,
SplitData,
fit_preprocessor,
load_attachment3_case,
load_official_splits,
transform_split,
)
Q2_DIR = Path(__file__).resolve().parent
RESULTS = Q2_DIR / "results"
TEXT_MODEL_ID = "google-bert/bert-base-uncased"
SEED = 20260924
MASK_RATES = (0.0, 0.1, 0.3, 0.5, 0.7)
MASK_MODES = ("single", "sync", "partial", "async")
DELTA_U = 0.01
DISTILL_TEMPERATURE = 2.0
LAMBDA_Y = 1.0
LAMBDA_DISTILL = 0.1
LAMBDA_RECON = 0.05
LAMBDA_GROUP = 0.1
GROUP_TEMPERATURE = 0.1
GROUP_RISK_CANDIDATES = ((0.05, 0.1), (0.1, 0.05), (0.1, 0.1), (0.1, 0.2), (0.2, 0.1))
LAMBDA_EMISSION = 1e-4
LAMBDA_TRANSITION = 1e-4
DEFAULT_RELIABILITY = (0.5, 0.05, 0.05, 0.05)
RELIABILITY_CANDIDATES = (
(0.5, 0.0, 0.0, 0.0),
(0.5, 0.05, 0.05, 0.05),
(0.5, 0.1, 0.0, 0.0),
(0.3, 0.05, 0.05, 0.05),
(0.7, 0.05, 0.05, 0.05),
)
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
@contextmanager
def fixed_torch_seed(seed: int, device: torch.device):
devices = [device.index if device.index is not None else torch.cuda.current_device()] if device.type == "cuda" else []
with torch.random.fork_rng(devices=devices):
torch.manual_seed(int(seed))
if device.type == "cuda":
torch.cuda.manual_seed_all(int(seed))
yield
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as stream:
for block in iter(lambda: stream.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def label_resolution_from_train(labels: np.ndarray) -> float:
"""Use half the smallest positive nonzero magnitude spacing in fit labels."""
magnitudes = np.unique(np.round(np.abs(np.asarray(labels, dtype=np.float64)) / 3.0, 6))
spacing = np.diff(magnitudes)
spacing = spacing[spacing > 1e-5]
return float(np.clip(0.5 * spacing.min(), 1e-4, 0.1)) if len(spacing) else 0.01
def fit_magnitude_priors(labels: np.ndarray) -> np.ndarray:
"""Method-of-moments sign-specific Beta priors from train labels only."""
result = np.zeros((2, 2), dtype=np.float64)
for slot, selected in enumerate((np.asarray(labels) < 0, np.asarray(labels) > 0)):
values = np.abs(np.asarray(labels, dtype=np.float64)[selected]) / 3.0
values = np.clip(values, 1e-4, 1.0 - 1e-4)
if len(values) < 2:
mean, concentration = 0.5, 4.0
else:
mean, variance = float(values.mean()), float(values.var(ddof=1))
concentration = mean * (1.0 - mean) / max(variance, 1e-5) - 1.0
concentration = float(np.clip(concentration, 2.0, 100.0))
result[slot] = (max(1e-3, mean * concentration), max(1e-3, (1.0 - mean) * concentration))
return result.astype(np.float32)
def assert_group_disjoint(splits: dict[str, SplitData]) -> dict[str, int]:
overlap: dict[str, int] = {}
for left, right in (("train", "valid"), ("train", "test"), ("valid", "test")):
shared = set(splits[left].groups) & set(splits[right].groups)
overlap[f"{left}_{right}"] = len(shared)
if shared:
raise ValueError(f"source-video leakage across {left}/{right}: {sorted(shared)[:5]}")
return overlap
def select_rows(split: SplitData, indices: np.ndarray, name: str) -> SplitData:
idx = np.asarray(indices, dtype=np.int64)
return SplitData(
name=name,
x={m: split.x[m][idx].copy() for m in MODALITIES},
mask=split.mask[idx].copy(),
class_y=split.class_y[idx].copy() if split.class_y is not None else None,
regression_y=split.regression_y[idx].copy() if split.regression_y is not None else None,
ids=[split.ids[int(i)] for i in idx],
groups=split.groups[idx].copy(),
)
def split_calibration(train: SplitData, seed: int, fraction: float = 0.1) -> tuple[SplitData, SplitData]:
splitter = GroupShuffleSplit(n_splits=1, test_size=fraction, random_state=seed)
fit_idx, cal_idx = next(splitter.split(np.zeros(train.n), train.class_y, train.groups))
fit, cal = select_rows(train, fit_idx, "fit"), select_rows(train, cal_idx, "calibration")
if set(fit.groups) & set(cal.groups):
raise AssertionError("internal fit/calibration source videos overlap")
return fit, cal
def _best_interval(
visible: np.ndarray,
wanted: int,
cap: int,
location: str,
rng: np.random.Generator,
) -> tuple[int, int] | None:
steps = len(visible)
candidates: list[tuple[int, int, int, int]] = []
for left in range(steps):
hits = 0
for right in range(left, steps):
hits += int(visible[right])
count = min(hits, cap)
if count:
candidates.append((abs(count - wanted), right - left + 1, left, right))
if not candidates:
return None
best = min((error, span) for error, span, _, _ in candidates)
tied = [(left, right) for error, span, left, right in candidates if (error, span) == best]
if location == "start":
return min(tied, key=lambda pair: (pair[0], pair[1]))
if location == "end":
return max(tied, key=lambda pair: (pair[1], pair[0]))
if location == "middle":
center = (steps - 1) / 2
return min(tied, key=lambda pair: (abs((pair[0] + pair[1]) / 2 - center), pair[0]))
if location != "random":
raise ValueError(f"unknown interval location: {location}")
return tied[int(rng.integers(0, len(tied)))]
def _spread_short_spans(visible: np.ndarray, wanted: int, cap: int) -> np.ndarray:
positions = np.flatnonzero(visible)
count = min(int(wanted), int(cap), len(positions))
chosen = np.zeros(len(visible), dtype=bool)
if count <= 0:
return chosen
n_spans = min(3, count)
chunks = np.array_split(positions, n_spans)
allocations = [count // n_spans + int(i < count % n_spans) for i in range(n_spans)]
for chunk, amount in zip(chunks, allocations):
if amount <= 0 or len(chunk) == 0:
continue
amount = min(amount, len(chunk))
start = max(0, (len(chunk) - amount) // 2)
chosen[chunk[start:start + amount]] = True
return chosen
def continuous_mask(
original: np.ndarray,
rate: float,
mode: str,
rng: np.random.Generator,
*,
modalities: tuple[int, ...] | None = None,
location: str = "random",
span_structure: str = "long",
kind: str = "continuous",
) -> np.ndarray:
"""Hide contiguous feature rows while preserving at least 20% per selected source.
``sync``, ``partial`` and ``async`` use a shared, shifted-overlap, or
staggered span layout. Evaluation controls may pin the affected modalities,
gap location, and one-long versus several-short structure.
"""
observed = np.asarray(original, dtype=bool)
result = observed.copy()
if rate <= 0 or mode == "none":
return result
steps, modality_count = observed.shape
present = [m for m in range(modality_count) if observed[:, m].any()]
if not present:
return result
if modalities is not None:
selected = [int(m) for m in modalities if int(m) in present]
if not selected:
return result
elif mode == "single":
selected = [int(rng.choice(present))]
elif mode in {"sync", "partial", "async"}:
if len(present) == 1:
selected = present
else:
count = int(rng.integers(2, min(3, len(present)) + 1))
selected = sorted(int(v) for v in rng.choice(present, size=count, replace=False))
else:
raise ValueError(f"unknown mask mode: {mode}")
def max_hide(modality: int) -> int:
count = int(observed[:, modality].sum())
keep = max(1, int(math.ceil(0.2 * count)))
return max(0, count - keep)
target = {m: min(max_hide(m), int(round(rate * int(observed[:, m].sum())))) for m in selected}
if kind == "point":
for modality in selected:
candidates = np.flatnonzero(observed[:, modality])
count = target[modality]
if count > 0:
hidden = rng.choice(candidates, size=count, replace=False)
result[hidden, modality] = False
return result
if kind != "continuous":
raise ValueError(f"unknown mask kind: {kind}")
if mode == "sync":
span = max(1, int(round(rate * steps)))
if location == "start":
left = 0
elif location == "end":
left = steps - span
elif location == "middle":
left = (steps - span) // 2
elif location == "random":
left = int(rng.integers(0, max(1, steps - span + 1)))
else:
raise ValueError(f"unknown interval location: {location}")
right = min(steps - 1, left + span - 1)
for m in selected:
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
amount = min(len(candidates), max_hide(m), target[m])
if amount:
offset = 0 if location != "end" else len(candidates) - amount
candidates = candidates[max(0, offset):max(0, offset) + amount]
result[candidates, m] = False
else:
common_span = max(1, int(round(rate * steps)))
for rank, m in enumerate(selected):
wanted = target[m]
if wanted <= 0:
continue
cap = max_hide(m)
if span_structure == "multi_short":
hide = _spread_short_spans(observed[:, m], wanted, cap)
elif span_structure != "long":
raise ValueError(f"unknown span structure: {span_structure}")
elif mode == "single" and location != "random":
span = max(1, int(round(rate * steps)))
if location == "start":
left = 0
elif location == "end":
left = steps - span
elif location == "middle":
left = (steps - span) // 2
else:
raise ValueError(f"unknown interval location: {location}")
right = min(steps - 1, left + span - 1)
hide = np.zeros(steps, dtype=bool)
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
hide[candidates[:cap]] = True
elif mode in {"partial", "async"}:
if mode == "partial":
base_left = int(rng.integers(0, max(1, steps - common_span + 1))) if location == "random" else (
0 if location == "start" else steps - common_span if location == "end" else (steps - common_span) // 2
)
offset = int(round(rank * common_span * 0.5))
else:
base_left = 0 if location == "random" else (
0 if location == "start" else steps - common_span if location == "end" else (steps - common_span) // 2
)
available = max(1, steps - common_span + 1)
offsets = np.rint(np.linspace(0, max(0, available - 1), len(selected))).astype(int)
if location == "random":
rng.shuffle(offsets)
offset = int(offsets[rank])
left = min(max(0, base_left + offset), max(0, steps - common_span))
right = min(steps - 1, left + common_span - 1)
hide = np.zeros(steps, dtype=bool)
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
amount = min(len(candidates), wanted, cap)
if amount:
hide[candidates[:amount]] = True
else:
interval = _best_interval(observed[:, m], wanted, cap, location, rng)
hide = np.zeros(steps, dtype=bool)
if interval is not None:
left, right = interval
candidates = np.flatnonzero(observed[left:right + 1, m]) + left
amount = min(len(candidates), wanted, cap)
if amount:
offset = 0 if location != "end" else len(candidates) - amount
hide[candidates[max(0, offset):max(0, offset) + amount]] = True
result[hide, m] = False
return result
def make_scenarios(split: SplitData, seed: int) -> dict[str, np.ndarray]:
scenarios = {"0.0/none": split.mask.copy()}
for rate in MASK_RATES[1:]:
for mode in MASK_MODES:
key = f"{rate:.1f}/{mode}"
masks = []
for sample_id, original in zip(split.ids, split.mask):
sample_seed = int.from_bytes(hashlib.sha256(f"{seed}:{sample_id}:{key}".encode()).digest()[:8], "little")
masks.append(continuous_mask(original, rate, mode, np.random.default_rng(sample_seed)))
scenarios[key] = np.stack(masks)
modality_sets = (
((0,), "T"), ((1,), "A"), ((2,), "V"),
((0, 1), "TA"), ((0, 2), "TV"), ((1, 2), "AV"), ((0, 1, 2), "TAV"),
)
for selected, label in modality_sets:
key = f"0.3/modality_{label}"
scenarios[key] = np.stack([
continuous_mask(
original, 0.3, "sync", np.random.default_rng(_scenario_seed(seed, sample_id, key)),
modalities=selected,
)
for sample_id, original in zip(split.ids, split.mask)
])
for modality_index, label in enumerate(("T", "A", "V")):
for location in ("start", "middle", "end"):
key = f"0.3/location_{location}_{label}"
scenarios[key] = np.stack([
continuous_mask(
original, 0.3, "single", np.random.default_rng(_scenario_seed(seed, sample_id, key)),
modalities=(modality_index,), location=location,
)
for sample_id, original in zip(split.ids, split.mask)
])
for structure in ("long", "multi_short"):
key = f"0.3/span_{structure}_{label}"
scenarios[key] = np.stack([
continuous_mask(
original, 0.3, "single", np.random.default_rng(_scenario_seed(seed, sample_id, key)),
modalities=(modality_index,), span_structure=structure,
)
for sample_id, original in zip(split.ids, split.mask)
])
for mode in ("sync", "partial", "async"):
key = f"0.3/synchrony_{mode}"
scenarios[key] = np.stack([
continuous_mask(
original, 0.3, mode, np.random.default_rng(_scenario_seed(seed, sample_id, key)),
modalities=(0, 1, 2),
)
for sample_id, original in zip(split.ids, split.mask)
])
return scenarios
def _scenario_seed(seed: int, sample_id: str, key: str) -> int:
return int.from_bytes(hashlib.sha256(f"{seed}:{sample_id}:{key}".encode()).digest()[:8], "little")
def make_reliability_scenarios(split: SplitData, seed: int) -> dict[str, np.ndarray]:
scenarios = {"0.0/natural": split.mask.copy()}
for rate, mode in ((0.3, "single"), (0.3, "sync"), (0.5, "async")):
key = f"{rate:.1f}/{mode}"
scenarios[key] = np.stack([
continuous_mask(original, rate, mode, np.random.default_rng(_scenario_seed(seed, sample_id, key)))
for sample_id, original in zip(split.ids, split.mask)
])
return scenarios
def to_device_batch(
arrays: dict[str, np.ndarray], masks: np.ndarray, indices: np.ndarray, device: torch.device,
) -> tuple[list[torch.Tensor], torch.Tensor]:
idx = np.asarray(indices, dtype=np.int64)
xs = [torch.from_numpy(arrays[m][idx]).to(device=device, dtype=torch.float32) for m in MODALITIES]
observed = torch.from_numpy(masks[idx].astype(bool)).to(device=device)
return xs, observed
def _betacf(a: torch.Tensor, b: torch.Tensor, x: torch.Tensor, iterations: int = 64) -> torch.Tensor:
"""Differentiable continued fraction for the regularized incomplete beta."""
tiny = 1e-12
qab, qap, qam = a + b, a + 1.0, a - 1.0
c = torch.ones_like(x)
d = 1.0 - qab * x / qap
d = 1.0 / torch.where(d.abs() < tiny, torch.full_like(d, tiny), d)
h = d
for m in range(1, iterations + 1):
mf = float(m)
aa = mf * (b - mf) * x / ((qam + 2.0 * mf) * (a + 2.0 * mf))
d = 1.0 + aa * d
d = 1.0 / torch.where(d.abs() < tiny, torch.full_like(d, tiny), d)
c = 1.0 + aa / torch.where(c.abs() < tiny, torch.full_like(c, tiny), c)
h = h * d * c
aa = -(a + mf) * (qab + mf) * x / ((a + 2.0 * mf) * (qap + 2.0 * mf))
d = 1.0 + aa * d
d = 1.0 / torch.where(d.abs() < tiny, torch.full_like(d, tiny), d)
c = 1.0 + aa / torch.where(c.abs() < tiny, torch.full_like(c, tiny), c)
h = h * d * c
return h
def regularized_beta(x: torch.Tensor, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
x_full, a_full, b_full = torch.broadcast_tensors(x, a, b)
safe_x = x_full.clamp(1e-7, 1.0 - 1e-7)
log_bt = torch.lgamma(a_full + b_full) - torch.lgamma(a_full) - torch.lgamma(b_full)
log_bt = log_bt + a_full * torch.log(safe_x) + b_full * torch.log1p(-safe_x)
bt = torch.exp(log_bt.clamp(-80.0, 30.0))
lower = safe_x < (a_full + 1.0) / (a_full + b_full + 2.0)
direct = bt * _betacf(a_full, b_full, safe_x) / a_full
complement = 1.0 - bt * _betacf(b_full, a_full, 1.0 - safe_x) / b_full
result = torch.where(lower, direct, complement).clamp(0.0, 1.0)
return torch.where(x_full <= 0.0, torch.zeros_like(result), torch.where(x_full >= 1.0, torch.ones_like(result), result))
def supervised_loss_per_sample(
output: dict[str, Any], class_y: torch.Tensor, regression_y: torch.Tensor,
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
probs = output["class_probs_by_path"].clamp_min(1e-8)
beta = output["beta_params"].clamp_min(1e-4)
path_count, batch = probs.shape[:2]
negative = class_y == 0
neutral = class_y == 1
positive = class_y == 2
u = (regression_y.abs() / 3.0).clamp(0.0, 1.0)
lo = (u - DELTA_U).clamp(0.0, 1.0)
hi = (u + DELTA_U).clamp(0.0, 1.0)
neutral_mass = probs[:, :, 1]
neg_params = beta[:, :, 0, :]
pos_params = beta[:, :, 1, :]
cdf_hi_neg = regularized_beta(hi.unsqueeze(0), neg_params[..., 0], neg_params[..., 1])
cdf_lo_neg = regularized_beta(lo.unsqueeze(0), neg_params[..., 0], neg_params[..., 1])
cdf_hi_pos = regularized_beta(hi.unsqueeze(0), pos_params[..., 0], pos_params[..., 1])
cdf_lo_pos = regularized_beta(lo.unsqueeze(0), pos_params[..., 0], pos_params[..., 1])
neg_mass = probs[:, :, 0] * (cdf_hi_neg - cdf_lo_neg).clamp_min(1e-12)
pos_mass = probs[:, :, 2] * (cdf_hi_pos - cdf_lo_pos).clamp_min(1e-12)
selected = torch.where(neutral.unsqueeze(0), neutral_mass, torch.where(negative.unsqueeze(0), neg_mass, pos_mass))
mixture_mass = selected.mean(dim=0).clamp_min(1e-12)
nll = -torch.log(mixture_mass)
beta_mean = beta[..., 0] / beta.sum(dim=-1)
conditional = 3.0 * (probs[:, :, 2] * beta_mean[:, :, 1] - probs[:, :, 0] * beta_mean[:, :, 0])
mean_score = conditional.mean(dim=0)
scaled_error = (regression_y - mean_score) / 3.0
huber = F.huber_loss(scaled_error, torch.zeros_like(scaled_error), reduction="none", delta=0.25)
total = nll + LAMBDA_Y * huber
return total, {"nll": nll, "huber": huber, "mean_score": mean_score}
def reconstruction_loss_per_sample(
output: dict[str, Any], xs: list[torch.Tensor], hidden: torch.Tensor,
) -> torch.Tensor:
batch = hidden.shape[0]
per_modal = []
for m, prediction in enumerate(output["reconstructions"]):
target = xs[m].unsqueeze(0)
error = (prediction - target).abs().mean(dim=-1)
mask = hidden[:, :, m].float().unsqueeze(0)
numerator = (error * mask).sum(dim=(0, 2))
denominator = mask.sum(dim=(0, 2)).clamp_min(1.0)
per_modal.append(numerator / denominator)
values = torch.stack(per_modal, dim=-1)
active = torch.stack([hidden[:, :, m].any(dim=1) for m in range(len(MODALITIES))], dim=-1).float()
return (values * active).sum(dim=-1) / active.sum(dim=-1).clamp_min(1.0)
def _regularized_logits(probabilities: np.ndarray, temperature: float) -> np.ndarray:
logp = np.log(np.clip(probabilities, 1e-12, 1.0)) / temperature
logp -= logp.max(axis=1, keepdims=True)
exp = np.exp(logp)
return exp / exp.sum(axis=1, keepdims=True)
def _decode_mixture(
probabilities_by_path: np.ndarray,
beta_params: np.ndarray,
temperature: float = 1.0,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
pbar = _regularized_logits(probabilities_by_path.mean(axis=0), temperature)
maximum = pbar.max(axis=1, keepdims=True)
ties = np.isclose(pbar, maximum, rtol=0.0, atol=1e-12)
predicted_class = np.asarray([next(c for c in (1, 0, 2) if row[c]) for row in ties], dtype=np.int64)
score = np.zeros(len(predicted_class), dtype=np.float32)
for i, cls in enumerate(predicted_class):
if cls == 1:
continue
sign_index = 0 if cls == 0 else 1
weights = probabilities_by_path[:, i, cls]
params = beta_params[:, i, sign_index]
denominator = float(weights.sum())
if denominator <= 1e-12:
magnitude = 0.5
else:
low, high = 0.0, 1.0
for _ in range(48):
middle = (low + high) / 2.0
cdf = float(np.dot(weights, betainc(params[:, 0], params[:, 1], middle)) / denominator)
if cdf < 0.5:
low = middle
else:
high = middle
magnitude = (low + high) / 2.0
score[i] = (-3.0 if cls == 0 else 3.0) * magnitude
return pbar, predicted_class, score
def _predictive_intervals(
probabilities_by_path: np.ndarray,
beta_params: np.ndarray,
temperature: float,
quantiles: tuple[float, float] = (0.05, 0.95),
) -> tuple[np.ndarray, np.ndarray]:
"""Central intervals of the calibrated signed point-mass/Beta mixture."""
pbar = _regularized_logits(probabilities_by_path.mean(axis=0), temperature)
lower = np.empty(len(pbar), dtype=np.float32)
upper = np.empty(len(pbar), dtype=np.float32)
for i, marginal in enumerate(pbar):
negative_weight = probabilities_by_path[:, i, 0].astype(np.float64)
positive_weight = probabilities_by_path[:, i, 2].astype(np.float64)
negative_weight /= negative_weight.sum()
positive_weight /= positive_weight.sum()
negative = beta_params[:, i, 0].astype(np.float64)
positive = beta_params[:, i, 1].astype(np.float64)
def cdf(value: float) -> float:
if value < 0.0:
magnitude_threshold = min(1.0, max(0.0, -value / 3.0))
conditional = np.dot(
negative_weight,
1.0 - betainc(negative[:, 0], negative[:, 1], magnitude_threshold),
)
return float(marginal[0] * conditional)
magnitude_threshold = min(1.0, max(0.0, value / 3.0))
conditional = np.dot(
positive_weight,
betainc(positive[:, 0], positive[:, 1], magnitude_threshold),
)
return float(marginal[0] + marginal[1] + marginal[2] * conditional)
for slot, quantile in enumerate(quantiles):
lo, hi = -3.0, 3.0
for _ in range(52):
mid = (lo + hi) / 2.0
if cdf(mid) >= quantile:
hi = mid
else:
lo = mid
if slot == 0:
lower[i] = hi
else:
upper[i] = hi
return lower, upper
def _trajectory_variance_components(
probabilities_by_path: np.ndarray,
beta_params: np.ndarray,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Verify Var(Y)=E_b Var(Y|b)+Var_b(E[Y|b]) before class calibration."""
probs = np.asarray(probabilities_by_path, dtype=np.float64)
beta = np.asarray(beta_params, dtype=np.float64)
mean = beta[..., 0] / beta.sum(axis=-1)
second = beta[..., 0] * (beta[..., 0] + 1.0) / (beta.sum(axis=-1) * (beta.sum(axis=-1) + 1.0))
path_mean = 3.0 * (probs[..., 2] * mean[..., 1] - probs[..., 0] * mean[..., 0])
path_second = 9.0 * (probs[..., 2] * second[..., 1] + probs[..., 0] * second[..., 0])
conditional_variance = np.maximum(0.0, path_second - np.square(path_mean))
within = conditional_variance.mean(axis=0)
between = path_mean.var(axis=0)
total = within + between
return total.astype(np.float32), within.astype(np.float32), between.astype(np.float32)
def _calibrated_mixture_moments(
probabilities_by_path: np.ndarray,
beta_params: np.ndarray,
temperature: float,
) -> tuple[np.ndarray, np.ndarray]:
"""Recompute mean and variance after temperature calibration of class mass."""
pcal = _regularized_logits(probabilities_by_path.mean(axis=0), temperature)
beta = np.asarray(beta_params, dtype=np.float64)
raw_probs = np.asarray(probabilities_by_path, dtype=np.float64)
beta_mean = beta[..., 0] / beta.sum(axis=-1)
beta_second = beta[..., 0] * (beta[..., 0] + 1.0) / (beta.sum(axis=-1) * (beta.sum(axis=-1) + 1.0))
conditional_mean = np.zeros((raw_probs.shape[1], 2), dtype=np.float64)
conditional_second = np.zeros((raw_probs.shape[1], 2), dtype=np.float64)
for sign_index, class_index in enumerate((0, 2)):
weights = raw_probs[:, :, class_index]
weights = weights / weights.sum(axis=0, keepdims=True)
conditional_mean[:, sign_index] = np.sum(weights * beta_mean[:, :, sign_index], axis=0)
conditional_second[:, sign_index] = np.sum(weights * beta_second[:, :, sign_index], axis=0)
mean = 3.0 * (pcal[:, 2] * conditional_mean[:, 1] - pcal[:, 0] * conditional_mean[:, 0])
second = 9.0 * (pcal[:, 2] * conditional_second[:, 1] + pcal[:, 0] * conditional_second[:, 0])
variance = np.maximum(0.0, second - np.square(mean))
return mean.astype(np.float32), variance.astype(np.float32)
def calculate_metrics(
y_cls: np.ndarray,
y_reg: np.ndarray,
probs: np.ndarray,
pred_cls: np.ndarray,
pred_reg: np.ndarray,
selection_nll: float | None = None,
interval_lower: np.ndarray | None = None,
interval_upper: np.ndarray | None = None,
variance_components: tuple[np.ndarray, np.ndarray, np.ndarray] | None = None,
calibrated_moments: tuple[np.ndarray, np.ndarray] | None = None,
) -> dict[str, Any]:
p = np.clip(probs, 1e-8, 1.0)
onehot = np.eye(3, dtype=np.float64)[y_cls]
nll = float(-np.log(p[np.arange(len(y_cls)), y_cls]).mean())
confidence = p.max(axis=1)
correct = (pred_cls == y_cls).astype(np.float64)
ece = 0.0
for left in np.linspace(0, 1, 16)[:-1]:
right = left + 1 / 15
selected = (confidence >= left) & (confidence < right if right < 1 else confidence <= right)
if selected.any():
ece += selected.mean() * abs(confidence[selected].mean() - correct[selected].mean())
pearson = float(pearsonr(y_reg, pred_reg).statistic) if np.std(y_reg) > 0 and np.std(pred_reg) > 0 else float("nan")
support = np.bincount(y_cls.astype(int), minlength=3)
result = {
"n": int(len(y_cls)),
"accuracy": float(accuracy_score(y_cls, pred_cls)),
"macro_f1": float(f1_score(y_cls, pred_cls, labels=[0, 1, 2], average="macro", zero_division=0)),
"negative_support": int(support[0]),
"neutral_support": int(support[1]),
"positive_support": int(support[2]),
"negative_recall": float(recall_score(y_cls, pred_cls, labels=[0], average="macro", zero_division=0)),
"middle_recall": float(recall_score(y_cls, pred_cls, labels=[1], average="macro", zero_division=0)),
"positive_recall": float(recall_score(y_cls, pred_cls, labels=[2], average="macro", zero_division=0)),
"regression_mae": float(mean_absolute_error(y_reg, pred_reg)),
"regression_rmse": float(np.sqrt(mean_squared_error(y_reg, pred_reg))),
"pearson": pearson,
"brier": float(np.square(probs - onehot).sum(axis=1).mean()),
"classification_nll": nll,
"ece_15": float(ece),
}
if selection_nll is not None:
result["selection_nll"] = float(selection_nll)
if interval_lower is not None and interval_upper is not None:
result["interval_90_coverage"] = float(np.mean((y_reg >= interval_lower) & (y_reg <= interval_upper)))
result["interval_90_mean_width"] = float(np.mean(interval_upper - interval_lower))
if variance_components is not None:
total, within, between = variance_components
result["predictive_variance_mean_uncalibrated"] = float(np.mean(total))
result["within_trajectory_variance_mean"] = float(np.mean(within))
result["between_trajectory_variance_mean"] = float(np.mean(between))
if calibrated_moments is not None:
calibrated_mean, calibrated_variance = calibrated_moments
result["predictive_mean_mean_calibrated"] = float(np.mean(calibrated_mean))
result["predictive_variance_mean_calibrated"] = float(np.mean(calibrated_variance))
return result
def evaluate(
model: CRG,
arrays: dict[str, np.ndarray],
split: SplitData,
device: torch.device,
batch_size: int,
*,
masks: np.ndarray | None = None,
temperature: float = 1.0,
collect_gate_diagnostics: bool = False,
) -> tuple[dict[str, Any], dict[str, np.ndarray]]:
model.eval()
source_masks = split.mask if masks is None else np.asarray(masks, dtype=bool)
prob_paths, beta_paths, loss_parts = [], [], []
gate_parts: dict[str, list[np.ndarray]] = {
"fusion_weights": [], "null_weights": [], "time_pool_weights": [], "reliability": [],
"imputation_uncertainty": [], "gap": [], "span": [], "distance_before": [], "distance_after": [],
}
with torch.inference_mode():
for start in range(0, split.n, batch_size):
idx = np.arange(start, min(split.n, start + batch_size))
xs, observed = to_device_batch(arrays, source_masks, idx, device)
out = model(xs, observed, paths=16 if model.use_joint_draws else 1, joint_draws=model.use_joint_draws)
prob_paths.append(out["class_probs_by_path"].cpu().numpy())
beta_paths.append(out["beta_params"].cpu().numpy())
if collect_gate_diagnostics:
for name, key in (("fusion_weights", "fusion_weights_by_path"),
("null_weights", "null_weights_by_path"),
("time_pool_weights", "time_pool_weights_by_path")):
gate_parts[name].append(out[key].mean(dim=0).cpu().numpy())
for name in ("reliability", "imputation_uncertainty", "gap", "span", "distance_before", "distance_after"):
gate_parts[name].append(out[name].cpu().numpy())
cy = torch.from_numpy(split.class_y[idx]).to(device)
ry = torch.from_numpy(split.regression_y[idx]).to(device)
nll, _ = supervised_loss_per_sample(out, cy, ry)
loss_parts.append(nll.cpu().numpy())
ppaths = np.concatenate(prob_paths, axis=1)
betas = np.concatenate(beta_paths, axis=1)
probs, pred_class, pred_score = _decode_mixture(ppaths, betas, temperature)
interval_lower, interval_upper = _predictive_intervals(ppaths, betas, temperature)
variance_components = _trajectory_variance_components(ppaths, betas)
calibrated_moments = _calibrated_mixture_moments(ppaths, betas, temperature)
metrics = calculate_metrics(
split.class_y, split.regression_y, probs, pred_class, pred_score,
float(np.mean(np.concatenate(loss_parts))), interval_lower, interval_upper, variance_components,
calibrated_moments,
)
predictions = {
"probabilities": probs,
"probabilities_by_path": ppaths,
"beta_params": betas,
"predicted_class": pred_class,
"predicted_score": pred_score,
"interval_lower": interval_lower,
"interval_upper": interval_upper,
"predictive_variance_uncalibrated": variance_components[0],
"within_trajectory_variance": variance_components[1],
"between_trajectory_variance": variance_components[2],
"predictive_mean_calibrated": calibrated_moments[0],
"predictive_variance_calibrated": calibrated_moments[1],
}
if collect_gate_diagnostics:
predictions.update({name: np.concatenate(values, axis=0) for name, values in gate_parts.items()})
return metrics, predictions
def gate_diagnostic_rows(split: SplitData, predictions: dict[str, np.ndarray]) -> list[dict[str, Any]]:
required = {"fusion_weights", "null_weights", "time_pool_weights", "reliability",
"imputation_uncertainty", "gap", "span", "distance_before", "distance_after"}
if not required.issubset(predictions):
raise ValueError(f"missing gate diagnostic arrays: {sorted(required - set(predictions))}")
rows = []
for sample_index, sample_id in enumerate(split.ids):
for step in range(split.mask.shape[1]):
for modality_index, modality in enumerate(MODALITIES):
rows.append({
"sample_id": sample_id,
"video_id": split.groups[sample_index],
"step": step,
"relative_position": step / max(1, split.mask.shape[1] - 1),
"modality": modality,
"observed": bool(split.mask[sample_index, step, modality_index]),
"fusion_weight_mean_over_paths": float(predictions["fusion_weights"][sample_index, step, modality_index]),
"null_weight_mean_over_paths": float(predictions["null_weights"][sample_index, step]),
"reliability": float(predictions["reliability"][sample_index, step, modality_index]),
"imputation_uncertainty": float(predictions["imputation_uncertainty"][sample_index, step, modality_index]),
"nearest_observation_gap": float(predictions["gap"][sample_index, step, modality_index]),
"continuous_missing_span": float(predictions["span"][sample_index, step, modality_index]),
"distance_before": float(predictions["distance_before"][sample_index, step, modality_index]),
"distance_after": float(predictions["distance_after"][sample_index, step, modality_index]),
"time_pool_weight_mean_over_paths": float(predictions["time_pool_weights"][sample_index, step]),
})
return rows
def fit_temperature(probabilities: np.ndarray, class_y: np.ndarray) -> float:
def objective(log_temperature: float) -> float:
p = _regularized_logits(probabilities, float(np.exp(log_temperature)))
return float(-np.log(np.clip(p[np.arange(len(class_y)), class_y], 1e-12, 1.0)).mean())
fitted = minimize_scalar(objective, bounds=(-2.0, 2.0), method="bounded", options={"xatol": 1e-5})
return float(np.exp(fitted.x))
def _group_ids(original_mask: np.ndarray, hidden_mask: np.ndarray) -> np.ndarray:
# Group by which sources received new artificial gaps, not only by sources
# that were erased completely (the mask design deliberately preserves 20%).
original = np.asarray(original_mask, dtype=bool)
current = np.asarray(hidden_mask, dtype=bool)
hidden_mod = (original & ~current).any(axis=1)
bits = hidden_mod[:, 0].astype(int) + 2 * hidden_mod[:, 1].astype(int) + 4 * hidden_mod[:, 2].astype(int)
# PDF (5.62): total missing rate is computed per source on the valid time
# axis, then averaged equally across T/A/V (never weighted by feature size
# or by the number of naturally observed rows).
final_missing_by_modality = 1.0 - current.mean(axis=1)
realized_rate = final_missing_by_modality.mean(axis=1)
coarse = np.where(realized_rate <= 0.2, 0, np.where(realized_rate <= 0.5, 1, 2))
return bits * 3 + coarse
def _missing_rate_summary(original: np.ndarray, current: np.ndarray) -> dict[str, Any]:
"""Return PDF (5.62)–(5.63) rates, with equal modality weighting."""
original = np.asarray(original, dtype=bool)
current = np.asarray(current, dtype=bool)
if original.shape != current.shape or original.ndim not in (2, 3):
raise ValueError("mask rate inputs must have matching [T,M] or [N,T,M] shapes")
newly_hidden = original & ~current
natural = 1.0 - original.mean(axis=-2)
final = 1.0 - current.mean(axis=-2)
observed_count = original.sum(axis=-2)
additional_count = newly_hidden.sum(axis=-2)
additional = np.divide(
additional_count, observed_count,
out=np.full(np.shape(additional_count), np.nan, dtype=np.float64),
where=observed_count > 0,
)
synchronous = (~current.any(axis=-1)).mean(axis=-1)
return {
"natural_by_modality": natural,
"additional_by_modality": additional,
"final_by_modality": final,
"natural_global": float(np.mean(natural)),
"additional_global": float(np.nanmean(additional)),
"final_global": float(np.mean(final)),
"synchronous_no_observation": synchronous,
}
def smooth_group_risk(
losses: torch.Tensor,
group_ids: np.ndarray,
lambda_group: float = LAMBDA_GROUP,
group_temperature: float = GROUP_TEMPERATURE,
) -> torch.Tensor:
gids = torch.as_tensor(group_ids, device=losses.device, dtype=torch.long)
unique = torch.unique(gids)
group_losses, priors = [], []
for group in unique:
selected = gids == group
group_losses.append(losses[selected].mean())
priors.append(selected.float().mean())
values = torch.stack(group_losses)
prior = torch.stack(priors).clamp_min(1e-8)
expected = (prior * values).sum()
worst = group_temperature * torch.logsumexp(torch.log(prior) + values / group_temperature, dim=0)
return (1.0 - lambda_group) * expected + lambda_group * worst
def _distillation_per_sample(
student: dict[str, Any], teacher: dict[str, Any], original: np.ndarray, current: np.ndarray,
) -> torch.Tensor:
temp = DISTILL_TEMPERATURE
p_teacher = teacher["tempered_probs_by_path"].mean(dim=0).detach().clamp_min(1e-8)
p_student = student["tempered_probs_by_path"].mean(dim=0).clamp_min(1e-8)
entropy = -(p_teacher * p_teacher.log()).sum(dim=-1)
confidence_weight = (1.0 - entropy / math.log(3.0)).clamp(0.0, 1.0)
retain_by_modality = []
orig_t = torch.as_tensor(original, device=p_teacher.device, dtype=torch.float32)
curr_t = torch.as_tensor(current, device=p_teacher.device, dtype=torch.float32)
for m in range(len(MODALITIES)):
denominator = orig_t[:, :, m].sum(dim=1)
retained = (orig_t[:, :, m] * curr_t[:, :, m]).sum(dim=1) / denominator.clamp_min(1.0)
retain_by_modality.append(torch.where(denominator > 0, retained, torch.ones_like(retained)))
retain = torch.stack(retain_by_modality, dim=-1).mean(dim=-1)
weight = confidence_weight * retain
kl = (p_teacher * (p_teacher.log() - p_student.log())).sum(dim=-1) * temp * temp
teacher_score = teacher["mixed_score"].detach()
student_score = student["mixed_score"]
reg = F.huber_loss((teacher_score - student_score) / 3.0, torch.zeros_like(teacher_score), reduction="none", delta=0.25)
return weight * (kl + reg)
def fit_imputer(
imputer: StructuredGaussianImputer,
arrays: dict[str, np.ndarray],
split: SplitData,
device: torch.device,
epochs: int,
batch_size: int,
seed: int,
) -> list[dict[str, float]]:
imputer.train()
optimizer = torch.optim.AdamW(imputer.parameters(), lr=3e-4, weight_decay=1e-4)
rng = np.random.default_rng(seed)
history = []
for epoch in range(1, epochs + 1):
order = rng.permutation(split.n)
losses = []
for start in range(0, split.n, batch_size):
idx = order[start:start + batch_size]
xs, observed = to_device_batch(arrays, split.mask, idx, device)
nll = imputer.observed_nll(xs, observed).mean()
observed_scalars = sum(observed[:, :, m].sum(dim=1).float() * INPUT_DIMS[m] for m in range(len(MODALITIES)))
emission_penalty = sum(value.square().mean() for value in imputer.emissions())
transition_penalty = imputer._transition().square().mean()
loss = nll / observed_scalars.mean().clamp_min(1.0)
loss = loss + LAMBDA_EMISSION * emission_penalty + LAMBDA_TRANSITION * transition_penalty
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(imputer.parameters(), 5.0)
optimizer.step()
losses.append(float(loss.detach().cpu()))
row = {"stage": "structured_imputer", "epoch": epoch, "train_observed_nll_per_scalar": float(np.mean(losses))}
history.append(row)
print(f"imputer {epoch}/{epochs}: observed_nll/scalar={row['train_observed_nll_per_scalar']:.4f}", flush=True)
imputer.eval()
for parameter in imputer.parameters():
parameter.requires_grad_(False)
return history
def _make_variant(
name: str,
imputer: StructuredGaussianImputer,
reliability_hparams: tuple[float, float, float, float] = DEFAULT_RELIABILITY,
) -> CRG:
if name == "C1":
flags = dict(use_imputer=False, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False)
elif name == "C2":
flags = dict(use_imputer=True, use_joint_draws=False, use_final_gate=False, use_source_attention=False, reliability_update=False, use_low_rank=False)
elif name == "C3":
flags = dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=False, reliability_update=False, use_low_rank=False)
elif name == "C4":
flags = dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True, reliability_update=False, use_low_rank=False)
else:
flags = dict(use_imputer=True, use_joint_draws=True, use_final_gate=True, use_source_attention=True,
reliability_update=True, use_low_rank=name == "C6" or name.startswith("C7"))
return CRG(copy.deepcopy(imputer), reliability_hparams=reliability_hparams, **flags)
def _set_reliability_hparams(model: CRG, values: tuple[float, float, float, float]) -> None:
rho_imp, lambda_u, lambda_gap, lambda_span = values
if not 0.0 < rho_imp < 1.0 or min(lambda_u, lambda_gap, lambda_span) < 0.0:
raise ValueError("invalid reliability hyperparameters")
with torch.no_grad():
model.rho_imp.fill_(rho_imp)
model.rel_u.fill_(lambda_u)
model.rel_gap.fill_(lambda_gap)
model.rel_span.fill_(lambda_span)
def tune_reliability_hparams(
model: CRG,
arrays: dict[str, np.ndarray],
split: SplitData,
scenarios: dict[str, np.ndarray],
device: torch.device,
batch_size: int,
model_name: str,
seed: int = SEED + 551,
candidate_values: tuple[tuple[float, float, float, float], ...] = RELIABILITY_CANDIDATES,
) -> tuple[tuple[float, float, float, float], list[dict[str, Any]]]:
if not (model.use_final_gate or model.use_source_attention or model.reliability_update):
return DEFAULT_RELIABILITY, [{"model": model_name, "selected": True,
"rho_imp": DEFAULT_RELIABILITY[0], "lambda_u": DEFAULT_RELIABILITY[1],
"lambda_gap": DEFAULT_RELIABILITY[2], "lambda_span": DEFAULT_RELIABILITY[3],
"inner_selection_nll": float("nan"), "note": "not used by this ablation"}]
rows = []
best_values, best_loss = DEFAULT_RELIABILITY, float("inf")
if not candidate_values:
raise ValueError("at least one reliability candidate is required")
for values in candidate_values:
_set_reliability_hparams(model, values)
scenario_losses: dict[str, float] = {}
for scenario, masks in scenarios.items():
with fixed_torch_seed(_scenario_seed(seed, split.name, scenario), device):
metrics, _ = evaluate(model, arrays, split, device, batch_size, masks=masks)
scenario_losses[scenario] = float(metrics["selection_nll"])
score = float(np.mean(list(scenario_losses.values())))
row = {"model": model_name, "rho_imp": values[0], "lambda_u": values[1],
"lambda_gap": values[2], "lambda_span": values[3], "inner_selection_nll": score,
"scenario_selection_nll": json.dumps(scenario_losses, sort_keys=True),
"selected": False, "selection_split": split.name}
rows.append(row)
if score < best_loss:
best_loss, best_values = score, values
_set_reliability_hparams(model, best_values)
for row in rows:
row["selected"] = (row["rho_imp"], row["lambda_u"], row["lambda_gap"], row["lambda_span"]) == best_values
return best_values, rows
def tune_group_risk_model(
imputer: StructuredGaussianImputer,
train: SplitData,
valid: SplitData,
arrays: dict[str, dict[str, np.ndarray]],
reliability_validation: SplitData,
reliability_arrays: dict[str, np.ndarray],
reliability_scenarios: dict[str, np.ndarray],
device: torch.device,
epochs: int,
batch_size: int,
patience: int,
seed: int,
) -> tuple[CRG, list[dict[str, float]], list[dict[str, Any]], list[dict[str, Any]], tuple[float, float, float, float], tuple[float, float], float]:
candidates = []
all_history: list[dict[str, float]] = []
all_reliability_rows: list[dict[str, Any]] = []
risk_rows: list[dict[str, Any]] = []
for candidate_index, (group_lambda, group_temperature) in enumerate(GROUP_RISK_CANDIDATES):
candidate_name = f"C7_group_lambda{group_lambda:.2f}_tau{group_temperature:.2f}"
model = _make_variant("C7_group", imputer)
with fixed_torch_seed(seed + 303, device):
model, history = _fit_neural(
model, candidate_name, train, valid, arrays, device, epochs, batch_size, patience,
np.random.default_rng(seed + 303), use_group_risk=True,
group_lambda=group_lambda, group_temperature=group_temperature,
selection_split=reliability_validation, selection_arrays=reliability_arrays,
selection_scenarios=reliability_scenarios,
)
all_history.extend(history)
selected_reliability, tuning_rows = tune_reliability_hparams(
model, reliability_arrays, reliability_validation, reliability_scenarios,
device, batch_size, candidate_name, seed=seed + 551,
)
selected_row = next(row for row in tuning_rows if row.get("selected"))
inner_score = float(selected_row["inner_selection_nll"])
for row in tuning_rows:
row["group_lambda"] = group_lambda
row["group_temperature"] = group_temperature
row["risk_candidate_selected"] = False
row["candidate"] = row["model"]
row["model"] = "C7_group"
all_reliability_rows.extend(tuning_rows)
candidates.append((inner_score, model, selected_reliability, (group_lambda, group_temperature), candidate_name))
risk_rows.append({"model": "C7_group", "candidate": candidate_name,
"lambda_group": group_lambda, "group_temperature": group_temperature,
"selected_reliability": selected_reliability,
"inner_selection_nll": inner_score, "selected": False,
"selection_split": reliability_validation.name})
best = min(candidates, key=lambda row: row[0])
score, model, selected_reliability, selected_group, candidate_name = best
for row in all_reliability_rows:
row["risk_candidate_selected"] = row["candidate"] == candidate_name
row["selected"] = bool(row.get("selected") and row["risk_candidate_selected"])
for row in risk_rows:
row["selected"] = row["candidate"] == candidate_name
if row["selected"]:
row["selected"] = True
return model, all_history, all_reliability_rows, risk_rows, selected_reliability, selected_group, score
def _fit_neural(
model: CRG,
name: str,
train: SplitData,
valid: SplitData,
arrays: dict[str, dict[str, np.ndarray]],
device: torch.device,
epochs: int,
batch_size: int,
patience: int,
rng: np.random.Generator,
*,
teacher: CRG | None = None,
use_group_risk: bool = False,
group_lambda: float = LAMBDA_GROUP,
group_temperature: float = GROUP_TEMPERATURE,
selection_split: SplitData | None = None,
selection_arrays: dict[str, np.ndarray] | None = None,
selection_scenarios: dict[str, np.ndarray] | None = None,
mask_kind: str = "continuous",
use_reconstruction: bool = True,
) -> tuple[CRG, list[dict[str, float]]]:
model.to(device)
model.imputer.eval()
for parameter in model.imputer.parameters():
parameter.requires_grad_(False)
if teacher is not None:
teacher.eval()
for parameter in teacher.parameters():
parameter.requires_grad_(False)
optimizer = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=3e-4, weight_decay=1e-3)
best_loss, best_state, stale = float("inf"), None, 0
history: list[dict[str, float]] = []
rate_choices = np.asarray(MASK_RATES, dtype=np.float64)
pattern_choices = np.asarray(MASK_MODES, dtype=object)
for epoch in range(1, epochs + 1):
model.train()
order = rng.permutation(train.n)
epoch_losses = []
for start in range(0, train.n, batch_size):
idx = order[start:start + batch_size]
xs, natural = to_device_batch(arrays["fit"], train.mask, idx, device)
if name == "teacher":
rates = np.zeros(len(idx), dtype=np.float64)
modes = ["none"] * len(idx)
else:
rates = rng.choice(rate_choices, size=len(idx), p=np.asarray([0.2] * 5))
modes = rng.choice(pattern_choices, size=len(idx)).tolist()
masks_np = np.stack([
continuous_mask(train.mask[int(i)], float(rate), str(mode), rng, kind=mask_kind)
for i, rate, mode in zip(idx, rates, modes)
])
current = torch.as_tensor(masks_np, device=device, dtype=torch.bool)
class_y = torch.as_tensor(train.class_y[idx], device=device, dtype=torch.long)
regression_y = torch.as_tensor(train.regression_y[idx], device=device, dtype=torch.float32)
output = model(xs, current, paths=4 if model.use_joint_draws else 1, joint_draws=model.use_joint_draws)
supervised, _ = supervised_loss_per_sample(output, class_y, regression_y)
hidden_np = train.mask[idx] & ~masks_np
hidden = torch.as_tensor(hidden_np, device=device, dtype=torch.bool)
reconstruction = reconstruction_loss_per_sample(output, xs, hidden)
per_sample = supervised + (LAMBDA_RECON * reconstruction if use_reconstruction else 0.0)
if teacher is not None:
with torch.no_grad():
teacher_output = teacher(xs, natural, paths=4, joint_draws=True)
distill = _distillation_per_sample( output, teacher_output, train.mask[idx], masks_np)
per_sample = per_sample + LAMBDA_DISTILL * distill
if use_group_risk:
group_ids = _group_ids(train.mask[idx], masks_np)
loss = smooth_group_risk(per_sample, group_ids, group_lambda, group_temperature)
else:
loss = per_sample.mean()
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_([p for p in model.parameters() if p.requires_grad], 1.0)
optimizer.step()
epoch_losses.append(float(loss.detach().cpu()))
if selection_split is not None and selection_arrays is not None and selection_scenarios:
selected_metrics = {}
for scenario, masks in selection_scenarios.items():
scenario_metrics, _ = evaluate(model, selection_arrays, selection_split, device, batch_size, masks=masks)
selected_metrics[scenario] = scenario_metrics
val_loss = float(np.mean([metrics["selection_nll"] for metrics in selected_metrics.values()]))
natural_metrics = next((metrics for key, metrics in selected_metrics.items()
if key.startswith("0.0/") or key.endswith("natural")),
next(iter(selected_metrics.values())))
selection_source = "group_disjoint_internal_scenarios"
else:
natural_metrics, _ = evaluate(model, arrays["valid"], valid, device, batch_size)
val_loss = float(natural_metrics["selection_nll"])
selection_source = "official_valid_fallback"
row = {"stage": name, "epoch": epoch, "train_loss": float(np.mean(epoch_losses)),
"inner_selection_nll": val_loss, "inner_natural_accuracy": natural_metrics["accuracy"],
"inner_natural_macro_f1": natural_metrics["macro_f1"],
"inner_natural_mae": natural_metrics["regression_mae"],
"selection_source": selection_source}
history.append(row)
print(f"{name} {epoch}/{epochs}: train={row['train_loss']:.4f} innerNLL={val_loss:.4f} "
f"acc={row['inner_natural_accuracy']:.4f} macroF1={row['inner_natural_macro_f1']:.4f}", flush=True)
if val_loss < best_loss:
best_loss, best_state, stale = val_loss, copy.deepcopy(model.state_dict()), 0
else:
stale += 1
if stale >= patience:
break
if best_state is not None:
model.load_state_dict(best_state)
model.eval()
return model, history
def _sample_statistics(
split: SplitData, arrays: dict[str, np.ndarray], indices: np.ndarray, masks: np.ndarray | None = None,
) -> np.ndarray:
parts = []
effective_mask = split.mask if masks is None else np.asarray(masks, dtype=bool)
for modality_index, modality in enumerate(MODALITIES):
x = arrays[modality][indices]
mask = effective_mask[indices, :, modality_index]
count = mask.sum(axis=1, keepdims=True)
mean = (x * mask[:, :, None]).sum(axis=1) / np.maximum(count, 1)
variance = (((x - mean[:, None, :]) ** 2) * mask[:, :, None]).sum(axis=1) / np.maximum(count, 1)
missing = 1.0 - mask.mean(axis=1, keepdims=True)
max_gap = []
for row in mask:
longest = current = 0
for visible in row:
current = 0 if visible else current + 1
longest = max(longest, current)
max_gap.append(longest / max(1, len(row)))
parts.extend((mean, np.sqrt(variance), missing, np.asarray(max_gap, np.float32)[:, None]))
return np.concatenate(parts, axis=1).astype(np.float32)
def fit_c0(
train: SplitData,
valid: SplitData,
arrays: dict[str, dict[str, np.ndarray]],
) -> tuple[dict[str, Any], dict[str, Any]]:
train_x = _sample_statistics(train, arrays["fit"], np.arange(train.n))
classifier = LogisticRegression(C=0.05, max_iter=2500, random_state=SEED)
classifier.fit(train_x, train.class_y)
regressor = Ridge(alpha=25.0)
regressor.fit(train_x, train.regression_y)
return {}, {"classifier": classifier, "regressor": regressor}
def calibrate_c0_interval(
state: dict[str, Any],
calibration: SplitData,
arrays: dict[str, np.ndarray],
) -> None:
features = _sample_statistics(calibration, arrays, np.arange(calibration.n))
residual = calibration.regression_y - state["regressor"].predict(features)
state["residual_q05"], state["residual_q95"] = (
float(np.quantile(residual, 0.05)), float(np.quantile(residual, 0.95))
)
def evaluate_c0(
state: dict[str, Any], split: SplitData, arrays: dict[str, np.ndarray],
temperature: float = 1.0, masks: np.ndarray | None = None,
) -> tuple[dict[str, Any], dict[str, np.ndarray]]:
features = _sample_statistics(split, arrays, np.arange(split.n), masks)
classifier, regressor = state["classifier"], state["regressor"]
raw = np.zeros((split.n, 3), np.float64)
raw[:, classifier.classes_] = classifier.predict_proba(features)
probs = _regularized_logits(raw, temperature)
maximum = probs.max(axis=1, keepdims=True)
ties = np.isclose(probs, maximum, rtol=0.0, atol=1e-12)
pred_class = np.asarray([next(c for c in (1, 0, 2) if row[c]) for row in ties], dtype=np.int64)
score = np.clip(regressor.predict(features), -3.0, 3.0)
interval_lower = np.clip(score + state["residual_q05"], -3.0, 3.0)
interval_upper = np.clip(score + state["residual_q95"], -3.0, 3.0)
metrics = calculate_metrics(
split.class_y, split.regression_y, probs, pred_class, score,
interval_lower=interval_lower, interval_upper=interval_upper,
)
scaled_error = torch.as_tensor((split.regression_y - score) / 3.0, dtype=torch.float32)
metrics["selection_nll"] = metrics["classification_nll"] + float(
F.huber_loss(scaled_error, torch.zeros_like(scaled_error), delta=0.25)
)
return metrics, {"probabilities": probs, "predicted_class": pred_class, "predicted_score": score,
"interval_lower": interval_lower, "interval_upper": interval_upper}
def controlled_c0(
state: dict[str, Any], split: SplitData, arrays: dict[str, np.ndarray],
scenarios: dict[str, np.ndarray], temperature: float,
) -> tuple[list[dict[str, Any]], dict[str, dict[str, np.ndarray]]]:
rows = []
predictions = {}
for scenario, mask in scenarios.items():
metrics, prediction = evaluate_c0(state, split, arrays, temperature, mask)
predictions[scenario] = prediction
rates = _missing_rate_summary(split.mask, mask)
additional_by_modality = np.nanmean(rates["additional_by_modality"], axis=0)
rate, mode = scenario.split("/", 1)
rows.append({"evaluation_split": split.name, "model": "C0", "rate_requested_per_selected_source": float(rate),
"mask_pattern": mode, "rate_realized_global": rates["additional_global"],
"rate_realized_additional_global": rates["additional_global"],
"rate_realized_additional_by_modality": json.dumps([None if not np.isfinite(x) else float(x) for x in additional_by_modality]),
"natural_missing_rate_global": rates["natural_global"],
"natural_missing_rate_by_modality": json.dumps(np.mean(rates["natural_by_modality"], axis=0).tolist()),
"rate_final_total_missing_global": rates["final_global"],
"rate_final_total_missing_by_modality": json.dumps(np.mean(rates["final_by_modality"], axis=0).tolist()),
"synchronous_no_observation_rate": float(np.mean(rates["synchronous_no_observation"])), **metrics})
return rows, predictions
def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8-sig", newline="") as stream:
fieldnames = list(dict.fromkeys(key for row in rows for key in row))
writer = csv.DictWriter(stream, fieldnames=fieldnames, extrasaction="raise")
writer.writeheader()
writer.writerows(rows)
def mask_audit_rows(split: SplitData, scenarios: dict[str, np.ndarray], seed: int) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
for scenario, masks in scenarios.items():
requested_rate, pattern = scenario.split("/", 1)
for index, (sample_id, original, current) in enumerate(zip(split.ids, split.mask, masks)):
newly_hidden = np.asarray(original, dtype=bool) & ~np.asarray(current, dtype=bool)
intervals: dict[str, list[list[int]]] = {}
for modality_index, modality in enumerate(MODALITIES):
positions = np.flatnonzero(newly_hidden[:, modality_index])
spans: list[list[int]] = []
if len(positions):
start = previous = int(positions[0])
for position in positions[1:]:
position = int(position)
if position != previous + 1:
spans.append([start, previous + 1])
start = position
previous = position
spans.append([start, previous + 1])
intervals[modality] = spans
rate_summary = _missing_rate_summary(original, current)
rows.append({
"natural_missing_rate": rate_summary["natural_global"],
"natural_missing_rate_by_modality": json.dumps(rate_summary["natural_by_modality"].tolist()),
"realized_additional_rate": rate_summary["additional_global"],
"realized_additional_rate_by_modality": json.dumps([
None if not np.isfinite(x) else float(x)
for x in rate_summary["additional_by_modality"]
]),
"final_total_missing_rate": rate_summary["final_global"],
"final_total_missing_rate_by_modality": json.dumps(rate_summary["final_by_modality"].tolist()),
"synchronous_no_observation_rate": float(rate_summary["synchronous_no_observation"]),
"sample_id": sample_id,
"video_id": split.groups[index],
"scenario": scenario,
"requested_rate_per_selected_source": float(requested_rate),
"mask_pattern": pattern,
"mask_seed": _scenario_seed(seed, sample_id, scenario),
"selected_modalities": ",".join(m for m in MODALITIES if intervals[m]),
"newly_hidden_intervals_step_half_open": json.dumps(intervals, separators=(",", ":")),
"original_observed_steps": int(np.asarray(original, dtype=bool).sum()),
"newly_hidden_steps": int(newly_hidden.sum()),
})
return rows
def group_bootstrap(
split: SplitData, prediction: dict[str, np.ndarray], reps: int = 1000, seed: int = SEED + 44,
) -> list[dict[str, Any]]:
rng = np.random.default_rng(seed)
groups = np.unique(split.groups)
group_indices = {g: np.flatnonzero(split.groups == g) for g in groups}
rows = {name: [] for name in ("accuracy", "macro_f1", "mae", "rmse", "pearson", "interval_90_coverage", "interval_90_mean_width")}
for _ in range(reps):
chosen = rng.choice(groups, size=len(groups), replace=True)
idx = np.concatenate([group_indices[g] for g in chosen])
cls, score = prediction["predicted_class"][idx], prediction["predicted_score"][idx]
ycls, yreg = split.class_y[idx], split.regression_y[idx]
rows["accuracy"].append(accuracy_score(ycls, cls))
rows["macro_f1"].append(f1_score(ycls, cls, labels=[0, 1, 2], average="macro", zero_division=0))
rows["mae"].append(mean_absolute_error(yreg, score))
rows["rmse"].append(np.sqrt(mean_squared_error(yreg, score)))
rows["pearson"].append(pearsonr(yreg, score).statistic if np.std(score) and np.std(yreg) else np.nan)
rows["interval_90_coverage"].append(np.mean((yreg >= prediction["interval_lower"][idx]) & (yreg <= prediction["interval_upper"][idx])))
rows["interval_90_mean_width"].append(np.mean(prediction["interval_upper"][idx] - prediction["interval_lower"][idx]))
result = []
for name, values in rows.items():
values = np.asarray(values, dtype=np.float64)
result.append({"metric": name, "estimate": float(np.nanmedian(values)),
"ci_2_5": float(np.nanpercentile(values, 2.5)),
"ci_97_5": float(np.nanpercentile(values, 97.5)),
"replicates": reps, "unit": "source video group"})
return result
def paired_group_bootstrap_deltas(
split: SplitData,
predictions: dict[str, dict[str, np.ndarray]],
reps: int = 1000,
seed: int = SEED + 88,
) -> list[dict[str, Any]]:
"""Paired validation-set model deltas using one shared source-video resample."""
groups = np.unique(split.groups)
group_indices = {group: np.flatnonzero(split.groups == group) for group in groups}
compare_models = [name for name in predictions if name != "C0"]
metrics = ("accuracy", "macro_f1", "mae", "rmse", "pearson", "interval_90_coverage", "interval_90_mean_width")
def value(name: str, metric: str, indices: np.ndarray) -> float:
pred = predictions[name]
y_cls, y_reg = split.class_y[indices], split.regression_y[indices]
if metric == "accuracy":
return float(accuracy_score(y_cls, pred["predicted_class"][indices]))
if metric == "macro_f1":
return float(f1_score(y_cls, pred["predicted_class"][indices], labels=[0, 1, 2], average="macro", zero_division=0))
if metric == "mae":
return float(mean_absolute_error(y_reg, pred["predicted_score"][indices]))
if metric == "rmse":
return float(np.sqrt(mean_squared_error(y_reg, pred["predicted_score"][indices])))
if metric == "pearson":
estimate = pred["predicted_score"][indices]
return float(pearsonr(y_reg, estimate).statistic) if np.std(y_reg) and np.std(estimate) else float("nan")
if metric == "interval_90_coverage":
return float(np.mean((y_reg >= pred["interval_lower"][indices]) & (y_reg <= pred["interval_upper"][indices])))
if metric == "interval_90_mean_width":
return float(np.mean(pred["interval_upper"][indices] - pred["interval_lower"][indices]))
raise ValueError(metric)
point = {
(model, metric): value(model, metric, np.arange(split.n)) - value("C0", metric, np.arange(split.n))
for model in compare_models for metric in metrics
}
draws = {key: [] for key in point}
rng = np.random.default_rng(seed)
for _ in range(reps):
chosen = rng.choice(groups, size=len(groups), replace=True)
indices = np.concatenate([group_indices[group] for group in chosen])
baseline = {metric: value("C0", metric, indices) for metric in metrics}
for model in compare_models:
for metric in metrics:
draws[(model, metric)].append(value(model, metric, indices) - baseline[metric])
rows = []
for (model, metric), values in draws.items():
values = np.asarray(values, dtype=np.float64)
rows.append({
"model": model, "baseline": "C0", "metric": metric,
"point_delta": point[(model, metric)],
"bootstrap_median_delta": float(np.nanmedian(values)),
"ci_2_5": float(np.nanpercentile(values, 2.5)),
"ci_97_5": float(np.nanpercentile(values, 97.5)),
"replicates": reps, "unit": "paired source-video group resample",
})
return rows
def controlled_metrics(
model: CRG,
split: SplitData,
arrays: dict[str, np.ndarray],
scenarios: dict[str, np.ndarray],
device: torch.device,
batch_size: int,
model_name: str,
temperature: float,
seed: int = SEED + 552,
) -> tuple[list[dict[str, Any]], dict[str, dict[str, np.ndarray]]]:
rows = []
predictions = {}
for scenario, mask in scenarios.items():
scenario_seed = _scenario_seed(seed, split.name, scenario)
with fixed_torch_seed(scenario_seed, device):
metrics, prediction = evaluate(model, arrays, split, device, batch_size, masks=mask, temperature=temperature)
predictions[scenario] = prediction
rates = _missing_rate_summary(split.mask, mask)
additional_by_modality = np.nanmean(rates["additional_by_modality"], axis=0)
rate, mode = scenario.split("/", 1)
rows.append({"evaluation_split": split.name, "model": model_name, "rate_requested_per_selected_source": float(rate),
"mask_pattern": mode, "rate_realized_global": rates["additional_global"],
"rate_realized_additional_global": rates["additional_global"],
"rate_realized_additional_by_modality": json.dumps([None if not np.isfinite(x) else float(x) for x in additional_by_modality]),
"natural_missing_rate_global": rates["natural_global"],
"natural_missing_rate_by_modality": json.dumps(np.mean(rates["natural_by_modality"], axis=0).tolist()),
"rate_final_total_missing_global": rates["final_global"],
"rate_final_total_missing_by_modality": json.dumps(np.mean(rates["final_by_modality"], axis=0).tolist()),
"synchronous_no_observation_rate": float(np.mean(rates["synchronous_no_observation"])), **metrics})
print(f"validation mask {scenario}: additional={rates['additional_global']:.3f} final={rates['final_global']:.3f} "
f"macroF1={metrics['macro_f1']:.4f} MAE={metrics['regression_mae']:.4f}", flush=True)
return rows, predictions
def controlled_group_bootstrap(
split: SplitData,
predictions: dict[str, dict[str, dict[str, np.ndarray]]],
scenario_masks: dict[str, np.ndarray],
repeats: int,
seed: int,
) -> list[dict[str, Any]]:
"""Paired source-video bootstrap for each fixed mask and the MAE-rate AURC."""
if "C0" not in predictions:
raise ValueError("controlled bootstrap requires C0 predictions")
scenarios = list(predictions["C0"])
if any(set(model_predictions) != set(scenarios) for model_predictions in predictions.values()):
raise ValueError("all models must use identical controlled-mask scenarios")
if set(scenario_masks) != set(scenarios):
raise ValueError("controlled-mask audit and model predictions must use identical scenarios")
groups = np.unique(split.groups)
group_indices = {group: np.flatnonzero(split.groups == group) for group in groups}
metric_names = ("accuracy", "macro_f1", "mae", "rmse", "pearson")
keys = [(model, scenario, metric) for model in predictions for scenario in scenarios for metric in metric_names]
point = {}
draws = {key: [] for key in keys}
delta_draws = {key: [] for key in keys if key[0] != "C0"}
curve_modes = tuple(MASK_MODES)
curve_rates = np.asarray((0.0, 0.1, 0.3, 0.5, 0.7), dtype=np.float64)
curve_scenarios = {
mode: tuple("0.0/none" if rate == 0.0 else f"{rate:.1f}/{mode}" for rate in curve_rates)
for mode in curve_modes
}
for mode, curve in curve_scenarios.items():
if any(scenario not in scenarios for scenario in curve):
raise ValueError(f"missing fixed rate-curve scenario for {mode}")
curve_point = {}
curve_draws = {}
curve_actual_rates: dict[tuple[str, str], np.ndarray] = {}
additional_rate_by_sample: dict[str, np.ndarray] = {}
for scenario, mask in scenario_masks.items():
by_modality = _missing_rate_summary(split.mask, mask)["additional_by_modality"]
counts = np.isfinite(by_modality).sum(axis=1)
additional_rate_by_sample[scenario] = np.divide(
np.nansum(by_modality, axis=1), counts,
out=np.full(split.n, np.nan, dtype=np.float64), where=counts > 0,
)
for model in predictions:
for mode, curve in curve_scenarios.items():
key = (model, mode)
errors = [np.abs(split.regression_y - predictions[model][scenario]["predicted_score"]) for scenario in curve]
point_curve = np.asarray([float(error.mean()) for error in errors])
actual_rates = np.asarray([
0.0 if scenario == "0.0/none" else float(np.nanmean(additional_rate_by_sample[scenario]))
for scenario in curve
])
curve_actual_rates[model, mode] = actual_rates
order = np.argsort(actual_rates, kind="stable")
x = actual_rates[order]
y = point_curve[order]
unique_x, inverse = np.unique(x, return_inverse=True)
unique_y = np.asarray([y[inverse == index].mean() for index in range(len(unique_x))])
curve_point[key] = (float(np.trapezoid(unique_y, unique_x) / unique_x[-1])
if len(unique_x) > 1 and unique_x[-1] > 0.0 else float(point_curve[0]))
curve_draws[key] = []
def metric_values(model: str, scenario: str, indices: np.ndarray) -> dict[str, float]:
prediction = predictions[model][scenario]
y_class, y_score = split.class_y[indices], split.regression_y[indices]
predicted_class = prediction["predicted_class"][indices]
predicted_score = prediction["predicted_score"][indices]
return {
"accuracy": float(accuracy_score(y_class, predicted_class)),
"macro_f1": float(f1_score(y_class, predicted_class, labels=[0, 1, 2], average="macro", zero_division=0)),
"mae": float(mean_absolute_error(y_score, predicted_score)),
"rmse": float(np.sqrt(mean_squared_error(y_score, predicted_score))),
"pearson": float(pearsonr(y_score, predicted_score).statistic)
if np.std(y_score) and np.std(predicted_score) else float("nan"),
}
for model in predictions:
for scenario in scenarios:
values = metric_values(model, scenario, np.arange(split.n))
for metric, value in values.items():
point[(model, scenario, metric)] = value
within_model_mae_delta_draws = {
(model, scenario): [] for model in predictions for scenario in scenarios if scenario != "0.0/none"
}
rng = np.random.default_rng(seed)
for _ in range(repeats):
chosen = rng.choice(groups, size=len(groups), replace=True)
indices = np.concatenate([group_indices[group] for group in chosen])
replicate_values = {}
for model in predictions:
for scenario in scenarios:
values = metric_values(model, scenario, indices)
for metric, value in values.items():
key = (model, scenario, metric)
draws[key].append(value)
replicate_values[key] = value
for model in predictions:
if model == "C0":
continue
for scenario in scenarios:
for metric in metric_names:
key = (model, scenario, metric)
delta_draws[key].append(replicate_values[key] - replicate_values[("C0", scenario, metric)])
for model in predictions:
natural_mae = replicate_values[(model, "0.0/none", "mae")]
for scenario in scenarios:
if scenario != "0.0/none":
within_model_mae_delta_draws[(model, scenario)].append(
replicate_values[(model, scenario, "mae")] - natural_mae
)
for model in predictions:
for mode, curve in curve_scenarios.items():
curve_mae = [replicate_values[(model, scenario, "mae")] for scenario in curve]
sample_rates = np.asarray([
0.0 if scenario == "0.0/none" else float(np.nanmean(additional_rate_by_sample[scenario][indices]))
for scenario in curve
])
order = np.argsort(sample_rates, kind="stable")
x = sample_rates[order]
y = np.asarray(curve_mae)[order]
unique_x, inverse = np.unique(x, return_inverse=True)
unique_y = np.asarray([y[inverse == index].mean() for index in range(len(unique_x))])
auc = (float(np.trapezoid(unique_y, unique_x) / unique_x[-1])
if len(unique_x) > 1 and unique_x[-1] > 0.0 else float(curve_mae[0]))
curve_draws[(model, mode)].append(auc)
def interval(values: list[float]) -> tuple[float, float, float]:
samples = np.asarray(values, dtype=np.float64)
return float(np.nanmedian(samples)), float(np.nanpercentile(samples, 2.5)), float(np.nanpercentile(samples, 97.5))
rows = []
for model, scenario, metric in keys:
median, lower, upper = interval(draws[(model, scenario, metric)])
row = {"model": model, "scenario": scenario, "metric": metric,
"estimate": point[(model, scenario, metric)], "bootstrap_median": median,
"ci_2_5": lower, "ci_97_5": upper, "replicates": repeats,
"unit": "paired source-video group resample"}
if model != "C0":
delta_median, delta_lower, delta_upper = interval(delta_draws[(model, scenario, metric)])
row.update({"baseline": "C0", "delta_estimate": point[(model, scenario, metric)] - point[("C0", scenario, metric)],
"delta_bootstrap_median": delta_median, "delta_ci_2_5": delta_lower,
"delta_ci_97_5": delta_upper})
if scenario != "0.0/none" and metric == "mae":
natural = point[(model, "0.0/none", "mae")]
within_median, within_lower, within_upper = interval(within_model_mae_delta_draws[(model, scenario)])
row.update({"delta_to_natural_mae": point[(model, scenario, "mae")] - natural,
"delta_to_natural_bootstrap_median": within_median,
"delta_to_natural_ci_2_5": within_lower,
"delta_to_natural_ci_97_5": within_upper})
rows.append(row)
for model in predictions:
for mode in curve_modes:
median, lower, upper = interval(curve_draws[(model, mode)])
row = {"model": model, "scenario": f"MAE_rate_curve/{mode}", "metric": "AURC_MAE",
"estimate": curve_point[(model, mode)], "bootstrap_median": median,
"ci_2_5": lower, "ci_97_5": upper,
"curve_additional_rates_realized": json.dumps(curve_actual_rates[(model, mode)].tolist()),
"replicates": repeats, "unit": "paired source-video group resample"}
if model != "C0":
baseline_values = curve_draws[("C0", mode)]
deltas = np.asarray(curve_draws[(model, mode)]) - np.asarray(baseline_values)
delta_median, delta_lower, delta_upper = interval(deltas.tolist())
row.update({"baseline": "C0", "delta_estimate": curve_point[(model, mode)] - curve_point[("C0", mode)],
"delta_bootstrap_median": delta_median, "delta_ci_2_5": delta_lower,
"delta_ci_97_5": delta_upper})
rows.append(row)
return rows
def reencode_attachment3(device: torch.device) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
from data import restricted_load
files = sorted(ATTACHMENT3_ALIGNED.glob("附件3_*.pkl"), key=lambda p: int(p.stem.split("_")[-1]))
if len(files) != 30:
raise FileNotFoundError(f"expected 30 aligned attachment-3 files, found {len(files)} under {ATTACHMENT3_ALIGNED}")
file_ids = [path.stem for path in files]
if len(set(file_ids)) != len(file_ids):
raise ValueError("attachment-3 input file IDs are not unique")
bert = AutoModel.from_pretrained(TEXT_MODEL_ID, local_files_only=True).to(device).eval()
cases, audit = [], []
with torch.inference_mode():
for path in files:
case = load_attachment3_case(path)
ids = torch.from_numpy(case["input_ids"][None]).to(device)
attention = torch.from_numpy(case["attention_mask"][None].astype(np.int64)).to(device)
segments = torch.from_numpy(case["token_type_ids"][None]).to(device)
text = bert(input_ids=ids, attention_mask=attention, token_type_ids=segments).last_hidden_state[0].float().cpu().numpy()
text[~case["attention_mask"]] = 0.0
case_id = path.stem
cases.append({"case_id": case_id, "text": text, "audio": case["audio"], "vision": case["vision"], "text_mask": case["attention_mask"]})
audio_mask = np.any(case["audio"] != 0, axis=1)
vision_mask = np.any(case["vision"] != 0, axis=1)
audit.append({"case_id": case_id, "source_file": path.name,
"text_visible_steps": int(case["attention_mask"].sum()),
"audio_visible_steps": int(audio_mask.sum()), "vision_visible_steps": int(vision_mask.sum()),
"audio_missing_fraction": float(1.0 - audio_mask.mean()),
"vision_missing_fraction": float(1.0 - vision_mask.mean()),
"unknown_quality_flag": True, "labels_available": False,
"input_note": "official aligned_50 only; unaligned rows and raw transcript are not used",
"source_sha256": sha256(path)})
del bert
return cases, audit
def infer_attachment3(
model: CRG,
cases: list[dict[str, Any]],
fitted: dict[str, dict[str, np.ndarray]],
device: torch.device,
temperature: float,
prior_probs: np.ndarray,
magnitude_priors: np.ndarray,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
model.eval()
predictions, audit = [], []
for case in cases:
mask = np.stack((case["text_mask"], np.any(case["audio"] != 0, axis=1), np.any(case["vision"] != 0, axis=1)), axis=-1)[None]
xs_np = {}
for modality in MODALITIES:
values = case[modality].astype(np.float32)
values = np.clip((values - fitted[modality]["mean"]) / fitted[modality]["std"], -10.0, 10.0)
values[~mask[0, :, MODALITIES.index(modality)]] = 0.0
xs_np[modality] = values[None]
low_information = not bool(mask.any())
if low_information:
p = np.asarray(prior_probs, dtype=np.float64)
p = p / p.sum()
max_probability = float(p.max())
predicted_class = next(c for c in (1, 0, 2) if math.isclose(float(p[c]), max_probability, rel_tol=0.0, abs_tol=1e-12))
score = 0.0
beta = np.full((1, 1, 2, 2), np.nan, np.float32)
if predicted_class != 1:
sign_index = 0 if predicted_class == 0 else 1
magnitude = float(betaincinv(magnitude_priors[sign_index, 0], magnitude_priors[sign_index, 1], 0.5))
score = (-3.0 if predicted_class == 0 else 3.0) * magnitude
beta[0, 0] = magnitude_priors
ppaths = p[None, None, :]
else:
arrays = {m: xs_np[m] for m in MODALITIES}
xs, observed = to_device_batch(arrays, mask, np.asarray([0]), device)
with torch.inference_mode():
out = model(xs, observed, paths=16 if model.use_joint_draws else 1, joint_draws=model.use_joint_draws)
ppaths = out["class_probs_by_path"].cpu().numpy()[:, 0:1]
beta = out["beta_params"].cpu().numpy()[:, 0:1]
p, classes, scores = _decode_mixture(ppaths, beta, temperature)
p, predicted_class, score = p[0], int(classes[0]), float(scores[0])
interval_temperature = 1.0 if low_information else temperature
interval_lower, interval_upper = _predictive_intervals(ppaths, beta, interval_temperature)
variance_components = _trajectory_variance_components(ppaths, beta)
calibrated_moments = _calibrated_mixture_moments(ppaths, beta, interval_temperature)
row = {"case_id": case["case_id"], "predicted_class": predicted_class,
"predicted_class_name": ("negative", "neutral", "positive")[predicted_class],
"predicted_sentiment": float(score), "p_negative": float(p[0]),
"p_neutral": float(p[1]), "p_positive": float(p[2]),
"interval_90_lower": float(interval_lower[0]), "interval_90_upper": float(interval_upper[0]),
"predictive_variance_mean_uncalibrated": float(variance_components[0][0]),
"within_trajectory_variance": float(variance_components[1][0]),
"between_trajectory_variance": float(variance_components[2][0]),
"predictive_mean_calibrated": float(calibrated_moments[0][0]),
"predictive_variance_calibrated": float(calibrated_moments[1][0]),
"beta_negative_alpha": float(beta[0, 0, 0, 0]), "beta_negative_beta": float(beta[0, 0, 0, 1]),
"beta_positive_alpha": float(beta[0, 0, 1, 0]), "beta_positive_beta": float(beta[0, 0, 1, 1]),
"low_information_prior_fallback": low_information,
"output_note": "unlabeled attachment-3 case; no accuracy/F1 is defined"}
predictions.append(row)
audit.append({"case_id": case["case_id"], "visible_text_steps": int(mask[0, :, 0].sum()),
"visible_audio_steps": int(mask[0, :, 1].sum()), "visible_vision_steps": int(mask[0, :, 2].sum()),
"low_information_prior_fallback": low_information,
"calibration_temperature": interval_temperature,
"interval_90_lower": float(interval_lower[0]), "interval_90_upper": float(interval_upper[0]),
"predictive_variance_mean_uncalibrated": float(variance_components[0][0]),
"predictive_variance_calibrated": float(calibrated_moments[1][0])})
return predictions, audit
def validate_attachment3_predictions(expected_ids: list[str], predictions: list[dict[str, Any]]) -> None:
"""Check the unlabeled submission contract and sign/strength consistency."""
predicted_ids = [str(row["case_id"]) for row in predictions]
if len(predictions) != len(expected_ids) or len(set(expected_ids)) != len(expected_ids):
raise ValueError("attachment-3 prediction count differs from unique expected IDs")
if len(set(predicted_ids)) != len(predicted_ids) or set(predicted_ids) != set(expected_ids):
raise ValueError("attachment-3 predictions must contain every expected ID exactly once")
for row in predictions:
predicted_class = int(row["predicted_class"])
score = float(row["predicted_sentiment"])
probabilities = np.asarray([row["p_negative"], row["p_neutral"], row["p_positive"]], dtype=np.float64)
if predicted_class not in (0, 1, 2) or not np.isfinite(score) or not -3.0 <= score <= 3.0:
raise ValueError(f"invalid attachment-3 class/strength for {row['case_id']}")
if (predicted_class == 1 and score != 0.0) or (predicted_class == 0 and not score < 0.0) or (
predicted_class == 2 and not score > 0.0
):
raise ValueError(f"attachment-3 class/strength polarity mismatch for {row['case_id']}")
if not np.isfinite(probabilities).all() or np.any(probabilities < 0.0) or not np.isclose(probabilities.sum(), 1.0, atol=1e-6):
raise ValueError(f"invalid attachment-3 class probabilities for {row['case_id']}")
low, high = float(row["interval_90_lower"]), float(row["interval_90_upper"])
if not np.isfinite([low, high]).all() or low > high or low < -3.0 or high > 3.0:
raise ValueError(f"invalid attachment-3 prediction interval for {row['case_id']}")
def main() -> None:
global DELTA_U, RESULTS
parser = argparse.ArgumentParser()
parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="aligned_50")
parser.add_argument("--output-dir", type=Path, default=None,
help="Write this run to a new directory instead of the default results directory")
parser.add_argument("--epochs", type=int, default=12)
parser.add_argument("--imputer-epochs", type=int, default=8)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--patience", type=int, default=3)
parser.add_argument("--seed", type=int, default=SEED)
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
parser.add_argument("--bootstrap-repeats", type=int, default=1000)
parser.add_argument("--skip-attachment3", action="store_true")
args = parser.parse_args()
if args.input_version == "unaligned_50" and not args.skip_attachment3:
parser.error("unaligned attachment 3 has no numerical text or trusted lengths; use --skip-attachment3 for the training comparison")
RESULTS = args.output_dir or Q2_DIR / ("results_unaligned" if args.input_version == "unaligned_50" else "results")
if args.output_dir is not None and RESULTS.exists() and any(RESULTS.iterdir()):
parser.error(f"refusing to overwrite non-empty result directory: {RESULTS}")
seed_everything(args.seed)
rng = np.random.default_rng(args.seed)
device = torch.device(args.device)
RESULTS.mkdir(parents=True, exist_ok=True)
if device.type == "cuda":
print(f"device={device} ({torch.cuda.get_device_name(device)})", flush=True)
input_path = ROOT / "E题数据" / "附件2-数据集特征文件" / f"{args.input_version}.pkl"
official = load_official_splits(input_path, version=args.input_version)
overlaps = assert_group_disjoint(official)
fit, heldout_train = split_calibration(official["train"], args.seed)
reliability_validation, temperature_calibration = split_calibration(heldout_train, args.seed + 1, fraction=0.5)
reliability_validation.name = "reliability_validation"
temperature_calibration.name = "temperature_calibration"
if (set(fit.groups) & set(reliability_validation.groups)
or set(fit.groups) & set(temperature_calibration.groups)
or set(reliability_validation.groups) & set(temperature_calibration.groups)):
raise AssertionError("fit, reliability-selection, and temperature-calibration videos must be disjoint")
DELTA_U = label_resolution_from_train(fit.regression_y)
magnitude_priors = fit_magnitude_priors(fit.regression_y)
class_counts_fit = np.bincount(fit.class_y, minlength=3).astype(np.float64)
class_prior_probs = (class_counts_fit + 1.0) / (class_counts_fit.sum() + 3.0)
# Freeze the exact validation masks before fitting any model.
validation_scenarios = make_scenarios(official["valid"], args.seed + 909)
reliability_scenarios = make_reliability_scenarios(reliability_validation, args.seed + 906)
print("official splits:", {k: (v.n, len(np.unique(v.groups))) for k, v in official.items()},
"fit/reliability_validation/temperature_calibration:",
(fit.n, reliability_validation.n, temperature_calibration.n), "group_overlap:", overlaps, flush=True)
fitted = fit_preprocessor(fit)
transformed = {name: transform_split(split, fitted) for name, split in official.items()}
transformed["fit"] = transform_split(fit, fitted)
transformed["reliability_validation"] = transform_split(reliability_validation, fitted)
transformed["temperature_calibration"] = transform_split(temperature_calibration, fitted)
np.savez_compressed(RESULTS / "preprocessor.npz", **{f"{m}_{k}": v for m, stats in fitted.items() for k, v in stats.items()})
imputer = StructuredGaussianImputer(INPUT_DIMS).to(device)
imputer_history = fit_imputer(imputer, transformed["fit"], fit, device, args.imputer_epochs, args.batch_size, args.seed + 1)
torch.save({k: v.detach().cpu() for k, v in imputer.state_dict().items()}, RESULTS / "structured_imputer.pt")
teacher = _make_variant("C6", imputer).to(device)
teacher, history_teacher = _fit_neural(
teacher, "teacher", fit, official["valid"], transformed, device, args.epochs,
args.batch_size, args.patience, np.random.default_rng(args.seed + 2),
selection_split=reliability_validation, selection_arrays=transformed["reliability_validation"],
selection_scenarios={"0.0/natural": reliability_validation.mask.copy()},
)
teacher_reliability, teacher_tuning_rows = tune_reliability_hparams(
teacher, transformed["reliability_validation"], reliability_validation, reliability_scenarios,
device, args.batch_size, "teacher", seed=args.seed + 551,
)
torch.save({k: v.detach().cpu() for k, v in teacher.state_dict().items()}, RESULTS / "teacher.pt")
ablation_rows: list[dict[str, Any]] = []
history = list(imputer_history) + history_teacher
reliability_tuning_rows = list(teacher_tuning_rows)
group_risk_tuning_rows: list[dict[str, Any]] = []
models: dict[str, CRG] = {}
_, c0_state = fit_c0(fit, official["valid"], transformed)
calibrate_c0_interval(c0_state, temperature_calibration, transformed["temperature_calibration"])
_, c0_cal_pred = evaluate_c0(c0_state, temperature_calibration, transformed["temperature_calibration"])
c0_temperature = fit_temperature(c0_cal_pred["probabilities"], temperature_calibration.class_y)
c0_metrics, c0_valid_pred = evaluate_c0(c0_state, official["valid"], transformed["valid"], c0_temperature)
c0_metrics["temperature"] = c0_temperature
ablation_rows.append({"model": "C0", **c0_metrics, "description": "observed mean/std + masks + maximum gap; logistic/ridge"})
best_model_name: str | None = None
best_valid_loss = float("inf")
temperatures = {"C0": c0_temperature}
selected_group_risk: dict[str, tuple[float, float]] = {}
validation_predictions: dict[str, dict[str, np.ndarray]] = {"C0": c0_valid_pred}
definitions = {
"C1": "masked BiGRU; no posterior imputation, explicit reliability or source gate",
"C2": "exact Gaussian posterior mean; no joint trajectory integral",
"C3": "joint trajectory integral plus final reliability/content fusion gate",
"C4": "C3 plus bounded cross-time source attention and null source",
"C5": "C4 plus reliability-modulated BiGRU update",
"C6": "C5 plus optional rank-4 CP residual",
"C6_no_distance": "C6 with uncertainty retained but both distance/span reliability penalties fixed to zero",
"C6_no_reconstruction": "C6 trained without the auxiliary hidden-feature reconstruction loss",
"C6_pointmask": "C6 trained with independent point masking instead of contiguous spans",
"C7_distill": "C6 plus entropy/retention-weighted teacher distillation only",
"C7_group": "C6 plus smooth worst-group risk only",
}
diagnostic_ablation_names = {"C6_no_distance", "C6_no_reconstruction", "C6_pointmask"}
model_names = ("C1", "C2", "C3", "C4", "C5", "C6", *sorted(diagnostic_ablation_names), "C7_distill", "C7_group")
no_distance_candidates = tuple(value for value in RELIABILITY_CANDIDATES if value[2] == 0.0 and value[3] == 0.0)
for name in model_names:
if name == "C7_group":
(variant, rows, reliability_rows, risk_rows, selected_reliability,
selected_risk, _) = tune_group_risk_model(
imputer, fit, official["valid"], transformed,
reliability_validation, transformed["reliability_validation"], reliability_scenarios,
device, args.epochs, args.batch_size, args.patience, args.seed + 303,
)
reliability_tuning_rows.extend(reliability_rows)
group_risk_tuning_rows.extend(risk_rows)
selected_group_risk[name] = selected_risk
else:
base_name = "C6" if name in diagnostic_ablation_names else name
variant = _make_variant(base_name, imputer)
if name == "C6_no_distance":
_set_reliability_hparams(variant, (DEFAULT_RELIABILITY[0], DEFAULT_RELIABILITY[1], 0.0, 0.0))
kd_teacher = teacher if name == "C7_distill" else None
mask_kind = "point" if name == "C6_pointmask" else "continuous"
variant, rows = _fit_neural(
variant, name, fit, official["valid"], transformed, device, args.epochs,
args.batch_size, args.patience, np.random.default_rng(args.seed + 303),
teacher=kd_teacher,
selection_split=reliability_validation,
selection_arrays=transformed["reliability_validation"],
selection_scenarios=reliability_scenarios,
mask_kind=mask_kind,
use_reconstruction=name != "C6_no_reconstruction",
)
selected_reliability, tuning_rows = tune_reliability_hparams(
variant, transformed["reliability_validation"], reliability_validation, reliability_scenarios,
device, args.batch_size, name, seed=args.seed + 551,
candidate_values=no_distance_candidates if name == "C6_no_distance" else RELIABILITY_CANDIDATES,
)
reliability_tuning_rows.extend(tuning_rows)
selected_group_risk[name] = (float("nan"), float("nan"))
history.extend(rows)
models[name] = variant
_, calibration_pred = evaluate(variant, transformed["temperature_calibration"], temperature_calibration,
device, args.batch_size)
model_temperature = fit_temperature(calibration_pred["probabilities"], temperature_calibration.class_y)
temperatures[name] = model_temperature
raw_metrics, _ = evaluate(variant, transformed["valid"], official["valid"], device, args.batch_size)
metrics, valid_prediction = evaluate(variant, transformed["valid"], official["valid"], device, args.batch_size,
temperature=model_temperature)
validation_predictions[name] = valid_prediction
metrics["validation_selection_loss"] = raw_metrics["selection_nll"]
metrics["temperature"] = model_temperature
metrics.update({"rho_imp": selected_reliability[0], "lambda_u": selected_reliability[1],
"lambda_gap": selected_reliability[2], "lambda_span": selected_reliability[3]})
metrics["lambda_group"], metrics["group_temperature"] = selected_group_risk[name]
ablation_rows.append({"model": name, **metrics, "description": definitions[name]})
if name not in diagnostic_ablation_names and raw_metrics["selection_nll"] < best_valid_loss:
best_valid_loss, best_model_name = raw_metrics["selection_nll"], name
write_csv(RESULTS / "ablation_validation.csv", ablation_rows)
write_csv(RESULTS / "reliability_hparam_tuning.csv", reliability_tuning_rows)
write_csv(RESULTS / "group_risk_tuning.csv", group_risk_tuning_rows)
write_csv(RESULTS / "validation_group_bootstrap_deltas.csv",
paired_group_bootstrap_deltas(official["valid"], validation_predictions, args.bootstrap_repeats, args.seed + 88))
if best_model_name is None:
raise RuntimeError("no neural ablation candidate completed")
best_model = models[best_model_name]
# Calibration is on a group-held-out slice of official training data, never on official test.
temperature = temperatures[best_model_name]
valid_metrics, valid_pred = evaluate(best_model, transformed["valid"], official["valid"], device, args.batch_size, temperature=temperature)
test_metrics, test_pred = evaluate(
best_model, transformed["test"], official["test"], device, args.batch_size,
temperature=temperature, collect_gate_diagnostics=True,
)
torch.save({k: v.detach().cpu() for k, v in best_model.state_dict().items()}, RESULTS / "crg_student.pt")
(RESULTS / "validation_metrics.json").write_text(json.dumps({**valid_metrics, "selected_model": best_model_name, "temperature": temperature}, indent=2), encoding="utf-8")
(RESULTS / "test_metrics.json").write_text(json.dumps({**test_metrics, "selected_model": best_model_name, "temperature": temperature}, indent=2), encoding="utf-8")
test_rows = []
for i, sample_id in enumerate(official["test"].ids):
p = test_pred["probabilities"][i]
test_rows.append({"sample_id": sample_id, "source_video_id": official["test"].groups[i],
"true_class": int(official["test"].class_y[i]), "predicted_class": int(test_pred["predicted_class"][i]),
"true_sentiment": float(official["test"].regression_y[i]), "predicted_sentiment": float(test_pred["predicted_score"][i]),
"p_negative": float(p[0]), "p_neutral": float(p[1]), "p_positive": float(p[2]),
"interval_90_lower": float(test_pred["interval_lower"][i]),
"interval_90_upper": float(test_pred["interval_upper"][i]),
"predictive_variance_mean_uncalibrated": float(test_pred["predictive_variance_uncalibrated"][i]),
"within_trajectory_variance": float(test_pred["within_trajectory_variance"][i]),
"between_trajectory_variance": float(test_pred["between_trajectory_variance"][i])})
write_csv(RESULTS / "test_predictions.csv", test_rows)
write_csv(RESULTS / "test_gate_diagnostics.csv", gate_diagnostic_rows(official["test"], test_pred))
write_csv(RESULTS / "group_bootstrap_ci.csv",
group_bootstrap(official["test"], test_pred, args.bootstrap_repeats, args.seed + 44))
write_csv(RESULTS / "controlled_mask_audit.csv",
mask_audit_rows(official["valid"], validation_scenarios, args.seed + 909))
controlled_rows, c0_scenario_predictions = controlled_c0(
c0_state, official["valid"], transformed["valid"], validation_scenarios, c0_temperature,
)
controlled_predictions: dict[str, dict[str, dict[str, np.ndarray]]] = {"C0": c0_scenario_predictions}
for name, candidate in models.items():
candidate_rows, candidate_predictions = controlled_metrics(
candidate, official["valid"], transformed["valid"], validation_scenarios,
device, args.batch_size, name, temperatures[name], seed=args.seed + 552,
)
controlled_rows.extend(candidate_rows)
controlled_predictions[name] = candidate_predictions
write_csv(RESULTS / "controlled_missingness.csv", controlled_rows)
write_csv(RESULTS / "controlled_group_bootstrap.csv",
controlled_group_bootstrap(official["valid"], controlled_predictions, validation_scenarios,
args.bootstrap_repeats, args.seed + 553))
write_csv(RESULTS / "training_history.csv", history)
attachment_count = 0
if not args.skip_attachment3:
cases, attachment_audit = reencode_attachment3(device)
attachment_predictions, inference_audit = infer_attachment3(best_model, cases, fitted, device, temperature, class_prior_probs, magnitude_priors)
validate_attachment3_predictions([case["case_id"] for case in cases], attachment_predictions)
write_csv(RESULTS / "attachment3_predictions.csv", attachment_predictions)
write_csv(RESULTS / "attachment3_audit.csv", [dict(a, **next(x for x in inference_audit if x["case_id"] == a["case_id"])) for a in attachment_audit])
attachment_count = len(cases)
manifest = {
"seed": args.seed,
"text_encoder": ("official precomputed text field; encoder revision not supplied"
if args.input_version == "unaligned_50" else TEXT_MODEL_ID),
"training_configuration": {"student_epoch_limit": args.epochs, "imputer_epochs": args.imputer_epochs,
"batch_size": args.batch_size, "early_stopping_patience": args.patience,
"device": str(device),
"device_name": torch.cuda.get_device_name(device) if device.type == "cuda" else "CPU",
"optimizer": "AdamW", "student_learning_rate": 3e-4,
"student_weight_decay": 1e-3, "imputer_learning_rate": 3e-4,
"imputer_weight_decay": 1e-4,
"early_stopping_metric": "mean untempered selection_nll over fixed group-disjoint internal training scenarios",
"inner_selection_scenarios": list(reliability_scenarios),
"inner_selection_source_video_groups": int(len(np.unique(reliability_validation.groups)))},
"training_input": str(input_path.relative_to(ROOT)),
"input_version": args.input_version,
"q1_alignment_adapter": {name: split.alignment_audit for name, split in official.items()}
if args.input_version == "unaligned_50" else None,
"training_sha256": sha256(input_path),
"official_group_overlap": overlaps,
"official_splits": {name: {"n": split.n, "source_video_groups": int(len(np.unique(split.groups)))} for name, split in official.items()},
"internal_train_holdouts": {
"fit": {"n": fit.n, "video_groups": int(len(np.unique(fit.groups)))},
"reliability_selection": {"n": reliability_validation.n, "video_groups": int(len(np.unique(reliability_validation.groups)))},
"temperature_calibration": {"n": temperature_calibration.n, "video_groups": int(len(np.unique(temperature_calibration.groups)))},
"all_group_disjoint": True,
},
"feature_standardization": "fit-only observed rows, per-dimension; fixed for valid/test/attachment3",
"missing_mask": ("official text attention and source lengths plus row observation; normalized-progress overlap preserves empty bins; q*=1 only where visible, J_Q=0"
if args.input_version == "unaligned_50" else
"row-level all-zero convention; q*=1 only where currently visible, J_Q=0; hidden metadata is zeroed"),
"observation_quality": {"quality_score_fields_present": False, "quality_available_flag_present": False,
"fallback": "q*=1 and J_Q=0 for visible rows; R_eff=R",
"quality_noise_mapping_ablation": f"not identifiable on {args.input_version} because no row quality score varies"},
"imputer": {"type": "structured linear Gaussian shared-private state space", "state_dims": {"shared": 8, "private_each": 4},
"posterior": "block-tridiagonal equivalent Kalman information filter + RTS smoother",
"sampling": "joint latent trajectories and missing emissions; observed features copied exactly",
"fit_objective": "train-only observed Gaussian marginal likelihood including log determinants",
"epochs": args.imputer_epochs, "frozen_before_teacher_student": True},
"architecture": {"projection": 32, "bigru_hidden_each_direction": 16, "cross_source_layers": 1,
"cross_time_read": True, "rank": 4, "reliability_gru": "directional hidden decay; reset applied before candidate map; update gate multiplied by rho",
"final_gate": "rho times bounded content score plus positive null prior",
"output": "neutral point mass plus sign-specific Beta magnitudes; K-path probabilities mixed before decoding"},
"reliability_hyperparameters": {"selected_per_model_on": "group-disjoint internal training reliability-validation slice",
"candidate_values": [list(v) for v in RELIABILITY_CANDIDATES],
"validation_scenarios": list(reliability_scenarios),
"selected_by_model": {row["model"]: [row["rho_imp"], row["lambda_u"], row["lambda_gap"], row["lambda_span"]]
for row in reliability_tuning_rows if row.get("selected")}},
"group_risk_hyperparameters": {
"selection": "lambda_group and group_temperature jointly selected with reliability hyperparameters on fixed group-disjoint internal training scenarios",
"candidate_values": [list(v) for v in GROUP_RISK_CANDIDATES],
"selected": list(selected_group_risk["C7_group"]),
"selection_split": reliability_validation.name,
},
"loss": {"supervision": "negative log mixture of Beta interval masses plus scaled Huber mean term",
"delta_u": DELTA_U, "delta_u_source": "half the minimum positive spacing of nonzero absolute labels in fit only",
"lambda_y": LAMBDA_Y, "lambda_distill": LAMBDA_DISTILL,
"lambda_reconstruction": LAMBDA_RECON,
"lambda_group_default": LAMBDA_GROUP, "group_temperature_default": GROUP_TEMPERATURE,
"selected_group_risk": list(selected_group_risk["C7_group"]),
"distill_temperature": DISTILL_TEMPERATURE, "distill_retention_exponent": 1.0,
"imputer_regularization": {"emission_l2": LAMBDA_EMISSION, "transition_l2": LAMBDA_TRANSITION},
"group_and_distill_separate": True},
"calibration": {"method": "temperature scaling on a group-disjoint internal official-train holdout, separated from reliability selection",
"temperature": temperature, "valid_used_for_selection": True, "test_used_for_selection_or_calibration": False},
"selected_model": best_model_name,
"attachment3_low_information_priors": {"class_probability_method": "fit counts + one pseudocount per class",
"class_probability_values": class_prior_probs.tolist(),
"negative_beta": magnitude_priors[0].tolist(),
"positive_beta": magnitude_priors[1].tolist()},
"ablation_definitions": definitions,
"masking": {"rates": MASK_RATES, "patterns": list(MASK_MODES), "preserve_at_least_fraction_per_selected_modality": 0.2,
"controlled_sweep_split": "official validation", "identical_masks_across_models": True,
"scenario_count": len(validation_scenarios), "scenario_seed": args.seed + 909,
"training_mask_rng_seed": args.seed + 303,
"reliability_scenario_seed": args.seed + 906,
"controlled_torch_sampling_seed": args.seed + 552,
"paired_control_bootstrap_seed": args.seed + 553,
"mask_audit_file": "controlled_mask_audit.csv",
"additional_one_factor_controls": ["modality T/A/V and combinations", "start/middle/end", "one-long/multiple-short", "sync/partial/async"],
"semantic_position_control": f"not run: {args.input_version} does not provide audited semantic boundary indices; raw text is prohibited in student inputs"},
"final_test_metrics": test_metrics,
"test_gate_diagnostics_file": "test_gate_diagnostics.csv",
"attachment3_cases": attachment_count,
"attachment3_labeled_metrics": None,
"completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
}
(RESULTS / "run_manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
print("Q2 complete:", json.dumps({"selected_model": best_model_name, "test_accuracy": test_metrics["accuracy"],
"test_macro_f1": test_metrics["macro_f1"], "test_mae": test_metrics["regression_mae"],
"temperature": temperature, "attachment3_cases": attachment_count}, ensure_ascii=False), flush=True)
if __name__ == "__main__":
main()