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