Add Q3 MoFE router visualizations and explanations

This commit is contained in:
2026-09-25 23:46:20 +08:00
parent adc9c2064b
commit a86560da64
258 changed files with 17686 additions and 585 deletions
+67
View File
@@ -0,0 +1,67 @@
# Q3:分层反事实证据归因
## 第一轮比较
第一轮直接复用 Q2 官方 unaligned_50 实验中的 EarlyConcat + BiGRU、MoFE-7 + MLP Router 权重和 train-only robust scaler。这样 E0/E1/E2 使用固定预测器,主要比较解释方式。
| 方案 | 预测器 | 解释 |
|---|---|---|
| E0 | EarlyConcat + BiGRU | 三模态精确 Shapley、配对交互、多尺度局部遮蔽 |
| E1 | MoFE-7 + MLP Router | Router 模态/位置权重;用反事实删除检验其是否 faithful |
| E2 | 与 E1 相同的 MoFE 检查点 | 三模态精确 Shapley、配对交互、多尺度局部遮蔽 |
E1 与 E2 的预测逐样本相同。E1 的路由权重只描述融合机制,只有通过删除检验后才能说明它在这些样本上是否与预测行为一致。
## 运行
从项目根目录执行,附件目录由 FINAL_DATA_DIR 指定。它应包含 附件2-数据集特征文件/unaligned_50.pkl 和 附件4-可解释专项视频样本与特征文件/。
~~~bash
export FINAL_DATA_DIR="/path/to/task-data"
python -m final.q3.run_experiments \
--output-dir final/output/q3/first_round \
--device auto
~~~
检查点与 scaler 默认读取 final/experiments/q2/unaligned_deep_two_b128/。完整运行会计算官方验证集预测及误差归因。--skip-validation 跳过这一步;--no-frames 跳过候选帧抽取。每轮实验使用新的空输出目录;若上次运行中断且留下部分文件,可对该目录增加 --resume 重新生成结果。
## 方法定义
令三个模态为 (M={T,A,V})。对每个附件 4 样本完整计算 8 个 coalition。分类价值函数使用完整输入预测类别的 logit,并在所有 coalition 上固定该类别;回归价值函数使用模型情感强度输出。绝对 Shapley 贡献除以三模态绝对贡献之和,得到模态作用比例;带符号值保留支持或反对预测的方向。配对交互采用标准 Shapley interaction index 系数。
局部证据对每个可见模态位置计算窗口宽度 (win{1,3,5}) 的遮蔽前后差值,三个尺度等权平均。每个模态选取局部贡献绝对值最高的 10% 位置,合并相邻位置,并输出代表证据段。
Faithfulness 使用固定预测类别 logit。Comprehensiveness 比较完整输入与删去高排名证据后的分数;sufficiency 比较完整输入与只保留高排名证据后的分数;deletion AUC 汇总删除 0% 到 70% 的分数下降。另记录宽度 1/3/5 局部图之间的 Spearman 相关作为尺度稳定性诊断。MoFE Router 与精确 Shapley 按样本比较 Spearman 排序相关和主导模态一致率。
## 输出
- attachment4_predictions.csv:E0、E1、E2 对 20 个样本的完整预测与类别概率。
- attachment4_modal_shapley.csv:分类与强度的有符号贡献、绝对比例和完备性残差。
- attachment4_pairwise_interactions.csv:Text–Audio、Text–Vision、Audio–Vision 交互。
- attachment4_local_evidence.csv:E0/E2 的 1/3/5-bin 遮蔽差值和来源行。
- attachment4_router_profiles.csv、attachment4_router_local_evidence.csv:MoFE 的专家权重和逐位置 Router utility。
- attachment4_evidence_segments.csv:稀疏代表证据段、来源行、转写文本和候选视频时间。
- faithfulness_by_sample.csv、q3_method_comparison.csv:逐样本与方案汇总的删除/保留检查。
- validation_predictions.csv、validation_errors.csv、validation_error_attribution.csv:官方验证集指标、误差样本和分类 margin 归因。
- explanation_cards/、typical_explanation_card.md、evidence_profiles/:逐样本解释卡、代表卡和 3×50 热图。
- evidence_frames/:从原始视频抽取的候选帧。
- run_manifest.json:检查点、scaler、输入模式、公式、运行边界与产物清单。
## Router 热力图样例
在完成 Q3 第一轮后,还可以绘制附件 4 样本 02、03 的输入遮蔽与 Router 权重对照图。每个样本分别展示原始完整输入和受控残缺输入:样本 02 遮蔽 30% 文本与音频,样本 03 遮蔽 30% 视觉。蓝色斜线只标输入中被遮蔽的位置;紧接着一行把七个 expert(T、A、V、TA、TV、AV、TAV)横向排列,色块/数字显示样本平均 Router 权重,下方窄条显示 50 个 bin 上的 α[t,e]。
~~~bash
export FINAL_DATA_DIR="/path/to/task-data"
python -m final.q3.plot_router_heatmaps --output-dir final/output/q3
~~~
这两个残缺样例是人为遮蔽的对照,不是附件 4 的原生缺失数据。文本 token 按序列顺序显示,蓝色斜线标出对应遮蔽词段;音频波形和视频帧按相对进程显示蓝色遮蔽区。每个 bin 的七个 expert 权重在可用专家集合内归一化,未满足模态条件的 expert 权重为 0;色条使用原始 0–1 数值并通过平方根归一化提高低权重的可读性。样例顶部另列 T/A/V 的总体 Router exposure share。视频帧和音轨只按归一化进程投影,不能当作精确词/帧同步。Router 权重反映融合路由,不是预测贡献或情绪因果解释。
图表与原始数据输出为 `mofe_router_heatmap_examples.png`、`mofe_router_heatmap_examples.pdf`、`mofe_router_heatmap_scores.csv`(逐 bin 七个 α 分数及模态可见掩码)、`mofe_router_heatmap_summary.csv` 和 `mofe_router_heatmap_manifest.json`。
## 回溯边界
附件 4 的文本、音频和视觉序列为未对齐特征,没有逐词、逐音频帧或逐视频帧的真实时间戳。输出会从 adapter 的稀疏投影权重记录来源特征行和归一化进程。视频候选秒数由归一化进程乘视频时长估算,只供人工回看;它不是真实物理时间对齐。文本 token 需要本地缓存 google-bert/bert-base-uncased tokenizer;缓存不存在时,解释卡仍保留完整转写和来源行。
这些解释衡量的是当前预测器对输入遮蔽的响应,不是现实情绪成因。虽然模型在 Q2 训练时见过模态遮蔽,局部孤立遮蔽和只保留 10% 的输入仍可能偏离训练分布;相关数值按诊断结果报告,不称为解释准确率或因果效应。
+463
View File
@@ -0,0 +1,463 @@
"""Render input mask overlays followed by per-time weights for the seven MoFE experts.
Expert weights are internal routing signals, not causal or counterfactual
modality contribution scores.
"""
from __future__ import annotations
import argparse
import csv
import json
import subprocess
from pathlib import Path
from typing import Any
import matplotlib
matplotlib.use("Agg")
matplotlib.rcParams["font.family"] = "sans-serif"
matplotlib.rcParams["font.sans-serif"] = ["FandolHei", "DejaVu Sans"]
matplotlib.rcParams["axes.unicode_minus"] = False
import matplotlib.pyplot as plt
import numpy as np
import torch
from matplotlib.patches import Rectangle
from PIL import Image
from scipy import sparse
from . import run_experiments as q3
MODALITY_INDEX = {name: i for i, name in enumerate(q3.MODALITIES)}
CLASS_NAMES_ZH = {"negative": "负向", "neutral": "中性", "positive": "正向"}
MASK_COLOR = "#21a6df"
FRAME_PROGRESS = np.linspace(0.06, 0.94, 9)
SAMPLE_CONDITIONS = (
("02", "完整输入", {}),
("02", "残缺:Text / Audio 各遮蔽 30%", {"text": (5, 20), "audio": (18, 33)}),
("03", "完整输入", {}),
("03", "残缺:Vision 遮蔽 30%", {"vision": (18, 33)}),
)
def _get_tokenizer() -> Any:
if q3.AutoTokenizer is None:
return None
try:
return q3.AutoTokenizer.from_pretrained("google-bert/bert-base-uncased", local_files_only=True)
except (OSError, ValueError):
return None
def _load_audio(video_path: Path) -> tuple[np.ndarray, int] | None:
command = [
"ffmpeg", "-hide_banner", "-loglevel", "error", "-i", str(video_path),
"-map", "0:a:0", "-ac", "1", "-ar", "8000", "-f", "f32le", "pipe:1",
]
try:
proc = subprocess.run(command, check=True, capture_output=True, timeout=45)
wave = np.frombuffer(proc.stdout, dtype="<f4").copy()
if not len(wave):
return None
return wave, 8000
except (OSError, subprocess.SubprocessError):
return None
def _load_frame(video_path: Path, duration: float | None, progress: float) -> np.ndarray | None:
second = float(progress) * float(duration or 0.0)
if duration:
second = min(max(0.0, second), max(0.0, duration - 0.04))
command = [
"ffmpeg", "-hide_banner", "-loglevel", "error", "-ss", f"{second:.4f}",
"-i", str(video_path), "-frames:v", "1", "-vf", "scale=360:-2",
"-f", "image2pipe", "-vcodec", "mjpeg", "pipe:1",
]
try:
proc = subprocess.run(command, check=True, capture_output=True, timeout=30)
import io
return np.asarray(Image.open(io.BytesIO(proc.stdout)).convert("RGB"))
except (OSError, subprocess.SubprocessError, ValueError):
return None
def _source_words(case: dict[str, Any], tokenizer: Any) -> list[tuple[str, list[int]]]:
token_ids = np.asarray(case["text_bert"][0], dtype=np.int64)
attention = np.asarray(case["text_bert"][1], dtype=bool)
token_strings = tokenizer.convert_ids_to_tokens(token_ids.tolist())
words: list[tuple[str, list[int]]] = []
punctuation = set(".,!?;:%)]}’'\"…")
for source_index, token in enumerate(token_strings):
if not attention[source_index] or token in tokenizer.all_special_tokens:
continue
token = str(token)
if token.startswith("##") and words:
text, indices = words[-1]
words[-1] = (text + token[2:], indices + [source_index])
elif token in punctuation and words:
text, indices = words[-1]
words[-1] = (text + token, indices + [source_index])
else:
words.append((token, [source_index]))
return words
def _draw_text_mask(ax: plt.Axes, case: dict[str, Any], mask: np.ndarray, tokenizer: Any) -> None:
words = _source_words(case, tokenizer)
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.axis("off")
if not words:
ax.text(0.02, 0.52, "无可显示的有效文本 token", fontsize=8, color="#555")
return
x, y = 0.015, 0.78
char_width = 0.0125
line_height = 0.29
source_to_target = sparse.csc_matrix(case["provenance"]["text"].source_weights)
for word, source_indices in words:
label = word.replace("Ġ", "").replace("▁", "")
width = max(0.035, char_width * (len(label) + 1) + 0.016)
if x + width > 0.99:
x = 0.015
y -= line_height
if y < 0.05:
break
target_bins: set[int] = set()
for source_index in source_indices:
target_bins.update(
int(slot) for slot in source_to_target.getcol(source_index).tocoo().row
)
hidden = (
any(not mask[slot, MODALITY_INDEX["text"]] for slot in target_bins)
if target_bins
else False
)
face = MASK_COLOR if hidden else "#f0f2f5"
edge = "#0879ad" if hidden else "#c5cbd1"
ax.add_patch(
Rectangle((x, y - 0.12), width, 0.21, transform=ax.transAxes,
facecolor=face, edgecolor=edge, linewidth=0.45, clip_on=True,
hatch="//" if hidden else None)
)
ax.text(x + 0.008, y - 0.015, label, transform=ax.transAxes,
ha="left", va="center", fontsize=6.8, color="#fff" if hidden else "#111", clip_on=True)
x += width + 0.007
ax.text(0.015, 0.02, "文本按 token 顺序显示;蓝色斜线区域表示被遮蔽词段",
transform=ax.transAxes, fontsize=6.2, color="#555", va="bottom")
def _draw_wave_mask(ax: plt.Axes, wave: np.ndarray | None, mask: np.ndarray) -> None:
ax.set_xlim(0, 1)
ax.set_ylim(-1.08, 1.08)
if wave is None:
ax.text(0.5, 0.5, "视频中未读取到音轨", ha="center", va="center", transform=ax.transAxes, fontsize=8)
ax.set_xticks([])
ax.set_yticks([])
return
step = max(1, len(wave) // 14000)
values = wave[::step]
peak = float(np.percentile(np.abs(values), 99.5)) if len(values) else 0.0
if peak > 1e-8:
values = np.clip(values / peak, -1.0, 1.0)
x = np.linspace(0, 1, len(values), endpoint=False)
if len(values) >= 2:
ax.plot(x, values, color="#333a43", linewidth=0.45, rasterized=True, zorder=2)
for slot in np.flatnonzero(~mask[:, MODALITY_INDEX["audio"]]):
ax.axvspan(slot / 50, (slot + 1) / 50, facecolor=MASK_COLOR, alpha=0.20,
zorder=1, hatch="///", edgecolor="#0879ad", linewidth=0.0)
ax.axhline(0, color="#777", linewidth=0.45, alpha=0.5)
ax.set_yticks([])
ax.set_xticks([0, .25, .5, .75, 1.0])
ax.set_xticklabels(["0", ".25", ".50", ".75", "1"], fontsize=7)
ax.grid(axis="x", color="#bbb", alpha=0.25, linewidth=0.4)
def _draw_video_mask(
ax: plt.Axes,
frames: list[np.ndarray | None],
mask: np.ndarray,
) -> None:
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.set_yticks([])
ax.set_xticks([0, .25, .5, .75, 1.0])
ax.set_xticklabels(["0", ".25", ".50", ".75", "1"], fontsize=7)
ax.grid(axis="x", color="#bbb", alpha=0.25, linewidth=0.4)
n = len(FRAME_PROGRESS)
cell = 1.0 / n
for index, (progress, frame) in enumerate(zip(FRAME_PROGRESS, frames)):
left = index * cell + 0.012
width = cell - 0.024
slot = min(int(progress * 50), 49)
present = bool(mask[slot, MODALITY_INDEX["vision"]])
if frame is None:
ax.add_patch(Rectangle((left, .12), width, .76, transform=ax.transAxes,
facecolor="#d5d8dc", edgecolor="#888", linewidth=0.7))
ax.text(left + width / 2, .5, "无帧", transform=ax.transAxes, ha="center", va="center", fontsize=6)
else:
rgb = frame.astype(np.float32) / 255
if not present:
tint = np.asarray(matplotlib.colors.to_rgb(MASK_COLOR), dtype=np.float32)
rgb = np.clip(0.64 * rgb + 0.36 * tint, 0, 1)
ax.imshow(rgb, extent=(left, left + width, .12, .88), aspect="auto", origin="upper")
if not present:
ax.add_patch(Rectangle((left, .12), width, .76, transform=ax.transAxes,
facecolor=MASK_COLOR, alpha=0.16, edgecolor="none"))
ax.add_patch(Rectangle((left, .12), width, .76, transform=ax.transAxes,
fill=False, edgecolor="#6a7075" if present else "#0879ad",
linewidth=0.75 if present else 1.2,
hatch=None if present else "///"))
ax.text(left + width / 2, .04, f"{progress:.2f}", transform=ax.transAxes,
ha="center", va="bottom", fontsize=6, color="#555")
def _controlled_mask(base_mask: np.ndarray, missing_ranges: dict[str, tuple[int, int]]) -> np.ndarray:
mask = base_mask.copy()
for modality, (start, end) in missing_ranges.items():
index = MODALITY_INDEX[modality]
mask[start:end, index] = False
return mask
def _draw_router_weights(
ax: plt.Axes,
expert_weights: np.ndarray,
norm: Any,
cmap: Any,
) -> None:
ax.set_xlim(0, len(q3.EXPERT_NAMES))
ax.set_ylim(0, 1)
ax.axis("off")
for expert_index, expert_name in enumerate(q3.EXPERT_NAMES):
left = expert_index + 0.06
width = 0.88
mean_weight = float(expert_weights[:, expert_index].mean())
face = cmap(norm(mean_weight))
luminance = 0.2126 * face[0] + 0.7152 * face[1] + 0.0722 * face[2]
foreground = "#111" if luminance > 0.58 else "#fff"
ax.add_patch(Rectangle((left, 0.58), width, 0.31, facecolor=face,
edgecolor="#8a8f96", linewidth=0.55))
ax.text(left + width / 2, 0.79, expert_name, ha="center", va="center",
fontsize=8, color=foreground)
ax.text(left + width / 2, 0.65, f"{mean_weight:.3f}", ha="center", va="center",
fontsize=7, color=foreground)
for slot, value in enumerate(expert_weights[:, expert_index]):
cell_left = left + width * slot / expert_weights.shape[0]
ax.add_patch(Rectangle(
(cell_left, 0.22), width / expert_weights.shape[0], 0.18,
facecolor=cmap(norm(float(value))), edgecolor="white", linewidth=0.1,
))
def _mask_summary(mask: np.ndarray) -> dict[str, Any]:
return {
name: {
"visible_bins": int(mask[:, index].sum()),
"total_bins": int(mask.shape[0]),
"missing_fraction": float(1 - mask[:, index].mean()),
}
for index, name in enumerate(q3.MODALITIES)
}
def create_plot(args: argparse.Namespace) -> None:
torch.set_num_threads(4)
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available()
else ("cpu" if args.device == "auto" else args.device))
if device.type == "cuda":
torch.set_float32_matmul_precision("high")
output_dir = args.output_dir.expanduser().resolve()
output_dir.mkdir(parents=True, exist_ok=True)
centers, scales = q3._load_scaler(args.scaler.expanduser().resolve())
cases, _ = q3._read_attachment4("unaligned_50")
by_id = {case["case_id"]: case for case in cases}
needed = sorted({case_id for case_id, _, _ in SAMPLE_CONDITIONS})
missing = [case_id for case_id in needed if case_id not in by_id]
if missing:
raise KeyError(f"sample IDs not found in Attachment 4: {missing}")
dims = tuple(int(x.shape[-1]) for x in by_id[needed[0]]["features"])
model = q3._build_model("mofe", dims, args.checkpoint.expanduser().resolve(), device)
router_cmap = plt.get_cmap("magma")
tokenizer = _get_tokenizer()
if tokenizer is None:
raise RuntimeError("local google-bert/bert-base-uncased tokenizer is required to label text tokens")
case_cache: dict[str, dict[str, Any]] = {}
audio_cache: dict[str, np.ndarray | None] = {}
frame_cache: dict[tuple[str, int], np.ndarray | None] = {}
columns: list[dict[str, Any]] = []
score_rows: list[dict[str, Any]] = []
summaries: list[dict[str, Any]] = []
for case_id, condition, missing_ranges in SAMPLE_CONDITIONS:
case = case_cache.setdefault(case_id, by_id[case_id])
base_mask = case["mask"]
mask = _controlled_mask(base_mask, missing_ranges)
features = q3._scale_features(case["features"], base_mask, centers, scales)
output = q3._model_output(
model,
tuple(torch.as_tensor(x[None], dtype=torch.float32, device=device) for x in features),
torch.as_tensor(mask[None], dtype=torch.bool, device=device),
)
logits = output["logits"][0].float().cpu().numpy()
expert_weights = output["alpha"][0].float().cpu().numpy()
pred = int(logits.argmax())
confidence = float(torch.softmax(output["logits"][0].float(), dim=-1)[pred].cpu().item())
profile_row, utility = q3._router_profile(model, features, mask, device)
prediction_name = q3.CLASS_NAMES[pred]
modalities_summary = _mask_summary(mask)
if case_id not in audio_cache:
loaded_audio = _load_audio(case["video_file"]) if case["video_file"] else None
audio_cache[case_id] = loaded_audio[0] if loaded_audio is not None else None
wave = audio_cache[case_id]
frames = []
for frame_index, progress in enumerate(FRAME_PROGRESS):
key = (case_id, frame_index)
if key not in frame_cache:
frame_cache[key] = _load_frame(case["video_file"], case["video_duration_sec"], float(progress)) \
if case["video_file"] else None
frames.append(frame_cache[key])
columns.append({
"case": case,
"condition": condition,
"mask": mask,
"utility": utility,
"expert_weights": expert_weights,
"wave": wave,
"frames": frames,
"prediction": prediction_name,
"intensity": float(output["intensity"][0].float().cpu().item()),
"confidence": confidence,
"profile": profile_row,
})
for slot in range(mask.shape[0]):
score_row = {
"case_id": case_id,
"condition": condition,
"bin_0based": slot,
"relative_progress_start": slot / mask.shape[0],
"relative_progress_end": (slot + 1) / mask.shape[0],
}
for modality_index, modality in enumerate(q3.MODALITIES):
score_row[f"{modality}_visible"] = bool(mask[slot, modality_index])
score_row[f"router_utility_{modality}"] = (
float(utility[modality_index, slot]) if mask[slot, modality_index] else ""
)
for expert_index, expert_name in enumerate(q3.EXPERT_NAMES):
score_row[f"alpha_{expert_name}"] = float(expert_weights[slot, expert_index])
score_rows.append(score_row)
summary_row = {
"case_id": case_id,
"condition": condition,
"prediction": prediction_name,
"intensity": float(output["intensity"][0].float().cpu().item()),
"confidence": confidence,
"visible_bin_counts": modalities_summary,
"masked_ranges_0based_start_end_exclusive": missing_ranges,
"router_text_share": profile_row["router_text_share"],
"router_audio_share": profile_row["router_audio_share"],
"router_vision_share": profile_row["router_vision_share"],
}
summaries.append(summary_row)
router_norm = matplotlib.colors.PowerNorm(gamma=0.5, vmin=0.0, vmax=1.0, clip=True)
fig = plt.figure(figsize=(19.5, 12.8), facecolor="white")
grid = fig.add_gridspec(
nrows=4, ncols=4, left=0.065, right=0.92, top=0.86, bottom=0.14,
height_ratios=(0.95, 1.1, 1.35, 1.25), hspace=0.34, wspace=0.12,
)
for column_index, item in enumerate(columns):
case_id = item["case"]["case_id"]
shares = [item["profile"][f"router_{name}_share"] for name in q3.MODALITIES]
modality_label = " / ".join(f"{name[0].upper()} {share:.0%}" for name, share in zip(q3.MODALITIES, shares))
title = f"样本 {case_id}|{item['condition']}\n{CLASS_NAMES_ZH[item['prediction']]},置信度 {item['confidence']:.2f};Router 暴露 {modality_label}"
text_ax = fig.add_subplot(grid[0, column_index])
_draw_text_mask(text_ax, item["case"], item["mask"], tokenizer)
text_ax.set_title(title, loc="left", fontsize=9, pad=6, fontweight="normal")
wave_ax = fig.add_subplot(grid[1, column_index])
_draw_wave_mask(wave_ax, item["wave"], item["mask"])
if column_index == 0:
wave_ax.set_ylabel("Audio waveform\n归一化振幅", fontsize=7)
video_ax = fig.add_subplot(grid[2, column_index])
_draw_video_mask(video_ax, item["frames"], item["mask"])
if column_index == 0:
video_ax.set_ylabel("Vision frames", fontsize=7)
video_ax.set_xlabel("相对进程(非真实时间戳)", fontsize=7)
router_ax = fig.add_subplot(grid[3, column_index])
_draw_router_weights(router_ax, item["expert_weights"], router_norm, router_cmap)
for ax in (wave_ax, video_ax, router_ax):
ax.tick_params(axis="x", labelsize=7, pad=2)
fig.suptitle(
"先看输入遮蔽,再看 MoFE 七个 Router 专家的权重",
fontsize=16, fontweight="normal", y=0.96,
)
fig.text(
0.5, 0.91,
"蓝色斜线 = 输入被遮蔽;下方 T、A、V、TA、TV、AV、TAV 七张卡横排:大色块/数字是样本平均 α,窄条从左到右显示 50 个 bin 的 α[t,e]。",
ha="center", fontsize=9, color="#444",
)
fig.text(
0.5, 0.075,
"样本 02 残缺列遮蔽 Text bins 6–20 与 Audio bins 19–33;样本 03 残缺列遮蔽 Vision bins 19–33。每段均为 15/50 = 30%,完整列保留原始输入。",
ha="center", fontsize=8, color="#444",
)
fig.text(
0.5, 0.048,
"文本高亮按 token 顺序对应 bin;音频波形与视频帧按归一化进程展示。附件 4 无可信逐帧时间戳,Router 权重只描述专家路由。",
ha="center", fontsize=8, color="#666",
)
legend = Rectangle((0, 0), 1, 1, facecolor=MASK_COLOR, alpha=0.28,
edgecolor="#0879ad", hatch="///", label="输入被遮蔽")
fig.legend(handles=[legend], loc="lower center", bbox_to_anchor=(0.5, 0.105),
frameon=False, ncol=1, fontsize=8)
scalar = matplotlib.cm.ScalarMappable(norm=router_norm, cmap=router_cmap)
cax = fig.add_axes([0.935, 0.24, 0.018, 0.42])
colorbar = fig.colorbar(scalar, cax=cax)
colorbar.set_label("Expert router weight α[t,e]", fontsize=8)
colorbar.set_ticks(np.linspace(0, 1, 5))
colorbar.ax.tick_params(labelsize=7)
png_path = output_dir / "mofe_router_heatmap_examples.png"
pdf_path = output_dir / "mofe_router_heatmap_examples.pdf"
fig.savefig(png_path, dpi=300, bbox_inches="tight")
fig.savefig(pdf_path, bbox_inches="tight")
plt.close(fig)
def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
with path.open("w", newline="", encoding="utf-8-sig") as stream:
writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
writer.writeheader()
writer.writerows(rows)
write_csv(output_dir / "mofe_router_heatmap_summary.csv", summaries)
write_csv(output_dir / "mofe_router_heatmap_scores.csv", score_rows)
manifest = {
"source": "Attachment 4 unaligned_50; Q2 MoFE seed_20260924 checkpoint and train-only robust scaler",
"device": str(device),
"router_weight_definition": "alpha[t,e], the MoFE softmax weight assigned to each of the seven experts at a relative-progress bin; weights normalize over eligible experts for each bin",
"router_heatmap_color_scale_min": 0.0,
"router_heatmap_color_scale_max": 1.0,
"router_heatmap_color_norm": "PowerNorm gamma=0.5; colorbar ticks retain raw alpha values",
"conditions": summaries,
"controlled_missing_note": "The partial examples are controlled masks, not native Attachment 4 missingness. The complete examples retain all 50 bins in all modalities.",
"relative_time_limit": "The unaligned rows have no physical timestamps; frames and waveform are shown against normalized progress only.",
"outputs": [png_path.name, pdf_path.name, "mofe_router_heatmap_summary.csv", "mofe_router_heatmap_scores.csv"],
}
(output_dir / "mofe_router_heatmap_manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
print(json.dumps({"output_dir": str(output_dir), "summary": summaries, "files": manifest["outputs"]},
ensure_ascii=False, indent=2))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--output-dir", type=Path, default=q3.PROJECT_ROOT / "output" / "q3")
parser.add_argument("--checkpoint", type=Path, default=q3.DEFAULT_MOFE)
parser.add_argument("--scaler", type=Path, default=q3.DEFAULT_SCALER)
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
return parser.parse_args()
if __name__ == "__main__":
create_plot(parse_args())
File diff suppressed because it is too large Load Diff
+52
View File
@@ -0,0 +1,52 @@
from __future__ import annotations
import itertools
import unittest
import numpy as np
from .run_experiments import COALITIONS, exact_pair_interactions, exact_shapley, _scale_features
class ExactAttributionTests(unittest.TestCase):
def test_additive_game_shapley_and_interaction(self) -> None:
weights = (1.25, -0.5, 2.0)
values = {
coalition: 3.0 + sum(weights[player] for player in coalition)
for coalition in COALITIONS
}
np.testing.assert_allclose(exact_shapley(values), weights, atol=1e-12)
for interaction in exact_pair_interactions(values).values():
self.assertAlmostEqual(interaction, 0.0, places=12)
def test_pair_interaction_is_reported_with_standard_half_weight(self) -> None:
values = {}
for coalition in COALITIONS:
value = float(len(coalition))
if 0 in coalition and 1 in coalition:
value += 2.0
values[coalition] = value
interactions = exact_pair_interactions(values)
self.assertAlmostEqual(interactions[(0, 1)], 1.0, places=12)
self.assertAlmostEqual(interactions[(0, 2)], 0.0, places=12)
self.assertAlmostEqual(interactions[(1, 2)], 0.0, places=12)
def test_robust_scaling_supports_single_cases_and_validation_batches(self) -> None:
features = (np.full((2, 4, 2), 3.0, np.float32),)
centers = (np.asarray([1.0, 1.0], np.float32),)
scales = (np.asarray([2.0, 2.0], np.float32),)
single_mask = np.ones((4, 1), dtype=bool)
single = _scale_features((features[0][0],), single_mask, centers, scales)[0]
self.assertEqual(single.shape, (4, 2))
np.testing.assert_allclose(single, 1.0)
batch_mask = np.ones((2, 4, 1), dtype=bool)
batch_mask[1, 2:, 0] = False
batch = _scale_features(features, batch_mask, centers, scales)[0]
self.assertEqual(batch.shape, (2, 4, 2))
np.testing.assert_allclose(batch[0], 1.0)
np.testing.assert_allclose(batch[1, :2], 1.0)
np.testing.assert_allclose(batch[1, 2:], 0.0)
if __name__ == "__main__":
unittest.main()
+2 -565
View File
@@ -1,569 +1,6 @@
"""Train Q3 on the official training split and explain Attachment 4 cases."""
from __future__ import annotations
"""Compatibility entry point for the first Q3 explanation experiment."""
import argparse
import csv
import hashlib
import json
import math
import random
import time
from pathlib import Path
from typing import Any
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import torch.nn.functional as F
from sklearn.metrics import accuracy_score, confusion_matrix, f1_score, mean_absolute_error, mean_squared_error
from transformers import AutoTokenizer
from ..adapter import Q1AlignmentAdapter
from ..data_paths import ATTACHMENT4, DATA_ROOT, PROJECT_ROOT
from ..model.early_concat import AlignedFusionModel
from ..q2.deep_learning.q2.evaluate_math_protocol import continuous_mask, scenario_seed
from ..q2.math.data import (
MODALITIES,
fit_preprocessor,
load_official_splits,
restricted_load,
transform_split,
)
SEED = 20260924
TEXT_MODEL_ID = "google-bert/bert-base-uncased"
CLASS_NAMES = ("negative", "neutral", "positive")
MODALITY_NAMES = ("text", "audio", "vision")
def _sha256(path: Path) -> str:
h = hashlib.sha256()
with path.open("rb") as stream:
for block in iter(lambda: stream.read(1024 * 1024), b""):
h.update(block)
return h.hexdigest()
def _decode(value: Any) -> str:
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
if isinstance(value, np.bytes_):
return bytes(value).decode("utf-8", errors="replace")
if isinstance(value, np.ndarray):
if value.shape == ():
return _decode(value.item())
return " ".join(_decode(x) for x in value.reshape(-1))
return str(value)
def _scalar_int(value: Any, field: str) -> int:
arr = np.asarray(value).reshape(-1)
if not len(arr):
raise ValueError(f"Attachment 4 {field} is empty")
return int(arr[0])
def _attachment4_location(version: str) -> tuple[Path, Path]:
inner = ATTACHMENT4 / "附件4-可解释专项视频样本与特征文件"
version_dir = inner / ("未对齐版本" if version == "unaligned_50" else "对齐版本")
video_dir = inner / "videos"
if not version_dir.is_dir():
raise FileNotFoundError(f"Attachment 4 {version} directory not found: {version_dir}")
return version_dir, video_dir
def _read_attachment4(version: str) -> tuple[list[dict[str, Any]], dict[str, str]]:
if version != "unaligned_50":
raise ValueError("Q3 explanation currently uses the official unaligned_50 Attachment 4 features")
version_dir, video_dir = _attachment4_location(version)
paths = sorted(version_dir.glob("*.pkl"), key=lambda p: p.name)
if len(paths) != 20:
raise FileNotFoundError(f"expected 20 Attachment 4 cases, found {len(paths)} under {version_dir}")
video_by_stem = {p.stem: p for p in video_dir.rglob("*.mp4")} if video_dir.is_dir() else {}
adapter = Q1AlignmentAdapter(target_steps=50)
cases: list[dict[str, Any]] = []
for path in paths:
raw = restricted_load(path)
case_id = _decode(raw.get("id", path.stem)).strip() or path.stem
text_bert = np.asarray(raw["text_bert"], dtype=np.int64)
if text_bert.ndim == 3 and text_bert.shape[0] == 1:
text_bert = text_bert[0]
if text_bert.shape != (3, 50):
raise ValueError(f"{path.name}: expected text_bert (3,50), got {text_bert.shape}")
record = {
"id": case_id,
"sequence_order_verified": True,
"attention_mask": text_bert[1].astype(bool),
"text": np.asarray(raw["text"], dtype=np.float32),
"audio": np.asarray(raw["audio"], dtype=np.float32),
"vision": np.asarray(raw["vision"], dtype=np.float32),
"audio_length": _scalar_int(raw["audio_lengths"], "audio_lengths"),
"vision_length": _scalar_int(raw["vision_lengths"], "vision_lengths"),
}
aligned = adapter.align(record, mode="relative")
mask = np.stack([aligned.observed[m] for m in MODALITIES], axis=-1)
features = {m: aligned.features[m].astype(np.float32) for m in MODALITIES}
transcript = _decode(raw.get("raw_text", ""))
video_path = video_by_stem.get(path.stem) or video_by_stem.get(case_id)
media = ""
if video_path is not None:
try:
media = video_path.resolve().relative_to(DATA_ROOT).as_posix()
except ValueError:
media = str(video_path.resolve())
cases.append({
"case_id": case_id,
"source_file": path,
"source_sha256": _sha256(path),
"transcript": transcript,
"text_bert": text_bert,
"raw": raw,
"features": features,
"mask": mask,
"target_intervals": aligned.target_intervals.astype(np.float32),
"provenance": aligned.provenance,
"video_path": media,
"coordinate_mode": aligned.metadata["coordinate_mode"],
"input_audit": {
"case_id": case_id,
"source_file": path.name,
"source_sha256": _sha256(path),
"coordinate_mode": aligned.metadata["coordinate_mode"],
"physical_time_alignment": False,
"audio_reported_length": record["audio_length"],
"vision_reported_length": record["vision_length"],
"audio_length_conflict": bool(aligned.provenance["audio"].length_conflict),
"vision_length_conflict": bool(aligned.provenance["vision"].length_conflict),
"text_visible_target_slots": int(aligned.observed["text"].sum()),
"audio_visible_target_slots": int(aligned.observed["audio"].sum()),
"vision_visible_target_slots": int(aligned.observed["vision"].sum()),
"source_video": media,
},
})
return cases, {"version_dir": str(version_dir), "video_dir": str(video_dir)}
def _split_arrays(split: Any, transformed: dict[str, np.ndarray]) -> tuple[tuple[np.ndarray, ...], np.ndarray]:
return tuple(transformed[m] for m in MODALITIES), np.asarray(split.mask, dtype=bool)
def _predict(
model: torch.nn.Module,
xs: tuple[np.ndarray, ...],
masks: np.ndarray,
device: torch.device,
batch_size: int,
) -> dict[str, np.ndarray]:
model.eval()
logits: list[np.ndarray] = []
intensity: list[np.ndarray] = []
with torch.inference_mode():
for start in range(0, len(masks), batch_size):
end = min(start + batch_size, len(masks))
batch_x = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs)
batch_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
output = model(batch_x, batch_mask)
logits.append(output["logits"].float().cpu().numpy())
intensity.append(output["intensity"].float().cpu().numpy())
return {"logits": np.concatenate(logits), "intensity": np.concatenate(intensity)}
def _metrics(y_cls: np.ndarray, y_reg: np.ndarray, prediction: dict[str, np.ndarray]) -> dict[str, Any]:
logits = np.asarray(prediction["logits"])
score = np.clip(np.asarray(prediction["intensity"]).reshape(-1), -3.0, 3.0)
predicted = logits.argmax(axis=-1)
pearson = float(np.corrcoef(y_reg, score)[0, 1]) if np.std(y_reg) > 0 and np.std(score) > 0 else None
return {
"n": int(len(y_cls)),
"accuracy": float(accuracy_score(y_cls, predicted)),
"macro_f1": float(f1_score(y_cls, predicted, labels=[0, 1, 2], average="macro", zero_division=0)),
"mae": float(mean_absolute_error(y_reg, score)),
"rmse": float(math.sqrt(mean_squared_error(y_reg, score))),
"pearson": pearson,
"confusion_matrix_rows_true_columns_predicted": confusion_matrix(y_cls, predicted, labels=[0, 1, 2]).tolist(),
"per_class_support": {CLASS_NAMES[i]: int(np.sum(y_cls == i)) for i in range(3)},
}
def _loss(logits: torch.Tensor, intensity: torch.Tensor, y_cls: torch.Tensor, y_reg: torch.Tensor) -> torch.Tensor:
return F.cross_entropy(logits, y_cls) + 0.5 * F.smooth_l1_loss(intensity / 3.0, y_reg / 3.0)
def _validation_loss(
model: torch.nn.Module,
xs: tuple[np.ndarray, ...],
masks: list[np.ndarray],
y_cls: np.ndarray,
y_reg: np.ndarray,
device: torch.device,
batch_size: int,
) -> float:
values: list[float] = []
model.eval()
with torch.inference_mode():
for scenario in masks:
total, count = 0.0, 0
for start in range(0, len(y_cls), batch_size):
end = min(start + batch_size, len(y_cls))
bx = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs)
bm = torch.as_tensor(scenario[start:end], dtype=torch.bool, device=device)
by = torch.as_tensor(y_cls[start:end], dtype=torch.long, device=device)
br = torch.as_tensor(y_reg[start:end], dtype=torch.float32, device=device)
out = model(bx, bm)
total += float(_loss(out["logits"], out["intensity"], by, br).item()) * (end - start)
count += end - start
values.append(total / max(1, count))
return float(np.mean(values))
def _train(args: argparse.Namespace, out_dir: Path) -> tuple[AlignedFusionModel, dict[str, Any], dict[str, Any]]:
feature_path: Path
if args.data_path is not None:
feature_path = args.data_path.expanduser().resolve()
else:
from ..data_paths import ATTACHMENT2
feature_path = ATTACHMENT2 / f"{args.input_version}.pkl"
if not feature_path.is_file():
raise FileNotFoundError(f"Q3 training feature file not found: {feature_path}")
raw_splits = load_official_splits(feature_path, version=args.input_version)
train = raw_splits["train"]
valid = raw_splits["valid"]
fitted = fit_preprocessor(train)
transformed = {name: transform_split(split, fitted) for name, split in raw_splits.items()}
train_x, train_mask = _split_arrays(train, transformed["train"])
valid_x, valid_mask = _split_arrays(valid, transformed["valid"])
dims = tuple(int(x.shape[-1]) for x in train_x)
np.savez_compressed(out_dir / "preprocessor.npz", **{
f"{modality}_{key}": value for modality, state in fitted.items() for key, value in state.items()
})
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else ("cpu" if args.device == "auto" else args.device))
if device.type == "cuda":
torch.cuda.manual_seed_all(SEED)
random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)
torch.set_num_threads(4)
model = AlignedFusionModel("concat", dims=dims).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay)
y_cls = np.asarray(train.class_y, dtype=np.int64)
y_reg = np.asarray(train.regression_y, dtype=np.float32)
vy_cls = np.asarray(valid.class_y, dtype=np.int64)
vy_reg = np.asarray(valid.regression_y, dtype=np.float32)
valid_rng_masks: list[np.ndarray] = [valid_mask.copy()]
for rate, mode in ((0.3, "single"), (0.3, "sync"), (0.5, "async")):
key = f"{rate:.1f}/{mode}"
valid_rng_masks.append(np.stack([
continuous_mask(mask, rate, mode, np.random.default_rng(scenario_seed(SEED + 177, sid, key)))
for sid, mask in zip(valid.ids, valid_mask)
]))
best = float("inf")
best_epoch = 0
stale = 0
history: list[dict[str, Any]] = []
for epoch in range(1, args.epochs + 1):
model.train()
train_corruption = np.stack([
continuous_mask(
mask,
float(np.random.choice((0.0, 0.1, 0.3, 0.5, 0.7))),
str(np.random.choice(("single", "sync", "partial", "async"))),
np.random.default_rng(scenario_seed(SEED + epoch, sid, f"train/{epoch}")),
)
for sid, mask in zip(train.ids, train_mask)
])
order = np.random.permutation(len(y_cls))
losses: list[float] = []
for start in range(0, len(order), args.batch_size):
ix = order[start:start + args.batch_size]
bx = tuple(torch.as_tensor(x[ix], dtype=torch.float32, device=device) for x in train_x)
bm = torch.as_tensor(train_corruption[ix], dtype=torch.bool, device=device)
by = torch.as_tensor(y_cls[ix], dtype=torch.long, device=device)
br = torch.as_tensor(y_reg[ix], dtype=torch.float32, device=device)
optimizer.zero_grad(set_to_none=True)
output = model(bx, bm)
loss = _loss(output["logits"], output["intensity"], by, br)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
losses.append(float(loss.item()))
validation = _validation_loss(model, valid_x, valid_rng_masks, vy_cls, vy_reg, device, args.batch_size)
history.append({"epoch": epoch, "train_loss": float(np.mean(losses)), "selection_loss": validation})
print(f"Q3 epoch {epoch}/{args.epochs}: train={np.mean(losses):.5f}, validation={validation:.5f}", flush=True)
if validation < best - 1e-7:
best, best_epoch, stale = validation, epoch, 0
torch.save({"state_dict": model.state_dict(), "dims": dims, "seed": SEED, "best_epoch": epoch}, out_dir / "model_best.pt")
else:
stale += 1
if stale >= args.patience:
break
checkpoint = torch.load(out_dir / "model_best.pt", map_location=device, weights_only=True)
model.load_state_dict(checkpoint["state_dict"])
model.eval()
prediction = _predict(model, valid_x, valid_mask, device, args.batch_size)
metric = _metrics(vy_cls, vy_reg, prediction)
metric["best_epoch"] = best_epoch
metric["selection_loss_clean_plus_fixed_missing_scenarios"] = best
metric["input_version"] = args.input_version
metric["adapter"] = "Q1AlignmentAdapter relative normalized progress"
metric["physical_time_alignment"] = False
_write_csv(out_dir / "training_history.csv", history)
_write_json(out_dir / "validation_metrics.json", metric)
validation_rows = []
for i, sid in enumerate(valid.ids):
prob = torch.softmax(torch.as_tensor(prediction["logits"][i]), dim=-1).numpy()
validation_rows.append({
"sample_id": sid,
"true_class": int(vy_cls[i]),
"true_class_name": CLASS_NAMES[int(vy_cls[i])],
"true_sentiment": float(vy_reg[i]),
"predicted_class": int(prob.argmax()),
"predicted_class_name": CLASS_NAMES[int(prob.argmax())],
"predicted_sentiment": float(prediction["intensity"][i]),
"p_negative": float(prob[0]), "p_neutral": float(prob[1]), "p_positive": float(prob[2]),
"absolute_error": float(abs(vy_reg[i] - prediction["intensity"][i])),
})
_write_csv(out_dir / "validation_predictions.csv", validation_rows)
errors = sorted(
(row for row in validation_rows if row["true_class"] != row["predicted_class"] or row["absolute_error"] >= metric["mae"]),
key=lambda row: (-row["absolute_error"], row["sample_id"]),
)
_write_csv(out_dir / "validation_errors.csv", errors[:100])
return model, {"metrics": metric, "feature_sha256": _sha256(feature_path), "feature_path": str(feature_path)}, {"x": valid_x, "mask": valid_mask, "y_cls": vy_cls, "y_reg": vy_reg, "prediction": prediction}
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
path.parent.mkdir(parents=True, exist_ok=True)
fields = list(dict.fromkeys(key for row in rows for key in row))
with path.open("w", encoding="utf-8-sig", newline="") as stream:
writer = csv.DictWriter(stream, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def _write_json(path: Path, payload: Any) -> None:
path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8")
def _model_output(model: torch.nn.Module, xs: tuple[torch.Tensor, ...], mask: torch.Tensor) -> dict[str, torch.Tensor]:
model.eval()
with torch.inference_mode():
return model(xs, mask)
def _span_evidence(case: dict[str, Any], modality_index: int, slot: int, tokenizer: Any) -> dict[str, Any]:
modality = MODALITY_NAMES[modality_index]
weights = case["provenance"][modality].source_weights.getrow(slot)
source_rows = weights.indices.tolist()
if source_rows:
low, high = min(source_rows), max(source_rows) + 1
else:
low = high = 0
start, end = case["target_intervals"][slot].astype(float).tolist()
text = ""
if modality == "text" and source_rows:
ids = np.asarray(case["text_bert"][0], dtype=np.int64)
token_ids = [int(ids[i]) for i in source_rows if i < len(ids) and int(ids[i]) not in tokenizer.all_special_ids]
text = " ".join(tokenizer.convert_ids_to_tokens(token_ids))
elif modality == "audio":
text = f"audio feature rows {low}–{high - 1}; inspect the same relative span in the linked source video/audio"
else:
text = f"video feature rows {low}–{high - 1}; inspect the same relative span in the linked source video"
return {
"modality": modality,
"slot": int(slot),
"relative_start": float(start),
"relative_end": float(end),
"source_row_start": int(low),
"source_row_end_exclusive": int(high),
"evidence": text,
}
def _explain_case(
model: torch.nn.Module,
case: dict[str, Any],
stats: dict[str, dict[str, np.ndarray]],
tokenizer: Any,
device: torch.device,
batch_size: int,
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
values: dict[str, np.ndarray] = {}
for modality_index, modality in enumerate(MODALITIES):
arr = case["features"][modality].astype(np.float32)
arr = np.clip((arr - stats[modality]["mean"]) / stats[modality]["std"], -10.0, 10.0)
arr[~case["mask"][:, modality_index]] = 0.0
values[modality] = arr
xs = tuple(torch.as_tensor(values[m][None], dtype=torch.float32, device=device) for m in MODALITIES)
mask = torch.as_tensor(case["mask"][None], dtype=torch.bool, device=device)
full = _model_output(model, xs, mask)
probs = torch.softmax(full["logits"], dim=-1)[0].cpu().numpy()
pred = int(np.argmax(probs))
contributions: dict[str, float] = {}
local_rows: list[dict[str, Any]] = []
for m, modality in enumerate(MODALITY_NAMES):
ablated_mask = mask.clone()
ablated_mask[:, :, m] = False
ablated = _model_output(model, xs, ablated_mask)
ablated_p = torch.softmax(ablated["logits"], dim=-1)[0, pred].item()
contributions[modality] = float(probs[pred] - ablated_p)
observed_slots = np.flatnonzero(case["mask"][:, m])
if not len(observed_slots):
continue
impacts: list[tuple[int, float]] = []
for start in range(0, len(observed_slots), batch_size):
chosen = observed_slots[start:start + batch_size]
bx = tuple(x.repeat(len(chosen), 1, 1) for x in xs)
bm = mask.repeat(len(chosen), 1, 1)
row_idx = torch.arange(len(chosen), device=device)
slot_idx = torch.as_tensor(chosen, dtype=torch.long, device=device)
bm[row_idx, slot_idx, m] = False
output = _model_output(model, bx, bm)
hidden_p = torch.softmax(output["logits"], dim=-1)[:, pred].cpu().numpy()
impacts.extend((int(slot), float(probs[pred] - p)) for slot, p in zip(chosen, hidden_p))
for slot, impact in sorted(impacts, key=lambda row: (-row[1], row[0]))[:3]:
evidence = _span_evidence(case, m, slot, tokenizer)
evidence["probability_drop"] = impact
evidence["case_id"] = case["case_id"]
evidence["source_video"] = case["video_path"]
local_rows.append(evidence)
principal = max(contributions, key=contributions.get)
intensity = float(full["intensity"][0].cpu().item())
explanation = {
"case_id": case["case_id"],
"predicted_class": pred,
"predicted_class_name": CLASS_NAMES[pred],
"predicted_sentiment": intensity,
"p_negative": float(probs[0]), "p_neutral": float(probs[1]), "p_positive": float(probs[2]),
"principal_modality": principal,
"text_contribution": contributions["text"],
"audio_contribution": contributions["audio"],
"vision_contribution": contributions["vision"],
"transcript": case["transcript"],
"source_video": case["video_path"],
"coordinate_mode": case["coordinate_mode"],
"interpretation_method": "single-modality and single-slot occlusion; probability drops measure model sensitivity",
}
return explanation, local_rows
def _write_cards(out_dir: Path, case_by_id: dict[str, dict[str, Any]], explanations: list[dict[str, Any]], local_rows: list[dict[str, Any]]) -> str:
cards = out_dir / "explanation_cards"
cards.mkdir(parents=True, exist_ok=True)
rows_by_id: dict[str, list[dict[str, Any]]] = {}
for row in local_rows:
rows_by_id.setdefault(str(row["case_id"]), []).append(row)
for item in explanations:
evidence = rows_by_id.get(str(item["case_id"]), [])
lines = [f"# Q3 Explanation: {item['case_id']}", "", f"- Prediction: **{item['predicted_class_name']}**", f"- Sentiment score: {item['predicted_sentiment']:.3f}", f"- Probabilities (negative / neutral / positive): {item['p_negative']:.3f} / {item['p_neutral']:.3f} / {item['p_positive']:.3f}", f"- Main modality by occlusion: **{item['principal_modality']}**", f"- Source video/audio: `{item['source_video'] or 'not found in the supplied video folder'}`", f"- Coordinate: normalized progress `[0,1]`; no physical timestamps are inferred from the unaligned feature rows.", "", "## Modality contribution", "", "Removing one modality changes the predicted-class probability by the values below. Positive values mean that modality supports the prediction under this model.", "", "| Modality | Probability drop |", "|---|---:|"]
for modality in MODALITY_NAMES:
lines.append(f"| {modality} | {item[f'{modality}_contribution']:.4f} |")
lines.extend(["", "## Local evidence", "", "Local values are single-slot occlusion sensitivity. Audio/video spans are relative positions in the supplied source clip; text is shown as BERT tokens and the full transcript is retained below.", ""])
for evidence_row in evidence:
lines.append(f"- **{evidence_row['modality']}**, slots {evidence_row['slot']} `[0-based]`, relative {evidence_row['relative_start']:.3f}–{evidence_row['relative_end']:.3f}, probability drop {evidence_row['probability_drop']:.4f}: {evidence_row['evidence']}")
lines.extend(["", "## Transcript", "", item["transcript"] or "(not supplied)", "", "## Interpretation note", "", "Occlusion scores describe how this trained model responds to removing features. They are not causal effects or proof that the signal expresses the named emotion.", ""])
safe = "".join(c if c.isalnum() or c in "-_" else "_" for c in str(item["case_id"]))
(cards / f"{safe}.md").write_text("\n".join(lines), encoding="utf-8")
confidence = np.asarray([max(row["p_negative"], row["p_neutral"], row["p_positive"]) for row in explanations])
representative = explanations[int(np.argmin(np.abs(confidence - np.median(confidence))))]
source = cards / ("".join(c if c.isalnum() or c in "-_" else "_" for c in str(representative["case_id"])) + ".md")
representative_card = out_dir / "typical_explanation_card.md"
representative_card.write_text(source.read_text(encoding="utf-8"), encoding="utf-8")
return str(representative["case_id"])
def _plot_validation(out_dir: Path, y_cls: np.ndarray, prediction: dict[str, np.ndarray]) -> None:
pred_cls = prediction["logits"].argmax(axis=-1)
matrix = confusion_matrix(y_cls, pred_cls, labels=[0, 1, 2])
fig, axes = plt.subplots(1, 2, figsize=(10, 4), constrained_layout=True)
image = axes[0].imshow(matrix, cmap="Blues")
axes[0].set_xticks(range(3), CLASS_NAMES, rotation=15)
axes[0].set_yticks(range(3), CLASS_NAMES)
axes[0].set_xlabel("Predicted")
axes[0].set_ylabel("True")
axes[0].set_title("Validation confusion matrix")
for (i, j), value in np.ndenumerate(matrix):
axes[0].text(j, i, str(value), ha="center", va="center")
fig.colorbar(image, ax=axes[0], fraction=0.046)
axes[1].scatter(prediction["intensity"], prediction["true_sentiment"], s=12, alpha=0.55)
axes[1].plot([-3, 3], [-3, 3], color="gray", linestyle="--", linewidth=1)
axes[1].set(xlim=(-3, 3), ylim=(-3, 3), xlabel="Predicted sentiment", ylabel="True sentiment", title="Validation intensity")
fig.savefig(out_dir / "validation_diagnostics.png", dpi=180)
plt.close(fig)
def run(args: argparse.Namespace) -> None:
out_dir = args.output_dir.expanduser().resolve()
out_dir.mkdir(parents=True, exist_ok=True)
started = time.time()
model, training_info, validation = _train(args, out_dir)
_plot_validation(out_dir, validation["y_cls"], {**validation["prediction"], "true_sentiment": validation["y_reg"]})
tokenizer = AutoTokenizer.from_pretrained(TEXT_MODEL_ID, use_fast=True)
with np.load(out_dir / "preprocessor.npz", allow_pickle=False) as saved:
stats = {m: {key: saved[f"{m}_{key}"].astype(np.float32) for key in ("mean", "std")} for m in MODALITIES}
cases, input_locations = _read_attachment4(args.attachment4_version)
device = next(model.parameters()).device
explanation_rows, all_local = [], []
prediction_rows = []
for case in cases:
explanation, local = _explain_case(model, case, stats, tokenizer, device, args.explanation_batch_size)
explanation_rows.append(explanation)
all_local.extend(local)
prediction_rows.append({key: explanation[key] for key in (
"case_id", "predicted_class", "predicted_class_name", "predicted_sentiment",
"p_negative", "p_neutral", "p_positive", "source_video",
)})
_write_csv(out_dir / "attachment4_predictions.csv", prediction_rows)
_write_csv(out_dir / "attachment4_explanations.csv", explanation_rows)
_write_csv(out_dir / "attachment4_local_evidence.csv", all_local)
_write_csv(out_dir / "attachment4_input_audit.csv", [case["input_audit"] for case in cases])
typical_id = _write_cards(out_dir, {case["case_id"]: case for case in cases}, explanation_rows, all_local)
manifest = {
"created_at_unix": time.time(),
"elapsed_seconds": time.time() - started,
"seed": SEED,
"training_input": training_info,
"attachment4": input_locations,
"attachment4_version": args.attachment4_version,
"attachment4_cases": len(cases),
"adapter": "Q1AlignmentAdapter shared relative-progress projection",
"coordinate_limit": "source-time stamps are absent; local audio/video positions are normalized progress, not seconds",
"model": "EarlyConcat + BiGRU",
"explanation": "single-modality and single-slot occlusion probability drops; model sensitivity, not causal attribution",
"validation_metrics": validation["metrics"],
"typical_explanation_case": typical_id,
"outputs": [
"model_best.pt", "preprocessor.npz", "validation_metrics.json", "validation_predictions.csv",
"validation_errors.csv", "validation_diagnostics.png", "attachment4_predictions.csv",
"attachment4_explanations.csv", "attachment4_local_evidence.csv", "attachment4_input_audit.csv", "typical_explanation_card.md",
],
}
_write_json(out_dir / "run_manifest.json", manifest)
print(f"Q3 complete: {len(cases)} Attachment 4 predictions saved under {out_dir}", flush=True)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="unaligned_50")
parser.add_argument("--attachment4-version", choices=("unaligned_50",), default="unaligned_50")
parser.add_argument("--data-path", type=Path, default=None, help="Optional explicit Attachment 2 pickle path")
parser.add_argument("--output-dir", type=Path, default=PROJECT_ROOT / "output" / "q3")
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--epochs", type=int, default=12)
parser.add_argument("--patience", type=int, default=3)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--learning-rate", type=float, default=3e-4)
parser.add_argument("--weight-decay", type=float, default=1e-3)
parser.add_argument("--explanation-batch-size", type=int, default=32)
args = parser.parse_args()
run(args)
from .run_experiments import main
if __name__ == "__main__":