ONNX
onnxruntime
onnx-mlir
quantization
fp32
ONNX_Models / scripts /registry_common.py
purejomo's picture
Finalize public ONNX/ONNX-MLIR validation release
ed3aeeb
Raw
History Blame Contribute Delete
7.93 kB
#!/usr/bin/env python3
"""Canonical model-registry columns and invariant checks."""
from __future__ import annotations
import csv
import hashlib
import json
import re
from pathlib import Path
from typing import Any, Iterable
REPO_ROOT = Path(__file__).resolve().parents[1]
SCHEMA = json.loads((REPO_ROOT / "schemas" / "model_registry.columns.json").read_text(encoding="utf-8"))
COLUMNS: list[str] = SCHEMA["columns"]
ELIGIBILITY = set(SCHEMA["eligibility_values"])
QUANTIZED_STATUSES = set(SCHEMA["public_quantized_status_values"])
HEX64 = re.compile(r"^[0-9a-f]{64}$")
EMPTY_MARKERS = {"", "UNKNOWN", "UNVERIFIED", "NOT_APPLICABLE", "N/A", "NONE", "NULL"}
def normalize_boolean(value: str) -> str:
normalized = value.strip().upper()
if normalized in {"TRUE", "YES", "Y", "1", "AVAILABLE", "PASS", "VERIFIED"}:
return "TRUE"
if normalized in {"FALSE", "NO", "N", "0", "UNAVAILABLE", "NOT_APPLICABLE", "N/A"}:
return "FALSE"
if normalized.startswith(("TRUE_", "YES_", "PASS:", "VERIFIED_", "PUBLIC_QUANTIZED_VERIFIED")):
return "TRUE"
if normalized.startswith(("FALSE_", "NO_", "UNAVAILABLE_", "NOT_APPLICABLE_", "N/A_")):
return "FALSE"
return normalized or "UNKNOWN"
def is_known(value: str) -> bool:
return value.strip().upper() not in EMPTY_MARKERS
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
while chunk := handle.read(1024 * 1024):
digest.update(chunk)
return digest.hexdigest()
def read_registry(path: Path) -> list[dict[str, str]]:
with path.open("r", encoding="utf-8-sig", newline="") as handle:
reader = csv.DictReader(handle)
if reader.fieldnames is None:
raise ValueError(f"CSV has no header: {path}")
return [{key: (value or "").strip() for key, value in row.items()} for row in reader]
def validate_row(row: dict[str, str], row_number: int, verify_files: bool) -> list[str]:
prefix = f"row {row_number} model_id={row.get('model_id', '')!r}"
errors: list[str] = []
missing = [column for column in COLUMNS if column not in row]
if missing:
errors.append(f"{prefix}: missing columns: {', '.join(missing)}")
return errors
blank = [column for column in COLUMNS if not row[column].strip()]
if blank:
errors.append(f"{prefix}: blank values: {', '.join(blank)}")
if row["eligibility"] not in ELIGIBILITY:
errors.append(f"{prefix}: invalid eligibility {row['eligibility']!r}")
if row["public_quantized_status"] not in QUANTIZED_STATUSES:
errors.append(f"{prefix}: invalid public_quantized_status {row['public_quantized_status']!r}")
public_available = normalize_boolean(row["public_quantized_available"])
paired_available = normalize_boolean(row["paired_fp32_available"])
if public_available == "FALSE" and row["eligibility"] != "DISCOVERY_ONLY":
errors.append(f"{prefix}: no public quantized artifact must be DISCOVERY_ONLY")
incomplete_public_artifact = (
row["eligibility"] == "DISCOVERY_ONLY"
and public_available == "TRUE"
and row["public_quantized_status"] == "PUBLIC_QUANTIZED_FOUND_UNVERIFIED"
)
if row["eligibility"] == "DISCOVERY_ONLY" and public_available == "TRUE" and not incomplete_public_artifact:
errors.append(
f"{prefix}: DISCOVERY_ONLY can retain an available artifact only as "
"PUBLIC_QUANTIZED_FOUND_UNVERIFIED"
)
if incomplete_public_artifact:
if paired_available != "TRUE" or row["pair_compatibility"].upper() != "VERIFIED":
errors.append(f"{prefix}: incomplete public artifact retention requires a verified FP32 pair")
if not HEX64.fullmatch(row["public_quantized_checksum"].lower()):
errors.append(f"{prefix}: invalid retained public quantized SHA-256")
if not HEX64.fullmatch(row["paired_fp32_checksum"].lower()):
errors.append(f"{prefix}: invalid retained paired FP32 SHA-256")
if not row["runtime_validation"].upper().startswith("BLOCKED_ARTIFACT_INCOMPLETE"):
errors.append(f"{prefix}: incomplete public artifact must preserve its runtime blocker")
if row["eligibility"] == "ELIGIBLE":
required_known = (
"artifact_id", "source_repository", "source_license",
"public_quantized_source", "public_quantized_artifact",
"public_quantized_format", "public_quantized_version",
"public_quantized_checksum", "public_quantized_license",
"runtime_validation", "runtime", "paired_fp32_source",
"paired_fp32_artifact", "paired_fp32_format", "paired_fp32_version",
"paired_fp32_checksum", "paired_fp32_license", "dataset", "input_shape",
"output_shape", "preprocessing", "label_space", "evidence_url",
"evidence_path", "download_command", "http_status", "status_reason",
)
unknown = [field for field in required_known if not is_known(row[field])]
if unknown:
errors.append(f"{prefix}: ELIGIBLE has unknown critical fields: {', '.join(unknown)}")
if public_available != "TRUE":
errors.append(f"{prefix}: ELIGIBLE requires public_quantized_available=TRUE")
if paired_available != "TRUE":
errors.append(f"{prefix}: ELIGIBLE requires paired_fp32_available=TRUE")
if row["pair_compatibility"].upper() != "VERIFIED":
errors.append(f"{prefix}: ELIGIBLE requires pair_compatibility=VERIFIED")
if row["public_quantized_status"] != "PUBLIC_QUANTIZED_VERIFIED":
errors.append(f"{prefix}: ELIGIBLE requires PUBLIC_QUANTIZED_VERIFIED")
if not HEX64.fullmatch(row["public_quantized_checksum"].lower()):
errors.append(f"{prefix}: invalid public quantized SHA-256")
if not HEX64.fullmatch(row["paired_fp32_checksum"].lower()):
errors.append(f"{prefix}: invalid paired FP32 SHA-256")
runtime = row["runtime_validation"].upper()
if not any(token in runtime for token in ("PASS", "VERIFIED", "INVOKE_SUCCESS", "LOAD_SUCCESS")):
errors.append(f"{prefix}: ELIGIBLE runtime_validation does not assert a successful load/invoke")
if verify_files and (row["eligibility"] == "ELIGIBLE" or incomplete_public_artifact):
for path_field, checksum_field in (
("artifact_local_path", "public_quantized_checksum"),
("paired_fp32_local_path", "paired_fp32_checksum"),
):
path = Path(row[path_field])
if not path.is_absolute():
path = REPO_ROOT / path
if not path.is_file():
errors.append(f"{prefix}: missing local artifact {path_field}={path}")
elif sha256_file(path) != row[checksum_field].lower():
errors.append(f"{prefix}: local artifact checksum mismatch for {path_field}")
return errors
def validate_registry(rows: list[dict[str, str]], verify_files: bool, require_complete: bool) -> list[str]:
errors: list[str] = []
identifiers: set[str] = set()
for index, row in enumerate(rows, start=2):
errors.extend(validate_row(row, index, verify_files))
identifier = row.get("model_id", "")
if identifier in identifiers:
errors.append(f"row {index}: duplicate model_id {identifier!r}")
identifiers.add(identifier)
if require_complete:
tasks = {row.get("task", "") for row in rows}
if len(rows) != 21:
errors.append(f"registry has {len(rows)} selected models; exactly 21 required")
if any(row.get("eligibility") != "ELIGIBLE" for row in rows):
errors.append("selected registry may contain only ELIGIBLE models")
if len(tasks) < 5:
errors.append(f"registry has {len(tasks)} tasks; at least 5 required")
return errors