etomoscow/mff_lora / code /scripts /export_acl_arr_supplement.py
etomoscow's picture
download
raw
10.8 kB
#!/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.