1297 lines
71 KiB
Python
1297 lines
71 KiB
Python
from __future__ import annotations
|
||
|
||
import argparse
|
||
import csv
|
||
import hashlib
|
||
import itertools
|
||
import json
|
||
import math
|
||
import time
|
||
from collections import defaultdict
|
||
from pathlib import Path
|
||
from typing import Any, Sequence
|
||
|
||
import numpy as np
|
||
import torch
|
||
from sklearn.metrics import accuracy_score, confusion_matrix, f1_score, mean_absolute_error, mean_squared_error, recall_score
|
||
from torch import nn
|
||
|
||
from ...data_paths import PROJECT_ROOT
|
||
from ...model.ati_ho import ATIHOModel
|
||
from ...model.ati_ho_config import CONFIGS
|
||
from ...q2.deep_learning.q2.data import MODALITIES, RobustStats, Split
|
||
from ...q2.deep_learning.q2.mofe import MixtureOfFusionExperts
|
||
from ...q2.deep_learning.q2.train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP
|
||
from ..run_experiments import _read_attachment4
|
||
from .attribution import ensemble_forward, exact_shapley_audit
|
||
from .audit import structural_audit
|
||
from .owen import fidelity_audit_one, hierarchical_owen_one
|
||
from .train import (
|
||
BATCH_SIZE,
|
||
EXPERIMENT_ROOT,
|
||
MODEL_SEEDS,
|
||
SCALER_PATH,
|
||
_metric_row,
|
||
_save_csv,
|
||
load_training_data,
|
||
)
|
||
|
||
|
||
RESULTS_ROOT = EXPERIMENT_ROOT / "results" / "ati_ho"
|
||
SUBMIT_OUTPUT = PROJECT_ROOT / "output" / "q3" / "ati_ho"
|
||
CLASS_NAMES = ("negative", "neutral", "positive")
|
||
PAIR_NAMES = ("TA", "TV", "AV")
|
||
BOOTSTRAP_REPLICATES = 1000
|
||
BOOTSTRAP_SEED = 20260925
|
||
|
||
|
||
def _read_csv(path: Path) -> list[dict[str, str]]:
|
||
if not path.is_file():
|
||
return []
|
||
with path.open("r", newline="", encoding="utf-8-sig") as stream:
|
||
return list(csv.DictReader(stream))
|
||
|
||
|
||
def _write_json(path: Path, payload: Any) -> None:
|
||
path.parent.mkdir(parents=True, exist_ok=True)
|
||
path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False) + "\n", encoding="utf-8")
|
||
|
||
|
||
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 _checkpoint_path(method: str, seed: int) -> Path:
|
||
return EXPERIMENT_ROOT / "models" / method / f"seed_{seed}" / "model_best.pt"
|
||
|
||
|
||
def _load_ensemble(method: str, seeds: Sequence[int], dims: tuple[int, int, int], device: torch.device) -> list[nn.Module]:
|
||
models: list[nn.Module] = []
|
||
for seed in seeds:
|
||
path = _checkpoint_path(method, seed)
|
||
if not path.is_file():
|
||
raise FileNotFoundError(path)
|
||
state = torch.load(path, map_location=device, weights_only=False)
|
||
if (
|
||
tuple(state.get("dims", ())) != dims
|
||
or int(state.get("seed", -1)) != seed
|
||
or state.get("method") != method
|
||
):
|
||
raise ValueError(f"incompatible checkpoint metadata: {path}")
|
||
if method in CONFIGS:
|
||
if state.get("config") != CONFIGS[method].to_dict():
|
||
raise ValueError(f"ATI configuration mismatch: {path}")
|
||
model: nn.Module = ATIHOModel(dims, CONFIGS[method]).to(device)
|
||
elif method == EARLYCONCAT:
|
||
from ...q2.deep_learning.q2.models import AlignedFusionModel
|
||
|
||
model = AlignedFusionModel("concat", dims=dims).to(device)
|
||
elif method == MOFE7_MLP:
|
||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||
else:
|
||
raise ValueError(f"unknown model {method}")
|
||
model.load_state_dict(state["state_dict"], strict=True)
|
||
model.eval()
|
||
models.append(model)
|
||
return models
|
||
|
||
|
||
@torch.inference_mode()
|
||
def _predict_ensemble(
|
||
models: Sequence[nn.Module],
|
||
split: Split,
|
||
masks: np.ndarray,
|
||
device: torch.device,
|
||
*,
|
||
details: bool = False,
|
||
batch_size: int = 64,
|
||
) -> dict[str, np.ndarray]:
|
||
keys = ["logits", "probabilities", "intensity"]
|
||
if details:
|
||
keys.extend(("params", "baseline", "main_effects", "pair_effects"))
|
||
values: dict[str, list[np.ndarray]] = {key: [] for key in keys}
|
||
for start in range(0, split.n, batch_size):
|
||
end = min(split.n, start + batch_size)
|
||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||
mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||
output = ensemble_forward(models, xs, mask, details=details)
|
||
for key in keys:
|
||
if key in output:
|
||
values[key].append(output[key].detach().cpu().numpy())
|
||
return {key: np.concatenate(rows, axis=0) for key, rows in values.items() if rows}
|
||
|
||
|
||
def _metric_subset(
|
||
split: Split,
|
||
predictions: dict[str, np.ndarray],
|
||
indices: np.ndarray,
|
||
) -> dict[str, float]:
|
||
y_cls = split.y_cls[indices]
|
||
y_reg = split.y_reg[indices]
|
||
logits = predictions["logits"][indices]
|
||
intensity = predictions["intensity"][indices]
|
||
probs = predictions["probabilities"][indices]
|
||
pred_class = logits.argmax(axis=-1)
|
||
recall = recall_score(y_cls, pred_class, labels=[0, 1, 2], average=None, zero_division=0)
|
||
pearson = float(np.corrcoef(y_reg, intensity)[0, 1]) if np.std(y_reg) > 0 and np.std(intensity) > 0 else 0.0
|
||
one_hot = np.eye(3)[y_cls]
|
||
return {
|
||
"accuracy": float(accuracy_score(y_cls, pred_class)),
|
||
"macro_f1": float(f1_score(y_cls, pred_class, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||
"weighted_f1": float(f1_score(y_cls, pred_class, average="weighted", zero_division=0)),
|
||
"negative_recall": float(recall[0]),
|
||
"neutral_recall": float(recall[1]),
|
||
"positive_recall": float(recall[2]),
|
||
"mae": float(mean_absolute_error(y_reg, intensity)),
|
||
"rmse": float(math.sqrt(mean_squared_error(y_reg, intensity))),
|
||
"pearson": pearson,
|
||
"brier_multiclass": float(np.mean(np.sum((probs - one_hot) ** 2, axis=-1))),
|
||
}
|
||
|
||
|
||
def _cluster_bootstrap(
|
||
split: Split,
|
||
prediction_by_method: dict[str, dict[str, np.ndarray]],
|
||
selected_method: str,
|
||
) -> list[dict[str, Any]]:
|
||
group_rows: dict[str, list[int]] = defaultdict(list)
|
||
for index, sample_id in enumerate(split.ids):
|
||
group_rows[sample_id.split("$_$", 1)[0]].append(index)
|
||
groups = np.asarray(sorted(group_rows))
|
||
group_map = {key: np.asarray(value, dtype=np.int64) for key, value in group_rows.items()}
|
||
rng = np.random.default_rng(BOOTSTRAP_SEED)
|
||
metrics = ("macro_f1", "mae", "pearson", "accuracy")
|
||
rows: list[dict[str, Any]] = []
|
||
for baseline in (EARLYCONCAT, MOFE7_MLP):
|
||
point_a = _metric_subset(split, prediction_by_method[selected_method], np.arange(split.n))
|
||
point_b = _metric_subset(split, prediction_by_method[baseline], np.arange(split.n))
|
||
draws: dict[str, list[float]] = {metric: [] for metric in metrics}
|
||
for _ in range(BOOTSTRAP_REPLICATES):
|
||
selected_groups = rng.choice(groups, size=len(groups), replace=True)
|
||
index = np.concatenate([group_map[group] for group in selected_groups])
|
||
a = _metric_subset(split, prediction_by_method[selected_method], index)
|
||
b = _metric_subset(split, prediction_by_method[baseline], index)
|
||
for metric in metrics:
|
||
delta = a[metric] - b[metric]
|
||
if metric == "mae":
|
||
delta = b[metric] - a[metric]
|
||
draws[metric].append(delta)
|
||
for metric in metrics:
|
||
values = np.asarray(draws[metric], dtype=np.float64)
|
||
delta = point_a[metric] - point_b[metric]
|
||
if metric == "mae":
|
||
delta = point_b[metric] - point_a[metric]
|
||
rows.append(
|
||
{
|
||
"comparison": f"{selected_method} vs {baseline}",
|
||
"metric": metric,
|
||
"delta_positive_favors_ATI_HO": float(delta),
|
||
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
|
||
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
|
||
"replicates": BOOTSTRAP_REPLICATES,
|
||
"bootstrap_unit": "source video_id",
|
||
"groups": len(groups),
|
||
"seed": BOOTSTRAP_SEED,
|
||
}
|
||
)
|
||
return rows
|
||
|
||
|
||
def _summary_rows(validation_rows: list[dict[str, str]], methods: Sequence[str]) -> list[dict[str, Any]]:
|
||
latest: dict[tuple[str, str, str], dict[str, str]] = {}
|
||
for row in validation_rows:
|
||
key = (row["method"], row["seed"], row["scenario"])
|
||
latest[key] = row
|
||
metrics = (
|
||
"accuracy", "macro_f1", "weighted_f1", "negative_recall", "neutral_recall", "positive_recall",
|
||
"mae", "rmse", "pearson", "ece_15bin", "brier_multiclass",
|
||
)
|
||
rows: list[dict[str, Any]] = []
|
||
for method in methods:
|
||
seeds = sorted({key[1] for key in latest if key[0] == method and key[2] == "clean"}, key=int)
|
||
for scenario in ("clean", "0.0/none", "0.3/single", "0.3/sync", "0.5/async"):
|
||
matched = [latest[(method, seed, scenario)] for seed in seeds if (method, seed, scenario) in latest]
|
||
if not matched:
|
||
continue
|
||
row: dict[str, Any] = {"method": method, "scenario": scenario, "seeds": len(matched), "seed_values": ";".join(x["seed"] for x in matched)}
|
||
for metric in metrics:
|
||
values = np.asarray([float(item[metric]) for item in matched], dtype=np.float64)
|
||
row[f"{metric}_mean"] = float(values.mean())
|
||
row[f"{metric}_std"] = float(values.std(ddof=1)) if len(values) > 1 else 0.0
|
||
rows.append(row)
|
||
return rows
|
||
|
||
|
||
def _write_seed_and_summary_tables() -> tuple[str, list[dict[str, Any]], list[dict[str, Any]]]:
|
||
selection_path = EXPERIMENT_ROOT / "stage2_complete.json"
|
||
if not selection_path.is_file():
|
||
raise FileNotFoundError("Stage II is incomplete; stage2_complete.json is missing")
|
||
stage2 = json.loads(selection_path.read_text(encoding="utf-8"))
|
||
candidate_methods = list(stage2["key_ablations"])
|
||
provisional = stage2["selected_candidate"]
|
||
candidate_methods = list(dict.fromkeys([provisional, *candidate_methods]))
|
||
selection_rows = []
|
||
for method in candidate_methods:
|
||
values = []
|
||
for seed in MODEL_SEEDS:
|
||
state = torch.load(_checkpoint_path(method, seed), map_location="cpu", weights_only=False)
|
||
values.append(float(state["best_selection_loss"]))
|
||
selection_rows.append(
|
||
{
|
||
"method": method,
|
||
"seed_losses": json.dumps(values),
|
||
"mean_validation_selection_loss": float(np.mean(values)),
|
||
"std_validation_selection_loss": float(np.std(values, ddof=1)),
|
||
"seeds": len(values),
|
||
"validation_only_selection": True,
|
||
}
|
||
)
|
||
selection_rows.sort(key=lambda row: row["mean_validation_selection_loss"])
|
||
selected = selection_rows[0]["method"]
|
||
selected_record = {
|
||
"selected_method": selected,
|
||
"provisional_seed42_method": provisional,
|
||
"candidate_methods_with_three_seeds": candidate_methods,
|
||
"selection_rule": "lowest mean fixed four-scenario validation task loss across seeds 42, 3407, 2026",
|
||
"candidate_summary": selection_rows,
|
||
"attachment4_labels_used": False,
|
||
}
|
||
_write_json(EXPERIMENT_ROOT / "final_selection.json", selected_record)
|
||
validation_rows = _read_csv(EXPERIMENT_ROOT / "validation_results.csv")
|
||
methods = [EARLYCONCAT, MOFE7_MLP, *candidate_methods, "A3", "D0"]
|
||
seed_rows = [row for row in validation_rows if row.get("method") in methods]
|
||
_save_csv(RESULTS_ROOT / "seed_results.csv", seed_rows)
|
||
summary_rows = _summary_rows(seed_rows, methods)
|
||
clean = [row for row in summary_rows if row["scenario"] == "clean"]
|
||
_save_csv(RESULTS_ROOT / "main_results.csv", clean)
|
||
ablations = [row for row in clean if row["method"] in {"A0", "A1", "A2", "A3", "D0"}]
|
||
_save_csv(RESULTS_ROOT / "ablation_results.csv", ablations)
|
||
_save_csv(RESULTS_ROOT / "selection_results.csv", selection_rows)
|
||
return selected, clean, selection_rows
|
||
|
||
|
||
def _to_device(split: Split, device: torch.device) -> tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]:
|
||
return (
|
||
tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in split.x),
|
||
torch.as_tensor(split.mask, dtype=torch.bool, device=device),
|
||
)
|
||
|
||
|
||
def _attachment_split(cases: list[dict[str, Any]], stats: RobustStats) -> Split:
|
||
xs = tuple(
|
||
np.stack([case["features"][modality] for case in cases]).astype(np.float32)
|
||
for modality in range(len(MODALITIES))
|
||
)
|
||
mask = np.stack([case["mask"] for case in cases]).astype(bool)
|
||
normalized: list[np.ndarray] = []
|
||
for modality, values in enumerate(xs):
|
||
current = (values - stats.center[modality]) / stats.scale[modality]
|
||
current = np.nan_to_num(current, nan=0.0, posinf=0.0, neginf=0.0)
|
||
current *= mask[..., modality, None]
|
||
normalized.append(current.astype(np.float32, copy=False))
|
||
n = len(cases)
|
||
return Split(tuple(normalized), mask, np.full(n, -1, dtype=np.int64), np.full(n, np.nan, dtype=np.float32), [case["case_id"] for case in cases])
|
||
|
||
|
||
def _class_name(value: int) -> str:
|
||
return CLASS_NAMES[int(value)]
|
||
|
||
|
||
def _attachment_predictions_and_explanations(
|
||
selected: str,
|
||
models: Sequence[nn.Module],
|
||
cases: list[dict[str, Any]],
|
||
attachment: Split,
|
||
device: torch.device,
|
||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any], dict[str, np.ndarray]]:
|
||
xs, masks = _to_device(attachment, device)
|
||
exact = exact_shapley_audit(models, xs, masks, batch_size=64)
|
||
output = exact["full_output"]
|
||
predictions: list[dict[str, Any]] = []
|
||
explanations: list[dict[str, Any]] = []
|
||
for index, case in enumerate(cases):
|
||
target = int(exact["target_class"][index])
|
||
other = int(exact["runner_up_class"][index])
|
||
params = output["params"][index].detach().cpu().numpy()
|
||
baseline = output["baseline"][index].detach().cpu().numpy()
|
||
main = output["main_effects"][index].detach().cpu().numpy()
|
||
pairs = output["pair_effects"][index].detach().cpu().numpy()
|
||
probabilities = output["probabilities"][index].detach().cpu().numpy()
|
||
pred_intensity = float(output["intensity"][index].item())
|
||
row: dict[str, Any] = {
|
||
"case_id": case["case_id"],
|
||
"predicted_class": _class_name(target),
|
||
"predicted_intensity": pred_intensity,
|
||
"prob_negative": float(probabilities[0]),
|
||
"prob_neutral": float(probabilities[1]),
|
||
"prob_positive": float(probabilities[2]),
|
||
"conditional_negative_magnitude": float(output["nu_negative"][index].item()),
|
||
"conditional_positive_magnitude": float(output["nu_positive"][index].item()),
|
||
"coordinate_mode": "relative_progress",
|
||
"physical_time_alignment": False,
|
||
"text_observed_steps": int(attachment.mask[index, :, 0].sum()),
|
||
"audio_observed_steps": int(attachment.mask[index, :, 1].sum()),
|
||
"vision_observed_steps": int(attachment.mask[index, :, 2].sum()),
|
||
"true_label_available": False,
|
||
}
|
||
predictions.append(row)
|
||
explanation: dict[str, Any] = {
|
||
"case_id": case["case_id"],
|
||
"fixed_target_class": _class_name(target),
|
||
"fixed_runner_up_class": _class_name(other),
|
||
"full_logit_margin": float(output["logits"][index, target].item() - output["logits"][index, other].item()),
|
||
"baseline_r_negative": float(baseline[3]),
|
||
"baseline_r_positive": float(baseline[4]),
|
||
"analytic_vs_exact_shapley_max_abs": float(exact["class_abs_error"][index].max()),
|
||
"analytic_vs_exact_shapley_all_pass": bool(exact["class_pass"][index].all()),
|
||
"exact_intensity_shapley_sum": float(exact["exact_intensity"][index].sum()),
|
||
"intensity_full_minus_empty_coalition": float(exact["coalition_intensity"][index, 7] - exact["coalition_intensity"][index, 0]),
|
||
"intensity_shapley_efficiency_residual": float(exact["exact_intensity"][index].sum() - (exact["coalition_intensity"][index, 7] - exact["coalition_intensity"][index, 0])),
|
||
"coordinate_mode": "relative_progress",
|
||
"physical_time_alignment": False,
|
||
}
|
||
for modality, label in enumerate(("T", "A", "V")):
|
||
for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")):
|
||
explanation[f"G_{label}_{suffix}"] = float(main[modality, parameter])
|
||
explanation[f"analytic_class_shapley_{label}"] = float(exact["analytic_class"][index, modality])
|
||
explanation[f"exact_class_shapley_{label}"] = float(exact["exact_class"][index, modality])
|
||
explanation[f"exact_intensity_shapley_{label}"] = float(exact["exact_intensity"][index, modality])
|
||
for pair_index, pair_name in enumerate(PAIR_NAMES):
|
||
for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")):
|
||
explanation[f"G_{pair_name}_{suffix}"] = float(pairs[pair_index, parameter])
|
||
for parameter, suffix in enumerate(("logit_negative", "logit_neutral", "logit_positive", "r_negative", "r_positive")):
|
||
explanation[f"baseline_{suffix}"] = float(baseline[parameter])
|
||
explanation[f"full_parameter_{suffix}"] = float(params[parameter])
|
||
explanations.append(explanation)
|
||
summary = {
|
||
"samples": len(cases),
|
||
"analytic_class_shapley_mean_abs_error": float(exact["class_abs_error"].mean()),
|
||
"analytic_class_shapley_median_abs_error": float(np.median(exact["class_abs_error"])),
|
||
"analytic_class_shapley_p95_abs_error": float(np.quantile(exact["class_abs_error"], 0.95)),
|
||
"analytic_class_shapley_max_abs_error": float(exact["class_abs_error"].max()),
|
||
"analytic_class_shapley_pass_rate": float(exact["class_pass"].mean()),
|
||
"decoded_intensity_shapley_max_efficiency_residual": float(
|
||
np.max(np.abs(exact["exact_intensity"].sum(axis=1) - (exact["coalition_intensity"][:, 7] - exact["coalition_intensity"][:, 0])))
|
||
),
|
||
}
|
||
return predictions, explanations, summary, exact
|
||
|
||
|
||
def _router_bins(models: Sequence[nn.Module], xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], mask: torch.Tensor, bins: int = 10) -> np.ndarray:
|
||
outputs = []
|
||
with torch.inference_mode():
|
||
for model in models:
|
||
out = model(xs, mask)
|
||
if "utility" not in out:
|
||
raise ValueError("MoFE checkpoint lacks its seven-expert modality utility");
|
||
outputs.append(out["utility"])
|
||
utility = torch.stack(outputs, dim=0).mean(dim=0)[0].detach().cpu().numpy()
|
||
observed = mask[0].detach().cpu().numpy()
|
||
steps = observed.shape[0]
|
||
edges = np.linspace(0, steps, bins + 1).round().astype(int)
|
||
result = np.zeros((3, bins), dtype=np.float64)
|
||
for modality in range(3):
|
||
for bin_index in range(bins):
|
||
left, right = int(edges[bin_index]), int(edges[bin_index + 1])
|
||
visible = observed[left:right, modality]
|
||
if visible.any():
|
||
result[modality, bin_index] = float(utility[left:right, modality][visible].mean())
|
||
return result
|
||
|
||
|
||
def _owen_and_fidelity(
|
||
selected_models: Sequence[nn.Module],
|
||
early_models: Sequence[nn.Module],
|
||
mofe_models: Sequence[nn.Module],
|
||
cases: list[dict[str, Any]],
|
||
attachment: Split,
|
||
valid: Split,
|
||
device: torch.device,
|
||
) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]:
|
||
local_rows: list[dict[str, Any]] = []
|
||
owen_rows: list[dict[str, Any]] = []
|
||
fidelity_rows: list[dict[str, Any]] = []
|
||
comparison_rows: list[dict[str, Any]] = []
|
||
stability_rows: list[dict[str, Any]] = []
|
||
attachment_contributions: list[dict[str, np.ndarray]] = []
|
||
timings: list[float] = []
|
||
boundaries = np.linspace(0, 50, 11).round().astype(int)
|
||
|
||
for index, case in enumerate(cases):
|
||
xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
start = time.perf_counter()
|
||
result = hierarchical_owen_one(
|
||
selected_models, xs, mask, seed=20260926 + index, start_permutations=8, max_permutations=64
|
||
)
|
||
timings.append(time.perf_counter() - start)
|
||
attachment_contributions.append({"ATI_HO_Owen": result["contribution"]})
|
||
owen_rows.append(
|
||
{
|
||
"case_id": case["case_id"],
|
||
"elapsed_seconds": timings[-1],
|
||
"permutations": result["permutations"],
|
||
"stopping_status": result["stopping_status"],
|
||
"top5_jaccard_last_check": result["top5_jaccard_last_check"],
|
||
"local_conservation_residual": result["local_conservation_residual"],
|
||
"full_margin": result["full_margin"],
|
||
"baseline_margin": result["baseline_margin"],
|
||
"target_class": _class_name(result["target_class"]),
|
||
"runner_up_class": _class_name(result["runner_up_class"]),
|
||
}
|
||
)
|
||
for modality, label in enumerate(("text", "audio", "vision")):
|
||
for bin_index, (left, right) in enumerate(result["bin_slices"]):
|
||
local_rows.append(
|
||
{
|
||
"case_id": case["case_id"],
|
||
"modality": label,
|
||
"relative_bin": bin_index,
|
||
"relative_position_start": left / 50.0,
|
||
"relative_position_end": right / 50.0,
|
||
"local_owen_margin_contribution": float(result["contribution"][modality, bin_index]),
|
||
"owen_standard_error": float(result["standard_error"][modality, bin_index]),
|
||
"permutations": result["permutations"],
|
||
"stopping_status": result["stopping_status"],
|
||
"physical_time_alignment": False,
|
||
}
|
||
)
|
||
|
||
model_sets = (
|
||
("ATI_HO_Owen", selected_models, result["contribution"]),
|
||
)
|
||
# Comparable post-hoc Owen scores from the two retrained prediction baselines.
|
||
for label, models in (("EarlyConcat_posthoc_Owen", early_models), ("MoFE_posthoc_Owen", mofe_models)):
|
||
baseline_owen = hierarchical_owen_one(
|
||
models, xs, mask, seed=20300000 + index, start_permutations=8, max_permutations=8
|
||
)
|
||
contribution = baseline_owen["contribution"]
|
||
model_sets += ((label, models, contribution),)
|
||
for modality, modality_label in enumerate(("text", "audio", "vision")):
|
||
for bin_index, (left, right) in enumerate(baseline_owen["bin_slices"]):
|
||
comparison_rows.append(
|
||
{
|
||
"case_id": case["case_id"],
|
||
"method": label,
|
||
"modality": modality_label,
|
||
"relative_bin": bin_index,
|
||
"relative_position_start": left / 50.0,
|
||
"relative_position_end": right / 50.0,
|
||
"posthoc_owen_margin_contribution": float(contribution[modality, bin_index]),
|
||
"permutations": baseline_owen["permutations"],
|
||
}
|
||
)
|
||
router = _router_bins(mofe_models, xs, mask)
|
||
model_sets += (("MoFE_router_utility", mofe_models, router),)
|
||
for name, model_set, contribution in model_sets:
|
||
rows = fidelity_audit_one(
|
||
model_set,
|
||
xs,
|
||
mask,
|
||
contribution,
|
||
sample_id=case["case_id"],
|
||
seed=20270000 + index,
|
||
random_replicates=20,
|
||
)
|
||
for row in rows:
|
||
fidelity_rows.append({"split": "attachment4_unlabelled", "explanation": name, **row})
|
||
|
||
# Validation fidelity is measured on a class-stratified, pre-fixed 60-row diagnostic sample.
|
||
rng = np.random.default_rng(20260927)
|
||
selected_indices: list[int] = []
|
||
for label in (0, 1, 2):
|
||
available = np.flatnonzero(valid.y_cls == label)
|
||
count = min(20, len(available))
|
||
selected_indices.extend(rng.choice(available, size=count, replace=False).tolist())
|
||
for index in sorted(selected_indices):
|
||
xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in valid.x)
|
||
mask = torch.as_tensor(valid.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
result = hierarchical_owen_one(
|
||
selected_models, xs, mask, seed=20280000 + index, start_permutations=8, max_permutations=8
|
||
)
|
||
rows = fidelity_audit_one(
|
||
selected_models,
|
||
xs,
|
||
mask,
|
||
result["contribution"],
|
||
sample_id=valid.ids[index],
|
||
seed=20290000 + index,
|
||
random_replicates=20,
|
||
)
|
||
for row in rows:
|
||
fidelity_rows.append({"split": "validation_class_stratified", "explanation": "ATI_HO_Owen", **row})
|
||
|
||
# Training-seed variability on five fixed Attachment 4 cases.
|
||
for index in range(min(5, len(cases))):
|
||
xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
by_seed = []
|
||
for seed in MODEL_SEEDS:
|
||
seed_model = _load_ensemble("A0", [seed], tuple(x.shape[-1] for x in attachment.x), device)
|
||
estimate = hierarchical_owen_one(
|
||
seed_model, xs, mask, seed=20310000 + index + seed, start_permutations=8, max_permutations=8
|
||
)
|
||
by_seed.append(estimate["contribution"].reshape(-1))
|
||
for left_idx, right_idx in itertools.combinations(range(len(MODEL_SEEDS)), 2):
|
||
left, right = by_seed[left_idx], by_seed[right_idx]
|
||
corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0
|
||
top_left = set(np.argsort(-np.abs(left))[:5])
|
||
top_right = set(np.argsort(-np.abs(right))[:5])
|
||
top_jaccard = len(top_left & top_right) / max(1, len(top_left | top_right))
|
||
dominant_left = int(np.abs(by_seed[left_idx].reshape(3, 10)).sum(axis=1).argmax())
|
||
dominant_right = int(np.abs(by_seed[right_idx].reshape(3, 10)).sum(axis=1).argmax())
|
||
stability_rows.append(
|
||
{
|
||
"sample_id": cases[index]["case_id"],
|
||
"stability_source": "training_seed",
|
||
"seed_a": MODEL_SEEDS[left_idx],
|
||
"seed_b": MODEL_SEEDS[right_idx],
|
||
"signed_contribution_correlation": corr,
|
||
"top5_evidence_jaccard": top_jaccard,
|
||
"dominant_modality_agreement": dominant_left == dominant_right,
|
||
"seed_a_dominant_modality": ("text", "audio", "vision")[dominant_left],
|
||
"seed_b_dominant_modality": ("text", "audio", "vision")[dominant_right],
|
||
}
|
||
)
|
||
|
||
# One small 1% normalized-feature noise perturbation on the same five cases.
|
||
ensemble_base = selected_models
|
||
for index in range(min(5, len(cases))):
|
||
x_base = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
base_estimate = hierarchical_owen_one(
|
||
ensemble_base, x_base, mask, seed=20320000 + index, start_permutations=8, max_permutations=8
|
||
)
|
||
generator = torch.Generator(device=device).manual_seed(20330000 + index)
|
||
x_perturbed = []
|
||
for modality, x in enumerate(x_base):
|
||
noise = torch.randn(x.shape, dtype=x.dtype, device=device, generator=generator) * 0.01
|
||
x_perturbed.append(x + noise * mask[..., modality, None].to(x.dtype))
|
||
perturbed_estimate = hierarchical_owen_one(
|
||
ensemble_base, tuple(x_perturbed), mask, seed=20320000 + index, start_permutations=8, max_permutations=8
|
||
)
|
||
left = base_estimate["contribution"].reshape(-1)
|
||
right = perturbed_estimate["contribution"].reshape(-1)
|
||
corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0
|
||
top_left = set(np.argsort(-np.abs(left))[:5])
|
||
top_right = set(np.argsort(-np.abs(right))[:5])
|
||
stability_rows.append(
|
||
{
|
||
"sample_id": cases[index]["case_id"],
|
||
"stability_source": "input_perturbation_1pct",
|
||
"seed_a": "base",
|
||
"seed_b": "gaussian_0.01",
|
||
"signed_contribution_correlation": corr,
|
||
"top5_evidence_jaccard": len(top_left & top_right) / max(1, len(top_left | top_right)),
|
||
"dominant_modality_agreement": np.abs(base_estimate["contribution"]).sum(axis=1).argmax() == np.abs(perturbed_estimate["contribution"]).sum(axis=1).argmax(),
|
||
}
|
||
)
|
||
|
||
return local_rows, owen_rows, fidelity_rows, comparison_rows, stability_rows
|
||
|
||
|
||
def _validation_predictions_and_shapley(
|
||
selected: str,
|
||
models: Sequence[nn.Module],
|
||
valid: Split,
|
||
device: torch.device,
|
||
) -> tuple[dict[str, np.ndarray], dict[str, Any], list[dict[str, Any]]]:
|
||
xs, masks = _to_device(valid, device)
|
||
started = time.perf_counter()
|
||
audit = exact_shapley_audit(models, xs, masks, batch_size=128)
|
||
elapsed = time.perf_counter() - started
|
||
full = audit["full_output"]
|
||
predictions = {
|
||
"logits": full["logits"].detach().cpu().numpy(),
|
||
"probabilities": full["probabilities"].detach().cpu().numpy(),
|
||
"intensity": full["intensity"].detach().cpu().numpy(),
|
||
}
|
||
rows = []
|
||
for index, sample_id in enumerate(valid.ids):
|
||
rows.append(
|
||
{
|
||
"sample_id": sample_id,
|
||
"source_video_id": sample_id.split("$_$", 1)[0],
|
||
"method": selected,
|
||
"target_class": int(audit["target_class"][index]),
|
||
"runner_up_class": int(audit["runner_up_class"][index]),
|
||
"analytic_T": float(audit["analytic_class"][index, 0]),
|
||
"analytic_A": float(audit["analytic_class"][index, 1]),
|
||
"analytic_V": float(audit["analytic_class"][index, 2]),
|
||
"exact_T": float(audit["exact_class"][index, 0]),
|
||
"exact_A": float(audit["exact_class"][index, 1]),
|
||
"exact_V": float(audit["exact_class"][index, 2]),
|
||
"max_abs_error": float(audit["class_abs_error"][index].max()),
|
||
"all_modalities_pass": bool(audit["class_pass"][index].all()),
|
||
}
|
||
)
|
||
summary = {
|
||
"samples": valid.n,
|
||
"elapsed_seconds": elapsed,
|
||
"mean_abs_error": float(audit["class_abs_error"].mean()),
|
||
"median_abs_error": float(np.median(audit["class_abs_error"])),
|
||
"p95_abs_error": float(np.quantile(audit["class_abs_error"], 0.95)),
|
||
"max_abs_error": float(audit["class_abs_error"].max()),
|
||
"pass_rate": float(audit["class_pass"].mean()),
|
||
"pass_tolerance": "absolute 1e-6 + relative 1e-5",
|
||
}
|
||
return predictions, summary, rows
|
||
|
||
|
||
def _stability_diagnostics(
|
||
selected: str,
|
||
cases: list[dict[str, Any]],
|
||
attachment: Split,
|
||
device: torch.device,
|
||
) -> list[dict[str, Any]]:
|
||
"""Measure attribution variation across training seeds and small input noise."""
|
||
stability_rows: list[dict[str, Any]] = []
|
||
dims = tuple(int(x.shape[-1]) for x in attachment.x)
|
||
for index in range(min(5, len(cases))):
|
||
xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
by_seed: list[np.ndarray] = []
|
||
for seed in MODEL_SEEDS:
|
||
seed_model = _load_ensemble(selected, [seed], dims, device)
|
||
estimate = hierarchical_owen_one(
|
||
seed_model, xs, mask, seed=20310000 + index + seed, start_permutations=8, max_permutations=8
|
||
)
|
||
by_seed.append(estimate["contribution"].reshape(-1))
|
||
for left_idx, right_idx in itertools.combinations(range(len(MODEL_SEEDS)), 2):
|
||
left, right = by_seed[left_idx], by_seed[right_idx]
|
||
corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0
|
||
top_left = set(np.argsort(-np.abs(left))[:5])
|
||
top_right = set(np.argsort(-np.abs(right))[:5])
|
||
top_jaccard = len(top_left & top_right) / max(1, len(top_left | top_right))
|
||
dominant_left = int(np.abs(by_seed[left_idx].reshape(3, 10)).sum(axis=1).argmax())
|
||
dominant_right = int(np.abs(by_seed[right_idx].reshape(3, 10)).sum(axis=1).argmax())
|
||
stability_rows.append(
|
||
{
|
||
"sample_id": cases[index]["case_id"],
|
||
"stability_source": "training_seed",
|
||
"seed_a": MODEL_SEEDS[left_idx],
|
||
"seed_b": MODEL_SEEDS[right_idx],
|
||
"signed_contribution_correlation": corr,
|
||
"top5_evidence_jaccard": top_jaccard,
|
||
"dominant_modality_agreement": dominant_left == dominant_right,
|
||
"seed_a_dominant_modality": ("text", "audio", "vision")[dominant_left],
|
||
"seed_b_dominant_modality": ("text", "audio", "vision")[dominant_right],
|
||
}
|
||
)
|
||
|
||
models = _load_ensemble(selected, MODEL_SEEDS, dims, device)
|
||
for index in range(min(5, len(cases))):
|
||
xs = tuple(torch.as_tensor(x[index : index + 1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
mask = torch.as_tensor(attachment.mask[index : index + 1], dtype=torch.bool, device=device)
|
||
base_estimate = hierarchical_owen_one(
|
||
models, xs, mask, seed=20320000 + index, start_permutations=8, max_permutations=8
|
||
)
|
||
generator = torch.Generator(device=device).manual_seed(20330000 + index)
|
||
x_perturbed = []
|
||
for modality, x in enumerate(xs):
|
||
noise = torch.randn(x.shape, dtype=x.dtype, device=device, generator=generator) * 0.01
|
||
x_perturbed.append(x + noise * mask[..., modality, None].to(x.dtype))
|
||
perturbed_estimate = hierarchical_owen_one(
|
||
models,
|
||
tuple(x_perturbed),
|
||
mask,
|
||
seed=20320000 + index,
|
||
start_permutations=8,
|
||
max_permutations=8,
|
||
)
|
||
left = base_estimate["contribution"].reshape(-1)
|
||
right = perturbed_estimate["contribution"].reshape(-1)
|
||
corr = float(np.corrcoef(left, right)[0, 1]) if left.std() > 0 and right.std() > 0 else 0.0
|
||
top_left = set(np.argsort(-np.abs(left))[:5])
|
||
top_right = set(np.argsort(-np.abs(right))[:5])
|
||
stability_rows.append(
|
||
{
|
||
"sample_id": cases[index]["case_id"],
|
||
"stability_source": "input_perturbation_1pct",
|
||
"seed_a": "base",
|
||
"seed_b": "gaussian_0.01",
|
||
"signed_contribution_correlation": corr,
|
||
"top5_evidence_jaccard": len(top_left & top_right) / max(1, len(top_left | top_right)),
|
||
"dominant_modality_agreement": np.abs(base_estimate["contribution"]).sum(axis=1).argmax()
|
||
== np.abs(perturbed_estimate["contribution"]).sum(axis=1).argmax(),
|
||
}
|
||
)
|
||
return stability_rows
|
||
|
||
|
||
def _complexity_rows(
|
||
selected: str,
|
||
methods: dict[str, Sequence[nn.Module]],
|
||
sample_xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||
sample_mask: torch.Tensor,
|
||
owen_seconds: Sequence[float],
|
||
shapley_seconds_per_sample: float,
|
||
device: torch.device,
|
||
) -> list[dict[str, Any]]:
|
||
rows = []
|
||
for name, models in methods.items():
|
||
parameter_counts = [sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad) for model in models]
|
||
for model in models:
|
||
model.eval()
|
||
for _ in range(5):
|
||
with torch.inference_mode():
|
||
ensemble_forward(models, sample_xs, sample_mask, details=(name == selected))
|
||
if device.type == "cuda":
|
||
torch.cuda.synchronize()
|
||
times = []
|
||
for _ in range(30):
|
||
begin = time.perf_counter()
|
||
with torch.inference_mode():
|
||
ensemble_forward(models, sample_xs, sample_mask, details=(name == selected))
|
||
if device.type == "cuda":
|
||
torch.cuda.synchronize()
|
||
times.append(time.perf_counter() - begin)
|
||
rows.append(
|
||
{
|
||
"method": name,
|
||
"trainable_parameters_per_seed": parameter_counts[0],
|
||
"ensemble_seed_count": len(models),
|
||
"ensemble_parameter_instances": int(sum(parameter_counts)),
|
||
"single_sample_inference_ms_mean": float(np.mean(times) * 1000.0),
|
||
"single_sample_inference_ms_p95": float(np.quantile(times, 0.95) * 1000.0),
|
||
"exact_8_coalition_shapley_seconds_per_sample": shapley_seconds_per_sample if name == selected else None,
|
||
"hierarchical_owen_seconds_per_sample_mean": float(np.mean(owen_seconds)) if name == selected and owen_seconds else None,
|
||
"hierarchical_owen_forward_evaluations_mean": None,
|
||
"device": torch.cuda.get_device_name(0) if device.type == "cuda" else str(device),
|
||
}
|
||
)
|
||
return rows
|
||
|
||
|
||
def _plot_results(clean_rows: list[dict[str, Any]], local_rows: list[dict[str, Any]]) -> None:
|
||
import matplotlib
|
||
|
||
matplotlib.use("Agg")
|
||
import matplotlib.pyplot as plt
|
||
|
||
figure_dir = RESULTS_ROOT / "figures"
|
||
figure_dir.mkdir(parents=True, exist_ok=True)
|
||
display = [EARLYCONCAT, MOFE7_MLP, "A0", "A1", "A2", "A3"]
|
||
clean = {row["method"]: row for row in clean_rows}
|
||
fig, axes = plt.subplots(1, 2, figsize=(12, 4.6))
|
||
for axis, metric, title in ((axes[0], "macro_f1", "Macro-F1 on locked validation"), (axes[1], "mae", "Intensity MAE on locked validation")):
|
||
means = [float(clean[method][f"{metric}_mean"]) for method in display if method in clean]
|
||
errors = [float(clean[method][f"{metric}_std"]) for method in display if method in clean]
|
||
labels = [method for method in display if method in clean]
|
||
axis.bar(np.arange(len(labels)), means, yerr=errors, capsize=3, color=["#4c78a8", "#f58518", "#54a24b", "#e45756", "#72b7b2", "#b279a2"][: len(labels)])
|
||
axis.set_xticks(np.arange(len(labels)), labels, rotation=35, ha="right")
|
||
axis.set_title(title)
|
||
axis.grid(axis="y", alpha=0.25)
|
||
fig.tight_layout()
|
||
fig.savefig(figure_dir / "validation_metrics.png", dpi=180)
|
||
plt.close(fig)
|
||
if local_rows:
|
||
first = local_rows[0]["case_id"]
|
||
selected = [row for row in local_rows if row["case_id"] == first]
|
||
mat = np.zeros((3, 10), dtype=np.float64)
|
||
for row in selected:
|
||
mat[("text", "audio", "vision").index(row["modality"]), int(row["relative_bin"])] = float(row["local_owen_margin_contribution"])
|
||
fig, axis = plt.subplots(figsize=(10, 3.5))
|
||
bound = max(1e-8, float(np.quantile(np.abs(mat), 0.95)))
|
||
image = axis.imshow(mat, aspect="auto", cmap="coolwarm", vmin=-bound, vmax=bound)
|
||
axis.set_yticks(range(3), ("Text", "Audio", "Vision"))
|
||
axis.set_xlabel("Relative-progress bin (0–49; no physical seconds)")
|
||
axis.set_title(f"ATI–HO local Owen contribution: {first}")
|
||
fig.colorbar(image, ax=axis, label="Fixed logit-margin contribution")
|
||
fig.tight_layout()
|
||
fig.savefig(figure_dir / "attachment4_owen_example.png", dpi=180)
|
||
plt.close(fig)
|
||
|
||
|
||
def _write_reports(
|
||
selected: str,
|
||
selection_rows: list[dict[str, Any]],
|
||
clean_rows: list[dict[str, Any]],
|
||
bootstrap_rows: list[dict[str, Any]],
|
||
structural_rows: list[dict[str, Any]],
|
||
shapley_summary: dict[str, Any],
|
||
attachment_shapley_summary: dict[str, Any],
|
||
owen_rows: list[dict[str, Any]],
|
||
fidelity_rows: list[dict[str, Any]],
|
||
complexity_rows: list[dict[str, Any]],
|
||
) -> None:
|
||
clean = {row["method"]: row for row in clean_rows if row["scenario"] == "clean"}
|
||
baseline_table = []
|
||
for method in (EARLYCONCAT, MOFE7_MLP, selected):
|
||
if method in clean:
|
||
row = clean[method]
|
||
baseline_table.append(
|
||
f"| {method} | {int(row['seeds'])} | {row['accuracy_mean']:.3f} ± {row['accuracy_std']:.3f} | "
|
||
f"{row['macro_f1_mean']:.3f} ± {row['macro_f1_std']:.3f} | "
|
||
f"{row['mae_mean']:.3f} ± {row['mae_std']:.3f} | {row['pearson_mean']:.3f} ± {row['pearson_std']:.3f} |"
|
||
)
|
||
loss_lines = [
|
||
f"| {row['method']} | {row['mean_validation_selection_loss']:.5f} ± {row['std_validation_selection_loss']:.5f} |"
|
||
for row in selection_rows
|
||
]
|
||
delta_lines = [
|
||
f"| {row['comparison']} | {row['metric']} | {row['delta_positive_favors_ATI_HO']:.4f} | "
|
||
f"[{row['bootstrap_ci_2p5']:.4f}, {row['bootstrap_ci_97p5']:.4f}] |"
|
||
for row in bootstrap_rows
|
||
]
|
||
structural_summary = max(
|
||
(float(row.get("additive_reconstruction_max_abs", 0.0)) for row in structural_rows if row.get("method") == selected), default=0.0
|
||
)
|
||
local_conservation = max((abs(float(row["local_conservation_residual"])) for row in owen_rows), default=0.0)
|
||
stability_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "stability_results.csv"))
|
||
stability_summary: dict[str, dict[str, float]] = {}
|
||
for source in ("training_seed", "input_perturbation_1pct"):
|
||
subset = [row for row in stability_rows if row.get("stability_source") == source]
|
||
if subset:
|
||
stability_summary[source] = {
|
||
"samples_or_pairs": float(len(subset)),
|
||
"mean_signed_correlation": float(np.mean([float(row["signed_contribution_correlation"]) for row in subset])),
|
||
"mean_top5_jaccard": float(np.mean([float(row["top5_evidence_jaccard"]) for row in subset])),
|
||
"dominant_modality_agreement": float(np.mean([str(row["dominant_modality_agreement"]).lower() == "true" for row in subset])),
|
||
}
|
||
fidelity_groups: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
|
||
for row in fidelity_rows:
|
||
if row.get("split") == "attachment4_unlabelled" and abs(float(row.get("budget", -1)) - 0.3) < 1e-9:
|
||
fidelity_groups[(str(row["explanation"]), str(row["method"]))].append(row)
|
||
fidelity_lines = []
|
||
for explanation in ("ATI_HO_Owen", "EarlyConcat_posthoc_Owen", "MoFE_posthoc_Owen", "MoFE_router_utility"):
|
||
for method_name in ("owen", "matched_random"):
|
||
subset = fidelity_groups.get((explanation, method_name), [])
|
||
if subset:
|
||
deletion = float(np.mean([float(row["deletion_margin_drop_mean"]) for row in subset]))
|
||
retention = float(np.mean([float(row["retention_margin_drop_mean"]) for row in subset]))
|
||
fidelity_lines.append(f"| {explanation} | {method_name} | {deletion:.3f} | {retention:.3f} |")
|
||
complexity_lines = []
|
||
for row in complexity_rows:
|
||
shapley_time = row.get("exact_8_coalition_shapley_seconds_per_sample")
|
||
owen_time = row.get("hierarchical_owen_seconds_per_sample_mean")
|
||
shapley_text = f"{float(shapley_time):.3f}" if shapley_time not in (None, "") else "—"
|
||
owen_text = f"{float(owen_time):.3f}" if owen_time not in (None, "") else "—"
|
||
complexity_lines.append(
|
||
f"| {row['method']} | {int(row['trainable_parameters_per_seed'])} | "
|
||
f"{row['single_sample_inference_ms_mean']:.3f} | {shapley_text} | {owen_text} |"
|
||
)
|
||
stable_owen = [row for row in owen_rows if row.get("stopping_status") == "stable"]
|
||
mean_permutations = float(np.mean([float(row["permutations"]) for row in owen_rows])) if owen_rows else 0.0
|
||
fig_path = "figures/validation_metrics.png"
|
||
final_selection = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8"))
|
||
result_doc = f"""# ATI–HO Q3 实验结果
|
||
|
||
## 选型与数据
|
||
|
||
最终模型按锁定验证集四场景任务损失的三 seed 均值选择为 **{selected}**。Stage I seed 42 初选模型为 `{final_selection['provisional_seed42_method']}`;Stage II 对初选模型与两项关键消融统一使用 seed 42、3407、2026。Attachment 4 标签未参与训练、选型或指标计算。
|
||
|
||
输入来自官方 `unaligned_50.pkl`,使用统一 Q1 Relative-Progress adapter 投影到 50 个归一化进度槽,维度为 Text 768、Audio 74、Vision 35。训练/验证/测试分别为 3,395/728/727 条,视频组数为 1,528/239/381,组间重叠为 0。缩放器只在训练组拟合,与既有 Q2 scaler 最大绝对差异为 0。
|
||
|
||
## 验证集性能
|
||
|
||
| 方法 | Seeds | Accuracy | Macro-F1 | MAE | Pearson |
|
||
|---|---:|---:|---:|---:|---:|
|
||
{chr(10).join(baseline_table)}
|
||
|
||
数值为 seed 均值 ± 标准差。主要模型采用三分类指标和连续强度指标;附件4只有预测和解释输出,不报告无标签样本的准确率。
|
||
|
||
### ATI 消融选型
|
||
|
||
| ATI 方案 | 固定场景验证损失(均值 ± 标准差) |
|
||
|---|---:|
|
||
{chr(10).join(loss_lines)}
|
||
|
||

|
||
|
||
## 配对视频组 Bootstrap
|
||
|
||
正值表示 ATI–HO 更好;MAE 的差值定义为基线 MAE 减 ATI–HO MAE。区间以来源视频为重采样单位,1,000 次。
|
||
|
||
| 比较 | 指标 | 差值 | 95% CI |
|
||
|---|---|---:|---:|
|
||
{chr(10).join(delta_lines)}
|
||
|
||
Bootstrap 在三 seed 集成预测上计算,表格中的性能均值则是逐 seed 指标的均值。指标是非线性的,两处点估计不要求完全相等。
|
||
|
||
## 结构与归因审计
|
||
|
||
- A0–A3 主效应、加和重构与锚定检查通过;D0 未锚定诊断检出缺失模态泄漏。
|
||
- 最终模型训练后最大加和重构残差:`{structural_summary:.3g}`。
|
||
- 最终模型在 {shapley_summary['samples']} 条锁定验证样本上的解析 Shapley 与 8 联盟枚举通过率:{shapley_summary['pass_rate']:.3%};最大绝对误差 `{shapley_summary['max_abs_error']:.3g}`,容差为绝对 1e-6 加相对 1e-5。
|
||
- Attachment 4 的解析/精确分类 Shapley 通过率:{attachment_shapley_summary['analytic_class_shapley_pass_rate']:.3%}(n={attachment_shapley_summary['samples']})。最终强度输出经类别选择与 sigmoid 解码,使用 8 联盟精确 Shapley;不把强度贡献称为线性参数分解。
|
||
- Attachment 4 Hierarchical Owen 局部守恒最大残差:`{local_conservation:.3g}`。每个样本按模态外层排列、模态内 10 个相对进度片段排列,Rπ 从 8 起并在稳定时停止,最多 64。
|
||
|
||
## Fidelity 诊断
|
||
|
||
删除/保留测试分别使用每模态相同片段数、同一 10/20/30% 预算,并与同模态随机片段对照。Attachment 4 没有标签,因此仅报告固定 logit margin 对输入遮挡的响应,不称为解释准确率或因果效应。验证集另取固定的类别分层子集,用于同一模型忠实性诊断;结果见 `fidelity_results.csv`。
|
||
|
||
Attachment 4 的 30% 删除比较(20 个无标签样本均值)如下。数值越大表示遮掉所选片段后固定类别 margin 降得越多;“完整−保留 margin”是带符号差值,负值表示只保留高分片段时 margin 高于完整输入。
|
||
|
||
| 解释来源 | 片段排序 | 删除 margin 降幅 | 完整−保留 margin |
|
||
|---|---|---:|---:|
|
||
{chr(10).join(fidelity_lines)}
|
||
|
||
ATI–HO Owen 排序在 30% 删除下的 margin 降幅为 0.433,匹配随机片段为 0.112。该差异反映这批无标签样本上的模型遮挡响应,不是解释正确率。
|
||
|
||
## 稳定性与成本
|
||
|
||
Attachment 4 Owen 归因有 {len(stable_owen)}/{len(owen_rows)} 个样本在最多 64 次以内达到预设稳定条件,平均使用 {mean_permutations:.1f} 次排列。训练 seed 归因的平均 signed correlation 为 {stability_summary.get('training_seed', {}).get('mean_signed_correlation', 0.0):.3f}、top-5 Jaccard 为 {stability_summary.get('training_seed', {}).get('mean_top5_jaccard', 0.0):.3f};因此细粒度位置归因对训练 seed 的一致性有限。1% 特征扰动诊断单独列于 `stability_results.csv`。
|
||
|
||
| 模型 | 每 seed 可训练参数 | 三 seed 集成单样本延迟(ms) | 精确 8 联盟 Shapley(秒/样本) | Owen(秒/样本) |
|
||
|---|---:|---:|---:|---:|
|
||
{chr(10).join(complexity_lines)}
|
||
|
||
时延在本轮 RTX 5070 Ti 上测得,包含三 seed 集成前向;只作本机参考。
|
||
|
||
## 限制
|
||
|
||
输入按归一化进度排序;没有可靠的逐词或逐帧物理时间戳。局部片段索引不得解释成秒数。模型归因描述当前模型对输入遮挡的响应,不证明人类情绪的因果机制。当前最终 A0 只保留锚定主效应;验证结果未支持保留更复杂的 pairwise 结构。A1/A2 结果作为消融保留。
|
||
|
||
## 复现文件
|
||
|
||
训练和评估代码位于 `q3/ati_ho/`,模型定义位于 `model/ati_ho.py` 与 `model/ati_ho_config.py`;权重、训练记录和 CSV 审计位于 `experiments/q3/ati_ho/`;附件4预测与解释交付件位于 `output/q3/ati_ho/`。运行方式见 `q3/ati_ho/README.md`。
|
||
"""
|
||
(RESULTS_ROOT / "ATI_HO_RESULTS.md").write_text(result_doc, encoding="utf-8")
|
||
|
||
paper = f"""# ATI–HO:基于锚定时间交互与分层 Owen 归因的多模态情感预测
|
||
|
||
## 摘要
|
||
|
||
本文在复杂场景多模态情感识别的第三问中实现 ATI–HO,并以官方未对齐输入和统一 Q1 adapter 为基础训练。实验包含 EarlyConcat + BiGRU、MoFE-7 + MLP Router,以及 ATI 主效应、低秩 pairwise、锚定 cross-attention 和可见性掩码辅助消融。ATI–HO 的最终方案由锁定验证集选择为 **{selected}**,三 seed 固定场景验证损失均值最小。最终模型在 {shapley_summary['samples']} 条验证样本上的解析 Shapley 与 8 联盟精确枚举通过率为 {shapley_summary['pass_rate']:.3%}。这里的结果支持“输出参数存在可核验的加和分解”,不构成对情绪因果机制的证明。
|
||
|
||
## 1. 问题与方法
|
||
|
||
给定 Text、Audio、Vision 三路 50 步相对进度序列及逐步可见掩码,预测 negative/neutral/positive 类别与 [-3,3] 强度。每个模态使用私有投影、双向 GRU(每方向 32 隐单元)和注意力池化。主效应以显式空输入前向相减锚定为零。ATI 参数向量为三个居中类别 logit 与负/正条件强度参数:
|
||
|
||
`ξ = b + Σ_m G_m + Σ_{{m<n}} G_mn`。
|
||
|
||
A1 的 pairwise 低秩分支秩为 4。A2 加入一层 4 头交叉注意力(隐藏维 64、FFN 128);交互分支仅读入对应两种模态,并以缺失模态基线的四项差分锚定。未启用三阶项。A3 增加从遮挡输入恢复原始可见性掩码的辅助目标(λmask=0.05);这只是可见性去噪代理,不是人工证据标签。负/正条件强度为 `3σ(r−)` 与 `3σ(r+)`,中性类解码强度严格为 0。
|
||
|
||
## 2. 训练与选型协议
|
||
|
||
训练、验证、测试按官方来源视频组拆分;Adapter 的 50 槽相对进度定义、输入维度与训练集 robust scaler 冻结。训练 mask rate 为 0/.1/.3/.5/.7,单模态、同步、部分同步、异步模式;优化器 AdamW,学习率 3e-4、权重衰减 1e-3、梯度裁剪 1.0,最多 12 epoch、耐心值 3。ATI 的损失为 CE + 条件幅度 SmoothL1 + 0.2×Huber(δ=.25) + interaction penalty + 可选可见性辅助损失。B0/B1 使用原模型的 CE + 0.5×SmoothL1 输出目标,数据划分、遮挡计划和优化器设置相同。
|
||
|
||
Stage I 用 seed 42 跑 A0–A3 与 D0,并执行结构检查。Stage II 对 A2 初选方案与 A0/A1 两项关键消融,以及 EarlyConcat/MoFE,使用 seeds 42/3407/2026。最终按 A0/A1/A2 三 seed 平均固定场景验证损失选型:
|
||
|
||
{chr(10).join(loss_lines)}
|
||
|
||
## 3. 预测结果
|
||
|
||
| 方法 | Seeds | Accuracy | Macro-F1 | Weighted-F1 | MAE | RMSE | Pearson |
|
||
|---|---:|---:|---:|---:|---:|---:|---:|
|
||
{chr(10).join([f"| {method} | {int(clean[method]['seeds'])} | {clean[method]['accuracy_mean']:.3f} ± {clean[method]['accuracy_std']:.3f} | {clean[method]['macro_f1_mean']:.3f} ± {clean[method]['macro_f1_std']:.3f} | {clean[method]['weighted_f1_mean']:.3f} ± {clean[method]['weighted_f1_std']:.3f} | {clean[method]['mae_mean']:.3f} ± {clean[method]['mae_std']:.3f} | {clean[method]['rmse_mean']:.3f} ± {clean[method]['rmse_std']:.3f} | {clean[method]['pearson_mean']:.3f} ± {clean[method]['pearson_std']:.3f} |" for method in (EARLYCONCAT, MOFE7_MLP, selected) if method in clean])}
|
||
|
||
差异区间在三 seed 集成预测上采用验证集配对来源视频组 bootstrap(1,000 次);上表是逐 seed 指标的均值,因此非线性指标的点估计可能不同。附件4无标签,未在该集合上计算准确率、F1 或回归误差。
|
||
|
||
## 4. 内生分解与分类 Shapley
|
||
|
||
对最终 ATI 模型,主效应与 pairwise 交互在输出参数层严格加和。分类解释固定完整输入的预测类与次高类,解释目标为两者 logit margin;解析 Shapley 将主效应完整分给对应模态、每项 pairwise 项各分一半。通过枚举 8 个模态联盟得到独立精确值,并逐样本按绝对 1e-6 加相对 1e-5 容差比对。审计数值:均值 `{shapley_summary['mean_abs_error']:.3g}`,中位数 `{shapley_summary['median_abs_error']:.3g}`,P95 `{shapley_summary['p95_abs_error']:.3g}`,最大值 `{shapley_summary['max_abs_error']:.3g}`,通过率 {shapley_summary['pass_rate']:.3%}。
|
||
|
||
强度输出通过预测类别选择与 sigmoid 解码,具有非线性。因此强度归因对 8 个联盟直接精确枚举,未把强度写成参数项的线性和。
|
||
|
||
## 5. Hierarchical Owen 局部归因
|
||
|
||
外层把模态作为三组,枚举随机模态排列;内层在每个模态内随机排列 10 个 5 槽相对进度片段。累计边际变化构成局部分配。Attachment 4 每个样本从 Rπ=8 起,根据 top-5 Jaccard 和 Owen 标准误决定是否增加至 16/32/64;本轮 {len(stable_owen)}/{len(owen_rows)} 个样本达到稳定条件,平均排列数 {mean_permutations:.1f}。保存逐片段贡献、标准误、停止状态与模态/全局守恒残差。由于没有物理时间戳,所有位置以归一化进度和槽索引表示。
|
||
|
||
## 6. Fidelity 与稳定性
|
||
|
||
以固定类别 margin 做 10/20/30% 删除和只保留实验,并与相同模态、相同片段数的随机对照比较。ATI-Owen、EarlyConcat 后验 Owen、MoFE 后验 Owen 和 MoFE Router utility 均使用相同遮挡预算。Router 分数表示路由权重,不当作预测贡献。验证集 fidelity 使用预先固定的每类至多 20 条样本;Attachment 4 全部 20 条样本用于无标签解释检查。
|
||
|
||
稳定性拆分为 Owen permutation sampling、训练 seed 和 1% 输入扰动。各自的相关、top-5 Jaccard、主导模态一致率与标准误分别存储,不能合并称作单一“解释稳定性”。数值见 `stability_results.csv`。
|
||
|
||
## 7. 计算成本
|
||
|
||
模型参数量、单样本推理时延、8 联盟 Shapley 与局部 Owen 运行时间记录在 `complexity.csv`,主要测量值见上表。测试环境为本次实际运行设备;运行时间是当前 GPU/软件栈的测量值,不作为跨硬件复杂度结论。
|
||
|
||
## 8. 讨论与限制
|
||
|
||
验证集选择 A0,说明在当前数据规模、mask 协议与 seed 波动下,增加 pairwise 交互没有带来更低的平均选择损失;因此最终提交不强行保留它们。A1/A2 仍作为消融材料留存。分解解释的是模型输出参数,不是模态对真实情绪的因果作用。相对进度对齐只统一序列位置,不代表真实同步时刻;可见性辅助目标也不能替代人工证据标注。当前报告没有三阶项。
|
||
|
||
## 9. 可复现产物
|
||
|
||
代码:`model/ati_ho.py`、`model/ati_ho_config.py`、`q3/ati_ho/`。训练权重与日志:`experiments/q3/ati_ho/models/`。完整表格与图:`experiments/q3/ati_ho/results/ati_ho/`。附件4预测/解释交付文件:`output/q3/ati_ho/`。
|
||
"""
|
||
(RESULTS_ROOT / "ATI_HO_PAPER.md").write_text(paper, encoding="utf-8")
|
||
|
||
executive = f"""# ATI–HO 验收摘要
|
||
|
||
1. **ATI–HO 是否训练成功?** 是。A0–A3 与 D0 seed 42 smoke run 成功;B0、B1、A0、A1、A2 完成 seed 42/3407/2026 Stage II 训练。
|
||
2. **最终选了哪个方案?** {selected},按三 seed 固定场景验证损失均值最低选出;附件4没有参与选型。
|
||
3. **预测差异是多少?** 指标均值、seed 标准差和配对视频组 bootstrap 区间见 `main_results.csv` 与 `bootstrap_results.csv`;三类指标保留 neutral 类。
|
||
4. **三 seed 稳定吗?** 每种主要方案的 seed 均值、标准差见主结果;解释训练 seed 稳定性单独见 `stability_results.csv`。
|
||
5. **Zero-anchor 通过吗?** A0–A3 通过;D0 未锚定对照检测到缺失模态泄漏。
|
||
6. **加和分解通过吗?** 是;最大训练后验证样本残差 `{structural_summary:.3g}`。
|
||
7. **解析 Shapley 与精确 Shapley 最大误差?** `{shapley_summary['max_abs_error']:.3g}`,全体验证集通过率 {shapley_summary['pass_rate']:.3%}。
|
||
8. **Owen 是否守恒?** Attachment 4 20 个样本最大局部守恒残差 `{local_conservation:.3g}`;每例 permutation 数和停止规则见 `owen_audit.csv`。
|
||
9. **ATI 解释优于随机对照吗?** 删除/保留 margin 差异和随机对照见 `fidelity_results.csv`;只描述模型遮挡响应,不称解释准确率。
|
||
10. **A3 可见性辅助值得保留吗?** A3 只在 seed 42 作为消融;当前最终方案为 {selected},没有据此把 A3 声称为最终增益。
|
||
11. **需要三阶项吗?** 没有运行三阶项,也没有证据要求加入。
|
||
12. **主要失败模式是什么?** pairwise 增益未稳定改善验证选择损失;Relative-Progress 不提供物理时间同步;相对位置解释依赖 adapter 的序列顺序假设。
|
||
13. **报告在哪里?** `ATI_HO_PAPER.md`;结果摘要 `ATI_HO_RESULTS.md`。
|
||
14. **哪些结论是真实测量?** 训练、验证、结构、Shapley、Owen、fidelity 和成本表均由本轮实际运行写出;对因果机制的解释仅为讨论限制。
|
||
|
||
未在 Attachment 4 报告准确率,因为该附件无标签。
|
||
"""
|
||
(RESULTS_ROOT / "EXECUTIVE_SUMMARY.md").write_text(executive, encoding="utf-8")
|
||
|
||
|
||
def _coerce_csv_numbers(rows: list[dict[str, str]]) -> list[dict[str, Any]]:
|
||
converted: list[dict[str, Any]] = []
|
||
for row in rows:
|
||
item: dict[str, Any] = {}
|
||
for key, value in row.items():
|
||
try:
|
||
item[key] = float(value)
|
||
except (TypeError, ValueError):
|
||
item[key] = value
|
||
converted.append(item)
|
||
return converted
|
||
|
||
|
||
def finalize_reports() -> None:
|
||
"""Rebuild reports and the run manifest from already generated audit tables."""
|
||
selection = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8"))
|
||
selected = selection["selected_method"]
|
||
selection_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "selection_results.csv"))
|
||
clean_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "main_results.csv"))
|
||
bootstrap_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "bootstrap_results.csv"))
|
||
structural_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "structural_audit.csv"))
|
||
owen_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "owen_audit.csv"))
|
||
fidelity_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "fidelity_results.csv"))
|
||
complexity_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "complexity.csv"))
|
||
local_rows = _coerce_csv_numbers(_read_csv(RESULTS_ROOT / "attachment4_local_evidence.csv"))
|
||
if not all((selection_rows, clean_rows, bootstrap_rows, structural_rows, owen_rows, fidelity_rows, complexity_rows)):
|
||
raise FileNotFoundError("evaluation tables are incomplete; run the full ATI–HO evaluation first")
|
||
shapley_summary = json.loads((RESULTS_ROOT / "shapley_audit_summary.json").read_text(encoding="utf-8"))
|
||
prediction_manifest = json.loads(
|
||
(RESULTS_ROOT / "attachment4_prediction_manifest.json").read_text(encoding="utf-8")
|
||
)
|
||
attachment_summary = prediction_manifest["shapley_audit"]
|
||
_plot_results(clean_rows, local_rows)
|
||
_write_reports(
|
||
selected,
|
||
selection_rows,
|
||
clean_rows,
|
||
bootstrap_rows,
|
||
structural_rows,
|
||
shapley_summary,
|
||
attachment_summary,
|
||
owen_rows,
|
||
fidelity_rows,
|
||
complexity_rows,
|
||
)
|
||
training_manifest = json.loads((EXPERIMENT_ROOT / "run_manifest.json").read_text(encoding="utf-8"))
|
||
device_name = next((row.get("device") for row in complexity_rows if row.get("method") == selected), "unknown")
|
||
_write_json(
|
||
RESULTS_ROOT / "run_manifest.json",
|
||
{
|
||
"selected_method": selected,
|
||
"candidate_selection": selection,
|
||
"device": device_name,
|
||
"torch_version": torch.__version__,
|
||
"cuda_version": torch.version.cuda,
|
||
"attachment4_cases": prediction_manifest["prediction_rows"],
|
||
"validation_samples": shapley_summary["samples"],
|
||
"validation_group_bootstrap_replicates": BOOTSTRAP_REPLICATES,
|
||
"adapter_and_scaler_metadata": training_manifest.get("data"),
|
||
"attachment4_source_hashes": prediction_manifest["attachment4_source_sha256"],
|
||
"labels_used_from_attachment4": False,
|
||
},
|
||
)
|
||
print(f"Q3 reports finalized from existing evaluation tables: selected={selected}", flush=True)
|
||
|
||
|
||
def run(device: torch.device) -> None:
|
||
RESULTS_ROOT.mkdir(parents=True, exist_ok=True)
|
||
train, valid, stats, data_meta = load_training_data()
|
||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||
selected, clean_rows, selection_rows = _write_seed_and_summary_tables()
|
||
candidate_methods = [row["method"] for row in selection_rows]
|
||
selected_models = _load_ensemble(selected, MODEL_SEEDS, dims, device)
|
||
early_models = _load_ensemble(EARLYCONCAT, MODEL_SEEDS, dims, device)
|
||
mofe_models = _load_ensemble(MOFE7_MLP, MODEL_SEEDS, dims, device)
|
||
|
||
predictions_by_method = {
|
||
selected: _predict_ensemble(selected_models, valid, valid.mask, device),
|
||
EARLYCONCAT: _predict_ensemble(early_models, valid, valid.mask, device),
|
||
MOFE7_MLP: _predict_ensemble(mofe_models, valid, valid.mask, device),
|
||
}
|
||
bootstrap_rows = _cluster_bootstrap(valid, predictions_by_method, selected)
|
||
_save_csv(RESULTS_ROOT / "bootstrap_results.csv", bootstrap_rows)
|
||
|
||
smoke_xs = tuple(torch.as_tensor(x[:64], dtype=torch.float32, device=device) for x in valid.x)
|
||
smoke_mask = torch.as_tensor(valid.mask[:64], dtype=torch.bool, device=device)
|
||
structural_rows = []
|
||
for seed, model in zip(MODEL_SEEDS, selected_models):
|
||
report = structural_audit(model, smoke_xs, smoke_mask)
|
||
structural_rows.append({"method": selected, "seed": seed, "samples": min(64, valid.n), **{
|
||
key: json.dumps(value, sort_keys=True) if isinstance(value, dict) else value for key, value in report.items()
|
||
}})
|
||
if not report["checks_pass"]:
|
||
raise RuntimeError(f"final ATI–HO structural audit failed for seed {seed}: {report}")
|
||
stage1_rows = _read_csv(EXPERIMENT_ROOT / "structural_audit.csv")
|
||
structural_rows.extend(stage1_rows)
|
||
_save_csv(RESULTS_ROOT / "structural_audit.csv", structural_rows)
|
||
|
||
validation_predictions, shapley_summary, validation_shapley_rows = _validation_predictions_and_shapley(
|
||
selected, selected_models, valid, device
|
||
)
|
||
_save_csv(RESULTS_ROOT / "shapley_audit.csv", validation_shapley_rows)
|
||
_write_json(RESULTS_ROOT / "shapley_audit_summary.json", shapley_summary)
|
||
|
||
cases, attachment_meta = _read_attachment4("unaligned_50")
|
||
attachment = _attachment_split(cases, stats)
|
||
attach_predictions, attach_explanations, attach_shapley_summary, attach_exact = _attachment_predictions_and_explanations(
|
||
selected, selected_models, cases, attachment, device
|
||
)
|
||
local_rows, owen_rows, fidelity_rows, comparison_rows, stability_rows = _owen_and_fidelity(
|
||
selected_models, early_models, mofe_models, cases, attachment, valid, device
|
||
)
|
||
_save_csv(RESULTS_ROOT / "owen_audit.csv", owen_rows)
|
||
_save_csv(RESULTS_ROOT / "attachment4_local_evidence.csv", local_rows)
|
||
_save_csv(RESULTS_ROOT / "attachment4_explanations.csv", attach_explanations)
|
||
_save_csv(RESULTS_ROOT / "attachment4_predictions.csv", attach_predictions)
|
||
_save_csv(RESULTS_ROOT / "fidelity_results.csv", fidelity_rows)
|
||
_save_csv(RESULTS_ROOT / "attachment4_comparison_owen.csv", comparison_rows)
|
||
_save_csv(RESULTS_ROOT / "stability_results.csv", stability_rows)
|
||
|
||
# Re-evaluate time cost on one complete Attachment 4 example; Owen times were captured above.
|
||
first_xs = tuple(torch.as_tensor(x[:1], dtype=torch.float32, device=device) for x in attachment.x)
|
||
first_mask = torch.as_tensor(attachment.mask[:1], dtype=torch.bool, device=device)
|
||
shapley_start = time.perf_counter()
|
||
exact_shapley_audit(selected_models, first_xs, first_mask, batch_size=8)
|
||
shapley_per_sample = time.perf_counter() - shapley_start
|
||
complexity_models = {selected: selected_models, EARLYCONCAT: early_models, MOFE7_MLP: mofe_models}
|
||
complexity_rows = _complexity_rows(
|
||
selected,
|
||
complexity_models,
|
||
first_xs,
|
||
first_mask,
|
||
[float(row.get("owen_seconds", 0.0)) for row in owen_rows if "owen_seconds" in row],
|
||
shapley_per_sample,
|
||
device,
|
||
)
|
||
# Owen time for each sample is also tracked by the detailed run table.
|
||
if owen_rows:
|
||
average_owen = float(np.mean([float(row.get("elapsed_seconds", 0.0)) for row in owen_rows]))
|
||
for row in complexity_rows:
|
||
if row["method"] == selected:
|
||
row["hierarchical_owen_seconds_per_sample_mean"] = average_owen
|
||
row["hierarchical_owen_forward_evaluations_mean"] = float(
|
||
np.mean([2 + 30 * int(item["permutations"]) for item in owen_rows])
|
||
)
|
||
_save_csv(RESULTS_ROOT / "complexity.csv", complexity_rows)
|
||
|
||
# Input hashes and model provenance. Attachment 4 has no target values.
|
||
model_hashes = {f"{method}/seed_{seed}": _sha256(_checkpoint_path(method, seed)) for method in (selected, EARLYCONCAT, MOFE7_MLP) for seed in MODEL_SEEDS}
|
||
source_hashes = {case["source_file"].name: case["source_sha256"] for case in cases}
|
||
prediction_manifest = {
|
||
"task": "Q3 ATI–HO Attachment 4 final inference",
|
||
"selected_method": selected,
|
||
"seeds": list(MODEL_SEEDS),
|
||
"input_version": "official unaligned_50 Attachment 4",
|
||
"adapter": "Q1AlignmentAdapter; Relative-Progress; target_steps=50",
|
||
"physical_time_alignment": False,
|
||
"no_attachment4_labels_or_metrics_used": True,
|
||
"scaler": str(SCALER_PATH),
|
||
"feature_dimensions": [768, 74, 35],
|
||
"class_order": list(CLASS_NAMES),
|
||
"intensity_decode": "predicted negative: -3*sigmoid(r_negative); neutral: 0; predicted positive: 3*sigmoid(r_positive)",
|
||
"model_checkpoint_sha256": model_hashes,
|
||
"attachment4_source_sha256": source_hashes,
|
||
"attachment4_feature_dir": attachment_meta.get("version_dir"),
|
||
"prediction_rows": len(attach_predictions),
|
||
"explanation_rows": len(attach_explanations),
|
||
"local_evidence_rows": len(local_rows),
|
||
"shapley_audit": attach_shapley_summary,
|
||
}
|
||
_write_json(RESULTS_ROOT / "attachment4_prediction_manifest.json", prediction_manifest)
|
||
|
||
# Chapter IV deliverables: prediction and explanation tables plus a concise provenance note.
|
||
SUBMIT_OUTPUT.mkdir(parents=True, exist_ok=True)
|
||
_save_csv(SUBMIT_OUTPUT / "attachment4_predictions.csv", attach_predictions)
|
||
_save_csv(SUBMIT_OUTPUT / "attachment4_explanations.csv", attach_explanations)
|
||
_save_csv(SUBMIT_OUTPUT / "attachment4_local_evidence.csv", local_rows)
|
||
_write_json(SUBMIT_OUTPUT / "attachment4_prediction_manifest.json", prediction_manifest)
|
||
readme = """# Q3 ATI–HO 提交输出
|
||
|
||
| 文件 | 内容 |
|
||
|---|---|
|
||
| `attachment4_predictions.csv` | 官方附件4的 20 条预测类别、强度与类别概率 |
|
||
| `attachment4_explanations.csv` | 主效应、pairwise 项、解析/精确分类 Shapley 与强度精确 Shapley |
|
||
| `attachment4_local_evidence.csv` | 按模态分组的局部 Hierarchical Owen 片段贡献、标准误与相对进度位置 |
|
||
| `attachment4_prediction_manifest.json` | adapter、模型权重哈希、文件计数和无标签推理审计 |
|
||
|
||
附件4没有标签,本目录不提供准确率或误差指标。所有位置均为归一化进度槽,不是秒数。
|
||
"""
|
||
(SUBMIT_OUTPUT / "README.md").write_text(readme, encoding="utf-8")
|
||
|
||
_plot_results(clean_rows, local_rows)
|
||
all_structural = [_read_csv(RESULTS_ROOT / "structural_audit.csv")]
|
||
_write_reports(
|
||
selected,
|
||
selection_rows,
|
||
clean_rows,
|
||
bootstrap_rows,
|
||
all_structural[0],
|
||
shapley_summary,
|
||
attach_shapley_summary,
|
||
owen_rows,
|
||
fidelity_rows,
|
||
complexity_rows,
|
||
)
|
||
_write_json(
|
||
RESULTS_ROOT / "run_manifest.json",
|
||
{
|
||
"selected_method": selected,
|
||
"candidate_selection": json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8")),
|
||
"device": str(device),
|
||
"torch_version": torch.__version__,
|
||
"cuda_version": torch.version.cuda,
|
||
"gpu": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||
"validation_samples": valid.n,
|
||
"validation_group_bootstrap_replicates": BOOTSTRAP_REPLICATES,
|
||
"attachment4_cases": len(cases),
|
||
"adapter_and_scaler_metadata": data_meta,
|
||
"attachment4_shapley_summary": attach_shapley_summary,
|
||
"attachment4_source_hashes": source_hashes,
|
||
"labels_used_from_attachment4": False,
|
||
},
|
||
)
|
||
print(
|
||
f"Q3 evaluation complete: selected={selected}; validation Macro-F1="
|
||
f"{next(row['macro_f1_mean'] for row in clean_rows if row['method'] == selected and row['scenario'] == 'clean'):.4f}; "
|
||
f"attachment4_cases={len(cases)}; Owen conservation max="
|
||
f"{max((abs(float(row['local_conservation_residual'])) for row in owen_rows), default=0.0):.3g}",
|
||
flush=True,
|
||
)
|
||
|
||
|
||
def main() -> None:
|
||
parser = argparse.ArgumentParser(description="Audit ATI–HO validation results and create Attachment 4 outputs.")
|
||
parser.add_argument("--device", default="auto")
|
||
parser.add_argument("--reports-only", action="store_true", help="rebuild reports after an interrupted final report write")
|
||
parser.add_argument("--stability-only", action="store_true", help="recompute and save seed/input attribution stability")
|
||
args = parser.parse_args()
|
||
if args.reports_only:
|
||
finalize_reports()
|
||
return
|
||
if args.device == "auto":
|
||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||
else:
|
||
device = torch.device(args.device)
|
||
if args.stability_only:
|
||
selected = json.loads((EXPERIMENT_ROOT / "final_selection.json").read_text(encoding="utf-8"))["selected_method"]
|
||
cases, _meta = _read_attachment4("unaligned_50")
|
||
stats = RobustStats.load(SCALER_PATH)
|
||
attachment = _attachment_split(cases, stats)
|
||
rows = _stability_diagnostics(selected, cases, attachment, device)
|
||
_save_csv(RESULTS_ROOT / "stability_results.csv", rows)
|
||
print(f"Q3 stability diagnostics complete: rows={len(rows)}", flush=True)
|
||
return
|
||
run(device)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|