Files
modeling_zhaocui/final/q3/plot_router_heatmaps.py

464 lines
21 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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())