File size: 6,766 Bytes
ed3aeeb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | #!/usr/bin/env python3
"""Validate the 42 published ONNX Dialect execution-dependency graphs."""
from __future__ import annotations
import argparse
import csv
import json
import sys
import xml.etree.ElementTree as ET
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from scripts.mlir_graph_common import (
VARIANTS,
atomic_json,
resolve_coverage_path,
sha256,
)
def utc_now() -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds").replace("+00:00", "Z")
def normalize_variant(value: str) -> str:
return "public_quantized" if value in {"quantized", "public_quantized"} else value
def read_rows(path: Path) -> list[dict[str, str]]:
with path.open(newline="", encoding="utf-8") as stream:
return list(csv.DictReader(stream))
def validate(args: argparse.Namespace) -> tuple[dict[str, Any], int]:
root = args.repo_root.resolve()
report_dir = (root / args.report_dir).resolve() if not args.report_dir.is_absolute() else args.report_dir.resolve()
matrix = (root / args.coverage_matrix).resolve() if not args.coverage_matrix.is_absolute() else args.coverage_matrix.resolve()
output = (root / args.output).resolve() if not args.output.is_absolute() else args.output.resolve()
for path in (report_dir, matrix, output.parent):
try:
path.relative_to(root)
except ValueError as error:
raise SystemExit(f"path outside repository root: {path}") from error
checks: list[dict[str, Any]] = []
def check(name: str, passed: bool, detail: Any = None) -> None:
checks.append({"name": name, "status": "PASS" if passed else "FAIL", "detail": detail})
rows = read_rows(matrix)
keyed: dict[tuple[str, str], dict[str, str]] = {}
for row in rows:
variant = normalize_variant(row.get("variant", ""))
key = (row.get("model_id", ""), variant)
check("unique_matrix_row", key not in keyed, {"key": key})
keyed[key] = row
model_ids = sorted({model_id for model_id, _ in keyed})
check("matrix_variant_count", len(keyed) == 42, len(keyed))
check("matrix_model_count", len(model_ids) == 21, len(model_ids))
check(
"matrix_pair_completeness",
all((model_id, variant) in keyed for model_id in model_ids for variant in VARIANTS),
)
expected_svg: set[Path] = set()
expected_png: set[Path] = set()
for model_id in model_ids:
for variant in VARIANTS:
row = keyed.get((model_id, variant))
if row is None:
continue
context = f"{model_id}:{variant}"
check("onnx_dialect_status", row.get("onnx_status") == "PASS", {"graph": context, "status": row.get("onnx_status")})
try:
source = resolve_coverage_path(row.get("onnx_artifact", ""), root)
source_ok = source.is_file() and sha256(source) == row.get("onnx_sha256")
except (OSError, ValueError) as error:
source_ok = False
source = Path(row.get("onnx_artifact", ""))
check("onnx_source_resolution", False, {"graph": context, "error": str(error)})
else:
check("onnx_source_checksum", source_ok, {"graph": context, "source": str(source)})
graph_dir = report_dir / "graphs" / model_id / variant / "onnx"
svg = graph_dir / "execution_dependency_graph.svg"
png = graph_dir / "execution_dependency_graph.png"
expected_svg.add(svg)
expected_png.add(png)
try:
svg_text = svg.read_text(encoding="utf-8")
ET.fromstring(svg_text)
svg_ok = "STATIC_MLIR_PROGRAM_ORDER" in svg_text and f"{model_id}:{variant}:ONNX" in svg_text
except (OSError, ET.ParseError) as error:
svg_ok = False
check("svg_read", False, {"graph": context, "error": str(error)})
else:
check("svg_semantics", svg_ok, context)
try:
with png.open("rb") as stream:
png_ok = stream.read(8) == b"\x89PNG\r\n\x1a\n" and png.stat().st_size > 8
except OSError as error:
png_ok = False
check("png_read", False, {"graph": context, "error": str(error)})
else:
check("png_signature", png_ok, context)
actual_svg = set(report_dir.glob("graphs/*/*/*/execution_dependency_graph.svg"))
actual_png = set(report_dir.glob("graphs/*/*/*/execution_dependency_graph.png"))
check("exact_svg_set", actual_svg == expected_svg, {"expected": 42, "actual": len(actual_svg)})
check("exact_png_set", actual_png == expected_png, {"expected": 42, "actual": len(actual_png)})
check("no_affine_graphs", not any("affine_scf_memref" in path.parts for path in report_dir.rglob("*")))
check("no_internal_graph_records", not any(report_dir.rglob("graph_record*.json")))
failed = [row for row in checks if row["status"] != "PASS"]
result = {
"schema_version": "1.0",
"stage": "ONNX_DIALECT_PRIMARY_GRAPH_VALIDATION",
"status": "PASS" if not failed else "FAIL",
"generated_at": utc_now(),
"counts": {
"models": len(model_ids),
"variants": len(keyed),
"svg": len(actual_svg),
"png": len(actual_png),
"checks": len(checks),
"failed": len(failed),
},
"checks": checks,
"policy": {
"primary_graph_stage": "ONNX_DIALECT",
"order_semantics": "STATIC_MLIR_PROGRAM_ORDER",
"supplemental_lower_graphs_included": False,
"model_runtime_run": False,
"mlir_toolchain_run": False,
},
}
output.parent.mkdir(parents=True, exist_ok=True)
atomic_json(output, result)
return result, 0 if result["status"] == "PASS" else 1
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--repo-root", type=Path, default=REPO_ROOT)
parser.add_argument("--report-dir", type=Path, default=Path("reports/graphs/mlir"))
parser.add_argument("--coverage-matrix", type=Path, default=Path("reports/conversion/ir_stage_coverage.csv"))
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
result, exit_code = validate(args)
print(json.dumps({"status": result["status"], "counts": result["counts"]}, ensure_ascii=False, sort_keys=True))
return exit_code
if __name__ == "__main__":
raise SystemExit(main())
|