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()