623 lines
32 KiB
Python
623 lines
32 KiB
Python
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()
|