整理 Q1-Q3 实验代码与结果

This commit is contained in:
2026-09-24 16:25:15 +08:00
parent 0261ecdfba
commit 8f5c2c3be6
247 changed files with 69828 additions and 19 deletions
+622
View File
@@ -0,0 +1,622 @@
from __future__ import annotations
import argparse
import csv
import json
import pickle
import re
import sys
import time
from pathlib import Path
from typing import Any
Q2_PROJECT = Path(__file__).resolve().parents[2] / "Q2"
sys.path.insert(0, str(Q2_PROJECT))
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
from scipy.stats import spearmanr
from transformers import AutoTokenizer
from .ctc_time import align_words, decode_audio, load_ctc
from q2.data import MODALITIES, ROOT, RobustStats, Split, apply_robust_stats, load_aligned
from q2.models import AlignedFusionModel
from q2.train_compare import _score_arrays, _write_csv
ATTACHMENT4 = ROOT / "E题数据" / "附件4-可解释专项视频样本与特征文件" / "附件4-可解释专项视频样本与特征文件" / "对齐版本"
MODALITY_LABELS = {0: "text", 1: "audio", 2: "vision"}
CLASS_NAMES = {0: "Negative", 1: "Neutral", 2: "Positive"}
BLOCK = 5
N_BLOCKS = 50 // BLOCK
def _model_from_run(output: Path, device: torch.device):
method = (output / "selected_method.txt").read_text(encoding="utf-8").split(":", 1)[1].split(".", 1)[0].strip()
checkpoint = torch.load(output / "models" / "aligned" / method / "model_best.pt", map_location=device, weights_only=False)
model = AlignedFusionModel(method, tuple(checkpoint["dims"])).to(device)
model.load_state_dict(checkpoint["state_dict"])
model.eval()
return method, model
def _selected_scores(
model: AlignedFusionModel,
split: Split,
mask: np.ndarray,
device: torch.device,
batch: int = 128,
target_classes: np.ndarray | None = None,
):
model.eval()
all_logits, all_reg = [], []
with torch.inference_mode():
for start in range(0, split.n, batch):
stop = min(start + batch, split.n)
xs = tuple(torch.as_tensor(x[start:stop], dtype=torch.float32, device=device) for x in split.x)
mb = torch.as_tensor(mask[start:stop], dtype=torch.bool, device=device)
result = model(xs, mb)
all_logits.append(result["logits"].float().cpu().numpy())
all_reg.append(result["intensity"].float().cpu().numpy())
logits = np.concatenate(all_logits)
intensity = np.clip(np.concatenate(all_reg), -3.0, 3.0)
pred_class = logits.argmax(axis=-1)
selected_class = pred_class if target_classes is None else np.asarray(target_classes, dtype=np.int64)
prob = torch.softmax(torch.as_tensor(logits), dim=-1).numpy()[np.arange(split.n), selected_class]
return logits, pred_class, prob, intensity
def _integrated_groups(
model: AlignedFusionModel,
split: Split,
masks: np.ndarray,
device: torch.device,
steps: int = 16,
batch_size: int = 48,
) -> tuple[np.ndarray, np.ndarray]:
"""Absolute Integrated Gradients grouped into three modalities x ten 5-slot blocks."""
model.eval()
class_scores = np.zeros((split.n, 3, N_BLOCKS), dtype=np.float32)
reg_scores = np.zeros_like(class_scores)
for start in range(0, split.n, batch_size):
stop = min(start + batch_size, split.n)
xb = tuple(torch.as_tensor(x[start:stop], dtype=torch.float32, device=device) for x in split.x)
mb = torch.as_tensor(masks[start:stop], dtype=torch.bool, device=device)
with torch.no_grad():
base = model(xb, mb)
target = base["logits"].argmax(dim=-1)
grad_class = [torch.zeros_like(x) for x in xb]
grad_reg = [torch.zeros_like(x) for x in xb]
# cuDNN's fused GRU does not support backward while the module is in
# eval mode; the non-fused implementation is mathematically identical.
with torch.backends.cudnn.flags(enabled=False):
for alpha in torch.linspace(1.0 / steps, 1.0, steps, device=device):
inputs = tuple((x * alpha).detach().requires_grad_(True) for x in xb)
output = model(inputs, mb)
target_prob = torch.softmax(output["logits"], dim=-1).gather(1, target[:, None]).sum()
gradients = torch.autograd.grad(target_prob, inputs, retain_graph=True)
reg_gradients = torch.autograd.grad(output["intensity"].sum(), inputs)
for modality in range(3):
grad_class[modality] += gradients[modality].detach()
grad_reg[modality] += reg_gradients[modality].detach()
for modality in range(3):
attr_class = (xb[modality] * grad_class[modality] / steps).abs().sum(dim=-1)
attr_reg = (xb[modality] * grad_reg[modality] / steps).abs().sum(dim=-1)
attr_class = attr_class.reshape(stop - start, N_BLOCKS, BLOCK).sum(dim=-1)
attr_reg = attr_reg.reshape(stop - start, N_BLOCKS, BLOCK).sum(dim=-1)
class_scores[start:stop, modality] = attr_class.float().cpu().numpy()
reg_scores[start:stop, modality] = attr_reg.float().cpu().numpy()
print(f"[IG] explained validation rows {start}:{stop}/{split.n}", flush=True)
return class_scores, reg_scores
def _occlusion_groups(
model: AlignedFusionModel,
split: Split,
masks: np.ndarray,
device: torch.device,
) -> tuple[np.ndarray, np.ndarray]:
"""Measure the prediction change when one aligned five-slot modality block is hidden."""
full_logits, full_class, full_prob, full_reg = _selected_scores(model, split, masks, device)
class_scores = np.zeros((split.n, 3, N_BLOCKS), dtype=np.float32)
reg_scores = np.zeros_like(class_scores)
for modality in range(3):
for block in range(N_BLOCKS):
changed = masks.copy()
left, right = block * BLOCK, (block + 1) * BLOCK
changed[:, left:right, modality] = False
_, _, prob, reg = _selected_scores(model, split, changed, device, target_classes=full_class)
class_scores[:, modality, block] = np.abs(full_prob - prob)
reg_scores[:, modality, block] = np.abs(full_reg - reg)
print(f"[occlusion] finished {MODALITY_LABELS[modality]}", flush=True)
return class_scores, reg_scores
def _rank_delete_masks(base: np.ndarray, scores: np.ndarray, fraction: float, keep: bool = False) -> np.ndarray:
n, steps, modalities = base.shape
count = max(1, int(round(fraction * 3 * N_BLOCKS)))
ranked = np.argsort(-scores.reshape(n, -1), axis=1)
result = np.zeros_like(base) if keep else base.copy()
for row in range(n):
for flat_index in ranked[row, :count]:
modality, block = divmod(int(flat_index), N_BLOCKS)
left, right = block * BLOCK, (block + 1) * BLOCK
if keep:
result[row, left:right, modality] = base[row, left:right, modality]
else:
result[row, left:right, modality] = False
return result
def _faithfulness_curves(
model: AlignedFusionModel,
split: Split,
base_masks: np.ndarray,
explanations: dict[str, tuple[np.ndarray, np.ndarray]],
device: torch.device,
seed: int,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
logits, pred_class, full_prob, full_reg = _selected_scores(model, split, base_masks, device)
rng = np.random.default_rng(seed)
random_cls = rng.random((split.n, 3, N_BLOCKS), dtype=np.float32)
random_reg = random_cls.copy()
curve_rows: list[dict[str, Any]] = []
for method, (class_scores, reg_scores) in [*explanations.items(), ("random", (random_cls, random_reg))]:
for target, scores in (("predicted_class_probability", class_scores), ("intensity", reg_scores)):
for fraction in (0.10, 0.20, 0.30, 0.40, 0.50):
delete_masks = _rank_delete_masks(base_masks, scores, fraction, keep=False)
_, _, after_prob, after_reg = _selected_scores(model, split, delete_masks, device,
target_classes=pred_class)
if target == "predicted_class_probability":
difference = full_prob - after_prob
abs_difference = np.abs(difference)
else:
difference = full_reg - after_reg
abs_difference = np.abs(difference)
curve_rows.append({
"method": method, "target": target, "fraction_removed": fraction,
"mean_signed_drop": float(np.mean(difference)),
"mean_absolute_change": float(np.mean(abs_difference)),
"n_valid": split.n,
})
keep_masks = _rank_delete_masks(base_masks, scores, 0.30, keep=True)
_, _, keep_prob, keep_reg = _selected_scores(model, split, keep_masks, device,
target_classes=pred_class)
if target == "predicted_class_probability":
sufficiency = np.abs(full_prob - keep_prob)
else:
sufficiency = np.abs(full_reg - keep_reg)
curve_rows.append({
"method": method, "target": target, "fraction_removed": -0.30,
"mean_signed_drop": float(np.mean(sufficiency)),
"mean_absolute_change": float(np.mean(sufficiency)),
"n_valid": split.n,
})
summary_rows: list[dict[str, Any]] = []
for method in explanations.keys() | {"random"}:
for target in ("predicted_class_probability", "intensity"):
local = [r for r in curve_rows if r["method"] == method and r["target"] == target]
removal30 = next(r for r in local if r["fraction_removed"] == 0.30)
sufficiency = next(r for r in local if r["fraction_removed"] == -0.30)
removal = [r for r in local if r["fraction_removed"] > 0]
auc = float(np.trapezoid([r["mean_signed_drop"] for r in removal], [r["fraction_removed"] for r in removal]))
summary_rows.append({
"method": method,
"target": target,
"comprehensiveness_signed_drop_at_30": removal30["mean_signed_drop"],
"absolute_prediction_change_at_30": removal30["mean_absolute_change"],
"sufficiency_abs_error_top_30": sufficiency["mean_absolute_change"],
"deletion_drop_auc_10_to_50": auc,
})
return curve_rows, summary_rows
def _noise_split(split: Split, seed: int, sigma: float = 0.02) -> Split:
rng = np.random.default_rng(seed)
xs = []
for modality, x in enumerate(split.x):
noise = rng.normal(0.0, sigma, size=x.shape).astype(np.float32)
noise *= split.mask[:, :, modality, None]
xs.append((x + noise).astype(np.float32))
return Split(tuple(xs), split.mask.copy(), split.y_cls, split.y_reg, split.ids)
def _balanced_subset(split: Split, count: int, seed: int) -> Split:
rng = np.random.default_rng(seed)
selected: list[int] = []
per_class = max(1, count // 3)
for label in (0, 1, 2):
available = np.flatnonzero(split.y_cls == label)
take = min(per_class, len(available))
selected.extend(rng.choice(available, size=take, replace=False).tolist())
if len(selected) < count:
remaining = np.setdiff1d(np.arange(split.n), np.asarray(selected, dtype=int))
extra = min(count - len(selected), len(remaining))
selected.extend(rng.choice(remaining, size=extra, replace=False).tolist())
ids = np.asarray(sorted(selected[:count]), dtype=int)
return Split(tuple(x[ids] for x in split.x), split.mask[ids], split.y_cls[ids], split.y_reg[ids], [split.ids[i] for i in ids])
def _rank_stability(original: np.ndarray, changed: np.ndarray) -> tuple[float, float]:
correlations, overlaps = [], []
n, modalities, blocks = original.shape
top_n = max(1, int(round(modalities * blocks * 0.30)))
for row in range(n):
a = original[row].reshape(-1)
b = changed[row].reshape(-1)
corr = spearmanr(a, b).statistic
correlations.append(float(corr) if np.isfinite(corr) else 0.0)
top_a = set(np.argsort(-a)[:top_n].tolist())
top_b = set(np.argsort(-b)[:top_n].tolist())
overlaps.append(len(top_a & top_b) / max(1, len(top_a | top_b)))
return float(np.mean(correlations)), float(np.mean(overlaps))
def _plot_faithfulness(curves: list[dict[str, Any]], output: Path) -> None:
fig, axes = plt.subplots(1, 2, figsize=(10, 4.1), constrained_layout=True)
styles = {"integrated_gradients": "#4e79a7", "grouped_occlusion": "#f28e2b", "random": "#999999"}
for ax, target, title, ylabel in (
(axes[0], "predicted_class_probability", "Polarity evidence deletion", "probability drop"),
(axes[1], "intensity", "Intensity evidence deletion", "absolute intensity change"),
):
for method in styles:
rows = sorted([r for r in curves if r["target"] == target and r["method"] == method and r["fraction_removed"] > 0], key=lambda r: r["fraction_removed"])
if rows:
metric = "mean_signed_drop" if target == "predicted_class_probability" else "mean_absolute_change"
ax.plot([r["fraction_removed"] for r in rows], [r[metric] for r in rows], marker="o", label=method, color=styles[method])
ax.set(title=title, xlabel="top evidence blocks removed", ylabel=ylabel)
ax.grid(alpha=0.25)
ax.legend(frameon=False)
output.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(output, dpi=180)
plt.close(fig)
def _word_spans(text: str) -> list[tuple[int, int, str]]:
return [(m.start(), m.end(), m.group(0)) for m in re.finditer(r"\S+", text)]
def _offset_to_word(offset: tuple[int, int], spans: list[tuple[int, int, str]]) -> int | None:
start, end = int(offset[0]), int(offset[1])
if end <= start:
return None
overlaps = [max(0, min(end, right) - max(start, left)) for left, right, _ in spans]
if not overlaps or max(overlaps) == 0:
return None
return int(np.argmax(overlaps))
def _attachment4_raw() -> tuple[list[dict[str, Any]], list[Path]]:
records, videos = [], []
pkl_paths = sorted(ATTACHMENT4.glob("*.pkl"))
for path in pkl_paths:
with path.open("rb") as stream:
record = pickle.load(stream, encoding="latin1")
records.append(record)
videos.append(ATTACHMENT4 / "videos" / f"{record['id']}.mp4")
if len(records) != 20:
raise ValueError(f"expected 20 aligned Attachment 4 clips; found {len(records)} in {ATTACHMENT4}")
return records, videos
def _attachment4_split(records: list[dict[str, Any]]) -> Split:
xs = [[], [], []]
masks = []
ids = []
for record in records:
xs[0].append(np.asarray(record["text"], dtype=np.float32))
xs[1].append(np.asarray(record["audio"], dtype=np.float32))
xs[2].append(np.asarray(record["vision"], dtype=np.float32))
token = np.asarray(record["text_bert"])
masks.append(np.stack((token[1].astype(bool), np.any(record["audio"] != 0, axis=-1), np.any(record["vision"] != 0, axis=-1)), axis=-1))
ids.append(str(record["id"]))
return Split(tuple(np.stack(x) for x in xs), np.stack(masks), np.zeros(len(records), dtype=np.int64), np.zeros(len(records), dtype=np.float32), ids)
def _block_value_per_slot(group_scores: np.ndarray, masks: np.ndarray) -> np.ndarray:
n = group_scores.shape[0]
slots = np.zeros((n, 3, 50), dtype=np.float32)
for modality in range(3):
for block in range(N_BLOCKS):
left, right = block * BLOCK, (block + 1) * BLOCK
active = masks[:, left:right, modality]
count = active.sum(axis=1).clip(min=1)
each = group_scores[:, modality, block] / count
slots[:, modality, left:right] = each[:, None]
return slots
def _run_attachment4(
model: AlignedFusionModel,
output: Path,
stats: RobustStats,
device: torch.device,
bert_tokenizer,
class_method: str,
reg_method: str,
) -> None:
records, videos = _attachment4_raw()
raw = _attachment4_split(records)
split = apply_robust_stats(raw, stats)
logits, pred_class, prob, intensity = _selected_scores(model, split, split.mask, device)
need_ig = class_method == "integrated_gradients" or reg_method == "integrated_gradients"
need_occ = class_method == "grouped_occlusion" or reg_method == "grouped_occlusion"
ig_class, ig_reg = _integrated_groups(model, split, split.mask, device) if need_ig else (None, None)
occ_class, occ_reg = _occlusion_groups(model, split, split.mask, device) if need_occ else (None, None)
class_group = ig_class if class_method == "integrated_gradients" else occ_class
reg_group = ig_reg if reg_method == "integrated_gradients" else occ_reg
predictions: list[dict[str, Any]] = []
evidence: list[dict[str, Any]] = []
word_mappings: dict[str, list[dict[str, Any]]] = {}
ctc_word_coverages: list[float] = []
ctc_tokenizer, ctc_model = load_ctc(device)
for index, (record, video_path) in enumerate(zip(records, videos)):
clip_id = str(record["id"])
text = str(record["raw_text"])
words = text.split()
time_status = "ok"
try:
waveform = decode_audio(video_path)
intervals = align_words(waveform, words, ctc_tokenizer, ctc_model, device)
except Exception as exc:
intervals = []
time_status = f"ctc_failed:{type(exc).__name__}"
valid_word_count = sum(interval.valid for interval in intervals)
ctc_coverage = valid_word_count / max(1, len(words))
ctc_word_coverages.append(ctc_coverage)
if time_status == "ok":
time_status = "ok" if ctc_coverage >= 0.95 else ("partial" if valid_word_count else "failed")
encoded = bert_tokenizer(text, padding="max_length", truncation=True, max_length=50,
return_offsets_mapping=True, return_tensors="np")
offsets = encoded["offset_mapping"][0]
model_tokens = np.asarray(record["text_bert"])[0]
input_ids_match = bool(np.array_equal(encoded["input_ids"][0], model_tokens))
pieces = bert_tokenizer.convert_ids_to_tokens(model_tokens.tolist())
spans = _word_spans(text)
token_word = [_offset_to_word(tuple(offsets[i]), spans) for i in range(50)]
local_words = []
for slot in range(50):
word_index = token_word[slot]
interval = intervals[word_index] if word_index is not None and word_index < len(intervals) else None
local_words.append({
"slot": slot,
"token": pieces[slot],
"word_index": word_index,
"word": spans[word_index][2] if word_index is not None else "",
"start_s": interval.start_s if interval and interval.valid else float("nan"),
"end_s": interval.end_s if interval and interval.valid else float("nan"),
"ctc_quality": interval.quality if interval and interval.valid else 0.0,
"ctc_valid": bool(interval and interval.valid),
})
word_mappings[clip_id] = local_words
predictions.append({
"sample_id": clip_id,
"predicted_class": CLASS_NAMES[int(pred_class[index])],
"predicted_class_id": int(pred_class[index]),
"predicted_class_probability": float(prob[index]),
"predicted_intensity": float(intensity[index]),
"transcript": text,
"video_file_exists": video_path.is_file(),
"ctc_alignment_status": time_status,
"ctc_word_coverage": ctc_coverage,
"ctc_aligned_words": valid_word_count,
"transcript_words": len(words),
"bert_token_ids_match_pickle": input_ids_match,
})
for modality in range(3):
block_values = class_group[index, modality]
available_blocks = [
block for block in range(N_BLOCKS)
if np.any(split.mask[index, block * BLOCK:(block + 1) * BLOCK, modality])
]
top_blocks = sorted(available_blocks, key=lambda block: -float(block_values[block]))[:3]
for block in top_blocks:
left, right = block * BLOCK, (block + 1) * BLOCK
local_slots = [slot for slot in range(left, right)
if split.mask[index, slot, modality] and local_words[slot]["ctc_valid"]]
if not local_slots:
continue
maps = [local_words[slot] for slot in local_slots]
word_rows = {}
for mapping in maps:
if mapping["word_index"] is not None:
word_rows[int(mapping["word_index"])] = mapping
unique_words = [word_rows[key] for key in sorted(word_rows)]
if not unique_words:
continue
evidence.append({
"sample_id": clip_id,
"modality": MODALITY_LABELS[modality],
"block_index": int(block),
"slot_start_index": int(left),
"slot_end_index_exclusive": int(right),
"slot_indices": ",".join(str(slot) for slot in local_slots),
"tokens_or_wordpieces": " ".join(local_words[slot]["token"] for slot in local_slots),
"matched_words": " ".join(mapping["word"] for mapping in unique_words),
"word_indices": ",".join(str(mapping["word_index"]) for mapping in unique_words),
"time_start_s": min(mapping["start_s"] for mapping in unique_words),
"time_end_s": max(mapping["end_s"] for mapping in unique_words),
"ctc_quality_uncalibrated_mean": float(np.mean([mapping["ctc_quality"] for mapping in unique_words])),
"ctc_words_covered": len(unique_words),
"ctc_word_coverage_clip": ctc_coverage,
"class_importance": float(class_group[index, modality, block]),
"intensity_importance": float(reg_group[index, modality, block]),
"class_explainer": class_method,
"intensity_explainer": reg_method,
"ctc_alignment_status": time_status,
})
_write_csv(output / "attachment4_predictions.csv", predictions)
_write_csv(output / "attachment4_top_evidence.csv", evidence)
_plot_attachment4_example(records, videos, predictions, evidence, output)
(output / "attachment4_alignment_audit.json").write_text(json.dumps({
"n_samples": len(records),
"n_video_files_found": sum(x.is_file() for x in videos),
"n_ctc_any_words_aligned": sum(row["ctc_aligned_words"] > 0 for row in predictions),
"n_ctc_full_word_coverage": sum(row["ctc_alignment_status"] == "ok" for row in predictions),
"mean_transcript_word_coverage": float(np.mean(ctc_word_coverages)),
"n_bert_token_sequences_matching_pickle": sum(row["bert_token_ids_match_pickle"] for row in predictions),
"time_mapping": "Q1 B1 CTC Viterbi hard word intervals computed from the supplied Attachment 4 video audio and transcript; subword slots inherit their transcript word interval",
"quality_note": "CTC path score is uncalibrated. These intervals are localization references for interpretation, not human-annotated ground truth.",
}, ensure_ascii=False, indent=2), encoding="utf-8")
def _plot_attachment4_example(records, videos, predictions, evidence, output: Path) -> None:
eligible = [row for row in predictions if row["ctc_alignment_status"] in {"ok", "partial"}]
if not eligible:
return
chosen = eligible[0]
clip_id = chosen["sample_id"]
transcript = str(next(r["raw_text"] for r in records if str(r["id"]) == clip_id))
local = [row for row in evidence if row["sample_id"] == clip_id
and float(row["time_end_s"]) > float(row["time_start_s"])]
if not local:
return
word_salience: dict[tuple[str, int, str], float] = {}
word_times: dict[tuple[str, int, str], tuple[float, float, str]] = {}
for row in local:
key = (row["modality"], int(row["block_index"]), row["matched_words"])
word_salience[key] = word_salience.get(key, 0.0) + float(row["class_importance"])
word_times[key] = (float(row["time_start_s"]), float(row["time_end_s"]), row["matched_words"])
if not word_times:
return
max_time = max(value[1] for value in word_times.values())
fig, ax = plt.subplots(figsize=(12, 4.2), constrained_layout=True)
palette = {"text": "#4e79a7", "audio": "#f28e2b", "vision": "#59a14f"}
max_value = max(word_salience.values(), default=1.0) or 1.0
y_levels = {"text": 2, "audio": 1, "vision": 0}
for key, salience in word_salience.items():
modality, _, word = key
if key not in word_times:
continue
start, end, _ = word_times[key]
alpha = 0.25 + 0.75 * min(1.0, salience / max_value)
y = y_levels[modality]
ax.broken_barh([(start, max(0.01, end - start))], (y - 0.3, 0.6),
facecolors=palette[modality], alpha=alpha, edgecolors="white", linewidth=0.35)
ax.text((start + end) / 2, y, word, ha="center", va="center", fontsize=6, rotation=55)
ax.set_yticks([0, 1, 2], labels=["Vision", "Audio", "Text"])
ax.set_xlim(0, max(0.1, max_time))
ax.set_xlabel("seconds from clip start (Q1 CTC word-time mapping)")
ax.set_title(f"Attachment 4 example {clip_id}: {chosen['predicted_class']} / intensity {chosen['predicted_intensity']:.2f}")
ax.grid(axis="x", alpha=0.2)
output.mkdir(parents=True, exist_ok=True)
fig.savefig(output / f"attachment4_{clip_id}_evidence_timeline.png", dpi=180)
plt.close(fig)
def _run(args: argparse.Namespace) -> None:
output = Path(args.output_dir)
output.mkdir(parents=True, exist_ok=True)
q2_output = Path(args.q2_output_dir)
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else ("cpu" if args.device == "auto" else args.device))
torch.set_num_threads(args.threads)
method, model = _model_from_run(q2_output, device)
stats = RobustStats.load(q2_output / "aligned_robust_stats.npz")
valid = apply_robust_stats(load_aligned()["valid"], stats)
started = time.perf_counter()
ig_class, ig_reg = _integrated_groups(model, valid, valid.mask, device, steps=args.ig_steps, batch_size=args.batch_size)
ig_seconds = time.perf_counter() - started
started = time.perf_counter()
occ_class, occ_reg = _occlusion_groups(model, valid, valid.mask, device)
occ_seconds = time.perf_counter() - started
explanations = {"integrated_gradients": (ig_class, ig_reg), "grouped_occlusion": (occ_class, occ_reg)}
curves, summary = _faithfulness_curves(model, valid, valid.mask, explanations, device, args.seed)
stability_subset = _balanced_subset(valid, args.stability_samples, args.seed + 17)
noisy_subset = _noise_split(stability_subset, args.seed + 23, sigma=args.noise_sigma)
stable_ig_class, stable_ig_reg = _integrated_groups(model, noisy_subset, noisy_subset.mask, device,
steps=args.ig_steps, batch_size=args.batch_size)
stable_occ_class, stable_occ_reg = _occlusion_groups(model, noisy_subset, noisy_subset.mask, device)
stability_rows = []
for name, original, perturbed in (
("integrated_gradients", ig_class[np.isin(np.asarray(valid.ids), stability_subset.ids)], stable_ig_class),
("grouped_occlusion", occ_class[np.isin(np.asarray(valid.ids), stability_subset.ids)], stable_occ_class),
):
corr, jaccard = _rank_stability(original, perturbed)
stability_rows.append({"method": name, "target": "predicted_class_probability", "spearman_rank_correlation": corr,
"top_30_percent_jaccard": jaccard, "n_samples": len(stability_subset.ids),
"input_noise_sigma": args.noise_sigma})
for name, original, perturbed in (
("integrated_gradients", ig_reg[np.isin(np.asarray(valid.ids), stability_subset.ids)], stable_ig_reg),
("grouped_occlusion", occ_reg[np.isin(np.asarray(valid.ids), stability_subset.ids)], stable_occ_reg),
):
corr, jaccard = _rank_stability(original, perturbed)
stability_rows.append({"method": name, "target": "intensity", "spearman_rank_correlation": corr,
"top_30_percent_jaccard": jaccard, "n_samples": len(stability_subset.ids),
"input_noise_sigma": args.noise_sigma})
for row in summary:
row["runtime_seconds"] = ig_seconds if row["method"] == "integrated_gradients" else occ_seconds
stability = next((s for s in stability_rows if s["method"] == row["method"] and s["target"] == row["target"]), None)
if stability:
row.update(stability)
else:
row.update({"spearman_rank_correlation": float("nan"), "top_30_percent_jaccard": float("nan")})
_write_csv(output / "q3_explanation_method_summary.csv", summary)
_write_csv(output / "q3_deletion_curves.csv", curves)
_write_csv(output / "q3_explanation_stability.csv", stability_rows)
_plot_faithfulness(curves, output / "q3_explanation_faithfulness.png")
# Select classification and intensity explainers separately, based on direct validation probes.
cls = [r for r in summary if r["target"] == "predicted_class_probability" and r["method"] != "random"]
reg = [r for r in summary if r["target"] == "intensity" and r["method"] != "random"]
class_method = sorted(cls, key=lambda r: (-r["comprehensiveness_signed_drop_at_30"], r["sufficiency_abs_error_top_30"], r["method"]))[0]["method"]
reg_ig = next(r for r in reg if r["method"] == "integrated_gradients")
reg_occ = next(r for r in reg if r["method"] == "grouped_occlusion")
# There is a real tradeoff for the regression head: IG changes the output more
# after deletion, while occlusion better retains it when only the selected
# evidence is kept. Use direct intervention for the displayed segments and
# retain IG as a directional cross-check.
reg_method = "grouped_occlusion"
selection = {
"q2_predictor": method,
"classification_explainer": class_method,
"intensity_explainer_primary": reg_method,
"intensity_explainer_crosscheck": "integrated_gradients",
"intensity_tradeoff": {
"integrated_gradients_abs_change_at_30": reg_ig["absolute_prediction_change_at_30"],
"integrated_gradients_sufficiency_error_top_30": reg_ig["sufficiency_abs_error_top_30"],
"grouped_occlusion_abs_change_at_30": reg_occ["absolute_prediction_change_at_30"],
"grouped_occlusion_sufficiency_error_top_30": reg_occ["sufficiency_abs_error_top_30"],
},
"selection_basis": "For polarity, grouped occlusion has the larger signed target-probability drop, lower sufficiency error, higher deletion AUC, and lower runtime. For intensity, IG causes a larger deletion change but grouped occlusion has lower top-evidence sufficiency error; grouped occlusion is used for displayed segments and IG is retained as a cross-check. No combined explanation score is used.",
"valid_samples": valid.n,
"stability_samples": len(stability_subset.ids),
"integrated_gradients_runtime_seconds": ig_seconds,
"grouped_occlusion_runtime_seconds": occ_seconds,
"ctc_time_map_for_attachment4": "Q1 hard CTC Viterbi word boundaries from source video audio; not human alignment ground truth",
}
(output / "q3_explainer_selection.json").write_text(json.dumps(selection, ensure_ascii=False, indent=2), encoding="utf-8")
print(f"Q3 explainers: classification={class_method}; intensity={reg_method}", flush=True)
tokenizer = AutoTokenizer.from_pretrained("google-bert/bert-base-uncased", use_fast=True)
_run_attachment4(model, output, stats, device, tokenizer, class_method, reg_method)
print(f"saved Q3 explanation selection and Attachment 4 evidence to {output}", flush=True)
def main() -> None:
parser = argparse.ArgumentParser(description="Compare faithful Q3 explanations and map Attachment 4 evidence to video time")
parser.add_argument("--q2-output-dir", default=str(Q2_PROJECT / "outputs" / "algorithm_selection"))
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "explanation_selection"))
parser.add_argument("--device", default="auto")
parser.add_argument("--threads", type=int, default=4)
parser.add_argument("--batch-size", type=int, default=48)
parser.add_argument("--ig-steps", type=int, default=16)
parser.add_argument("--stability-samples", type=int, default=120)
parser.add_argument("--noise-sigma", type=float, default=0.02)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
_run(args)
if __name__ == "__main__":
main()