| #!/usr/bin/env python3 | |
| """Export the existing ACL/ARR cells, aggregates, and factor metadata.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import re | |
| from collections import defaultdict | |
| from pathlib import Path | |
| from statistics import mean | |
| from typing import Any | |
| import numpy as np | |
| from safetensors import safe_open | |
| from scripts.build_acl_arr_figures import ( | |
| BROAD_BLOCKS, | |
| FRAGILE_BLOCKS, | |
| _alpha_cells, | |
| _cells_from_results, | |
| ) | |
| METHOD_MULTIPLIER = { | |
| "fpeft_low": 0.5, | |
| "fpeft_bi_low": 0.5, | |
| "fpeft_bi_high": 0.5, | |
| "peft_default": 1.0, | |
| "filet": 1.0, | |
| } | |
| CALIBRATION_METHODS = ("fpeft_low", "fpeft_bi_low", "fpeft_bi_high") | |
| ALPHA_METHODS = ("fpeft_low", "fpeft_bi_high") | |
| ALPHA_CONTROLS = ("random_orthogonal", "peft_scaled_random") | |
| SEEDS = (42, 1337, 2024, 7, 123) | |
| FACTOR_ROOTS = { | |
| "llama": ("outputs/e1_0_additional_factors/en", "unsloth/Llama-3.1-8B"), | |
| "qwen3": ("outputs/factors/qwen3/en", "Qwen/Qwen3-8B"), | |
| "ministral": ("outputs/factors/ministral/en", "mistralai/Ministral-3-8B-Base-2512"), | |
| "gemma4": ("outputs/factors/gemma4/en", "google/gemma-4-E4B"), | |
| } | |
| def _write(path: Path, fields: list[str], rows: list[dict[str, object]]) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", newline="") as handle: | |
| writer = csv.DictWriter(handle, fieldnames=fields) | |
| writer.writeheader() | |
| writer.writerows(rows) | |
| def _bootstrap(values: list[float], rng: np.random.Generator) -> list[float]: | |
| array = np.asarray(values) | |
| draws = array[rng.integers(0, len(array), (5000, len(array)))].mean(axis=1) | |
| return np.quantile(draws * 100, (0.025, 0.975)).tolist() | |
| def _reported_aggregates( | |
| repo: Path, | |
| broad: list[dict[str, object]], | |
| alpha: list[dict[str, object]], | |
| ) -> dict[str, Any]: | |
| broad_values: dict[tuple[str, str], list[float]] = defaultdict(list) | |
| for row in broad: | |
| broad_values[(str(row["block"]), str(row["method"]))].append(float(row["accuracy"])) | |
| all_blocks = sorted(block for block, method in broad_values if method == "peft_default") | |
| def fixed(method: str, blocks: list[str]) -> dict[str, Any]: | |
| selected = [block for block in blocks if (block, method) in broad_values] | |
| deltas = [ | |
| mean(broad_values[(block, method)]) - mean(broad_values[(block, "peft_default")]) | |
| for block in selected | |
| ] | |
| return { | |
| "blocks": len(selected), | |
| "wins": sum(delta > 0 for delta in deltas), | |
| "mean_delta_pp": mean(deltas) * 100, | |
| "calibration_low_cells": sum( | |
| value < 0.70 for block in selected for value in broad_values[(block, method)] | |
| ), | |
| "peft_low_cells": sum( | |
| value < 0.70 for block in selected for value in broad_values[(block, "peft_default")] | |
| ), | |
| } | |
| def oracle(blocks: list[str]) -> dict[str, Any]: | |
| chosen = { | |
| block: max( | |
| (method for method in CALIBRATION_METHODS if (block, method) in broad_values), | |
| key=lambda method: mean(broad_values[(block, method)]), | |
| ) | |
| for block in blocks | |
| } | |
| deltas = [ | |
| mean(broad_values[(block, method)]) - mean(broad_values[(block, "peft_default")]) | |
| for block, method in chosen.items() | |
| ] | |
| return { | |
| "blocks": len(blocks), | |
| "mean_delta_pp": mean(deltas) * 100, | |
| "calibration_low_cells": sum( | |
| value < 0.70 | |
| for block, method in chosen.items() | |
| for value in broad_values[(block, method)] | |
| ), | |
| "peft_low_cells": sum( | |
| value < 0.70 for block in blocks for value in broad_values[(block, "peft_default")] | |
| ), | |
| "chosen_method_by_block": chosen, | |
| } | |
| fragile = [block for _, block, _ in FRAGILE_BLOCKS] | |
| alpha_values = { | |
| (float(row["alpha"]), str(row["block"]), str(row["method"]), int(row["seed"])): | |
| float(row["accuracy"]) | |
| for row in alpha | |
| } | |
| alpha_blocks = sorted({str(row["block"]) for row in alpha}) | |
| alpha_means = { | |
| f"{alpha_value:g}": { | |
| method: mean( | |
| value | |
| for (candidate_alpha, _, candidate_method, _), value in alpha_values.items() | |
| if candidate_alpha == alpha_value and candidate_method == method | |
| ) * 100 | |
| for method in (*ALPHA_METHODS, *ALPHA_CONTROLS) | |
| } | |
| for alpha_value in (0.25, 0.5, 1.0, 2.0) | |
| } | |
| rng = np.random.default_rng(20260717) | |
| matched_scale: dict[str, Any] = {} | |
| for method in ALPHA_METHODS: | |
| for control in ALPHA_CONTROLS: | |
| block_deltas = [ | |
| mean( | |
| alpha_values[(1.0, block, method, seed)] | |
| - alpha_values[(1.0, block, control, seed)] | |
| for seed in SEEDS | |
| ) | |
| for block in alpha_blocks | |
| ] | |
| matched_scale[f"{method}__vs__{control}"] = { | |
| "mean_delta_pp": mean(block_deltas) * 100, | |
| "block_bootstrap_95_pp": _bootstrap(block_deltas, rng), | |
| } | |
| interactions: dict[str, Any] = {} | |
| for method in ALPHA_METHODS: | |
| for control in ALPHA_CONTROLS: | |
| block_differences = [ | |
| mean( | |
| (alpha_values[(2.0, block, method, seed)] | |
| - alpha_values[(0.25, block, method, seed)]) | |
| - (alpha_values[(2.0, block, control, seed)] | |
| - alpha_values[(0.25, block, control, seed)]) | |
| for seed in SEEDS | |
| ) | |
| for block in alpha_blocks | |
| ] | |
| interactions[f"{method}__vs__{control}"] = { | |
| "mean_difference_in_differences_pp": mean(block_differences) * 100, | |
| "block_bootstrap_95_pp": _bootstrap(block_differences, rng), | |
| } | |
| projection = json.loads( | |
| (repo / "outputs/e1_grad_diag_bestinit/summary.json").read_text() | |
| )["groups"] | |
| return { | |
| "bootstrap": {"replicates": 5000, "rng": "numpy.default_rng", "seed": 20260717}, | |
| "fixed_broad": {method: fixed(method, all_blocks) for method in CALIBRATION_METHODS}, | |
| "selection_audit": { | |
| "all_fixed_input_tail": fixed("fpeft_low", all_blocks), | |
| "all_per_block_oracle": oracle(all_blocks), | |
| "fragile_fixed_input_tail": fixed("fpeft_low", fragile), | |
| "fragile_fixed_paired_low": fixed("fpeft_bi_low", fragile), | |
| "fragile_per_block_oracle": oracle(fragile), | |
| }, | |
| "matched_scale": matched_scale, | |
| "alpha_means_percent": alpha_means, | |
| "extreme_multiplier_interactions": interactions, | |
| "step_zero_projection": { | |
| key.removesuffix("__collapsed_False"): float(value["projection_score_mean"]) | |
| for key, value in projection.items() | |
| }, | |
| } | |
| def _factor_manifest(repo: Path) -> list[dict[str, object]]: | |
| rows = [] | |
| for model, (relative_root, model_id) in FACTOR_ROOTS.items(): | |
| for path in sorted((repo / relative_root).rglob("*.safetensors")): | |
| with safe_open(path, framework="pt", device="cpu") as handle: | |
| rows.append({ | |
| "model": model, | |
| "model_id": model_id, | |
| "factor_checkpoint_revision": "unknown_not_archived", | |
| "calibration_source": "FLORES+ dev/devtest English, 2009 texts", | |
| "dataset_revision": "unknown_not_archived", | |
| "factor_file": path.relative_to(repo).as_posix(), | |
| "layer": int(re.search(r"layers__(\d+)__", path.name).group(1)), | |
| "module": path.parent.name, | |
| "input_moment_shape": "x".join(map(str, handle.get_slice("A").get_shape())), | |
| "output_moment_shape": "x".join(map(str, handle.get_slice("B").get_shape())), | |
| "count_a": int(handle.get_tensor("count_a").reshape(-1)[0]), | |
| "count_b": int(handle.get_tensor("count_b").reshape(-1)[0]), | |
| "file_bytes": path.stat().st_size, | |
| }) | |
| return rows | |
| def export(repo: Path, output_dir: Path) -> tuple[Path, Path, Path, Path]: | |
| # ponytail: reuse the figure pipeline's validated result loaders. | |
| broad = [] | |
| for cell in _cells_from_results(repo): | |
| broad.append({ | |
| **cell, | |
| "multiplier": METHOD_MULTIPLIER[cell["method"]], | |
| "status": "complete", | |
| "source": f"{BROAD_BLOCKS[cell['block']]}/results.json", | |
| }) | |
| broad.sort(key=lambda row: ( | |
| str(row["task"]), str(row["model"]), str(row["method"]), | |
| int(row["rank"]), int(row["seed"]), | |
| )) | |
| alpha = [] | |
| for cell in _alpha_cells(repo): | |
| alpha.append({ | |
| **cell, | |
| "rank": 32, | |
| "status": "complete", | |
| "source": ( | |
| "outputs/paper_upgrade_alpha_grid_20260720/" | |
| f"alpha{cell['alpha']:g}/comparison/{cell['block']}/results.json" | |
| ), | |
| }) | |
| alpha.sort(key=lambda row: ( | |
| float(row["alpha"]), str(row["block"]), str(row["method"]), int(row["seed"]), | |
| )) | |
| broad_path = output_dir / "broad_grid_cells.csv" | |
| alpha_path = output_dir / "alpha_grid_cells.csv" | |
| aggregate_path = output_dir / "reported_aggregates.json" | |
| manifest_path = output_dir / "factor_manifest.csv" | |
| _write( | |
| broad_path, | |
| ["block", "model", "task", "method", "multiplier", "rank", "seed", | |
| "accuracy", "status", "source"], | |
| broad, | |
| ) | |
| _write( | |
| alpha_path, | |
| ["alpha", "block", "method", "rank", "seed", "accuracy", "status", "source"], | |
| alpha, | |
| ) | |
| aggregate_path.write_text( | |
| json.dumps(_reported_aggregates(repo, broad, alpha), indent=2, sort_keys=True) + "\n" | |
| ) | |
| manifest = _factor_manifest(repo) | |
| _write( | |
| manifest_path, | |
| ["model", "model_id", "factor_checkpoint_revision", "calibration_source", "dataset_revision", | |
| "factor_file", "layer", "module", "input_moment_shape", "output_moment_shape", | |
| "count_a", "count_b", "file_bytes"], | |
| manifest, | |
| ) | |
| return broad_path, alpha_path, aggregate_path, manifest_path | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--repo", type=Path, default=Path(".")) | |
| parser.add_argument( | |
| "--output-dir", | |
| type=Path, | |
| default=Path("paper/acl_arr/supplement"), | |
| ) | |
| args = parser.parse_args() | |
| for path in export(args.repo.resolve(), args.output_dir): | |
| print(path) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 10.8 kB
- Xet hash:
- 005fc916d49b3daf26c65468a8b8f09e2ec4c676120864ce7f01de704fb690f1
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.