File size: 13,359 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8a4a24
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
"""Fail-closed prerequisites for full Ouro paper-benchmark validation runs.

These checks establish traceability and bounded parity evidence, not identical
author settings, authentic authorship of a checkpoint, or all-input parity.
"""
import hashlib
import math
from collections.abc import Mapping
import torch

from .export import validate_export_artifact
from .evaluation_provenance import source_paths
from .objective import KL_MODES


def source_digest(root) -> str:
    digest = hashlib.sha256()
    for path in source_paths(root):
        digest.update(str(path.relative_to(root)).encode() + b"\0")
        digest.update(hashlib.sha256(path.read_bytes()).digest())
    return digest.hexdigest()


def require_ouro_paper_calibration(artifact: dict, *, study: bool = False) -> None:
    validate_export_artifact(artifact)
    model, c = artifact["model"], artifact["calibration"]
    if model.get("architecture") != "ouro" or model.get("physical_layers") != 24 or model.get("loop_count") != 4:
        raise ValueError("paper readiness currently supports Ouro-1.4B, 24 layers / 4 loops")
    if c.get("completed") is not True or (not study and c.get("paper_calibration") is not True):
        raise ValueError("full paper evaluation requires completed full calibration")
    from .experiment_contract import STUDY_SAMPLES, SLT_BUDGETS
    study_info = c.get("study") or {}
    samples = study_info.get("samples") if study else 1024
    if study and (c.get("run_kind") != "study" or c.get("paper_calibration") is not False
                  or samples not in STUDY_SAMPLES or not study_info.get("axes")
                  or not set(study_info["axes"]) <= {"calibration_size", "slt_budget"}):
        raise ValueError("study evaluation requires explicit, non-paper study identity")
    if c.get("numerical_contract") != artifact["quantization"]["numerical_contract"]:
        raise ValueError("calibration must record the current numerical contract")
    data = c.get("data", {})
    expected = {"dataset": "mit-han-lab/pile-val-backup", "split": "validation",
                "text_field": "text", "samples": samples, "max_length": 256,
                "dataset_revision": "2f5e46ae6a69cf0dce4b12f78241c408936ca0e4"}
    if any(data.get(k) != v for k, v in expected.items()):
        raise ValueError(f"calibration must record Pile validation/text, {samples} samples, max length 256")
    if len(data.get("dataset_revision", "")) != 40:
        raise ValueError("calibration dataset revision must be pinned")
    records = c.get("token_records", [])
    if len(records) != samples:
        raise ValueError(f"full calibration must retain {samples} actual token records")
    for index, record in enumerate(records):
        ids, mask = record.get("input_ids", []), record.get("attention_mask", [])
        if (record.get("sample_index") != index or len(ids) != 1 or not 0 < len(ids[0]) <= 256
                or any(type(i) is not int or i < 0 for i in ids[0]) or mask != [[1] * len(ids[0])]):
            raise ValueError("invalid unpadded calibration token record")
    choices = c.get("algorithm_choices", {})
    if choices.get("kl_mode") not in KL_MODES[:2]:
        raise ValueError("paper Table 6 uses top-1000; full-vocabulary KL is diagnostic only")
    if choices.get("statistics_coordinate") not in {"svd_parameters", "effective_factors"}:
        raise ValueError("missing statistics coordinate")
    if choices.get("fisher_estimator") not in {"trajectory_gradient_square", "model_score_mc"}:
        raise ValueError("missing Fisher estimator")
    expected_budget = 0 if c.get("ablation") == "no_slt" else 4
    if study:
        expected_budget = study_info.get("slt_budget")
        if expected_budget not in SLT_BUDGETS:
            raise ValueError("invalid study SLT budget")
        from .experiment_contract import calibration_identity
        expected_identity = calibration_identity(smoke=False, samples=samples,
            study_samples="calibration_size" in study_info["axes"],
            budget=expected_budget, ablation=c.get("ablation"))
        if expected_identity["study"] != study_info:
            raise ValueError("inconsistent study identity")
    if choices.get("slt_budget") != expected_budget or not c.get("resolved_optimizer"):
        raise ValueError("paper SLT budget or resolved optimizer missing")
    optimizer = c["resolved_optimizer"]
    if any(not isinstance(v, (int, float)) or not math.isfinite(v) for key in ("lr", "weight_decay", "eps")
           for v in [optimizer.get(key)]) or optimizer["lr"] <= 0 or optimizer["eps"] <= 0 or optimizer["weight_decay"] < 0:
        raise ValueError("resolved optimizer requires finite lr, weight_decay and eps")
    betas = optimizer.get("betas", ())
    if len(betas) != 2 or any(not isinstance(b, (float, int)) or not math.isfinite(b) or not 0 <= b < 1 for b in betas):
        raise ValueError("resolved optimizer requires two finite betas")
    if type(c.get("global_optimization_steps")) is not int or c["global_optimization_steps"] <= 0:
        raise ValueError("completed optimizer updates missing")
    components = artifact["components"]
    def check_finite(value):
        if isinstance(value, torch.Tensor) and not torch.isfinite(value).all():
            raise ValueError("non-finite tensor in calibrated components")
        if isinstance(value, Mapping):
            for item in value.values():
                check_finite(item)
        elif isinstance(value, (list, tuple)):
            for item in value:
                check_finite(item)
    check_finite(components)
    from adapters.ouro import PAPER_GROUPS
    from .las import LoopAwareActivationScales
    from .cta import CrossLoopTransitionAdapter
    from .transforms import SharedKroneckerTransform
    expected_groups = {f"model.layers.{layer}.{group}" for layer in range(24) for group in PAPER_GROUPS}
    shared, selected = components["shared_transforms"], components["selected_loop_transforms"]
    las = LoopAwareActivationScales.from_export_state(components["las"])
    if set(shared) != expected_groups or set(las.module_names) != expected_groups:
        raise ValueError("paper artifact requires all 24 x 4 transform and LAS groups")
    if not las.dynamic_clip or las.loop_count != 4 or las.group_size != 32:
        raise ValueError("current full-calibration profile requires dynamic module-loop LAS")
    if sum(p.numel() for p in las.parameters()) != (96 if c.get("ablation") == "no_las" else 384):
        raise ValueError("unexpected LAS parameter count")
    for key in las.module_names:
        for loop in range(4):
            scale = las.scales_for(key, loop)
            if not torch.isfinite(scale).all() or (scale <= 0).any():
                raise ValueError("LAS exponentiation produced invalid clipping multiplier")
    if len(selected) != expected_budget or not set(selected) <= expected_groups:
        raise ValueError("selected transform groups do not match the paper budget")
    def checked_transform(state):
        transform = SharedKroneckerTransform.from_export_state(state)
        for factor in (transform.left, transform.right):
            inverse, info = torch.linalg.inv_ex(factor.detach().double())
            if info.any() or not torch.isfinite(inverse).all():
                raise ValueError("calibrated transform is not invertible")
        return transform
    for key, state in shared.items():
        transform = checked_transform(state)
        width = 5632 if key.endswith("mlp_down") else 2048
        if transform.feature_size != width:
            raise ValueError("shared transform shape does not match Ouro projection")
    for key, states in selected.items():
        if set(states) != {"0", "1", "2", "3"}:
            raise ValueError("selected group must contain all four transforms")
        for state in states.values():
            if (state["left_size"], state["right_size"]) != (shared[key]["left_size"], shared[key]["right_size"]):
                raise ValueError("selected/shared factor dimensions differ")
            checked_transform(state)
    cta = CrossLoopTransitionAdapter.from_export_state(components["cta"])
    if (cta.hidden_size, cta.rank, cta.transition_count) != (2048, 8, 3):
        raise ValueError("paper Ouro CTA must have hidden 2048, rank 8, three transitions")


