""" scripts/plot_results.py ----------------------- Generate the four paper figures from saved results. Figure 1: Capability density heatmap (24 layers × 16 heads) Figure 2: Density vs ablation ΔPPL scatter (Pearson r, Spearman ρ) Figure 3: Density rank vs Wanda rank (orthogonality) Figure 4: PPL comparison bar chart Usage: python scripts/plot_results.py \\ --density_map results/density_map.npz \\ --results_dir results/ \\ --output_dir figures/ """ import argparse import json import os import sys import numpy as np sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from cgc.density import CapabilityDensityMap, N_LAYERS, N_HEADS def parse_args(): p = argparse.ArgumentParser(description="Generate CGC v1 paper figures") p.add_argument("--density_map", type=str, required=True) p.add_argument("--results_dir", type=str, required=True) p.add_argument("--output_dir", type=str, default="figures/") return p.parse_args() def plot_figure1(dm: CapabilityDensityMap, output_dir: str): """Figure 1: Capability density heatmap (24 × 16).""" import matplotlib.pyplot as plt import seaborn as sns fig, ax = plt.subplots(figsize=(14, 7)) sns.heatmap( dm.density, ax=ax, cmap="YlOrRd", vmin=0, vmax=1, xticklabels=[f"H{h}" for h in range(N_HEADS)], yticklabels=[f"L{l}" for l in range(N_LAYERS)], cbar_kws={"label": "Capability Density δ(c)"}, ) d = dm.density ax.set_title( f"Figure 1: Capability Density Map — GPT-2 Medium " f"({N_LAYERS} layers × {N_HEADS} heads)\n" f"Mean={d.mean():.4f} " f"Max={d.max():.4f} (L{d.max(1).argmax()} H{d.argmax(1)[d.max(1).argmax()]}) " f"Min={d.min():.4f}", fontsize=11, fontweight="bold", ) ax.set_xlabel("Attention Head", fontsize=11) ax.set_ylabel("Layer", fontsize=11) plt.tight_layout() path = os.path.join(output_dir, "figure1_density_heatmap.png") plt.savefig(path, dpi=150, bbox_inches="tight") plt.close() print(f"Saved: {path}") def plot_figure2(dm: CapabilityDensityMap, ablation: np.ndarray, output_dir: str): """Figure 2: Density vs ablation ΔPPL scatter.""" import matplotlib.pyplot as plt from scipy.stats import pearsonr, spearmanr if ablation is None: print("Skipping Figure 2 (no ablation data).") return density_flat = dm.density.flatten() ablation_flat = ablation.flatten() d_min = density_flat.min() d_max = density_flat.max() d_rsc = (density_flat - d_min) / (d_max - d_min + 1e-8) pr, pp = pearsonr(d_rsc, ablation_flat) sr, sp = spearmanr(d_rsc, ablation_flat) layer_colors = np.repeat(np.arange(N_LAYERS), N_HEADS) fig, ax = plt.subplots(figsize=(9, 6)) sc = ax.scatter( d_rsc, ablation_flat, c=layer_colors, cmap="RdYlBu_r", alpha=0.7, s=40, edgecolors="white", linewidths=0.3, ) z = np.polyfit(d_rsc, ablation_flat, 1) x_line = np.linspace(d_rsc.min(), d_rsc.max(), 100) ax.plot(x_line, np.poly1d(z)(x_line), "k--", linewidth=2, label="Linear fit") plt.colorbar(sc, ax=ax, label="Layer Index") ax.set_xlabel("Capability Density δ(c) [rescaled 0–1]", fontsize=12) ax.set_ylabel("Ablation Impact ΔPPL", fontsize=12) ax.set_title( f"Figure 2: Capability Density vs. Compression Vulnerability\n" f"Pearson r = {pr:.3f} (p = {pp:.2e}) | " f"Spearman ρ = {sr:.3f} (p = {sp:.2e}) | " f"n = {len(density_flat)} heads", fontsize=11, fontweight="bold", ) ax.legend(fontsize=10) plt.tight_layout() path = os.path.join(output_dir, "figure2_density_vs_ablation.png") plt.savefig(path, dpi=150, bbox_inches="tight") plt.close() print(f"Saved: {path} (r={pr:.4f}, p={pp:.2e})") def plot_figure3(dm: CapabilityDensityMap, wanda: np.ndarray, output_dir: str): """Figure 3: Density rank vs Wanda rank (orthogonality).""" import matplotlib.pyplot as plt from scipy.stats import rankdata, spearmanr if wanda is None: print("Skipping Figure 3 (no Wanda data).") return density_flat = dm.density.flatten() wanda_flat = wanda.flatten() d_ranks = rankdata(density_flat) w_ranks = rankdata(wanda_flat) rho, p = spearmanr(density_flat, wanda_flat) layer_colors = np.repeat(np.arange(N_LAYERS), N_HEADS) fig, ax = plt.subplots(figsize=(8, 6)) ax.scatter( w_ranks, d_ranks, c=layer_colors, cmap="RdYlBu_r", alpha=0.6, s=35, edgecolors="white", linewidths=0.3, ) ax.set_xlabel("Wanda Importance Rank", fontsize=12) ax.set_ylabel("Capability Density Rank", fontsize=12) ax.set_title( f"Figure 3: Capability Density vs. Wanda Importance — Signal Orthogonality\n" f"Spearman ρ = {rho:.3f} (p = {p:.2e}) | n = {len(density_flat)} heads", fontsize=12, fontweight="bold", ) plt.tight_layout() path = os.path.join(output_dir, "figure3_density_vs_wanda.png") plt.savefig(path, dpi=150, bbox_inches="tight") plt.close() print(f"Saved: {path} (ρ={rho:.4f}, p={p:.2e})") def plot_figure4(summary: dict, output_dir: str): """Figure 4: PPL comparison bar chart.""" import matplotlib.pyplot as plt methods = ["Dense", "Uniform", "CGC-L\n(ours)", "Inverted\n(wrong)"] ppls = [ summary["baseline_ppl"], summary["uniform"]["ppl"], summary["cgc"]["ppl"], summary["inverted"]["ppl"], ] colors = ["#2196F3", "#FF9800", "#4CAF50", "#F44336"] fig, ax = plt.subplots(figsize=(8, 5)) bars = ax.bar(methods, ppls, color=colors, edgecolor="white", linewidth=1.5) ax.set_ylabel("Perplexity (lower = better)", fontsize=12) ax.set_title( f"Figure 4: Compression PPL Comparison — GPT-2 Medium\n" f"(50% global attention head weight retention)", fontsize=12, fontweight="bold", ) for bar, ppl in zip(bars, ppls): ax.text( bar.get_x() + bar.get_width() / 2, ppl + 0.05, f"{ppl:.2f}", ha="center", fontsize=10, fontweight="bold", ) plt.tight_layout() path = os.path.join(output_dir, "figure4_compression_comparison.png") plt.savefig(path, dpi=150, bbox_inches="tight") plt.close() print(f"Saved: {path}") def main(): args = parse_args() os.makedirs(args.output_dir, exist_ok=True) dm = CapabilityDensityMap.load(args.density_map) print(dm.summary() + "\n") # Load optional arrays ablation_path = os.path.join(args.results_dir, "ablation_results.npy") ablation = np.load(ablation_path) if os.path.exists(ablation_path) else None if ablation is None: print("Note: ablation_results.npy not found — skipping Figure 2.") wanda_path = os.path.join(args.results_dir, "wanda_importance.npy") wanda = np.load(wanda_path) if os.path.exists(wanda_path) else None if wanda is None: print("Note: wanda_importance.npy not found — skipping Figure 3.") summary_path = os.path.join(args.results_dir, "compression_summary.json") with open(summary_path) as f: summary = json.load(f) plot_figure1(dm, args.output_dir) plot_figure2(dm, ablation, args.output_dir) plot_figure3(dm, wanda, args.output_dir) plot_figure4(summary, args.output_dir) print(f"\nAll figures written to: {args.output_dir}") if __name__ == "__main__": main()