Add Q3 MoFE router visualizations and explanations
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user