#!/usr/bin/env python3 """Orchestrate verified OCR-aware quantization, validation, and publication.""" from __future__ import annotations import argparse import importlib.metadata import json import platform import re import subprocess import sys import time from datetime import datetime, timezone from pathlib import Path QUANT_DIR = Path(__file__).resolve().parent PROJECT_ROOT = QUANT_DIR.parent sys.path.insert(0, str(PROJECT_ROOT)) sys.path.insert(0, str(PROJECT_ROOT / "src")) from quantization.release_gate import ( # noqa: E402 calibration_raw_evidence_paths, dataset_manifest, dataset_separation, generate_precision_map, load_json_object, sha256_file, validate_calibration_results, ) DEFAULT_SOURCE = PROJECT_ROOT / "reference" / "Unlimited-OCR" DEFAULT_SOURCE_ID = "baidu/Unlimited-OCR" DEFAULT_REFERENCE = "sahilchachra/unlimited-ocr-mxfp8-mlx" DEFAULT_OUTPUT = PROJECT_ROOT / "models" / "AX-Unlimited-OCR-3B-MoE-MLX-MXFP8" DEFAULT_ARTIFACTS = PROJECT_ROOT / "artifacts" / "ocr-aware-v1" DEFAULT_REPO = "AutomatosX/AX-Unlimited-OCR-3B-MoE-MLX-MXFP8" BASE_PRECISION_MAP = QUANT_DIR / "precision_map.json" def run_command(cmd: list[str], description: str, *, dry_run: bool = False) -> bool: """Run one subprocess and stop the pipeline on any non-zero result.""" print("\n" + "=" * 72) print(f"STEP: {description}") print("CMD: " + " ".join(str(part) for part in cmd)) print("=" * 72) if dry_run: print("[DRY RUN] Command not executed") return True started = time.perf_counter() result = subprocess.run(cmd, cwd=str(PROJECT_ROOT)) elapsed = time.perf_counter() - started print(f"[{'OK' if result.returncode == 0 else 'FAIL'}] {description} ({elapsed:.1f}s)") return result.returncode == 0 def _config_path(model_path: str) -> Path: local = Path(model_path) if local.is_dir(): return local / "config.json" from huggingface_hub import hf_hub_download return Path(hf_hub_download(model_path, filename="config.json")) def assert_unquantized_source(model_path: str) -> dict: """Reject an already quantized source before expensive work.""" config_path = _config_path(model_path) config = load_json_object(config_path) if config.get("quantization") or config.get("quantization_config"): raise ValueError("Source is already quantized; use the upstream BF16 checkpoint") architectures = config.get("architectures") if not isinstance(architectures, list) or "UnlimitedOCRForCausalLM" not in architectures: raise ValueError("Source is not an UnlimitedOCRForCausalLM checkpoint") return {"path": config_path.name, "sha256": sha256_file(config_path)} def verify_huggingface_revision(repo_id: str, revision: str) -> str: """Resolve an explicitly pinned Hub revision and require an exact commit.""" from huggingface_hub import HfApi resolved = HfApi().model_info(repo_id, revision=revision).sha if resolved != revision: raise ValueError( f"Revision for {repo_id} resolved to {resolved!r}, expected {revision!r}" ) return resolved def system_provenance() -> dict: versions = {} for distribution in ("mlx", "mlx-vlm", "huggingface-hub", "numpy", "Pillow"): try: versions[distribution] = importlib.metadata.version(distribution) except importlib.metadata.PackageNotFoundError: versions[distribution] = None return { "created_at": datetime.now(timezone.utc).isoformat(), "python": platform.python_version(), "platform": platform.platform(), "machine": platform.machine(), "processor": platform.processor(), "versions": versions, } def preflight(args: argparse.Namespace) -> bool: """Validate source, dataset, and release inputs before Metal allocation.""" try: source_config = assert_unquantized_source(args.model_path) verify_huggingface_revision(args.source_id, args.source_revision) verify_huggingface_revision( args.reference_model, args.reference_revision, ) calibration_dataset = dataset_manifest(args.calibration_dir) evaluation_dataset = dataset_manifest(args.eval_dir) separation = dataset_separation(calibration_dataset, evaluation_dataset) if not separation["passed"]: raise ValueError( "Calibration and final evaluation datasets must be disjoint; " f"details={separation}" ) if not args.image.is_file(): raise FileNotFoundError(f"Performance/R-SWA image not found: {args.image}") if args.output_dir.exists() and args.step in {"all", "convert"}: raise FileExistsError(f"Output model directory already exists: {args.output_dir}") args.artifacts_dir.mkdir(parents=True, exist_ok=True) provenance = { **system_provenance(), "source_model": args.source_id, "source_revision": args.source_revision, "source_local_name": Path(args.model_path).name, "source_config": source_config, "reference_model": args.reference_model, "reference_revision": args.reference_revision, "target_repo": args.repo_id, "calibration_dataset": calibration_dataset, "evaluation_dataset": evaluation_dataset, "smoke_image": { "name": args.image.name, "sha256": sha256_file(args.image), }, "parameters": { "accuracy_tokens": args.accuracy_tokens, "performance_tokens": args.performance_tokens, "performance_warmup": args.performance_warmup, "performance_runs": args.performance_runs, "rswa_lengths": args.rswa_lengths, }, } (args.artifacts_dir / "provenance.json").write_text( json.dumps(provenance, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) print(f"[OK] Unquantized source: {args.model_path}") print(f"[OK] Calibration samples: {calibration_dataset['num_samples']}") print(f"[OK] Calibration digest: {calibration_dataset['content_sha256']}") print(f"[OK] Held-out evaluation samples: {evaluation_dataset['num_samples']}") print(f"[OK] Evaluation digest: {evaluation_dataset['content_sha256']}") return True except Exception as exc: print(f"[FAIL] Preflight: {exc}") return False def sensitivity(args: argparse.Namespace) -> bool: cmd = [ sys.executable, str(QUANT_DIR / "layer_sensitivity.py"), "--model-path", args.model_path, "--source-id", args.source_id, "--source-revision", args.source_revision, "--eval-dir", str(args.calibration_dir), "--output", str(args.artifacts_dir / "sensitivity_results.json"), "--max-tokens", str(args.accuracy_tokens), ] return run_command(cmd, "Layer sensitivity analysis", dry_run=args.dry_run) def _accuracy_on_dir( model: str, eval_dir: Path, output: Path, args: argparse.Namespace, *, revision: str | None = None, ) -> list[str]: command = [ sys.executable, str(PROJECT_ROOT / "benchmarks" / "run_accuracy.py"), "--model-path", model, "--eval-dir", str(eval_dir), "--output", str(output), "--max-tokens", str(args.accuracy_tokens), "--profile", "accurate", ] if revision is not None: command.extend(["--served-revision", revision]) return command def calibrate(args: argparse.Namespace) -> bool: """Build LM-head candidates from sensitivity, measure them, select precision. Schema-3 release requires ``calibration_results.json`` with hashed raw evidence before the final precision map can be generated. """ if args.dry_run: print("[DRY RUN] Calibrate LM head (bf16 / mxfp8 / affine8 experiments)") return True try: sensitivity_path = args.artifacts_dir / "sensitivity_results.json" # Align group retention with the joint release CER/digit budgets so # cumulative MXFP8 error is less likely to blow the 0.01 absolute caps. interim_map = generate_precision_map( load_json_object(BASE_PRECISION_MAP), load_json_object(sensitivity_path), calibration_results=None, cer_threshold=0.01, digit_cer_threshold=0.01, table_degradation_threshold=0.01, ) interim_path = args.artifacts_dir / "interim_precision_map.json" interim_path.write_text( json.dumps(interim_map, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) print(f"[OK] Interim precision map (pre-calibration): {interim_path}") except Exception as exc: print(f"[FAIL] Interim precision map: {exc}") return False baseline_accuracy = args.artifacts_dir / "calibration_baseline_accuracy.json" reference_performance = ( args.artifacts_dir / "calibration_reference_performance.json" ) jobs: list[tuple[list[str], str]] = [ ( _accuracy_on_dir( args.model_path, args.calibration_dir, baseline_accuracy, args, revision=args.source_revision, ), "Calibration baseline BF16 accuracy", ), ( _performance_command( args.reference_model, reference_performance, args, revision=args.reference_revision, ), "Calibration reference performance", ), ] if not all(run_command(cmd, description) for cmd, description in jobs): return False experiment_specs = ( ("bf16-head", "bfloat16"), ("mxfp8-head", "mxfp8"), ("affine8-head", "affine8"), ) experiment_args: list[str] = [] for label, precision in experiment_specs: exp_map = dict(interim_map) exp_map["language_model.lm_head"] = precision exp_map_path = args.artifacts_dir / f"calibration_{precision}_precision_map.json" exp_map_path.write_text( json.dumps(exp_map, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) exp_model_dir = ( args.output_dir.parent / f"{args.output_dir.name}-cal-{precision}" ) if exp_model_dir.exists(): print(f"[FAIL] Calibration model directory already exists: {exp_model_dir}") return False convert_cmd = [ sys.executable, str(QUANT_DIR / "mixed_precision_convert.py"), "--model-path", args.model_path, "--source-revision", args.source_revision, "--precision-map", str(exp_map_path), "--output-dir", str(exp_model_dir), ] if not run_command(convert_cmd, f"Calibration convert ({label})"): return False # Filenames must match CALIBRATION_EVIDENCE_KEYS in release_gate.py. accuracy_path = args.artifacts_dir / f"calibration_{precision}_accuracy.json" performance_path = ( args.artifacts_dir / f"calibration_{precision}_performance.json" ) if not run_command( _accuracy_on_dir( str(exp_model_dir), args.calibration_dir, accuracy_path, args, ), f"Calibration accuracy ({label})", ): return False if not run_command( _performance_command(str(exp_model_dir), performance_path, args), f"Calibration performance ({label})", ): return False experiment_args.extend( [ "--experiment", label, precision, str(accuracy_path), str(performance_path), ] ) select_cmd = [ sys.executable, str(QUANT_DIR / "calibrate_precision.py"), "--bf16-accuracy", str(baseline_accuracy), "--reference-performance", str(reference_performance), "--calibration-dir", str(args.calibration_dir), "--source-revision", args.source_revision, "--reference-revision", args.reference_revision, *experiment_args, "--output", str(args.artifacts_dir / "calibration_results.json"), ] return run_command(select_cmd, "Select calibrated LM-head precision") def precision_map(args: argparse.Namespace) -> bool: sensitivity_path = args.artifacts_dir / "sensitivity_results.json" output_path = args.artifacts_dir / "generated_precision_map.json" if args.dry_run: print(f"[DRY RUN] Generate {output_path} from {sensitivity_path}") return True try: calibration_path = args.artifacts_dir / "calibration_results.json" calibration = load_json_object(calibration_path) evidence_paths = { "calibration_baseline_accuracy": ( args.artifacts_dir / "calibration_baseline_accuracy.json" ), "calibration_reference_performance": ( args.artifacts_dir / "calibration_reference_performance.json" ), **calibration_raw_evidence_paths(args.artifacts_dir, calibration), } evidence = { name: load_json_object(path) for name, path in evidence_paths.items() } validate_calibration_results( calibration, bf16_accuracy=evidence["calibration_baseline_accuracy"], reference_performance=evidence[ "calibration_reference_performance" ], calibration_dataset=dataset_manifest(args.calibration_dir), source_revision=args.source_revision, reference_revision=args.reference_revision, evidence=evidence, evidence_paths=evidence_paths, ) generated = generate_precision_map( load_json_object(BASE_PRECISION_MAP), load_json_object(sensitivity_path), calibration_results=calibration, # Must match the interim calibration map thresholds. cer_threshold=0.01, digit_cer_threshold=0.01, table_degradation_threshold=0.01, ) output_path.write_text( json.dumps(generated, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) print(f"[OK] Generated executable precision map: {output_path}") return True except Exception as exc: print(f"[FAIL] Precision-map generation: {exc}") return False def convert(args: argparse.Namespace) -> bool: cmd = [ sys.executable, str(QUANT_DIR / "mixed_precision_convert.py"), "--model-path", args.model_path, "--source-revision", args.source_revision, "--precision-map", str(args.artifacts_dir / "generated_precision_map.json"), "--output-dir", str(args.output_dir), ] return run_command(cmd, "BF16 to OCR-aware MXFP8 conversion", dry_run=args.dry_run) def _accuracy_command( model: str, output: Path, args: argparse.Namespace, *, revision: str | None = None, ) -> list[str]: command = [ sys.executable, str(PROJECT_ROOT / "benchmarks" / "run_accuracy.py"), "--model-path", model, "--eval-dir", str(args.eval_dir), "--output", str(output), "--max-tokens", str(args.accuracy_tokens), "--profile", "accurate", ] if revision is not None: command.extend(["--served-revision", revision]) return command def _performance_command( model: str, output: Path, args: argparse.Namespace, *, revision: str | None = None, ) -> list[str]: command = [ sys.executable, str(PROJECT_ROOT / "benchmarks" / "run_performance.py"), "--model-path", model, "--image", str(args.image), "--output", str(output), "--max-tokens", str(args.performance_tokens), "--warmup", str(args.performance_warmup), "--runs", str(args.performance_runs), ] if revision is not None: command.extend(["--served-revision", revision]) return command def validate(args: argparse.Namespace) -> bool: """Benchmark BF16, Sahil reference, candidate, then stress candidate R-SWA.""" jobs = [ (_accuracy_command(args.model_path, args.artifacts_dir / "bf16_accuracy.json", args, revision=args.source_revision), "BF16 accuracy"), (_accuracy_command(args.reference_model, args.artifacts_dir / "reference_accuracy.json", args, revision=args.reference_revision), "Sahil-reference accuracy"), (_accuracy_command(str(args.output_dir), args.artifacts_dir / "candidate_accuracy.json", args), "Candidate accuracy"), (_performance_command(args.reference_model, args.artifacts_dir / "reference_performance.json", args, revision=args.reference_revision), "Sahil-reference performance"), (_performance_command(str(args.output_dir), args.artifacts_dir / "candidate_performance.json", args), "Candidate performance"), ([ sys.executable, str(PROJECT_ROOT / "benchmarks" / "rswa_validation.py"), "--model-path", str(args.output_dir), "--image", str(args.image), "--output", str(args.artifacts_dir / "candidate_rswa.json"), "--lengths", *[str(length) for length in args.rswa_lengths], "--force-min-tokens", str(max(args.rswa_lengths)), ], "Candidate R-SWA stress validation"), ] return all(run_command(cmd, description, dry_run=args.dry_run) for cmd, description in jobs) def resolve_reference_dir(reference_model: str, revision: str) -> Path: local = Path(reference_model) if local.is_dir(): return local from huggingface_hub import snapshot_download return Path(snapshot_download( reference_model, revision=revision, allow_patterns=[ "*.safetensors", "model.safetensors.index.json", "config.json", "processor_config.json", "tokenizer*.json", "special_tokens_map.json", "chat_template.jinja", ], )) def gate(args: argparse.Namespace) -> bool: if args.dry_run: reference_dir = Path("") else: try: reference_dir = resolve_reference_dir( args.reference_model, args.reference_revision, ) except Exception as exc: print(f"[FAIL] Reference download: {exc}") return False cmd = [ sys.executable, str(QUANT_DIR / "release_gate.py"), "--candidate-dir", str(args.output_dir), "--reference-dir", str(reference_dir), "--source-dir", args.model_path, "--calibration-dir", str(args.calibration_dir), "--eval-dir", str(args.eval_dir), "--artifacts-dir", str(args.artifacts_dir), "--output", str(args.artifacts_dir / "release_manifest.json"), "--repo-id", args.repo_id, "--source-id", args.source_id, "--source-revision", args.source_revision, "--reference-id", args.reference_model, "--reference-revision", args.reference_revision, ] return run_command(cmd, "Fail-closed release gate", dry_run=args.dry_run) def publish(args: argparse.Namespace) -> bool: cmd = [ sys.executable, str(PROJECT_ROOT / "scripts" / "publish_optimized_model.py"), "--model-dir", str(args.output_dir), "--manifest", str(args.artifacts_dir / "release_manifest.json"), "--artifacts-dir", str(args.artifacts_dir), "--repo-id", args.repo_id, ] if args.dry_run: cmd.append("--dry-run") return run_command(cmd, "Publish approved candidate to Hugging Face") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Verified OCR-aware model release pipeline") parser.add_argument("--model-path", default=str(DEFAULT_SOURCE)) parser.add_argument("--source-id", default=DEFAULT_SOURCE_ID) parser.add_argument("--source-revision", required=True) parser.add_argument("--reference-model", default=DEFAULT_REFERENCE) parser.add_argument("--reference-revision", required=True) parser.add_argument( "--calibration-dir", required=True, type=Path, help="Selection-only OCR dataset used for sensitivity and head calibration", ) parser.add_argument( "--eval-dir", required=True, type=Path, help="Disjoint held-out OCR dataset used only for final release validation", ) parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT) parser.add_argument("--artifacts-dir", type=Path, default=DEFAULT_ARTIFACTS) parser.add_argument("--image", type=Path, default=PROJECT_ROOT / "test_data" / "test_invoice.png") parser.add_argument("--repo-id", default=DEFAULT_REPO) parser.add_argument("--accuracy-tokens", type=int, default=2048) parser.add_argument("--performance-tokens", type=int, default=256) parser.add_argument("--performance-warmup", type=int, default=1) parser.add_argument("--performance-runs", type=int, default=3) parser.add_argument("--rswa-lengths", type=int, nargs="+", default=[512, 2048, 8192]) parser.add_argument( "--step", choices=[ "all", "preflight", "sensitivity", "calibrate", "precision-map", "convert", "validate", "gate", "publish", ], default="all", ) parser.add_argument("--dry-run", action="store_true") args = parser.parse_args() for name in ("source_revision", "reference_revision"): if not re.fullmatch(r"[0-9a-f]{40}", getattr(args, name)): parser.error(f"--{name.replace('_', '-')} must be a 40-character lowercase commit SHA") positive = ( "accuracy_tokens", "performance_tokens", "performance_warmup", "performance_runs", ) for name in positive: if getattr(args, name) < 1: parser.error(f"--{name.replace('_', '-')} must be positive") if len(args.rswa_lengths) < 3 or args.rswa_lengths != sorted(set(args.rswa_lengths)) or 8192 not in args.rswa_lengths: parser.error("--rswa-lengths must be sorted, unique, include 8192, and contain at least three values") return args def main() -> None: args = parse_args() steps = { "preflight": lambda: preflight(args), "sensitivity": lambda: sensitivity(args), "calibrate": lambda: calibrate(args), "precision-map": lambda: precision_map(args), "convert": lambda: convert(args), "validate": lambda: validate(args), "gate": lambda: gate(args), "publish": lambda: publish(args), } # Full release order: selection data → calibrated map → convert → held-out gate. selected = list(steps) if args.step == "all" else [args.step] for step_name in selected: if not steps[step_name](): print(f"[ABORT] Pipeline stopped at: {step_name}") raise SystemExit(1) print("[DONE] Requested pipeline steps completed successfully") if __name__ == "__main__": main()