File size: 6,390 Bytes
e69b72a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 | #!/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()
|