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']}")
|