ravel / scripts /plot_mvsa_single.py
minhy112's picture
Upload RAVEL revision project without data or checkpoints
ea8bfa1 verified
Raw
History Blame Contribute Delete
5.18 kB
#!/usr/bin/env python3
"""Plot confusion matrix and ablation figure for MVSA-Single results."""
from __future__ import annotations
import argparse
import csv
from pathlib import Path
from typing import List, Tuple
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Plot MVSA-Single result figures")
parser.add_argument("--results-dir", default="results/mvsa_single")
parser.add_argument("--dpi", type=int, default=150)
return parser.parse_args()
def read_confusion_matrix(path: Path) -> Tuple[List[str], np.ndarray]:
with path.open("r", encoding="utf-8") as f:
rows = list(csv.reader(f))
if len(rows) < 2:
raise ValueError(f"Invalid confusion matrix file: {path}")
labels = [cell.strip() for cell in rows[0][1:] if cell.strip()]
values = []
for row in rows[1:]:
values.append([int(float(x)) for x in row[1 : 1 + len(labels)]])
cm = np.array(values, dtype=np.int64)
return labels, cm
def read_ablation(path: Path) -> List[Tuple[str, float, float]]:
rows: List[Tuple[str, float, float]] = []
with path.open("r", encoding="utf-8") as f:
reader = csv.DictReader(f)
for row in reader:
rows.append(
(
str(row["variant"]),
float(row["accuracy"]),
float(row["f1_weighted"]),
)
)
if not rows:
raise ValueError(f"Empty ablation file: {path}")
return rows
def plot_confusion_matrix(labels: List[str], cm: np.ndarray, out_path: Path, dpi: int) -> None:
fig, ax = plt.subplots(figsize=(5.2, 4.4))
row_sums = cm.sum(axis=1, keepdims=True).astype(np.float64)
with np.errstate(divide="ignore", invalid="ignore"):
cm_norm = np.divide(cm, row_sums, where=row_sums > 0)
cm_norm = np.nan_to_num(cm_norm)
im = ax.imshow(cm_norm, cmap="Blues", vmin=0.0, vmax=1.0)
for i in range(cm.shape[0]):
for j in range(cm.shape[1]):
pct = cm_norm[i, j] * 100.0
count = int(cm[i, j])
color = "white" if cm_norm[i, j] > 0.55 else "black"
ax.text(j, i, f"{count}\n({pct:.1f}%)", ha="center", va="center", fontsize=10, color=color)
ax.set_xticks(np.arange(len(labels)))
ax.set_yticks(np.arange(len(labels)))
ax.set_xticklabels([label.capitalize() for label in labels], fontsize=10)
ax.set_yticklabels([label.capitalize() for label in labels], fontsize=10)
ax.set_xlabel("Predicted", fontsize=11)
ax.set_ylabel("True", fontsize=11)
cbar = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
cbar.ax.tick_params(labelsize=9)
fig.tight_layout()
fig.savefig(out_path, dpi=dpi, bbox_inches="tight")
plt.close(fig)
def plot_ablation_table(rows: List[Tuple[str, float, float]], out_path: Path, dpi: int) -> None:
navy = "#2c3e6b"
orange = "#d4a017"
light_grey = "#f0f0f0"
fig, ax = plt.subplots(figsize=(6.2, 3.4))
ax.axis("off")
col_labels = ["Variant", "Acc", "F1-Weighted"]
table_data = [[name, f"{acc:.4f}", f"{f1w:.4f}"] for name, acc, f1w in rows]
table = ax.table(
cellText=table_data,
colLabels=col_labels,
loc="center",
cellLoc="center",
)
table.auto_set_font_size(False)
table.set_fontsize(10)
table.scale(1.3, 1.8)
for col_idx in range(len(col_labels)):
cell = table[0, col_idx]
cell.set_facecolor(navy)
cell.set_text_props(color="white", fontweight="bold")
cell.set_edgecolor("white")
full_f1 = rows[0][2]
for row_idx, (variant, _, f1w) in enumerate(rows, start=1):
is_full = variant == "Full"
for col_idx in range(len(col_labels)):
cell = table[row_idx, col_idx]
cell.set_edgecolor("#cccccc")
if is_full:
cell.set_facecolor("#dce6f1")
elif row_idx % 2 == 0:
cell.set_facecolor(light_grey)
if (f1w > full_f1) and (not is_full):
table[row_idx, 2].set_facecolor(orange)
table[row_idx, 2].set_text_props(fontweight="bold")
fig.tight_layout()
fig.savefig(out_path, dpi=dpi, bbox_inches="tight")
plt.close(fig)
def main() -> None:
args = parse_args()
results_dir = Path(args.results_dir)
cm_path = results_dir / "confusion_matrix.csv"
ablation_path = results_dir / "ablation_summary.csv"
if not cm_path.exists():
raise FileNotFoundError(f"Missing: {cm_path}")
if not ablation_path.exists():
raise FileNotFoundError(f"Missing: {ablation_path}")
labels, cm = read_confusion_matrix(cm_path)
ablation_rows = read_ablation(ablation_path)
out_cm = results_dir / "figure4a_confusion_matrix.png"
out_ablation = results_dir / "figure4b_ablation_study.png"
plot_confusion_matrix(labels, cm, out_cm, args.dpi)
plot_ablation_table(ablation_rows, out_ablation, args.dpi)
print(f"Saved -> {out_cm}")
print(f"Saved -> {out_ablation}")
if __name__ == "__main__":
main()