Download loopq_quantization/scripts/loopq/readiness.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 13.4 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/readiness.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/readiness.py
-
curl -L -o readiness.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/readiness.py
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']}") | |