JunYoungLee's picture
Archive anonymized LoopQ calibration artifacts from compute node 1
d8a4a24 verified
Raw History Blame Contribute Delete
13.4 kB
"""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']}")