Files

1414 lines
70 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Run the first Q3 explanation comparison on the official Attachment 4 cases.
The runner reuses the Q2 EarlyConcat and MoFE checkpoints and their train-only
robust scaler. E1 and E2 are two explanation views of the same MoFE model.
"""
from __future__ import annotations
import argparse
import csv
import gc
import hashlib
import itertools
import json
import math
import pickle
import re
import shutil
import subprocess
import time
from pathlib import Path
from typing import Any, Iterable, Mapping
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
from scipy.stats import spearmanr
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, mean_squared_error
from adapter import Q1AlignmentAdapter, adapt_official_split
from data_paths import ATTACHMENT2, ATTACHMENT4, DATA_ROOT, PROJECT_ROOT
from model.early_concat import AlignedFusionModel
from model.mofe import EXPERT_NAMES, MixtureOfFusionExperts
from q2.deep_learning.q2.data import MODALITIES
SEED = 20260924
CLASS_NAMES = ("negative", "neutral", "positive")
MODEL_VARIANTS = (
("E0_EarlyConcat", "early_concat", "exact Shapley + multiscale occlusion"),
("E1_MoFE_Router", "mofe", "router weights, tested by counterfactual deletion"),
("E2_MoFE_Shapley", "mofe", "exact Shapley + multiscale occlusion"),
)
COALITIONS = tuple(
frozenset(c)
for size in range(4)
for c in itertools.combinations(range(3), size)
)
WINDOWS = (1, 3, 5)
EXPLANATION_FRACTION = 0.10
DEFAULT_RUN = PROJECT_ROOT / "experiments" / "q2" / "unaligned_deep_two_b128"
DEFAULT_EARLY = DEFAULT_RUN / "models" / "B0_early_concat" / "seed_20260924" / "model_best.pt"
DEFAULT_MOFE = DEFAULT_RUN / "models" / "B5_mofe_mlp" / "seed_20260924" / "model_best.pt"
DEFAULT_SCALER = DEFAULT_RUN / "unaligned_50_robust_stats.npz"
try:
from transformers import AutoTokenizer
except ImportError: # Token highlighting degrades gracefully; inference has no HF dependency.
AutoTokenizer = None # type: ignore[assignment,misc]
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 _decode(value: Any) -> str:
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
if isinstance(value, np.bytes_):
return bytes(value).decode("utf-8", errors="replace")
if isinstance(value, np.ndarray):
if value.shape == ():
return _decode(value.item())
return " ".join(_decode(x) for x in value.reshape(-1))
return str(value)
def _safe_name(value: str) -> str:
name = re.sub(r"[^A-Za-z0-9_.-]+", "_", value).strip("_.")
return name or "sample"
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
fields = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", newline="", encoding="utf-8-sig") as stream:
writer = csv.DictWriter(stream, fieldnames=fields, extrasaction="ignore")
writer.writeheader()
writer.writerows(rows)
def _write_json(path: Path, payload: Any) -> None:
path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False) + "\n", encoding="utf-8")
def _optional_float(value: Any) -> float | None:
value = float(value)
return value if math.isfinite(value) else None
def exact_shapley(values: Mapping[frozenset[int], float], player_count: int = 3) -> np.ndarray:
"""Exact Shapley values for a small finite coalition game."""
players = set(range(player_count))
result = np.zeros(player_count, dtype=np.float64)
denom = math.factorial(player_count)
for player in range(player_count):
others = sorted(players - {player})
for size in range(player_count):
for coalition_tuple in itertools.combinations(others, size):
coalition = frozenset(coalition_tuple)
weight = math.factorial(size) * math.factorial(player_count - size - 1) / denom
result[player] += weight * (
values[coalition | {player}] - values[coalition]
)
return result
def exact_pair_interactions(
values: Mapping[frozenset[int], float], player_count: int = 3
) -> dict[tuple[int, int], float]:
"""Shapley interaction index with the standard one-half pair coefficient."""
result: dict[tuple[int, int], float] = {}
for first, second in itertools.combinations(range(player_count), 2):
remaining = sorted(set(range(player_count)) - {first, second})
total = 0.0
for size in range(len(remaining) + 1):
for coalition_tuple in itertools.combinations(remaining, size):
coalition = frozenset(coalition_tuple)
weight = (
math.factorial(size)
* math.factorial(player_count - size - 2)
/ (2 * math.factorial(player_count - 1))
)
total += weight * (
values[coalition | {first, second}]
- values[coalition | {first}]
- values[coalition | {second}]
+ values[coalition]
)
result[(first, second)] = float(total)
return result
def _load_scaler(path: Path) -> tuple[tuple[np.ndarray, ...], tuple[np.ndarray, ...]]:
with np.load(path, allow_pickle=False) as archive:
centers = tuple(archive[f"{name}_center"].astype(np.float32) for name in MODALITIES)
scales = tuple(archive[f"{name}_scale"].astype(np.float32) for name in MODALITIES)
if any(np.any(~np.isfinite(scale)) or np.any(scale <= 0) for scale in scales):
raise ValueError(f"invalid robust scaler: {path}")
return centers, scales
def _scale_features(
features: tuple[np.ndarray, ...],
mask: np.ndarray,
centers: tuple[np.ndarray, ...],
scales: tuple[np.ndarray, ...],
) -> tuple[np.ndarray, ...]:
result: list[np.ndarray] = []
for index, source in enumerate(features):
values = (np.asarray(source, dtype=np.float32) - centers[index]) / scales[index]
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
visible = mask[:, index] if mask.ndim == 2 else mask[:, :, index]
values *= visible[..., None]
result.append(values.astype(np.float32, copy=False))
return tuple(result)
def _attachment4_paths(version: str) -> tuple[Path, Path]:
if version != "unaligned_50":
raise ValueError("Q3 uses the official unaligned_50 Attachment 4 features")
inner = ATTACHMENT4 / "附件4-可解释专项视频样本与特征文件"
version_dir = inner / "未对齐版本"
# The submitted archive places videos inside the feature-version folder.
video_dir = version_dir / "videos"
if not video_dir.is_dir():
video_dir = inner / "videos"
if not version_dir.is_dir():
raise FileNotFoundError(f"Attachment 4 feature folder not found: {version_dir}")
return version_dir, video_dir
def _media_duration(path: Path | None) -> float | None:
if path is None or not path.is_file():
return None
try:
proc = subprocess.run(
[
"ffprobe", "-v", "error", "-show_entries", "format=duration",
"-of", "default=noprint_wrappers=1:nokey=1", str(path),
],
check=True,
capture_output=True,
text=True,
timeout=20,
)
duration = float(proc.stdout.strip())
return duration if duration > 0 else None
except (OSError, subprocess.SubprocessError, ValueError):
return None
def _read_attachment4(version: str) -> tuple[list[dict[str, Any]], dict[str, str]]:
version_dir, video_dir = _attachment4_paths(version)
paths = sorted(version_dir.glob("*.pkl"), key=lambda p: p.name)
if len(paths) != 20:
raise FileNotFoundError(f"expected 20 Attachment 4 cases, found {len(paths)} in {version_dir}")
video_by_stem = {p.stem: p for p in video_dir.glob("*.mp4")} if video_dir.is_dir() else {}
adapter = Q1AlignmentAdapter(target_steps=50)
cases: list[dict[str, Any]] = []
for path in paths:
with path.open("rb") as stream:
raw = pickle.load(stream, encoding="latin1")
case_id = _decode(raw.get("id", path.stem)).strip() or path.stem
text_bert = np.asarray(raw["text_bert"], dtype=np.int64)
if text_bert.ndim == 3 and text_bert.shape[0] == 1:
text_bert = text_bert[0]
if text_bert.shape != (3, 50):
raise ValueError(f"{path.name}: expected text_bert shape (3,50), got {text_bert.shape}")
record = {
"id": case_id,
"sequence_order_verified": True,
"attention_mask": text_bert[1].astype(bool),
"text": np.asarray(raw["text"], dtype=np.float32),
"audio": np.asarray(raw["audio"], dtype=np.float32),
"vision": np.asarray(raw["vision"], dtype=np.float32),
"audio_length": int(np.asarray(raw["audio_lengths"]).reshape(-1)[0]),
"vision_length": int(np.asarray(raw["vision_lengths"]).reshape(-1)[0]),
}
aligned = adapter.align(record, mode="relative")
mask = np.stack([aligned.observed[name] for name in MODALITIES], axis=-1)
features = tuple(aligned.features[name].astype(np.float32) for name in MODALITIES)
video = video_by_stem.get(path.stem) or video_by_stem.get(case_id)
try:
video_relative = video.resolve().relative_to(DATA_ROOT).as_posix() if video is not None else ""
except ValueError:
video_relative = str(video.resolve()) if video is not None else ""
duration = _media_duration(video)
cases.append(
{
"case_id": case_id,
"source_file": path,
"source_sha256": _sha256(path),
"transcript": _decode(raw.get("raw_text", "")),
"text_bert": text_bert,
"features": features,
"mask": mask,
"target_intervals": aligned.target_intervals.astype(np.float32),
"provenance": aligned.provenance,
"video_file": video,
"video_path": video_relative,
"video_duration_sec": duration,
"input_audit": {
"case_id": case_id,
"source_file": path.name,
"source_sha256": _sha256(path),
"coordinate_mode": aligned.metadata["coordinate_mode"],
"physical_time_alignment": False,
"audio_reported_length": record["audio_length"],
"vision_reported_length": record["vision_length"],
"audio_length_conflict": bool(aligned.provenance["audio"].length_conflict),
"vision_length_conflict": bool(aligned.provenance["vision"].length_conflict),
"text_visible_target_slots": int(aligned.observed["text"].sum()),
"audio_visible_target_slots": int(aligned.observed["audio"].sum()),
"vision_visible_target_slots": int(aligned.observed["vision"].sum()),
"source_video": video_relative,
"source_video_duration_sec": _optional_float(duration) if video else None,
},
}
)
return cases, {"version_dir": str(version_dir), "video_dir": str(video_dir)}
def _build_model(kind: str, dims: tuple[int, int, int], checkpoint_path: Path, device: torch.device) -> torch.nn.Module:
checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
stored_dims = tuple(int(x) for x in checkpoint.get("dims", ()))
if stored_dims != dims:
raise ValueError(f"{checkpoint_path} expects {stored_dims}, adapter produced {dims}")
if kind == "early_concat":
model: torch.nn.Module = AlignedFusionModel("concat", dims=dims)
elif kind == "mofe":
model = MixtureOfFusionExperts(dims=dims)
else:
raise ValueError(f"unknown model kind: {kind}")
model.load_state_dict(checkpoint["state_dict"], strict=True)
model.to(device).eval()
return model
def _model_output(
model: torch.nn.Module,
xs: tuple[torch.Tensor, ...],
mask: torch.Tensor,
) -> dict[str, torch.Tensor]:
model.eval()
with torch.inference_mode():
return model(xs, mask)
def _batched_output(
model: torch.nn.Module,
features: tuple[np.ndarray, ...],
masks: np.ndarray,
device: torch.device,
batch_size: int,
) -> tuple[np.ndarray, np.ndarray]:
logits: list[np.ndarray] = []
intensity: list[np.ndarray] = []
model.eval()
with torch.inference_mode():
for start in range(0, len(masks), batch_size):
end = min(start + batch_size, len(masks))
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in features)
batch_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
output = model(xs, batch_mask)
logits.append(output["logits"].float().cpu().numpy())
intensity.append(output["intensity"].float().cpu().numpy())
return np.concatenate(logits), np.concatenate(intensity)
def _single_input_mask_outputs(
model: torch.nn.Module,
features: tuple[np.ndarray, ...],
masks: np.ndarray,
device: torch.device,
batch_size: int,
) -> tuple[np.ndarray, np.ndarray]:
logits: list[np.ndarray] = []
intensity: list[np.ndarray] = []
model.eval()
with torch.inference_mode():
for start in range(0, len(masks), batch_size):
end = min(start + batch_size, len(masks))
count = end - start
xs = tuple(
torch.as_tensor(np.repeat(x[None], count, axis=0), dtype=torch.float32, device=device)
for x in features
)
output = model(xs, torch.as_tensor(masks[start:end], dtype=torch.bool, device=device))
logits.append(output["logits"].float().cpu().numpy())
intensity.append(output["intensity"].float().cpu().numpy())
return np.concatenate(logits), np.concatenate(intensity).reshape(-1)
def _coalition_outputs(
model: torch.nn.Module,
features: tuple[np.ndarray, ...],
base_mask: np.ndarray,
device: torch.device,
) -> tuple[np.ndarray, np.ndarray]:
masks = np.repeat(base_mask[None, :, :], len(COALITIONS), axis=0)
for index, coalition in enumerate(COALITIONS):
for modality in range(3):
if modality not in coalition:
masks[index, :, modality] = False
xs = tuple(
torch.as_tensor(np.repeat(x[None, :, :], len(COALITIONS), axis=0), dtype=torch.float32, device=device)
for x in features
)
with torch.inference_mode():
output = model(xs, torch.as_tensor(masks, dtype=torch.bool, device=device))
return output["logits"].float().cpu().numpy(), output["intensity"].float().cpu().numpy().reshape(-1)
def _values_for_task(
coalition_logits: np.ndarray,
coalition_intensity: np.ndarray,
predicted_class: int,
) -> tuple[dict[frozenset[int], float], dict[frozenset[int], float]]:
class_values = {coalition: float(coalition_logits[i, predicted_class]) for i, coalition in enumerate(COALITIONS)}
reg_values = {coalition: float(coalition_intensity[i]) for i, coalition in enumerate(COALITIONS)}
return class_values, reg_values
def _shares(values: np.ndarray) -> np.ndarray:
denominator = float(np.abs(values).sum())
return np.abs(values) / denominator if denominator > 1e-12 else np.zeros_like(values)
def _source_rows(case: dict[str, Any], modality_index: int, slots: Iterable[int]) -> list[int]:
provenance = case["provenance"][MODALITIES[modality_index]]
rows: set[int] = set()
for slot in slots:
rows.update(int(x) for x in provenance.source_weights.getrow(int(slot)).indices)
return sorted(rows)
def _decode_text_rows(case: dict[str, Any], source_rows: list[int], tokenizer: Any) -> tuple[str, str]:
ids = np.asarray(case["text_bert"][0], dtype=np.int64)
usable = [index for index in source_rows if 0 <= index < len(ids)]
if not usable:
return case["transcript"], "whole_transcript_no_token_overlap"
if tokenizer is None:
return case["transcript"], "whole_transcript_tokenizer_unavailable"
selected = [int(ids[index]) for index in usable if int(ids[index]) not in tokenizer.all_special_ids]
if not selected:
return case["transcript"], "whole_transcript_special_tokens_only"
return tokenizer.decode(selected, skip_special_tokens=True, clean_up_tokenization_spaces=True), "bert_token_ids"
def _local_occlusion(
model: torch.nn.Module,
features: tuple[np.ndarray, ...],
base_mask: np.ndarray,
full_logit: float,
full_intensity: float,
predicted_class: int,
device: torch.device,
batch_size: int,
) -> tuple[dict[tuple[int, int], dict[str, float]], dict[int, np.ndarray]]:
masks: list[np.ndarray] = []
meta: list[tuple[int, int, int]] = []
for modality in range(3):
for width in WINDOWS:
radius = width // 2
for slot in np.flatnonzero(base_mask[:, modality]).tolist():
masked = base_mask.copy()
left, right = max(0, slot - radius), min(base_mask.shape[0], slot + radius + 1)
masked[left:right, modality] = False
masks.append(masked)
meta.append((modality, slot, width))
class_drop: dict[tuple[int, int, int], float] = {}
reg_drop: dict[tuple[int, int, int], float] = {}
if masks:
logits, intensity = _single_input_mask_outputs(model, features, np.stack(masks), device, batch_size)
for key, logit_row, regression in zip(meta, logits, intensity):
class_drop[key] = float(full_logit - logit_row[predicted_class])
reg_drop[key] = float(full_intensity - regression)
local: dict[tuple[int, int], dict[str, float]] = {}
scale_maps = {width: np.full((3, base_mask.shape[0]), np.nan, dtype=np.float32) for width in WINDOWS}
for modality in range(3):
for slot in np.flatnonzero(base_mask[:, modality]).tolist():
row: dict[str, float] = {}
for width in WINDOWS:
row[f"class_logit_drop_w{width}"] = class_drop.get((modality, slot, width), 0.0)
row[f"intensity_drop_w{width}"] = reg_drop.get((modality, slot, width), 0.0)
scale_maps[width][modality, slot] = row[f"class_logit_drop_w{width}"]
row["class_logit_drop_multiscale"] = float(np.mean([row[f"class_logit_drop_w{w}"] for w in WINDOWS]))
row["intensity_drop_multiscale"] = float(np.mean([row[f"intensity_drop_w{w}"] for w in WINDOWS]))
local[(modality, slot)] = row
return local, scale_maps
def _router_profile(model: torch.nn.Module, features: tuple[np.ndarray, ...], mask: np.ndarray, device: torch.device) -> tuple[dict[str, Any], np.ndarray]:
xs = tuple(torch.as_tensor(x[None], dtype=torch.float32, device=device) for x in features)
tensor_mask = torch.as_tensor(mask[None], dtype=torch.bool, device=device)
output = _model_output(model, xs, tensor_mask)
alpha = output.get("alpha")
utility = output.get("utility")
if alpha is None or utility is None:
raise TypeError("router profile requested for a model without MoFE router outputs")
alpha_np = alpha[0].float().cpu().numpy()
utility_np = utility[0].float().cpu().numpy()
valid_slots = mask.any(axis=-1)
if valid_slots.any():
expert_mean = alpha_np[valid_slots].mean(axis=0)
exposure = utility_np[valid_slots].mean(axis=0)
else:
expert_mean = np.zeros(len(EXPERT_NAMES), dtype=np.float32)
exposure = np.zeros(3, dtype=np.float32)
exposure_share = _shares(exposure)
row: dict[str, Any] = {}
for name, value in zip(EXPERT_NAMES, expert_mean):
row[f"router_expert_{name}"] = float(value)
for index, name in enumerate(MODALITIES):
row[f"router_{name}_exposure"] = float(exposure[index])
row[f"router_{name}_share"] = float(exposure_share[index])
return row, utility_np.T
def _spearman(x: np.ndarray, y: np.ndarray) -> float | None:
if len(x) < 2 or np.allclose(x, x[0]) or np.allclose(y, y[0]):
return None
return _optional_float(spearmanr(x, y).statistic)
def _scale_stability(scale_maps: dict[int, np.ndarray], mask: np.ndarray) -> dict[str, float | None]:
results: dict[str, float | None] = {}
values = [scale_maps[width][mask.T] for width in WINDOWS]
for (left, right), a, b in zip(((1, 3), (1, 5), (3, 5)), (values[0], values[0], values[1]), (values[1], values[2], values[2])):
results[f"spearman_w{left}_w{right}"] = _spearman(a, b)
finite = [value for value in results.values() if value is not None]
results["mean_scale_spearman"] = float(np.mean(finite)) if finite else None
return results
def _ranked_positions(profile: np.ndarray, mask: np.ndarray) -> list[tuple[int, int]]:
candidates = [
(modality, slot, float(profile[modality, slot]))
for modality in range(3)
for slot in range(mask.shape[0])
if mask[slot, modality] and math.isfinite(float(profile[modality, slot]))
]
return [(m, t) for m, t, _ in sorted(candidates, key=lambda item: (-item[2], item[0], item[1]))]
def _faithfulness(
model: torch.nn.Module,
features: tuple[np.ndarray, ...],
base_mask: np.ndarray,
full_logit: float,
predicted_class: int,
importance: np.ndarray,
device: torch.device,
) -> dict[str, Any]:
positions = _ranked_positions(importance, base_mask)
n = len(positions)
deletion_fractions = (0.0, 0.1, 0.2, 0.3, 0.5, 0.7)
masks: list[np.ndarray] = []
tags: list[tuple[str, float, int]] = []
for fraction in deletion_fractions:
count = min(n, int(math.ceil(n * fraction))) if fraction else 0
masked = base_mask.copy()
for modality, slot in positions[:count]:
masked[slot, modality] = False
masks.append(masked)
tags.append(("delete", fraction, count))
for fraction in (0.1, 0.2, 0.3):
count = min(n, max(1, int(math.ceil(n * fraction)))) if n else 0
retained = np.zeros_like(base_mask)
for modality, slot in positions[:count]:
retained[slot, modality] = True
masks.append(retained)
tags.append(("retain", fraction, count))
logits, _ = _single_input_mask_outputs(model, features, np.stack(masks), device, batch_size=32)
row: dict[str, Any] = {"observed_cells": n, "ranking_cells": n}
drops: list[float] = []
for tag, logit_row in zip(tags, logits):
operation, fraction, count = tag
score = float(logit_row[predicted_class])
change = float(full_logit - score)
if operation == "delete":
row[f"comprehensiveness_delete_{int(fraction * 100)}pct"] = change
row[f"deleted_cells_{int(fraction * 100)}pct"] = count
drops.append(change)
else:
row[f"sufficiency_gap_retain_{int(fraction * 100)}pct"] = change
row[f"sufficiency_abs_gap_retain_{int(fraction * 100)}pct"] = abs(change)
row[f"retained_cells_{int(fraction * 100)}pct"] = count
row["deletion_auc_0_70_mean_logit_drop"] = float(
np.trapezoid(np.asarray(drops, dtype=np.float64), x=np.asarray(deletion_fractions)) / 0.7
)
return row
def _segment_profile(
profile: np.ndarray,
mask: np.ndarray,
case: dict[str, Any],
variant: str,
tokenizer: Any,
output_dir: Path,
extract_frames: bool,
) -> list[dict[str, Any]]:
segments: list[dict[str, Any]] = []
for modality in range(3):
visible = np.flatnonzero(mask[:, modality])
if not len(visible):
continue
k = max(1, int(math.ceil(len(visible) * EXPLANATION_FRACTION)))
chosen = sorted(
visible.tolist(),
key=lambda slot: (-abs(float(profile[modality, slot])), slot),
)[:k]
groups: list[list[int]] = []
for slot in sorted(chosen):
if groups and slot == groups[-1][-1] + 1:
groups[-1].append(slot)
else:
groups.append([slot])
groups.sort(key=lambda group: (-sum(abs(float(profile[modality, t])) for t in group), group[0]))
for rank, group in enumerate(groups[:2], start=1):
rows = _source_rows(case, modality, group)
start = float(case["target_intervals"][group[0], 0])
end = float(case["target_intervals"][group[-1], 1])
duration = case["video_duration_sec"]
start_sec = start * duration if duration is not None else None
end_sec = end * duration if duration is not None else None
if modality == 0:
evidence, text_method = _decode_text_rows(case, rows, tokenizer)
elif modality == 1:
evidence = f"unaligned audio feature rows {min(rows) if rows else 0}–{max(rows) if rows else -1}; review the linked source clip at the estimated relative span"
text_method = "feature_row_provenance"
else:
evidence = f"unaligned vision feature rows {min(rows) if rows else 0}–{max(rows) if rows else -1}; candidate frame time is estimated from relative progress"
text_method = "feature_row_provenance"
signed = float(sum(float(profile[modality, t]) for t in group))
strength = float(sum(abs(float(profile[modality, t])) for t in group))
if variant == "E1_MoFE_Router":
direction = "router_activity_not_signed_contribution"
score_semantics = "internal router utility; not a prediction effect"
else:
direction = "supports_predicted_class" if signed > 0 else ("opposes_predicted_class" if signed < 0 else "neutral")
score_semantics = "predicted-class logit drop after local occlusion"
frame_path = ""
if modality == 2 and extract_frames and case["video_file"] is not None and start_sec is not None and end_sec is not None:
midpoint = (start_sec + end_sec) / 2.0
if duration:
midpoint = min(max(midpoint, 0.0), max(0.0, duration - 0.05))
name = f"{_safe_name(variant)}_{_safe_name(case['case_id'])}_vision_{rank}.jpg"
target = output_dir / "evidence_frames" / name
target.parent.mkdir(parents=True, exist_ok=True)
try:
subprocess.run(
["ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-ss", f"{midpoint:.4f}",
"-i", str(case["video_file"]), "-frames:v", "1", "-vf", "scale=640:-2", str(target)],
check=True,
capture_output=True,
timeout=30,
)
frame_path = target.relative_to(output_dir).as_posix()
except (OSError, subprocess.SubprocessError):
frame_path = ""
segments.append(
{
"variant": variant,
"case_id": case["case_id"],
"modality": MODALITIES[modality],
"rank_within_modality": rank,
"slot_start_0based": group[0],
"slot_end_exclusive": group[-1] + 1,
"relative_progress_start": start,
"relative_progress_end": end,
"source_row_start": min(rows) if rows else 0,
"source_row_end_exclusive": max(rows) + 1 if rows else 0,
"local_score_sum_signed": signed,
"local_score_mass": strength,
"direction": direction,
"score_semantics": score_semantics,
"evidence_text": evidence,
"text_evidence_method": text_method,
"source_video": case["video_path"],
"video_time_start_sec_estimate": start_sec,
"video_time_end_sec_estimate": end_sec,
"video_time_basis": "relative progress times clip duration; approximate, not a physical feature timestamp",
"candidate_frame": frame_path,
}
)
return segments
def _plot_profile(
path: Path,
profile: np.ndarray,
mask: np.ndarray,
title: str,
router: bool = False,
) -> None:
values = np.asarray(profile, dtype=np.float32).copy()
values[~mask.T] = np.nan
fig, ax = plt.subplots(figsize=(11, 2.8), constrained_layout=True)
cmap = plt.get_cmap("viridis" if router else "coolwarm").copy()
cmap.set_bad("#d8d8d8")
if router:
vmax = max(float(np.nanmax(values)) if np.isfinite(values).any() else 0.0, 1e-6)
image = ax.imshow(np.ma.masked_invalid(values), aspect="auto", interpolation="nearest", cmap=cmap, vmin=0, vmax=vmax)
else:
bound = max(float(np.nanmax(np.abs(values))) if np.isfinite(values).any() else 0.0, 1e-6)
image = ax.imshow(np.ma.masked_invalid(values), aspect="auto", interpolation="nearest", cmap=cmap, vmin=-bound, vmax=bound)
ax.set_yticks(range(3), ("Text", "Audio", "Vision"))
ax.set_xticks(range(0, 50, 5), range(0, 50, 5))
ax.set_xlabel("Relative progress bin (0-based)")
ax.set_title(title)
fig.colorbar(image, ax=ax, label="router utility" if router else "predicted-class logit drop")
path.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(path, dpi=150)
plt.close(fig)
def _metrics(y_cls: np.ndarray, y_reg: np.ndarray, logits: np.ndarray, intensity: np.ndarray) -> dict[str, Any]:
predicted = logits.argmax(axis=-1)
score = np.clip(intensity.reshape(-1), -3.0, 3.0)
pearson = float(np.corrcoef(y_reg, score)[0, 1]) if np.std(y_reg) > 0 and np.std(score) > 0 else None
return {
"n": int(len(y_cls)),
"accuracy": float(accuracy_score(y_cls, predicted)),
"macro_f1": float(f1_score(y_cls, predicted, labels=[0, 1, 2], average="macro", zero_division=0)),
"mae": float(mean_absolute_error(y_reg, score)),
"rmse": float(math.sqrt(mean_squared_error(y_reg, score))),
"pearson": pearson,
}
def _load_validation(
path: Path,
centers: tuple[np.ndarray, ...],
scales: tuple[np.ndarray, ...],
) -> tuple[list[str], tuple[np.ndarray, ...], np.ndarray, np.ndarray, np.ndarray, dict[str, Any]]:
if not path.is_file():
raise FileNotFoundError(f"Attachment 2 unaligned_50 pickle not found: {path}")
print(f"Loading the official validation split from {path} ...", flush=True)
with path.open("rb") as stream:
raw = pickle.load(stream, encoding="latin1")
part = raw["valid"]
raw_ids = part["id"]
ids = [_decode(value) for value in raw_ids]
y_cls = np.asarray(part["classification_labels"], dtype=np.int64).reshape(-1)
y_reg = np.asarray(part["regression_labels"], dtype=np.float32).reshape(-1)
features_dict, mask, audit = adapt_official_split(part)
features = tuple(features_dict[name] for name in MODALITIES)
del part, raw, raw_ids, features_dict
gc.collect()
features = _scale_features(features, mask, centers, scales)
return ids, features, mask, y_cls, y_reg, audit
def _validation_errors(
ids: list[str],
features: tuple[np.ndarray, ...],
mask: np.ndarray,
y_cls: np.ndarray,
y_reg: np.ndarray,
model_rows: dict[str, tuple[torch.nn.Module, np.ndarray, np.ndarray]],
device: torch.device,
batch_size: int,
output_dir: Path,
) -> tuple[dict[str, Any], int]:
predictions: list[dict[str, Any]] = []
error_indices: set[int] = set()
validation_metrics: dict[str, Any] = {}
cache: dict[str, tuple[np.ndarray, np.ndarray]] = {}
for model_name, (model, _unused_logits, _unused_intensity) in model_rows.items():
logits, intensity = _batched_output(model, features, mask, device, batch_size)
cache[model_name] = logits, intensity
validation_metrics[model_name] = _metrics(y_cls, y_reg, logits, intensity)
predicted = logits.argmax(axis=-1)
abs_error = np.abs(y_reg - np.clip(intensity, -3.0, 3.0))
misses = np.flatnonzero(predicted != y_cls)
error_indices.update(int(i) for i in misses)
top_reg = np.argsort(-abs_error)[: min(20, len(abs_error))]
error_indices.update(int(i) for i in top_reg)
for i, sample_id in enumerate(ids):
probabilities = torch.softmax(torch.as_tensor(logits[i]), dim=-1).numpy()
predictions.append(
{
"model": model_name,
"sample_id": sample_id,
"true_class": int(y_cls[i]),
"true_class_name": CLASS_NAMES[int(y_cls[i])],
"predicted_class": int(predicted[i]),
"predicted_class_name": CLASS_NAMES[int(predicted[i])],
"classification_correct": bool(predicted[i] == y_cls[i]),
"true_intensity": float(y_reg[i]),
"predicted_intensity": float(intensity[i]),
"absolute_intensity_error": float(abs_error[i]),
"p_negative": float(probabilities[0]),
"p_neutral": float(probabilities[1]),
"p_positive": float(probabilities[2]),
}
)
_write_csv(output_dir / "validation_predictions.csv", predictions)
_write_json(output_dir / "validation_metrics.json", validation_metrics)
sorted_errors = sorted(error_indices, key=lambda i: (y_cls[i] == cache["early_concat"][0][i].argmax(), -abs(float(y_reg[i] - cache["early_concat"][1][i]))))
error_rows: list[dict[str, Any]] = []
for i in sorted_errors:
for model_name, (logits, intensity) in cache.items():
predicted = int(logits[i].argmax())
error_rows.append(
{
"model": model_name,
"sample_id": ids[i],
"true_class_name": CLASS_NAMES[int(y_cls[i])],
"predicted_class_name": CLASS_NAMES[predicted],
"classification_correct": bool(predicted == y_cls[i]),
"true_intensity": float(y_reg[i]),
"predicted_intensity": float(intensity[i]),
"absolute_intensity_error": float(abs(float(y_reg[i] - intensity[i]))),
"error_selection": "classification error or top-20 intensity error",
}
)
_write_csv(output_dir / "validation_errors.csv", error_rows)
attribution: list[dict[str, Any]] = []
for i in sorted_errors[:50]:
for model_name, (model, _, _) in model_rows.items():
coal_logits, coal_intensity = _coalition_outputs(model, tuple(x[i] for x in features), mask[i], device)
predicted = int(cache[model_name][0][i].argmax())
true = int(y_cls[i])
if predicted != true:
alternative = predicted
target_name = "predicted_logit_minus_true_logit"
class_values = {
coalition: float(coal_logits[j, alternative] - coal_logits[j, true])
for j, coalition in enumerate(COALITIONS)
}
else:
alternatives = [c for c in range(3) if c != true]
alternative = max(alternatives, key=lambda c: float(cache[model_name][0][i, c]))
target_name = "true_logit_minus_best_alternative"
class_values = {
coalition: float(coal_logits[j, true] - coal_logits[j, alternative])
for j, coalition in enumerate(COALITIONS)
}
reg_values = {coalition: float(coal_intensity[j]) for j, coalition in enumerate(COALITIONS)}
class_phi = exact_shapley(class_values)
reg_phi = exact_shapley(reg_values)
row: dict[str, Any] = {
"model": model_name,
"sample_id": ids[i],
"true_class_name": CLASS_NAMES[true],
"predicted_class_name": CLASS_NAMES[predicted],
"classification_correct": bool(predicted == true),
"true_intensity": float(y_reg[i]),
"predicted_intensity": float(cache[model_name][1][i]),
"absolute_intensity_error": float(abs(float(y_reg[i] - cache[model_name][1][i]))),
"classification_margin_target": target_name,
"classification_margin_phi_sum_residual": float(class_phi.sum() - (class_values[COALITIONS[-1]] - class_values[frozenset()])),
"regression_shapley_sum_residual": float(reg_phi.sum() - (reg_values[COALITIONS[-1]] - reg_values[frozenset()])),
}
for m, name in enumerate(MODALITIES):
row[f"class_margin_phi_{name}"] = float(class_phi[m])
row[f"class_margin_share_{name}"] = float(_shares(class_phi)[m])
row[f"regression_phi_{name}"] = float(reg_phi[m])
row[f"regression_share_{name}"] = float(_shares(reg_phi)[m])
row["class_margin_dominant_modality"] = MODALITIES[int(np.argmax(np.abs(class_phi)))]
attribution.append(row)
_write_csv(output_dir / "validation_error_attribution.csv", attribution)
return validation_metrics, len(sorted_errors)
def _q2_reference_metrics(path: Path) -> dict[str, dict[str, float]]:
if not path.is_file():
return {}
rows: dict[str, dict[str, float]] = {}
with path.open("r", newline="", encoding="utf-8-sig") as stream:
for row in csv.DictReader(stream):
if row.get("split") != "official_valid":
continue
if row.get("model") not in ("EarlyConcat", "MoFE-7"):
continue
rows[row["model"]] = {
key: float(row[key]) for key in ("accuracy", "macro_f1", "mae", "rmse", "pearson")
}
return rows
def _make_cards(
output_dir: Path,
cases: list[dict[str, Any]],
results: list[dict[str, Any]],
shapley_by_key: dict[tuple[str, str], dict[str, Any]],
interaction_by_key: dict[tuple[str, str], dict[str, Any]],
router_by_id: dict[str, dict[str, Any]],
segment_rows: list[dict[str, Any]],
faithfulness_rows: list[dict[str, Any]],
) -> str:
case_by_id = {case["case_id"]: case for case in cases}
segments_by: dict[tuple[str, str], list[dict[str, Any]]] = {}
for row in segment_rows:
segments_by.setdefault((str(row["variant"]), str(row["case_id"])), []).append(row)
faith_by: dict[tuple[str, str], dict[str, Any]] = {
(str(row["variant"]), str(row["case_id"])): row for row in faithfulness_rows
}
result_by = {(str(row["variant"]), str(row["case_id"])): row for row in results}
card_paths: dict[tuple[str, str], Path] = {}
confidences = [float(row["confidence"]) for row in results if row["variant"] == "E2_MoFE_Shapley"]
median_conf = float(np.median(confidences))
typical_id = min(
(str(row["case_id"]) for row in results if row["variant"] == "E2_MoFE_Shapley"),
key=lambda cid: abs(result_by[("E2_MoFE_Shapley", cid)]["confidence"] - median_conf),
)
for variant, _, _ in MODEL_VARIANTS:
folder = output_dir / "explanation_cards" / variant
folder.mkdir(parents=True, exist_ok=True)
for case in cases:
cid = str(case["case_id"])
item = result_by[(variant, cid)]
base_model = "early_concat" if variant == "E0_EarlyConcat" else "mofe"
shap = shapley_by_key[(base_model, cid)]
inter = interaction_by_key[(base_model, cid)]
router = router_by_id.get(cid, {})
faithful = faith_by[(variant, cid)]
lines = [
f"# Q3 explanation card — {variant} — {cid}",
"",
"## Prediction",
"",
f"- Polarity: **{item['predicted_class_name']}**",
f"- Intensity: {item['predicted_sentiment']:+.3f}",
f"- Predicted-class confidence: {item['confidence']:.3f}",
f"- Source clip: {case['video_path'] or 'video not found'}",
"",
]
if variant == "E1_MoFE_Router":
lines += [
"## Router profile (intrinsic routing signal, not prediction contribution)",
"",
"| Modality | Router exposure share | Exact Shapley absolute share |",
"|---|---:|---:|",
]
for name in MODALITIES:
lines.append(
f"| {name} | {router.get(f'router_{name}_share', 0.0):.3f} | {shap[f'class_share_{name}']:.3f} |"
)
lines += [
"",
f"Router–Shapley Spearman: {router.get('router_shapley_spearman')}; top modality agreement: {router.get('router_shapley_top1_agreement')}.",
"Router values describe mixture routing. The counterfactual scores below test whether that routing signal tracks model behavior.",
"",
]
else:
lines += [
"## Exact modality Shapley",
"",
"Positive values support the predicted class logit; negative values oppose it. Shares use absolute values and are model decision contributions, not real-world emotion importance.",
"",
"| Modality | Class logit contribution | Absolute share | Intensity contribution | Absolute share |",
"|---|---:|---:|---:|---:|",
]
for name in MODALITIES:
lines.append(
f"| {name} | {shap[f'class_phi_{name}']:+.4f} | {shap[f'class_share_{name}']:.3f} | {shap[f'regression_phi_{name}']:+.4f} | {shap[f'regression_share_{name}']:.3f} |"
)
lines += [
"",
f"Shapley completeness residuals: class {shap['class_completeness_residual']:.2e}, intensity {shap['regression_completeness_residual']:.2e}.",
"## Pairwise Shapley interaction",
"",
"| Pair | Class logit | Intensity |",
"|---|---:|---:|",
f"| Text + Audio | {inter['class_interaction_TA']:+.4f} | {inter['regression_interaction_TA']:+.4f} |",
f"| Text + Vision | {inter['class_interaction_TV']:+.4f} | {inter['regression_interaction_TV']:+.4f} |",
f"| Audio + Vision | {inter['class_interaction_AV']:+.4f} | {inter['regression_interaction_AV']:+.4f} |",
"",
]
lines += [
"## Local evidence segments",
"",
"Local counterfactual scores are the predicted-class logit difference after hiding a 1/3/5-bin window, averaged equally across the three scales. E1 ranks its router utility and is evaluated separately.",
"",
]
evidence = segments_by.get((variant, cid), [])
for row in evidence:
segment = (
f"{row['modality']} bins {row['slot_start_0based']}–{int(row['slot_end_exclusive']) - 1} "
f"(relative progress {row['relative_progress_start']:.3f}–{row['relative_progress_end']:.3f})"
)
if row.get("video_time_start_sec_estimate") is not None:
segment += f", estimated clip interval {row['video_time_start_sec_estimate']:.2f}–{row['video_time_end_sec_estimate']:.2f}s"
lines.append(f"- **{segment}** — {row['direction']}; evidence: {row['evidence_text']}")
if row.get("candidate_frame"):
lines.append(f" - Candidate frame: ![estimated visual evidence](../../{row['candidate_frame']})")
if not evidence:
lines.append("- No observed local evidence cells for this sample.")
image_rel = Path("..") / ".." / "evidence_profiles" / variant / f"{_safe_name(cid)}.png"
lines += [
"",
f"![Local evidence profile]({image_rel.as_posix()})",
"",
"## Faithfulness checks",
"",
f"- Comprehensiveness after deleting the top 10% / 30% cells: {faithful['comprehensiveness_delete_10pct']:+.4f} / {faithful['comprehensiveness_delete_30pct']:+.4f} predicted-class logit.",
f"- Sufficiency gap when retaining the top 10% / 30%: {faithful['sufficiency_gap_retain_10pct']:+.4f} / {faithful['sufficiency_gap_retain_30pct']:+.4f}. Smaller absolute gaps are better.",
f"- Mean deletion logit drop over 0–70% deletion: {faithful['deletion_auc_0_70_mean_logit_drop']:+.4f}.",
"",
"## Provenance limit",
"",
"Attachment 4 supplies unaligned feature sequences without word/audio/frame timestamps. Feature rows are traced to source rows and normalized progress. Clip-time estimates multiply that progress by the video duration; they are approximate review locations, not physical alignment timestamps.",
"",
"Occlusion and Shapley values describe this trained model's response to masked inputs. They do not establish causal effects or prove the emotion expressed by a person.",
"",
"## Transcript",
"",
case["transcript"] or "(not supplied)",
"",
]
path = folder / f"{_safe_name(cid)}.md"
path.write_text("\n".join(lines), encoding="utf-8")
card_paths[(variant, cid)] = path
source = card_paths[("E2_MoFE_Shapley", typical_id)]
(output_dir / "typical_explanation_card.md").write_text(source.read_text(encoding="utf-8"), encoding="utf-8")
return typical_id
def _evaluate_case(
case: dict[str, Any],
models: dict[str, torch.nn.Module],
device: torch.device,
explanation_batch_size: int,
tokenizer: Any,
output_dir: Path,
extract_frames: bool,
) -> dict[str, Any]:
xs = case["features"]
mask = case["mask"]
result: dict[str, Any] = {"case_id": case["case_id"], "models": {}, "variant_rows": []}
local_by_model: dict[str, dict[tuple[int, int], dict[str, float]]] = {}
scale_by_model: dict[str, dict[int, np.ndarray]] = {}
router_by_model: dict[str, dict[str, Any]] = {}
for model_name, model in models.items():
outputs = _model_output(
model,
tuple(torch.as_tensor(x[None], dtype=torch.float32, device=device) for x in xs),
torch.as_tensor(mask[None], dtype=torch.bool, device=device),
)
logits = outputs["logits"][0].float().cpu().numpy()
intensity = float(outputs["intensity"][0].float().cpu().item())
probabilities = torch.softmax(outputs["logits"][0].float(), dim=-1).cpu().numpy()
predicted = int(logits.argmax())
class_values, reg_values = _values_for_task(
*_coalition_outputs(model, xs, mask, device),
predicted,
)
class_phi = exact_shapley(class_values)
reg_phi = exact_shapley(reg_values)
class_share = _shares(class_phi)
reg_share = _shares(reg_phi)
class_interactions = exact_pair_interactions(class_values)
reg_interactions = exact_pair_interactions(reg_values)
modal_row: dict[str, Any] = {
"base_model": model_name,
"case_id": case["case_id"],
"predicted_class_name": CLASS_NAMES[predicted],
"class_value_empty": class_values[frozenset()],
"class_value_full": class_values[COALITIONS[-1]],
"regression_value_empty": reg_values[frozenset()],
"regression_value_full": reg_values[COALITIONS[-1]],
"class_completeness_residual": float(class_phi.sum() - (class_values[COALITIONS[-1]] - class_values[frozenset()])),
"regression_completeness_residual": float(reg_phi.sum() - (reg_values[COALITIONS[-1]] - reg_values[frozenset()])),
}
for index, name in enumerate(MODALITIES):
modal_row[f"class_phi_{name}"] = float(class_phi[index])
modal_row[f"class_share_{name}"] = float(class_share[index])
modal_row[f"regression_phi_{name}"] = float(reg_phi[index])
modal_row[f"regression_share_{name}"] = float(reg_share[index])
interaction_row: dict[str, Any] = {"base_model": model_name, "case_id": case["case_id"]}
for key, label in (((0, 1), "TA"), ((0, 2), "TV"), ((1, 2), "AV")):
interaction_row[f"class_interaction_{label}"] = class_interactions[key]
interaction_row[f"regression_interaction_{label}"] = reg_interactions[key]
local, scale_maps = _local_occlusion(
model, xs, mask, float(logits[predicted]), intensity, predicted, device, explanation_batch_size
)
local_by_model[model_name] = local
scale_by_model[model_name] = scale_maps
router_row: dict[str, Any] | None = None
if model_name == "mofe":
router_row, router_map = _router_profile(model, xs, mask, device)
router_row["case_id"] = case["case_id"]
router_row["router_shapley_spearman"] = _spearman(
np.asarray([router_row[f"router_{name}_share"] for name in MODALITIES]), class_share
)
router_row["router_shapley_top1_agreement"] = bool(
np.argmax([router_row[f"router_{name}_share"] for name in MODALITIES]) == np.argmax(class_share)
)
router_by_model[model_name] = router_row
result.setdefault("router_local", {})[model_name] = router_map
result["models"][model_name] = {
"logits": logits,
"probabilities": probabilities,
"predicted_class": predicted,
"intensity": intensity,
"modal_row": modal_row,
"interaction_row": interaction_row,
"router_row": router_row,
"local": local,
"scale_maps": scale_maps,
}
result["variant_rows"].append(
{
"variant": "E0_EarlyConcat" if model_name == "early_concat" else "E1_MoFE_Router",
"base_model": model_name,
"case_id": case["case_id"],
"predicted_class": predicted,
"predicted_class_name": CLASS_NAMES[predicted],
"predicted_sentiment": intensity,
"confidence": float(probabilities[predicted]),
"p_negative": float(probabilities[0]),
"p_neutral": float(probabilities[1]),
"p_positive": float(probabilities[2]),
"source_video": case["video_path"],
"explanation_method": "exact Shapley + multiscale occlusion" if model_name == "early_concat" else "router profile; not itself a counterfactual contribution",
}
)
if model_name == "mofe":
result["variant_rows"].append(
{
"variant": "E2_MoFE_Shapley",
"base_model": model_name,
"case_id": case["case_id"],
"predicted_class": predicted,
"predicted_class_name": CLASS_NAMES[predicted],
"predicted_sentiment": intensity,
"confidence": float(probabilities[predicted]),
"p_negative": float(probabilities[0]),
"p_neutral": float(probabilities[1]),
"p_positive": float(probabilities[2]),
"source_video": case["video_path"],
"explanation_method": "exact Shapley + multiscale occlusion",
}
)
result["local_by_model"] = local_by_model
result["scale_by_model"] = scale_by_model
result["router_by_model"] = router_by_model
return result
def run(args: argparse.Namespace) -> None:
out_dir = args.output_dir.expanduser().resolve()
out_dir.mkdir(parents=True, exist_ok=True)
remaining = [p for p in out_dir.iterdir() if p.name != ".gitkeep"]
if remaining and not args.resume:
raise FileExistsError(f"output directory is not empty; choose a new path: {out_dir}")
if remaining:
print(f"Reusing existing Q3 output directory after a partial run: {out_dir}", flush=True)
started = time.time()
for path in (args.early_checkpoint, args.mofe_checkpoint, args.scaler):
if not path.is_file():
raise FileNotFoundError(f"Q2 model artifact not found: {path}")
centers, scales = _load_scaler(args.scaler)
cases, input_locations = _read_attachment4(args.attachment4_version)
dims = tuple(int(x.shape[-1]) for x in cases[0]["features"])
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else ("cpu" if args.device == "auto" else args.device))
if device.type == "cuda":
torch.cuda.manual_seed_all(SEED)
torch.set_float32_matmul_precision("high")
torch.set_num_threads(4)
models = {
"early_concat": _build_model("early_concat", dims, args.early_checkpoint, device),
"mofe": _build_model("mofe", dims, args.mofe_checkpoint, device),
}
for case in cases:
case["features"] = _scale_features(case["features"], case["mask"], centers, scales)
tokenizer = None
if AutoTokenizer is not None:
try:
tokenizer = AutoTokenizer.from_pretrained("google-bert/bert-base-uncased", use_fast=True, local_files_only=True)
except Exception:
tokenizer = None
prediction_rows: list[dict[str, Any]] = []
shapley_rows: list[dict[str, Any]] = []
interaction_rows: list[dict[str, Any]] = []
local_rows: list[dict[str, Any]] = []
router_rows: list[dict[str, Any]] = []
router_local_rows: list[dict[str, Any]] = []
faithfulness_rows: list[dict[str, Any]] = []
segment_rows: list[dict[str, Any]] = []
profiles_for_summary: dict[str, list[np.ndarray]] = {variant[0]: [] for variant in MODEL_VARIANTS}
result_for_cards: list[dict[str, Any]] = []
shapley_by_key: dict[tuple[str, str], dict[str, Any]] = {}
interaction_by_key: dict[tuple[str, str], dict[str, Any]] = {}
router_by_id: dict[str, dict[str, Any]] = {}
validation_models: dict[str, tuple[torch.nn.Module, np.ndarray, np.ndarray]] = {}
for index, case in enumerate(cases, start=1):
print(f"Q3 Attachment 4: explaining {case['case_id']} ({index}/{len(cases)})", flush=True)
evaluated = _evaluate_case(
case, models, device, args.explanation_batch_size, tokenizer, out_dir, not args.no_frames
)
for model_name in ("early_concat", "mofe"):
model_result = evaluated["models"][model_name]
shap_row = model_result["modal_row"]
interaction_row = model_result["interaction_row"]
shapley_rows.append(shap_row)
interaction_rows.append(interaction_row)
shapley_by_key[(model_name, case["case_id"])] = shap_row
interaction_by_key[(model_name, case["case_id"])] = interaction_row
if model_name == "mofe":
router_by_id[case["case_id"]] = model_result["router_row"]
router_rows.append(model_result["router_row"])
for (modality, slot), scores in model_result["local"].items():
provenance_rows = _source_rows(case, modality, [slot])
interval = case["target_intervals"][slot]
duration = case["video_duration_sec"]
local_rows.append(
{
"base_model": model_name,
"case_id": case["case_id"],
"modality": MODALITIES[modality],
"slot_0based": slot,
"relative_progress_start": float(interval[0]),
"relative_progress_end": float(interval[1]),
**scores,
"source_row_start": min(provenance_rows) if provenance_rows else 0,
"source_row_end_exclusive": max(provenance_rows) + 1 if provenance_rows else 0,
"source_video": case["video_path"],
"video_time_start_sec_estimate": float(interval[0] * duration) if duration else None,
"video_time_end_sec_estimate": float(interval[1] * duration) if duration else None,
}
)
stability = _scale_stability(model_result["scale_maps"], case["mask"])
local_map = np.full((3, 50), np.nan, dtype=np.float32)
for (modality, slot), scores in model_result["local"].items():
local_map[modality, slot] = scores["class_logit_drop_multiscale"]
if model_name == "early_concat":
variant = "E0_EarlyConcat"
importance = np.nan_to_num(local_map, nan=0.0)
router_display = False
else:
variant = "E2_MoFE_Shapley"
importance = np.nan_to_num(local_map, nan=0.0)
router_display = False
faith = _faithfulness(
models[model_name], case["features"], case["mask"],
float(model_result["logits"][model_result["predicted_class"]]),
int(model_result["predicted_class"]), np.abs(importance), device,
)
faith.update(
{
"variant": variant,
"case_id": case["case_id"],
"scale_stability_mean_spearman": stability["mean_scale_spearman"],
"scale_stability_w1_w3": stability["spearman_w1_w3"],
"scale_stability_w1_w5": stability["spearman_w1_w5"],
"scale_stability_w3_w5": stability["spearman_w3_w5"],
}
)
faithfulness_rows.append(faith)
segment_rows.extend(
_segment_profile(importance, case["mask"], case, variant, tokenizer, out_dir, not args.no_frames)
)
profiles_for_summary[variant].append(importance)
if model_name == "mofe":
router_map = evaluated["router_local"]["mofe"]
router_profile = np.nan_to_num(router_map, nan=0.0)
router_faith = _faithfulness(
models["mofe"], case["features"], case["mask"],
float(model_result["logits"][model_result["predicted_class"]]),
int(model_result["predicted_class"]), router_profile, device,
)
router_faith.update({"variant": "E1_MoFE_Router", "case_id": case["case_id"]})
faithfulness_rows.append(router_faith)
segment_rows.extend(
_segment_profile(router_profile, case["mask"], case, "E1_MoFE_Router", tokenizer, out_dir, not args.no_frames)
)
profiles_for_summary["E1_MoFE_Router"].append(router_profile)
_plot_profile(
out_dir / "evidence_profiles" / "E1_MoFE_Router" / f"{_safe_name(case['case_id'])}.png",
router_profile, case["mask"], f"E1 MoFE router utility — {case['case_id']}", router=True,
)
for modality in range(3):
for slot in range(50):
if case["mask"][slot, modality]:
router_local_rows.append(
{
"case_id": case["case_id"],
"modality": MODALITIES[modality],
"slot_0based": slot,
"router_utility": float(router_map[modality, slot]),
"relative_progress_start": float(case["target_intervals"][slot, 0]),
"relative_progress_end": float(case["target_intervals"][slot, 1]),
}
)
display_profile = router_profile if model_name == "mofe" and router_display else importance
_plot_profile(
out_dir / "evidence_profiles" / variant / f"{_safe_name(case['case_id'])}.png",
display_profile, case["mask"], f"{variant} — {case['case_id']}", router=router_display,
)
result_for_cards.extend(evaluated["variant_rows"])
for row in evaluated["variant_rows"]:
prediction_rows.append(row)
validation_models["early_concat"] = (
models["early_concat"],
evaluated["models"]["early_concat"]["logits"],
np.asarray([evaluated["models"]["early_concat"]["intensity"]]),
)
validation_models["mofe"] = (
models["mofe"],
evaluated["models"]["mofe"]["logits"],
np.asarray([evaluated["models"]["mofe"]["intensity"]]),
)
_write_csv(out_dir / "attachment4_predictions.csv", prediction_rows)
_write_csv(out_dir / "attachment4_modal_shapley.csv", shapley_rows)
_write_csv(out_dir / "attachment4_pairwise_interactions.csv", interaction_rows)
_write_csv(out_dir / "attachment4_local_evidence.csv", local_rows)
_write_csv(out_dir / "attachment4_router_profiles.csv", router_rows)
_write_csv(out_dir / "attachment4_router_local_evidence.csv", router_local_rows)
_write_csv(out_dir / "attachment4_evidence_segments.csv", segment_rows)
_write_csv(out_dir / "faithfulness_by_sample.csv", faithfulness_rows)
typical_id = _make_cards(
out_dir, cases, result_for_cards, shapley_by_key, interaction_by_key, router_by_id, segment_rows, faithfulness_rows
)
if not args.skip_validation:
validation_path = args.validation_data.expanduser().resolve()
valid = _load_validation(validation_path, centers, scales)
validation_metrics, validation_error_count = _validation_errors(
valid[0], valid[1], valid[2], valid[3], valid[4],
validation_models, device, args.validation_batch_size, out_dir,
)
del valid
else:
validation_metrics, validation_error_count = {}, 0
reference = _q2_reference_metrics(args.q2_validation_reference.expanduser().resolve())
summary_rows: list[dict[str, Any]] = []
for variant, model_name, method in MODEL_VARIANTS:
variant_faith = [row for row in faithfulness_rows if row["variant"] == variant]
mean = lambda key: float(np.mean([float(row[key]) for row in variant_faith if row.get(key) is not None])) if any(row.get(key) is not None for row in variant_faith) else None
val_name = "EarlyConcat" if model_name == "early_concat" else "MoFE-7"
q2_metrics = reference.get(val_name, {})
q3_val = validation_metrics.get(model_name, {})
summary_rows.append(
{
"variant": variant,
"backbone": "EarlyConcat + BiGRU" if model_name == "early_concat" else "MoFE-7 + MLP Router",
"explanation_method": method,
"validation_accuracy": q3_val.get("accuracy", q2_metrics.get("accuracy")),
"validation_macro_f1": q3_val.get("macro_f1", q2_metrics.get("macro_f1")),
"validation_mae": q3_val.get("mae", q2_metrics.get("mae")),
"validation_rmse": q3_val.get("rmse", q2_metrics.get("rmse")),
"validation_pearson": q3_val.get("pearson", q2_metrics.get("pearson")),
"attachment4_cases": len(cases),
"comprehensiveness_delete_10pct_mean": mean("comprehensiveness_delete_10pct"),
"comprehensiveness_delete_30pct_mean": mean("comprehensiveness_delete_30pct"),
"sufficiency_abs_gap_retain_10pct_mean": mean("sufficiency_abs_gap_retain_10pct"),
"sufficiency_abs_gap_retain_30pct_mean": mean("sufficiency_abs_gap_retain_30pct"),
"deletion_auc_0_70_mean_logit_drop": mean("deletion_auc_0_70_mean_logit_drop"),
"scale_stability_mean_spearman": mean("scale_stability_mean_spearman"),
"router_shapley_mean_spearman": (
float(np.mean([row["router_shapley_spearman"] for row in router_rows if row.get("router_shapley_spearman") is not None]))
if model_name == "mofe" and any(row.get("router_shapley_spearman") is not None for row in router_rows)
else None
),
"router_shapley_top1_agreement_rate": (
float(np.mean([bool(row["router_shapley_top1_agreement"]) for row in router_rows]))
if model_name == "mofe" else None
),
}
)
_write_csv(out_dir / "q3_method_comparison.csv", summary_rows)
completeness = [
abs(float(row[key]))
for row in shapley_rows
for key in ("class_completeness_residual", "regression_completeness_residual")
]
if completeness and max(completeness) > 1e-4:
raise AssertionError(f"exact Shapley completeness check failed: max residual {max(completeness)}")
manifest = {
"experiment": "Q3 first-round hierarchical counterfactual evidence attribution",
"created_at_unix": time.time(),
"elapsed_seconds": time.time() - started,
"seed": SEED,
"device": str(device),
"attachment4": input_locations,
"attachment4_version": args.attachment4_version,
"attachment4_cases": len(cases),
"coordinate_mode": "relative normalized progress",
"physical_time_alignment": False,
"time_mapping_limit": "estimated clip seconds equal normalized progress times video duration; original unaligned rows have no physical timestamps",
"models": {
"E0_EarlyConcat": {"checkpoint": str(args.early_checkpoint), "sha256": _sha256(args.early_checkpoint)},
"E1_E2_MoFE": {"checkpoint": str(args.mofe_checkpoint), "sha256": _sha256(args.mofe_checkpoint)},
},
"scaler": {"path": str(args.scaler), "sha256": _sha256(args.scaler), "fit": "Q2 official training rows only"},
"shapley": {
"coalitions": [sorted(x) for x in COALITIONS],
"class_value": "full-input predicted-class logit, fixed class across coalitions",
"regression_value": "predicted intensity",
"exact_enumeration": True,
"max_completeness_residual": max(completeness) if completeness else None,
},
"interaction": "pairwise Shapley interaction index; positive values indicate synergistic logit/intensity interaction under this convention",
"local_evidence": {"method": "leave out a 1, 3, or 5-bin contiguous window from one modality", "windows": list(WINDOWS), "scale_weights": [1 / 3] * 3},
"router_note": "MoFE router exposure is an internal routing summary, not prediction contribution; compare its rank with exact Shapley and deletion faithfulness.",
"faithfulness": {
"score": "fixed predicted-class logit",
"comprehensiveness": "full score minus score after deleting top-ranked cells",
"sufficiency": "full score minus score with only top-ranked cells retained",
"deletion_auc_fraction_range": [0.0, 0.7],
"mask_training_rates": [0.0, 0.1, 0.3, 0.5, 0.7],
"limit": "isolated sparse masks may still differ from the contiguous masks used in training; 90% deletion/10% retention is not claimed as in-distribution",
},
"validation_metrics": validation_metrics,
"validation_error_examples": validation_error_count,
"q2_validation_reference": reference,
"tokenizer_available": tokenizer is not None,
"outputs": [
"attachment4_predictions.csv", "attachment4_modal_shapley.csv",
"attachment4_pairwise_interactions.csv", "attachment4_local_evidence.csv",
"attachment4_router_profiles.csv", "attachment4_router_local_evidence.csv",
"attachment4_evidence_segments.csv", "faithfulness_by_sample.csv",
"q3_method_comparison.csv", "explanation_cards/", "evidence_profiles/",
"evidence_frames/", "typical_explanation_card.md",
] + ([] if args.skip_validation else ["validation_metrics.json", "validation_predictions.csv", "validation_errors.csv", "validation_error_attribution.csv"]),
"typical_explanation_case": typical_id,
}
_write_json(out_dir / "run_manifest.json", manifest)
print(f"Q3 complete: {len(cases)} Attachment 4 cases; outputs saved under {out_dir}", flush=True)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--attachment4-version", choices=("unaligned_50",), default="unaligned_50")
parser.add_argument("--early-checkpoint", type=Path, default=DEFAULT_EARLY)
parser.add_argument("--mofe-checkpoint", type=Path, default=DEFAULT_MOFE)
parser.add_argument("--scaler", type=Path, default=DEFAULT_SCALER)
parser.add_argument("--validation-data", type=Path, default=ATTACHMENT2 / "unaligned_50.pkl")
parser.add_argument("--q2-validation-reference", type=Path, default=PROJECT_ROOT / "output" / "q2" / "comparison_validation.csv")
parser.add_argument("--output-dir", type=Path, default=PROJECT_ROOT / "output" / "q3")
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--explanation-batch-size", type=int, default=128)
parser.add_argument("--validation-batch-size", type=int, default=128)
parser.add_argument("--skip-validation", action="store_true", help="Skip the labeled official validation split.")
parser.add_argument("--no-frames", action="store_true", help="Do not extract approximate candidate frames from source videos.")
parser.add_argument("--resume", action="store_true", help="Rerun into an existing partial output directory, overwriting this run's outputs.")
args = parser.parse_args()
run(args)
if __name__ == "__main__":
main()