etomoscow/mff_lora / code /scripts /build_acl_main_figures.py
etomoscow's picture
download
raw
19.5 kB
#!/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.