#!/usr/bin/env python3 """Select a mixed-precision override from measured candidate experiments.""" from __future__ import annotations import argparse import json import math import sys import tempfile from datetime import datetime, timezone from pathlib import Path PROJECT_ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(PROJECT_ROOT)) from quantization.release_gate import ( # noqa: E402 CALIBRATION_EVIDENCE_KEYS, dataset_manifest, load_json_object, sha256_file, validate_calibration_results, validate_release_thresholds, ) SUPPORTED_HEAD_PRECISIONS = {"bfloat16", "mxfp8", "affine8"} def _number(payload: dict, key: str) -> float | None: value = payload.get(key) if ( not isinstance(value, (int, float)) or isinstance(value, bool) or not math.isfinite(float(value)) ): return None return float(value) def select_head_precision( bf16_accuracy: dict, reference_performance: dict, experiments: list[dict], thresholds: dict | None = None, ) -> dict: """Choose the fastest experiment that passes the existing release limits.""" limits = validate_release_thresholds(thresholds) bf16_cer = _number(bf16_accuracy, "mean_cer") bf16_digit = _number(bf16_accuracy, "mean_digit_cer") bf16_table = _number(bf16_accuracy, "mean_table_score") reference_tps = _number(reference_performance, "mean_tps") if None in (bf16_cer, bf16_digit, bf16_table, reference_tps) or reference_tps <= 0: raise ValueError("Baseline accuracy and reference throughput must be complete") if not isinstance(experiments, list) or not experiments: raise ValueError("At least one calibration experiment is required") if any(not isinstance(experiment, dict) for experiment in experiments): raise ValueError("Each calibration experiment must be an object") labels = [experiment.get("label") for experiment in experiments] precisions = [experiment.get("precision") for experiment in experiments] if any( not isinstance(label, str) or not label.strip() or label.strip() != label for label in labels ): raise ValueError("Each calibration experiment needs a normalized label") if any(not isinstance(precision, str) for precision in precisions): raise ValueError("Each calibration experiment needs a precision") if len(set(labels)) != len(labels): raise ValueError("Calibration experiment labels must be unique") if len(set(precisions)) != len(precisions): raise ValueError("Calibration experiment precisions must be unique") if set(precisions) != SUPPORTED_HEAD_PRECISIONS: raise ValueError( "Calibration requires exactly one experiment for each supported precision" ) evaluated = [] for experiment in experiments: label = experiment.get("label") precision = experiment.get("precision") accuracy = experiment.get("accuracy") performance = experiment.get("performance") if precision not in SUPPORTED_HEAD_PRECISIONS: raise ValueError(f"Unsupported head precision for {label}: {precision}") if not isinstance(accuracy, dict) or not isinstance(performance, dict): raise ValueError(f"Calibration metrics are missing for {label}") candidate_cer = _number(accuracy, "mean_cer") candidate_digit = _number(accuracy, "mean_digit_cer") candidate_table = _number(accuracy, "mean_table_score") candidate_tps = _number(performance, "mean_tps") metrics_complete = None not in ( candidate_cer, candidate_digit, candidate_table, candidate_tps, ) deltas = { "cer_vs_bf16": candidate_cer - bf16_cer if metrics_complete else None, "digit_cer_vs_bf16": candidate_digit - bf16_digit if metrics_complete else None, "table_degradation_vs_bf16": bf16_table - candidate_table if metrics_complete else None, "tps_ratio_vs_reference": candidate_tps / reference_tps if metrics_complete else None, } checks = { "cer": metrics_complete and deltas["cer_vs_bf16"] <= limits["max_cer_delta_vs_bf16"], "digit_cer": metrics_complete and deltas["digit_cer_vs_bf16"] <= limits["max_digit_cer_delta_vs_bf16"], "table_score": metrics_complete and deltas["table_degradation_vs_bf16"] <= limits["max_table_score_degradation_vs_bf16"], "throughput": metrics_complete and deltas["tps_ratio_vs_reference"] >= limits["min_tps_ratio_vs_reference"], } evaluated.append( { "label": label, "precision": precision, "passed": all(checks.values()), "checks": checks, "metrics": { "mean_cer": candidate_cer, "mean_digit_cer": candidate_digit, "mean_table_score": candidate_table, "mean_tps": candidate_tps, }, "deltas": deltas, } ) passing = [experiment for experiment in evaluated if experiment["passed"]] if not passing: raise RuntimeError( "No LM-head calibration experiment passed every release limit" ) selected = max(passing, key=lambda experiment: experiment["metrics"]["mean_tps"]) return { "schema_version": 3, "created_at": datetime.now(timezone.utc).isoformat(), "target_pattern": "language_model.lm_head", "selection_policy": "fastest candidate passing existing quality and throughput limits", "thresholds": limits, "experiments": evaluated, "selected": { "label": selected["label"], "precision": selected["precision"], }, "precision_overrides": { "language_model.lm_head": selected["precision"], }, } def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--bf16-accuracy", required=True, type=Path) parser.add_argument("--reference-performance", required=True, type=Path) parser.add_argument( "--calibration-dir", required=True, type=Path, help="Selection-only dataset used for every calibration accuracy run", ) parser.add_argument("--source-revision", required=True) parser.add_argument("--reference-revision", required=True) parser.add_argument( "--experiment", action="append", nargs=4, metavar=("LABEL", "PRECISION", "ACCURACY_JSON", "PERFORMANCE_JSON"), required=True, ) parser.add_argument("--output", required=True, type=Path) args = parser.parse_args() bf16_accuracy = load_json_object(args.bf16_accuracy) reference_performance = load_json_object(args.reference_performance) experiments = [] calibration_dataset = dataset_manifest(args.calibration_dir) evidence = { "calibration_baseline_accuracy": bf16_accuracy, "calibration_reference_performance": reference_performance, } evidence_paths: dict[str, Path] = { "calibration_baseline_accuracy": args.bf16_accuracy, "calibration_reference_performance": args.reference_performance, } for label, precision, accuracy_path, performance_path in args.experiment: if precision not in CALIBRATION_EVIDENCE_KEYS: raise ValueError(f"Unsupported calibration precision: {precision}") accuracy_key, performance_key = CALIBRATION_EVIDENCE_KEYS[precision] accuracy_path = Path(accuracy_path) performance_path = Path(performance_path) evidence[accuracy_key] = load_json_object(accuracy_path) evidence[performance_key] = load_json_object(performance_path) evidence_paths[accuracy_key] = accuracy_path evidence_paths[performance_key] = performance_path experiments.append( { "label": label, "precision": precision, "accuracy": evidence[accuracy_key], "performance": evidence[performance_key], } ) result = select_head_precision( bf16_accuracy, reference_performance, experiments, ) expected_input_keys = { "calibration_baseline_accuracy", "calibration_reference_performance", *(name for pair in CALIBRATION_EVIDENCE_KEYS.values() for name in pair), } if set(evidence_paths) != expected_input_keys: raise ValueError("Calibration requires one unique evidence pair per precision") resolved_paths = [path.resolve(strict=True) for path in evidence_paths.values()] if len(set(resolved_paths)) != len(resolved_paths): raise ValueError("Calibration input artifacts must be distinct files") result["input_artifacts"] = { name: { "filename": path.name, "size": path.stat().st_size, "sha256": sha256_file(path), } for name, path in evidence_paths.items() } result["dataset"] = calibration_dataset validate_calibration_results( result, bf16_accuracy=bf16_accuracy, reference_performance=reference_performance, calibration_dataset=calibration_dataset, source_revision=args.source_revision, reference_revision=args.reference_revision, evidence=evidence, evidence_paths=evidence_paths, ) if args.output.is_symlink() or (args.output.exists() and not args.output.is_file()): raise ValueError(f"Output must be a regular file: {args.output}") args.output.parent.mkdir(parents=True, exist_ok=True) temporary_path = None try: with tempfile.NamedTemporaryFile( mode="w", encoding="utf-8", dir=args.output.parent, prefix=f".{args.output.name}.", suffix=".tmp", delete=False, ) as temporary: temporary_path = Path(temporary.name) json.dump(result, temporary, indent=2, ensure_ascii=False, allow_nan=False) temporary.write("\n") temporary_path.replace(args.output) temporary_path = None finally: if temporary_path is not None: temporary_path.unlink(missing_ok=True) print(json.dumps(result["selected"], ensure_ascii=False)) if __name__ == "__main__": main()