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