| #!/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.