etomoscow/mff_lora / code /scripts /aggregate_alpha_grid.py
etomoscow's picture
download
raw
4.83 kB
#!/usr/bin/env python3
"""Summarize the all-block alpha grid without changing raw results."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from statistics import mean, pstdev
from typing import Any
import numpy as np
BLOCKS = [
("qwen3", "xnli"),
("ministral", "xnli"),
("gemma4", "xnli"),
("llama", "rte"),
("qwen3", "rte"),
("ministral", "rte"),
("gemma4", "rte"),
]
METHODS = ["fpeft_low", "fpeft_bi_high", "random_orthogonal", "peft_scaled_random"]
FISHER_METHODS = ["fpeft_low", "fpeft_bi_high"]
CONTROL_METHODS = ["random_orthogonal", "peft_scaled_random"]
SEEDS = [42, 1337, 2024, 7, 123]
RANK = 32
COLLAPSE_THRESHOLD = 0.70
def _block_name(model: str, task: str) -> str:
return f"{model}_{task}"
def _alpha_tag(alpha: float) -> str:
return f"alpha{alpha:g}"
def _summary(values: list[float]) -> dict[str, Any]:
return {
"mean": float(mean(values)) if values else None,
"std": float(pstdev(values)) if values else None,
"n": len(values),
"collapse_cells": sum(value < COLLAPSE_THRESHOLD for value in values),
}
def _bootstrap_interval(values: list[float], rng: np.random.Generator) -> list[float] | None:
if not values:
return None
array = np.asarray(values)
samples = array[rng.integers(0, len(array), (5000, len(array)))].mean(axis=1)
return np.quantile(samples, (0.025, 0.975)).tolist()
def _load_cells(root: Path, alpha: float, block: str) -> dict[tuple[str, int], float]:
path = root / _alpha_tag(alpha) / "comparison" / block / "results.json"
rows = json.loads(path.read_text()) if path.exists() else {}
cells: dict[tuple[str, int], float] = {}
for method in METHODS:
for seed in SEEDS:
key = f"{method}__rank{RANK}__seed{seed}"
row = rows.get(key)
if not isinstance(row, dict) or row.get("status") != "complete":
raise ValueError(f"incomplete cell: {path} {key}")
cells[(method, seed)] = float(row["eval_accuracy"])
return cells
def _paired(
cells_by_block: dict[str, dict[tuple[str, int], float]],
fisher: str,
control: str,
rng: np.random.Generator,
) -> dict[str, Any]:
deltas = []
block_means = []
for block, cells in cells_by_block.items():
block_deltas = [cells[(fisher, seed)] - cells[(control, seed)] for seed in SEEDS]
deltas.extend(block_deltas)
block_means.append(mean(block_deltas))
return {
**_summary(deltas),
"block_bootstrap_95": _bootstrap_interval(block_means, rng),
}
def aggregate(root: Path, alphas: tuple[float, ...] | list[float]) -> dict[str, Any]:
root = Path(root)
per_alpha: dict[str, Any] = {}
complete = 0
for alpha in alphas:
rng = np.random.default_rng(20260717)
blocks: dict[str, Any] = {}
cells_by_block: dict[str, dict[tuple[str, int], float]] = {}
for model, task in BLOCKS:
block = _block_name(model, task)
cells = _load_cells(root, alpha, block)
cells_by_block[block] = cells
blocks[block] = {
method: _summary([cells[(method, seed)] for seed in SEEDS])
for method in METHODS
}
complete += len(cells)
per_alpha[_alpha_tag(alpha).removeprefix("alpha")] = {
"blocks": blocks,
"paired": {
f"{fisher}__vs__{control}": _paired(cells_by_block, fisher, control, rng)
for fisher in FISHER_METHODS
for control in CONTROL_METHODS
},
}
expected = len(BLOCKS) * len(METHODS) * len(SEEDS) * len(alphas)
return {
"rank": RANK,
"seeds": list(SEEDS),
"collapse_threshold": COLLAPSE_THRESHOLD,
"alphas": list(alphas),
"methods": list(METHODS),
"blocks": [_block_name(*block) for block in BLOCKS],
"cell_counts": {"expected": expected, "complete": complete, "errors": 0},
"per_alpha": per_alpha,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--root", default="outputs/paper_upgrade_alpha_grid_20260720")
parser.add_argument("--output", default="outputs/paper_upgrade_alpha_grid_20260720/comparison_summary.json")
parser.add_argument("--alphas", nargs="+", type=float, default=[0.25, 0.5, 1.0, 2.0])
args = parser.parse_args()
result = aggregate(Path(args.root), args.alphas)
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(result, indent=2, sort_keys=True) + "\n")
print(f"wrote {output} ({result['cell_counts']['complete']}/{result['cell_counts']['expected']} cells)")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
4.83 kB
·
Xet hash:
939c113c2296a4003fd20c211abed3f332e2235fc29d5e9c10d66e08c7499973

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