etomoscow/mff_lora / code /scripts /build_acl_arr_figures.py
etomoscow's picture
download
raw
5.77 kB
#!/usr/bin/env python3
"""Load and validate ACL ARR source artifacts used by the current paper."""
from __future__ import annotations
import json
import re
from pathlib import Path
from statistics import mean
from typing import Any
import numpy as np
LOW_THRESHOLD = 0.70
KEY_RE = re.compile(r"^(?P<method>.+)__rank(?P<rank>\d+)__seed(?P<seed>\d+)$")
ALPHA_METHODS = ("fpeft_low", "fpeft_bi_high", "random_orthogonal", "peft_scaled_random")
BROAD_BLOCKS = {
"llama_boolq": "outputs/task1_boolq_llama",
"llama_rte": "outputs/task4_rte_llama",
"llama_hellaswag": "outputs/task5_hellaswag_llama",
"qwen3_xnli": "outputs/task3_xnli_qwen3",
"qwen3_boolq": "outputs/task4_boolq_qwen3",
"qwen3_rte": "outputs/task4_rte_qwen3",
"qwen3_hellaswag": "outputs/task5_hellaswag_qwen3",
"ministral_xnli": "outputs/task3_xnli_ministral_fixed_loader",
"ministral_boolq": "outputs/task4_boolq_ministral_fixed_loader",
"ministral_rte": "outputs/task4_rte_ministral_fixed_loader",
"ministral_hellaswag": "outputs/task5_hellaswag_ministral",
"gemma4_xnli": "outputs/task3_xnli_gemma4_fixed",
"gemma4_boolq": "outputs/task4_boolq_gemma4_retry_freegpu",
"gemma4_rte": "outputs/task4_rte_gemma4_fixed",
}
FRAGILE_BLOCKS = (
("Ministral RTE", "ministral_rte", "fpeft_bi_high"),
("Llama BoolQ", "llama_boolq", "fpeft_low"),
("Gemma4 BoolQ", "gemma4_boolq", "fpeft_low"),
("Ministral BoolQ", "ministral_boolq", "fpeft_bi_high"),
("Gemma4 RTE", "gemma4_rte", "fpeft_bi_low"),
)
def _read(path: Path) -> Any:
if not path.exists():
raise FileNotFoundError(path)
return json.loads(path.read_text())
def _cells_from_results(repo: Path) -> list[dict[str, Any]]:
cells = []
for block, relative_dir in BROAD_BLOCKS.items():
model, task = block.rsplit("_", 1)
for key, row in _read(repo / relative_dir / "results.json").items():
if not isinstance(row, dict) or row.get("status") != "complete":
continue
match = KEY_RE.match(key)
if not match or row.get("eval_accuracy") is None:
continue
cells.append({
"block": block,
"model": model,
"task": task,
"method": match.group("method"),
"rank": int(match.group("rank")),
"seed": int(match.group("seed")),
"accuracy": float(row["eval_accuracy"]),
})
return cells
def _alpha_cells(repo: Path) -> list[dict[str, Any]]:
root = repo / "outputs/paper_upgrade_alpha_grid_20260720"
summary = _read(root / "comparison_summary.json")
cells = []
for alpha in summary["alphas"]:
for block in summary["blocks"]:
rows = _read(root / f"alpha{alpha:g}" / "comparison" / block / "results.json")
for method in ALPHA_METHODS:
for seed in summary["seeds"]:
row = rows.get(f"{method}__rank32__seed{seed}")
if not isinstance(row, dict) or row.get("status") != "complete":
raise ValueError(
f"incomplete alpha cell: alpha={alpha}, block={block}, "
f"{method}, seed={seed}"
)
cells.append({
"alpha": float(alpha),
"block": block,
"method": method,
"seed": int(seed),
"accuracy": float(row["eval_accuracy"]),
})
return cells
def load_data(repo: Path) -> dict[str, object]:
"""Load only the result families consumed by the current main figures."""
repo = Path(repo)
projection_groups = _read(repo / "outputs/e1_grad_diag_bestinit/summary.json")["groups"]
projection = {
key.removesuffix("__collapsed_False"): float(value["projection_score_mean"])
for key, value in projection_groups.items()
}
return {
"broad_cells": _cells_from_results(repo),
"projection": projection,
"alpha_summary": _read(
repo / "outputs/paper_upgrade_alpha_grid_20260720/comparison_summary.json"
),
"alpha_cells": _alpha_cells(repo),
}
def _values(cells: list[dict[str, Any]], block: str, method: str) -> list[float]:
values = [
cell["accuracy"]
for cell in cells
if cell["block"] == block and cell["method"] == method
]
if len(values) != 15:
raise ValueError(f"expected 15 broad cells for {block}/{method}, found {len(values)}")
return values
def validate(data: dict[str, object]) -> None:
"""Fail early if current paper source artifacts cannot reproduce key values."""
broad_cells = data["broad_cells"]
assert isinstance(broad_cells, list)
fisher_low = sum(
sum(value < LOW_THRESHOLD for value in _values(broad_cells, block, fisher))
for _, block, fisher in FRAGILE_BLOCKS
)
peft_low = sum(
sum(value < LOW_THRESHOLD for value in _values(broad_cells, block, "peft_default"))
for _, block, _ in FRAGILE_BLOCKS
)
data["fragile_low_counts"] = (peft_low, fisher_low)
assert data["fragile_low_counts"] == (28, 7)
assert np.isclose(data["projection"]["fpeft"], 0.9166, atol=1e-4)
assert np.isclose(data["projection"]["peft_default"], 0.0353, atol=1e-4)
assert len(data["alpha_cells"]) == 560
assert data["alpha_summary"]["cell_counts"]["complete"] == 560
for _, block, fisher in FRAGILE_BLOCKS:
fisher_mean = mean(_values(broad_cells, block, fisher))
peft_mean = mean(_values(broad_cells, block, "peft_default"))
assert fisher_mean > peft_mean, block

Xet Storage Details

Size:
5.77 kB
·
Xet hash:
38b719c498d0604171fcd8ece193231bbf5ebcc25a55b661513738e4aea564d3

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.