# Established repository convention from the Track B parity work
# (docs/cowork/jun/0827_parity_gate_results.md): a reference top-2 gap at or
# below this is a near-tie that no backend can be required to break the same
# way, because the reference is not self-consistent there either.
PARITY_TIE_THRESHOLD = 0.25


def top1_verdict(row: dict, *, tie_threshold: float = PARITY_TIE_THRESHOLD) -> tuple[bool, bool]:
    """Top-1 equality, excusing only ties the reference itself cannot break.

    Returns (accepted, excused_as_tie). A row is excused only when the reference
    top-2 gap is within `tie_threshold` AND both backends propose the same
    unordered top-2 pair, i.e. they disagree on ordering, not on candidates.
    """
    if row.get("top1_agreement") is True:
        return True, False
    tokens = row.get("token_ids") or []
    reference = row.get("hf_logprobs") or []
    deployed = row.get("vllm_logprobs") or []
    if not (len(tokens) == len(reference) == len(deployed)) or len(tokens) < 2:
        return False, False
    reference_order = sorted(range(len(tokens)), key=lambda i: -reference[i])
    deployed_order = sorted(range(len(tokens)), key=lambda i: -deployed[i])
    gap = reference[reference_order[0]] - reference[reference_order[1]]
    same_pair = ({tokens[reference_order[0]], tokens[reference_order[1]]}
                 == {tokens[deployed_order[0]], tokens[deployed_order[1]]})
    excused = bool(gap <= tie_threshold and same_pair)
    return excused, excused


