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