464 lines
21 KiB
Python
464 lines
21 KiB
Python
"""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())
|