def parity_gate(report: dict, *, max_logprob_error: float, max_hidden_error: float) -> dict:
    """A bounded, unmodified six-prompt prefill gate; never an all-input proof."""
    reasons = []
    tie_excused = 0
    if any(not math.isfinite(x) or x < 0 for x in (max_logprob_error, max_hidden_error)):
        raise ValueError("parity tolerances must be finite and non-negative")
    interventions = ("<redacted-hf-token>", "hf_vllm_rope", "hf_vllm_attention", "hf_torch_flash_attention",
                     "hf_offline_weights", "hf_packed_linears", "attention_replay_artifact")
    if any(report.get(name) for name in interventions):
        reasons.append("diagnostic intervention is not production parity evidence")
    rows = report.get("rows", [])
    if len(rows) < 6:
        reasons.append("at least six prefill fixtures are required")
    if len({tuple(row.get("prompt_token_ids", [])) for row in rows}) != len(rows):
        reasons.append("parity fixtures must be distinct")
    for row in rows:
        error, hidden = row.get("max_abs_logprob_error"), row.get("loop_relative_hidden_error", [])
        accepted, excused = top1_verdict(row)
        tie_excused += excused
        if not accepted or row.get("compared_logprob_count", 0) < 1000:
            reasons.append("top1 or top-1000 comparison incomplete")
        if not isinstance(error, (int, float)) or not math.isfinite(error) or not 0 <= error <= max_logprob_error:
            reasons.append("logprob error exceeds tolerance")
        if len(hidden) != 4 or any(not isinstance(x, (int, float)) or not math.isfinite(x) or not 0 <= x <= max_hidden_error for x in hidden):
            reasons.append("four-loop hidden-state parity failed")
    if report.get("source_changed_during_run"):
        reasons.append("source changed during parity measurement")
    packed = report.get("packed_storage_diagnostic")
    if packed is not None:
        if packed.get('manifest_changed_during_run'):
            reasons.append("packed manifest changed during parity measurement")
        comparisons = packed.get("dense_vllm_comparison", [])
        if len(comparisons) != len(rows) or any(
            item.get("same_token") is not True or item.get("same_topk_indices") is not True
            or item.get("compared_logprobs", 0) < 1000 or item.get("max_abs_logprob_error") != 0.0
            for item in comparisons):
            reasons.append("packed/dense prefill mismatch")
    return dict(passed=not reasons, reasons=sorted(set(reasons)),
                max_logprob_error=max_logprob_error, max_hidden_error=max_hidden_error,
                tie_threshold=PARITY_TIE_THRESHOLD, rows_excused_as_reference_tie=tie_excused,
                scope="six_or_more_single_sequence_prefills; decode/batch/long_context still require A40 validation")


def require_parity_report(report: dict, *, artifact_sha256: str, current_source_digest: str) -> None:
    if report.get("artifact_sha256") != artifact_sha256 or report.get("source_digest") != current_source_digest:
        raise ValueError("parity evidence does not match this artifact and current code")
    gate = report.get("acceptance", {})
    if "max_logprob_error" not in gate or "max_hidden_error" not in gate:
        raise ValueError("parity acceptance thresholds missing")
    result = parity_gate(report, max_logprob_error=gate["max_logprob_error"], max_hidden_error=gate["max_hidden_error"])
    if not result["passed"]:
        raise ValueError(f"prefill parity gate failed: {result['reasons']}")