1414 lines
70 KiB
Python
1414 lines
70 KiB
Python
"""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: ")
|
||
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"})",
|
||
"",
|
||
"## 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()
|