2062 lines
111 KiB
Python
2062 lines
111 KiB
Python
"""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()
|