strata-headquotient-q25 / evaluate_release.py
nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
6.39 kB
#!/usr/bin/env python3
"""Re-run Q25/dense PPL and controlled semantic confirmation."""
from __future__ import annotations
import argparse
from collections import Counter
import json
import math
from pathlib import Path
import sys
import numpy as np
import torch
ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT / "source/src"))
sys.path.insert(0, str(ROOT / "source/scripts"))
import run_strata_headquotient_v1_1_frontier as frontier # noqa: E402
from strata.eval.head_quotient import bucket_loss_sums, dense_bucket_loss_sums # noqa: E402
from strata.eval.head_quotient_causal import ( # noqa: E402
evaluate_scoped_variants,
prepare_causal_cases,
)
from strata.experiments.compose_rf import load_dense_base # noqa: E402
from strata.training.lm_data import PackedLMDataset # noqa: E402
BUCKETS = ("0-2048", "2048-4096", "4096-8192")
LANGUAGES = ("en", "zh", "de", "es", "ar")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", default="cuda:0" if torch.cuda.is_available() else "cpu")
parser.add_argument("--documents", type=int, default=470)
parser.add_argument("--semantic-examples", type=int, default=2000)
parser.add_argument("--bootstrap-samples", type=int, default=10000)
parser.add_argument("--output", type=Path, default=ROOT / "reproduced_evaluation.json")
return parser.parse_args()
def loss_row(row: dict[str, tuple[float, int]]) -> tuple[dict[str, float], float]:
buckets = {key: total / count for key, (total, count) in row.items()}
return buckets, sum(value[0] for value in row.values()) / sum(
value[1] for value in row.values()
)
def interval(values: np.ndarray, samples: int, seed: int) -> dict[str, float | int]:
rng = np.random.default_rng(seed)
draws = np.empty(samples, dtype=np.float64)
for start in range(0, samples, 1000):
stop = min(samples, start + 1000)
indices = rng.integers(0, values.size, size=(stop - start, values.size))
draws[start:stop] = values[indices].mean(axis=1)
lower, upper = np.quantile(draws, (0.025, 0.975))
mean = float(values.mean())
return {
"documents": int(values.size),
"nll_difference": mean,
"ppl_ratio": math.exp(mean),
"ppl_ratio_lower": math.exp(float(lower)),
"ppl_ratio_upper": math.exp(float(upper)),
}
def main() -> None:
args = parse_args()
if not 1 <= args.documents <= 470:
raise ValueError("--documents must be between 1 and 470")
if not 1 <= args.semantic_examples <= 2000:
raise ValueError("--semantic-examples must be between 1 and 2000")
device = torch.device(args.device)
if device.type == "cuda":
torch.cuda.set_device(device)
torch.manual_seed(20260721)
config = json.loads((ROOT / "config/headquotient.json").read_text(encoding="utf-8"))
config["dense_checkpoint"] = str(ROOT / "checkpoints/dense_model.pt")
config["model_config"] = "config/model.json"
config["heldout_corpus"] = str(ROOT / "evaluation/heldout-5lang-1m")
plan = json.loads((ROOT / "selection/Q25/plan.json").read_text(encoding="utf-8"))
scoped = json.loads((ROOT / "scoped_adapter/result.json").read_text(encoding="utf-8"))
scoped["adapter"]["path"] = str(ROOT / "scoped_adapter/adapter.pt")
frontier.ROOT = ROOT
dense, _ = load_dense_base(
ROOT / "config/model.json", ROOT / "checkpoints/dense_model.pt", device
)
dense.eval()
model, _model_config, _selected, _graph_groups, _coverage, _compaction = (
frontier.construct_export(config, plan, scoped, device)
)
model.load_state_dict(
torch.load(ROOT / "checkpoints/q25_export.pt", map_location=device, weights_only=True),
strict=True,
)
model.eval()
with np.load(ROOT / "evaluation/q25_470_documents.npz") as bundle:
tokens = bundle["tokens"][: args.documents]
languages = bundle["languages"][: args.documents]
rows = []
with torch.inference_mode():
for index, values in enumerate(tokens):
ids = torch.from_numpy(values.astype(np.int64)).unsqueeze(0).to(device)
dense_buckets, dense_all = loss_row(dense_bucket_loss_sums(dense, ids))
q25_buckets, q25_all = loss_row(bucket_loss_sums(model, ids))
rows.append({
"language": str(languages[index]),
"difference": q25_all - dense_all,
"buckets": {
key: q25_buckets[key] - dense_buckets[key] for key in BUCKETS
},
})
if (index + 1) % 10 == 0:
print(json.dumps({"evaluated_documents": index + 1}), flush=True)
differences = np.asarray([row["difference"] for row in rows], dtype=np.float64)
overall = interval(differences, args.bootstrap_samples, 20260722)
position = {
bucket: interval(
np.asarray([row["buckets"][bucket] for row in rows]),
args.bootstrap_samples,
20260822 + offset,
)
for offset, bucket in enumerate(BUCKETS)
}
language = {}
for offset, name in enumerate(LANGUAGES):
values = np.asarray([
row["difference"] for row in rows if row["language"] == name
], dtype=np.float64)
if values.size:
language[name] = interval(values, args.bootstrap_samples, 20260922 + offset)
semantic_dataset = PackedLMDataset(config["heldout_corpus"], seq_len=64)
semantic_cases = prepare_causal_cases(
model, semantic_dataset, args.semantic_examples, device, start=8192
)
semantic, _margins, null_exact = evaluate_scoped_variants(
model, semantic_cases, batch_size=32
)
payload = {
"documents": len(rows),
"languages": dict(Counter(str(item) for item in languages[: len(rows)])),
"ppl": overall,
"position_buckets": position,
"per_language": language,
"semantic_examples": args.semantic_examples,
"semantic": semantic,
"null_controls_exact": null_exact,
"noninferiority_pass": float(overall["ppl_ratio_upper"]) < 1.03,
}
args.output.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
print(json.dumps(payload, indent=2, sort_keys=True))
if __name__ == "__main__":
main()