整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Q3 interpretation and evidence-localization experiments."""
|
||||
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import AutoModelForCTC, AutoTokenizer
|
||||
|
||||
|
||||
SAMPLE_RATE = 16_000
|
||||
MODEL_ID = "facebook/wav2vec2-base-960h"
|
||||
|
||||
|
||||
@dataclass
|
||||
class WordInterval:
|
||||
word: str
|
||||
start_s: float
|
||||
end_s: float
|
||||
quality: float
|
||||
valid: bool
|
||||
|
||||
|
||||
def decode_audio(video_path: Path) -> np.ndarray:
|
||||
result = subprocess.run(
|
||||
[
|
||||
"ffmpeg", "-v", "error", "-i", str(video_path), "-map", "0:a:0",
|
||||
"-ac", "1", "-ar", str(SAMPLE_RATE), "-f", "f32le", "pipe:1",
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
)
|
||||
waveform = np.frombuffer(result.stdout, dtype="<f4").copy()
|
||||
if not len(waveform):
|
||||
raise RuntimeError(f"no decoded audio in {video_path}")
|
||||
return waveform
|
||||
|
||||
|
||||
def _ctc_targets(words: list[str], tokenizer) -> tuple[list[int], list[list[int]]]:
|
||||
vocab = tokenizer.get_vocab()
|
||||
delimiter = int(tokenizer.convert_tokens_to_ids(tokenizer.word_delimiter_token or "|"))
|
||||
unknown = int(tokenizer.unk_token_id)
|
||||
targets: list[int] = []
|
||||
per_word: list[list[int]] = [[] for _ in words]
|
||||
for word_index, raw_word in enumerate(words):
|
||||
if word_index:
|
||||
targets.append(delimiter)
|
||||
normalized = re.sub(r"[^a-z']", "", raw_word.lower())
|
||||
for character in normalized:
|
||||
per_word[word_index].append(len(targets))
|
||||
targets.append(int(vocab.get(character, unknown)))
|
||||
return targets, per_word
|
||||
|
||||
|
||||
def _viterbi(log_probs: np.ndarray, targets: list[int], blank: int) -> np.ndarray | None:
|
||||
if not targets or log_probs.ndim != 2:
|
||||
return None
|
||||
states = np.full(2 * len(targets) + 1, blank, dtype=np.int64)
|
||||
states[1::2] = np.asarray(targets, dtype=np.int64)
|
||||
frames, count = log_probs.shape[0], len(states)
|
||||
if frames == 0 or frames < len(targets):
|
||||
return None
|
||||
previous = np.full(count, -np.inf, dtype=np.float64)
|
||||
previous[0] = float(log_probs[0, blank])
|
||||
previous[1] = float(log_probs[0, states[1]])
|
||||
back = np.zeros((frames, count), dtype=np.uint8)
|
||||
skip = np.zeros(count, dtype=bool)
|
||||
if count > 2:
|
||||
skip[2:] = (states[2:] != blank) & (states[2:] != states[:-2])
|
||||
for frame in range(1, frames):
|
||||
stay = previous
|
||||
one = np.full(count, -np.inf, dtype=np.float64)
|
||||
one[1:] = previous[:-1]
|
||||
two = np.full(count, -np.inf, dtype=np.float64)
|
||||
if skip.any():
|
||||
two[skip] = previous[np.flatnonzero(skip) - 2]
|
||||
candidates = np.stack((stay, one, two), axis=0)
|
||||
choice = candidates.argmax(axis=0).astype(np.uint8)
|
||||
previous = candidates[choice, np.arange(count)] + log_probs[frame, states]
|
||||
back[frame] = choice
|
||||
state = count - 1 if previous[-1] >= previous[-2] else count - 2
|
||||
path = np.empty(frames, dtype=np.int32)
|
||||
path[-1] = state
|
||||
for frame in range(frames - 1, 0, -1):
|
||||
state -= int(back[frame, state])
|
||||
path[frame - 1] = state
|
||||
return path
|
||||
|
||||
|
||||
def model_time_constants(model) -> tuple[float, float]:
|
||||
config = model.config
|
||||
stride = int(np.prod(config.conv_stride))
|
||||
receptive = 1
|
||||
jump = 1
|
||||
for kernel, local_stride in zip(config.conv_kernel, config.conv_stride):
|
||||
receptive += (int(kernel) - 1) * jump
|
||||
jump *= int(local_stride)
|
||||
return stride / SAMPLE_RATE, receptive / (2 * SAMPLE_RATE)
|
||||
|
||||
|
||||
def align_words(
|
||||
waveform: np.ndarray,
|
||||
words: list[str],
|
||||
tokenizer,
|
||||
model,
|
||||
device: torch.device,
|
||||
) -> list[WordInterval]:
|
||||
targets, word_targets = _ctc_targets(words, tokenizer)
|
||||
blank = int(tokenizer.pad_token_id)
|
||||
if not targets or not len(waveform):
|
||||
return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words]
|
||||
with torch.inference_mode():
|
||||
values = torch.as_tensor(waveform, dtype=torch.float32, device=device).unsqueeze(0)
|
||||
logits = model(input_values=values).logits[0].float()
|
||||
log_probs = torch.log_softmax(logits, dim=-1).cpu().numpy()
|
||||
path = _viterbi(log_probs, targets, blank)
|
||||
frame_step, center_s = model_time_constants(model)
|
||||
duration = len(waveform) / SAMPLE_RATE
|
||||
intervals: list[WordInterval] = []
|
||||
if path is None:
|
||||
return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words]
|
||||
for word, target_indices in zip(words, word_targets):
|
||||
states = np.asarray([2 * index + 1 for index in target_indices], dtype=np.int32)
|
||||
frame_indices = np.flatnonzero(np.isin(path, states)) if len(states) else np.empty(0, dtype=np.int64)
|
||||
if not len(frame_indices):
|
||||
intervals.append(WordInterval(word, float("nan"), float("nan"), 0.0, False))
|
||||
continue
|
||||
first, last = int(frame_indices[0]), int(frame_indices[-1])
|
||||
start = max(0.0, first * frame_step + center_s - frame_step / 2)
|
||||
end = min(duration, (last + 1) * frame_step + center_s - frame_step / 2)
|
||||
char_scores = []
|
||||
for target_index in target_indices:
|
||||
selected = np.flatnonzero(path == 2 * target_index + 1)
|
||||
if len(selected):
|
||||
char_scores.extend(log_probs[selected, targets[target_index]].tolist())
|
||||
quality = float(np.exp(np.mean(char_scores))) if char_scores else 0.0
|
||||
valid = end > start
|
||||
intervals.append(WordInterval(word, start, end, quality, valid))
|
||||
return intervals
|
||||
|
||||
|
||||
def load_ctc(device: torch.device):
|
||||
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
|
||||
model = AutoModelForCTC.from_pretrained(MODEL_ID).to(device).eval()
|
||||
return tokenizer, model
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user