| #!/usr/bin/env python3 | |
| """Build the three reader-facing ACL main-paper figures.""" | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| from statistics import mean | |
| from textwrap import dedent | |
| import numpy as np | |
| try: | |
| from .build_acl_arr_figures import ( | |
| FRAGILE_BLOCKS, | |
| LOW_THRESHOLD, | |
| _values, | |
| load_data, | |
| validate, | |
| ) | |
| except ImportError: | |
| from build_acl_arr_figures import ( | |
| FRAGILE_BLOCKS, | |
| LOW_THRESHOLD, | |
| _values, | |
| load_data, | |
| validate, | |
| ) | |
| BLUE = "#0077BB" | |
| ORANGE = "#EE7733" | |
| INK = "#1B2A41" | |
| LIGHT_GREY = "#D6D9DC" | |
| FISHER_LABELS = { | |
| "fpeft_low": "Input-Tail", | |
| "fpeft_bi_low": "Paired-Low", | |
| "fpeft_bi_high": "Paired-High", | |
| } | |
| def _pyplot(): | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| matplotlib.rcParams.update({ | |
| "font.family": "DejaVu Sans", | |
| "font.size": 8, | |
| "axes.labelsize": 8, | |
| "axes.titlesize": 8.5, | |
| "xtick.labelsize": 7, | |
| "ytick.labelsize": 7, | |
| "legend.fontsize": 7, | |
| "pdf.fonttype": 42, | |
| "ps.fonttype": 42, | |
| }) | |
| return plt | |
| def _save(fig, output_dir: Path, stem: str) -> Path: | |
| pdf = output_dir / f"{stem}.pdf" | |
| fig.savefig(pdf, dpi=300) | |
| fig.savefig(output_dir / f"{stem}.png", dpi=300) | |
| return pdf | |
| def _fragile_rows(data: dict[str, object]) -> list[dict[str, object]]: | |
| cells = data["broad_cells"] | |
| rows = [] | |
| for label, block, fisher_method in FRAGILE_BLOCKS: | |
| peft = _values(cells, block, "peft_default") | |
| fisher = _values(cells, block, fisher_method) | |
| peft_mean = mean(peft) * 100 | |
| fisher_mean = mean(fisher) * 100 | |
| rows.append({ | |
| "label": label, | |
| "variant": FISHER_LABELS[fisher_method], | |
| "peft_mean": peft_mean, | |
| "fisher_mean": fisher_mean, | |
| "gain": fisher_mean - peft_mean, | |
| "peft_low": sum(value < LOW_THRESHOLD for value in peft), | |
| "fisher_low": sum(value < LOW_THRESHOLD for value in fisher), | |
| }) | |
| return sorted(rows, key=lambda row: row["gain"], reverse=True) | |
| def write_method_svg(data: dict[str, object], output_dir: Path) -> Path: | |
| projection = data["projection"] | |
| fisher = float(projection["fpeft"]) | |
| peft = float(projection["peft_default"]) | |
| ratio = round(fisher / peft) | |
| rows = _fragile_rows(data) | |
| peft_low = sum(row["peft_low"] for row in rows) | |
| fisher_low = sum(row["fisher_low"] for row in rows) | |
| svg = dedent( | |
| f"""\ | |
| <svg xmlns="http://www.w3.org/2000/svg" width="1200" height="480" viewBox="0 0 1200 480" role="img" | |
| aria-labelledby="title desc"> | |
| <title id="title">Same function, different initial tangent spaces</title> | |
| <desc id="desc">PEFT and calibration-informed LoRA preserve the same step-zero function but produce | |
| different unnormalized gradient-access statistics.</desc> | |
| <defs> | |
| <marker id="arrow-ink" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="7" markerHeight="7" | |
| orient="auto"> | |
| <path d="M 0 0 L 10 5 L 0 10 z" fill="{INK}"/> | |
| </marker> | |
| <marker id="arrow-blue" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="7" markerHeight="7" | |
| orient="auto"> | |
| <path d="M 0 0 L 10 5 L 0 10 z" fill="{BLUE}"/> | |
| </marker> | |
| <marker id="arrow-grey" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="7" markerHeight="7" | |
| orient="auto"> | |
| <path d="M 0 0 L 10 5 L 0 10 z" fill="#69717A"/> | |
| </marker> | |
| <linearGradient id="fisher-band" x1="0" y1="1" x2="1" y2="0"> | |
| <stop offset="0" stop-color="{BLUE}" stop-opacity="0.08"/> | |
| <stop offset="1" stop-color="{BLUE}" stop-opacity="0.28"/> | |
| </linearGradient> | |
| <linearGradient id="peft-band" x1="0" y1="0" x2="1" y2="1"> | |
| <stop offset="0" stop-color="#69717A" stop-opacity="0.08"/> | |
| <stop offset="1" stop-color="#D55E5E" stop-opacity="0.22"/> | |
| </linearGradient> | |
| <style> | |
| text {{ font-family: Arial, Helvetica, sans-serif; fill: {INK}; }} | |
| .panel {{ font-size: 20px; font-weight: 700; }} | |
| .heading {{ font-size: 22px; font-weight: 700; }} | |
| .label {{ font-size: 18px; font-weight: 700; }} | |
| .body {{ font-size: 16px; }} | |
| .small {{ font-size: 13px; }} | |
| .formula {{ font-size: 18px; font-style: italic; }} | |
| .card {{ fill: #FAFBFC; stroke: #D6D9DC; stroke-width: 1.5; }} | |
| </style> | |
| </defs> | |
| <rect width="1200" height="480" fill="white"/> | |
| <text class="panel" x="24" y="28">A</text> | |
| <text class="heading" x="55" y="28">Same function, different initial tangent spaces</text> | |
| <rect x="24" y="45" width="800" height="405" rx="16" fill="#FBFCFD" stroke="#D6D9DC" stroke-width="1.5"/> | |
| <ellipse cx="684" cy="116" rx="112" ry="64" fill="#DDF0FA" stroke="{BLUE}" stroke-width="1.6"/> | |
| <ellipse cx="684" cy="116" rx="82" ry="44" fill="none" stroke="{BLUE}" stroke-opacity="0.45"/> | |
| <ellipse cx="700" cy="355" rx="102" ry="58" fill="#F8E2E0" stroke="#C95C54" stroke-width="1.6"/> | |
| <ellipse cx="700" cy="355" rx="73" ry="38" fill="none" stroke="#C95C54" stroke-opacity="0.45"/> | |
| <path d="M 88 390 C 230 330, 240 155, 410 105 C 570 58, 705 62, 790 110" | |
| fill="none" stroke="#D6D9DC" stroke-width="1.2"/> | |
| <path d="M 72 420 C 260 365, 260 180, 432 130 C 600 82, 724 92, 808 145" | |
| fill="none" stroke="#E3E6E8" stroke-width="1.2"/> | |
| <path d="M 92 330 C 230 300, 275 245, 390 225 C 550 198, 685 248, 805 335" | |
| fill="none" stroke="#E3E6E8" stroke-width="1.2"/> | |
| <polygon points="323,251 353,287 749,392 774,338" fill="url(#peft-band)"/> | |
| <polygon points="324,249 357,281 718,110 684,72" fill="url(#fisher-band)"/> | |
| <path d="M 349 266 C 455 280, 565 316, 704 354" fill="none" stroke="#69717A" | |
| stroke-width="3" marker-end="url(#arrow-grey)"/> | |
| <path d="M 349 263 C 450 229, 555 170, 677 108" fill="none" stroke="{BLUE}" | |
| stroke-width="3.5" marker-end="url(#arrow-blue)"/> | |
| <path d="M 349 263 L 716 82" fill="none" stroke="{INK}" stroke-width="2" | |
| stroke-dasharray="7 6" marker-end="url(#arrow-ink)"/> | |
| <circle cx="349" cy="263" r="12" fill="{INK}" stroke="white" stroke-width="4"/> | |
| <path d="M 190 217 L 332 254" fill="none" stroke="{INK}" stroke-width="1.5"/> | |
| <text class="label" x="55" y="188">Same step-0 function</text> | |
| <text class="formula" x="55" y="216">W_eff,0 = W</text> | |
| <text class="body" x="55" y="241">B₀ = 0; A₀ changes only tangent access</text> | |
| <text class="label" x="430" y="145" fill="{BLUE}">Fisher-informed tangent space</text> | |
| <text class="body" x="458" y="170" fill="{BLUE}">better task-gradient access</text> | |
| <text class="label" x="432" y="335" fill="#5E6670">PEFT tangent space</text> | |
| <text class="body" x="458" y="360" fill="#5E6670">can remain poorly aligned</text> | |
| <text class="body" x="585" y="63">full task gradient G</text> | |
| <text class="label" x="684" y="110" text-anchor="middle" fill="{BLUE}">Task-aligned</text> | |
| <text class="body" x="684" y="134" text-anchor="middle" fill="{BLUE}">descent region</text> | |
| <text class="label" x="700" y="350" text-anchor="middle" fill="#A83E37">Ineffective region</text> | |
| <text class="body" x="700" y="374" text-anchor="middle" fill="#A83E37">fragile early training</text> | |
| <text class="small" x="46" y="432" fill="#69717A">Conceptual geometry; not a measured 2-D loss surface.</text> | |
| <line x1="844" y1="45" x2="844" y2="450" stroke="#D6D9DC" stroke-width="1.5"/> | |
| <text class="panel" x="866" y="28">B</text> | |
| <text class="heading" x="897" y="28">Empirical signature</text> | |
| <rect class="card" x="866" y="45" width="310" height="155" rx="14"/> | |
| <text class="body" x="888" y="75">Initial gradient access</text> | |
| <text class="small" x="888" y="97">Input-Top mechanism probe</text> | |
| <text x="1021" y="139" text-anchor="middle" font-size="32" font-weight="700" fill="{BLUE}">{peft:.2f} → {fisher:.2f}</text> | |
| <text x="1021" y="174" text-anchor="middle" font-size="22" font-weight="700" fill="{BLUE}">≈ {ratio}×</text> | |
| <rect class="card" x="866" y="216" width="310" height="234" rx="14"/> | |
| <text class="body" x="888" y="247">Low-accuracy cells</text> | |
| <text class="small" x="888" y="269">Five fragile blocks; family envelope</text> | |
| <text class="small" x="888" y="304">PEFT</text> | |
| <rect x="932" y="290" width="214" height="19" rx="9" fill="#E8EAEC"/> | |
| <rect x="932" y="290" width="{214 * peft_low / 75:.1f}" height="19" rx="9" fill="#D55E5E"/> | |
| <text class="body" x="1146" y="306" text-anchor="end">{peft_low} / 75</text> | |
| <text class="small" x="888" y="350" fill="{BLUE}">Fisher</text> | |
| <rect x="932" y="336" width="214" height="19" rx="9" fill="#E8EAEC"/> | |
| <rect x="932" y="336" width="{214 * fisher_low / 75:.1f}" height="19" rx="9" fill="#D55E5E"/> | |
| <text class="body" x="1146" y="352" text-anchor="end">{fisher_low} / 75</text> | |
| <text x="1021" y="402" text-anchor="middle" font-size="31" font-weight="700" fill="{BLUE}">{peft_low} → {fisher_low}</text> | |
| <text class="label" x="1021" y="430" text-anchor="middle" fill="{BLUE}">−75% lower tail</text> | |
| </svg> | |
| """ | |
| ) | |
| path = output_dir / "fig1_method.svg" | |
| path.write_text(svg) | |
| return path | |
| def plot_accuracy_gain(data: dict[str, object], output_dir: Path) -> Path: | |
| plt = _pyplot() | |
| rows = _fragile_rows(data) | |
| x = np.arange(len(rows)) | |
| floor = 65.0 | |
| fig, ax = plt.subplots(figsize=(3.05, 2.65)) | |
| fig.subplots_adjust(left=0.16, right=0.985, bottom=0.29, top=0.82) | |
| fisher_bars = ax.bar( | |
| x, | |
| [row["fisher_mean"] - floor for row in rows], | |
| bottom=floor, | |
| width=0.68, | |
| color=BLUE, | |
| edgecolor=BLUE, | |
| linewidth=0.9, | |
| alpha=0.20, | |
| label="Calibration envelope", | |
| zorder=2, | |
| ) | |
| peft_bars = ax.bar( | |
| x, | |
| [row["peft_mean"] - floor for row in rows], | |
| bottom=floor, | |
| width=0.68, | |
| color=ORANGE, | |
| edgecolor=ORANGE, | |
| linewidth=0.9, | |
| alpha=0.48, | |
| label="PEFT", | |
| zorder=3, | |
| ) | |
| gain_bars = ax.bar( | |
| x, | |
| [row["gain"] for row in rows], | |
| bottom=[row["peft_mean"] for row in rows], | |
| width=0.68, | |
| facecolor="none", | |
| edgecolor=BLUE, | |
| linewidth=0.9, | |
| hatch="////", | |
| label="Difference", | |
| zorder=4, | |
| ) | |
| for gid, bars in (("fisher", fisher_bars), ("peft", peft_bars), ("gain", gain_bars)): | |
| for bar in bars: | |
| bar.set_gid(gid) | |
| ax.set_ylim(floor, 95) | |
| ax.set_ylabel("Mean accuracy (%)") | |
| ax.set_xticks( | |
| x, | |
| [f"{row['label'].replace(' ', chr(10), 1)}\n{row['variant']}" for row in rows], | |
| fontsize=5.8, | |
| ) | |
| ax.grid(axis="y", color="#E4E4E4", linewidth=0.6) | |
| ax.set_axisbelow(True) | |
| ax.spines[["top", "right"]].set_visible(False) | |
| handles, labels = ax.get_legend_handles_labels() | |
| ax.legend( | |
| [handles[index] for index in (1, 0, 2)], | |
| ["PEFT", "Envelope", "Gain"], | |
| frameon=False, | |
| ncol=3, | |
| loc="lower center", | |
| bbox_to_anchor=(0.5, 1.01), | |
| columnspacing=1.0, | |
| handlelength=1.6, | |
| ) | |
| for center, row in zip(x, rows): | |
| arrow_x = center + 0.27 | |
| ax.text( | |
| center, | |
| row["fisher_mean"] + 0.35, | |
| f"{row['fisher_mean']:.2f}", | |
| ha="center", | |
| va="bottom", | |
| color=BLUE, | |
| fontsize=6, | |
| weight="bold", | |
| ) | |
| ax.text( | |
| center, | |
| row["peft_mean"] - 0.35, | |
| f"{row['peft_mean']:.2f}", | |
| ha="center", | |
| va="top", | |
| color="#9A4219", | |
| fontsize=6, | |
| weight="bold", | |
| ) | |
| ax.annotate( | |
| "", | |
| xy=(arrow_x, row["fisher_mean"]), | |
| xytext=(arrow_x, row["peft_mean"]), | |
| arrowprops={"arrowstyle": "<->", "color": BLUE, "linewidth": 1.1}, | |
| zorder=6, | |
| ) | |
| ax.text( | |
| center, | |
| (row["peft_mean"] + row["fisher_mean"]) / 2, | |
| f"+{row['gain']:.2f} pp", | |
| ha="center", | |
| va="center", | |
| color=BLUE, | |
| fontsize=5.8, | |
| weight="bold", | |
| bbox={"facecolor": "white", "edgecolor": "none", "alpha": 0.88, "pad": 0.8}, | |
| zorder=7, | |
| ) | |
| for offset in (0.015, 0.045): | |
| ax.plot( | |
| (-0.008, 0.008), | |
| (offset - 0.012, offset + 0.012), | |
| transform=ax.transAxes, | |
| color=INK, | |
| linewidth=1.0, | |
| clip_on=False, | |
| ) | |
| path = _save(fig, output_dir, "fig2_accuracy_gain") | |
| plt.close(fig) | |
| return path | |
| def plot_threshold_sensitivity(data: dict[str, object], output_dir: Path) -> Path: | |
| plt = _pyplot() | |
| cells = data["broad_cells"] | |
| peft_values, fisher_values = [], [] | |
| for _, block, fisher_method in FRAGILE_BLOCKS: | |
| peft_values.extend(value * 100 for value in _values(cells, block, "peft_default")) | |
| fisher_values.extend(value * 100 for value in _values(cells, block, fisher_method)) | |
| thresholds = np.arange(55.0, 90.01, 0.25) | |
| peft_rates = np.searchsorted(np.sort(peft_values), thresholds, side="left") / len(peft_values) * 100 | |
| fisher_rates = np.searchsorted(np.sort(fisher_values), thresholds, side="left") / len(fisher_values) * 100 | |
| cutoff = LOW_THRESHOLD * 100 | |
| peft_count = sum(value < cutoff for value in peft_values) | |
| fisher_count = sum(value < cutoff for value in fisher_values) | |
| peft_rate = peft_count / len(peft_values) * 100 | |
| fisher_rate = fisher_count / len(fisher_values) * 100 | |
| fig, ax = plt.subplots(figsize=(3.05, 2.35)) | |
| fig.subplots_adjust(left=0.18, right=0.975, bottom=0.22, top=0.96) | |
| peft_line, = ax.step( | |
| thresholds, | |
| peft_rates, | |
| where="post", | |
| color=ORANGE, | |
| linewidth=2, | |
| label="PEFT", | |
| zorder=3, | |
| ) | |
| fisher_line, = ax.step( | |
| thresholds, | |
| fisher_rates, | |
| where="post", | |
| color=BLUE, | |
| linewidth=2, | |
| linestyle="--", | |
| label="Calibration envelope", | |
| zorder=3, | |
| ) | |
| peft_line.set_gid("peft-threshold") | |
| fisher_line.set_gid("fisher-threshold") | |
| ax.fill_between( | |
| thresholds, | |
| fisher_rates, | |
| peft_rates, | |
| step="post", | |
| color=BLUE, | |
| alpha=0.10, | |
| zorder=1, | |
| ) | |
| ax.axvline(cutoff, color="#777777", linestyle="--", linewidth=0.9, zorder=2) | |
| ax.scatter([cutoff], [peft_rate], color=ORANGE, s=30, zorder=4) | |
| ax.scatter([cutoff], [fisher_rate], color=BLUE, s=30, zorder=4) | |
| ax.text( | |
| cutoff + 0.6, | |
| peft_rate + 2, | |
| f"{peft_count}/75 ({peft_rate:.0f}%)", | |
| color=ORANGE, | |
| fontsize=7, | |
| weight="bold", | |
| va="bottom", | |
| ) | |
| ax.text( | |
| cutoff + 0.6, | |
| fisher_rate - 1, | |
| f"{fisher_count}/75 ({fisher_rate:.0f}%)", | |
| color=BLUE, | |
| fontsize=7, | |
| weight="bold", | |
| va="top", | |
| ) | |
| ax.text( | |
| cutoff - 0.6, | |
| 73, | |
| "70% cutoff", | |
| color="#555555", | |
| fontsize=7, | |
| ha="right", | |
| va="top", | |
| ) | |
| ax.set_xlim(55, 90) | |
| ax.set_ylim(0, 75) | |
| ax.set_xlabel("Accuracy threshold (%)") | |
| ax.set_ylabel("Cells below threshold (%)") | |
| ax.grid(axis="y", color="#E4E4E4", linewidth=0.6) | |
| ax.set_axisbelow(True) | |
| ax.spines[["top", "right"]].set_visible(False) | |
| ax.legend(frameon=False, ncol=1, loc="upper right", fontsize=6, handlelength=2.0) | |
| ax.text( | |
| 0.99, | |
| 0.04, | |
| "Lower is better", | |
| transform=ax.transAxes, | |
| ha="right", | |
| va="bottom", | |
| color="#555555", | |
| fontsize=7, | |
| ) | |
| path = _save(fig, output_dir, "fig3_threshold_sensitivity") | |
| plt.close(fig) | |
| return path | |
| def write_captions(output_dir: Path) -> Path: | |
| captions = dedent( | |
| r""" | |
| % Requires \usepackage{svg} | |
| \begin{figure*}[t] | |
| \centering | |
| \includesvg[width=\textwidth]{paper_figures/main/fig1_method} | |
| \caption{Calibration-informed initialization selects the initial LoRA basis while preserving the pretrained | |
| function at step zero. Input-factor variants use eigendirections of $S_X$; paired variants use | |
| Fisher-scored right singular directions. In the reported diagnostic probe, the raw unnormalized | |
| statistic is 0.0353 for PEFT and 0.9166 for Input-Top.} | |
| \label{fig:method} | |
| \end{figure*} | |
| \begin{figure}[t] | |
| \centering | |
| \includegraphics[width=\columnwidth]{paper_figures/main/fig2_accuracy_gain.pdf} | |
| \caption{Mean accuracy for PEFT and the calibration-moment envelope in five outcome-conditioned | |
| blocks. The translucent bars show absolute accuracy; blue hatching and arrows mark the observed | |
| difference. Each value aggregates 15 rank--seed cells. The envelope is selected within each block. | |
| Across the same 75 cells, 28 PEFT and 7 envelope cells are below 70\% accuracy. The y-axis is | |
| truncated at 65\%.} | |
| \label{fig:accuracy-gain} | |
| \end{figure} | |
| \begin{figure}[t] | |
| \centering | |
| \includegraphics[width=\columnwidth]{paper_figures/main/fig3_threshold_sensitivity.pdf} | |
| \caption{Threshold sensitivity of the same 75 rank--seed cells. | |
| Curves report the fraction of cells below each accuracy threshold; lower is better. The selected | |
| calibration-moment envelope remains below PEFT from 55\% to 90\%. At the primary 70\% cutoff, the | |
| rates are 28/75 (37\%) for PEFT and 7/75 (9\%) for the envelope.} | |
| \label{fig:threshold-sensitivity} | |
| \end{figure} | |
| """ | |
| ).lstrip() | |
| path = output_dir / "captions.tex" | |
| path.write_text(captions) | |
| return path | |
| def build(repo: Path, output_dir: Path) -> list[Path]: | |
| data = load_data(repo) | |
| validate(data) | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| paths = [ | |
| write_method_svg(data, output_dir), | |
| plot_accuracy_gain(data, output_dir), | |
| plot_threshold_sensitivity(data, output_dir), | |
| write_captions(output_dir), | |
| ] | |
| return paths | |
| def main(argv: list[str] | None = None) -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--repo", type=Path, default=Path(__file__).resolve().parents[1]) | |
| parser.add_argument("--output-dir", type=Path, default=Path("paper_figures/main")) | |
| args = parser.parse_args(argv) | |
| output_dir = args.output_dir if args.output_dir.is_absolute() else args.repo / args.output_dir | |
| build(args.repo, output_dir) | |
| print(f"wrote 3 main figures to {output_dir}") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 19.5 kB
- Xet hash:
- 750113c0e7335626f6faf7a2ba7683ee1f06b788bace397fc53053df8b797e81
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.