"""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=" 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())