Add Q3 MoFE router visualizations and explanations
This commit is contained in:
@@ -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% 的输入仍可能偏离训练分布;相关数值按诊断结果报告,不称为解释准确率或因果效应。
|
||||
@@ -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
@@ -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()
|
||||
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user