#!/usr/bin/env python3 """Independently validate all non-AD01 T60 batch records and artifacts.""" from __future__ import annotations import argparse import json import shlex import sys import re from collections import Counter 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 import run_mlir_batch as batch def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument( "--report", type=Path, default=batch.BATCH_LOG_DIR / "validation_report.json", ) parser.add_argument( "--artifact-manifest", type=Path, default=batch.BATCH_LOG_DIR / "artifact_manifest.json", ) return parser.parse_args() def add_error(errors: list[str], message: str) -> None: errors.append(message) def main() -> int: args = parse_args() rows = batch.eligible_rows() expected_ids = {row["model_id"] for row in rows} errors: list[str] = [] config_count = 0 result_count = 0 variant_count = 0 stage_count = 0 command_count = 0 log_count = 0 artifact_records: dict[str, dict[str, Any]] = {} required = Counter() variants = Counter() failures = Counter() optional_tosa = Counter() optional_stablehlo = Counter() import_compiler_attempts = 0 prerequisite_attempts = 0 forbidden_operations = 0 command_policy_passes = 0 quant_source_representation = Counter() quant_onnx_preservation = Counter() quant_low_level = Counter() required_failure_signatures = Counter() blockers: list[dict[str, Any]] = [] for model_id in sorted(expected_ids): config_path = batch.CONFIG_DIR / f"{model_id}_mlir.json" if not config_path.is_file(): add_error(errors, f"{model_id}: config missing: {config_path}") continue config = json.loads(config_path.read_text()) try: batch.validate(batch.CONFIG_SCHEMA, config) except Exception as error: add_error(errors, f"{model_id}: config schema: {error}") continue config_count += 1 result_path = batch.resolve_config_path(config["model_dir"]) / "mlir" / "mlir_batch_run_result.json" if not result_path.is_file(): add_error(errors, f"{model_id}: result missing: {result_path}") continue result = json.loads(result_path.read_text()) try: batch.validate(batch.RESULT_SCHEMA, result) except Exception as error: add_error(errors, f"{model_id}: result schema: {error}") continue result_count += 1 if result.get("config_sha256") != batch.sha256(config_path): add_error(errors, f"{model_id}: config checksum mismatch") if result.get("model_id") != model_id: add_error(errors, f"{model_id}: result model_id mismatch") forbidden_operations += len(result.get("forbidden_operations_performed", [])) flattened_expected = sum(len(item.get("stages", [])) for item in result["variants"].values()) if len(result.get("stages", [])) != flattened_expected: add_error(errors, f"{model_id}: flattened stage record count mismatch") for variant_name, variant in result["variants"].items(): variant_count += 1 variants[variant["status"]] += 1 required[variant["required_path_status"]] += 1 if variant.get("failure_code"): failures[variant["failure_code"]] += 1 if variant_name == "public_quantized": quant = variant.get("quantization", {}) quant_source_representation[quant.get("source", {}).get("representation", "MISSING")] += 1 quant_onnx_preservation[quant.get("onnx_dialect_preservation", "BLOCKED")] += 1 quant_low_level[quant.get("low_level_status", "BLOCKED")] += 1 input_record = variant["input"] input_path = Path(input_record["path"]) if input_record["exists"]: if not input_path.is_file() or batch.sha256(input_path) != input_record["sha256"]: add_error(errors, f"{model_id}/{variant_name}: input checksum mismatch") if variant.get("input_integrity_unchanged") is False: add_error(errors, f"{model_id}/{variant_name}: input mutation detected") stages = variant["stages"] import_stage = next((item for item in stages if item["stage"] == "onnx_to_onnx_dialect"), None) if import_stage is None: add_error(errors, f"{model_id}/{variant_name}: no T60 prerequisite/import attempt") else: prerequisite_attempts += 1 if import_stage.get("command_argv", [None])[0] == str(batch.ONNX_MLIR): import_compiler_attempts += 1 required_nonpass = next( ( item for item in stages if not item.get("optional") and item.get("status") != "PASS" ), None, ) if required_nonpass: stderr_path = Path(required_nonpass.get("stderr_log", "")) stderr_text = stderr_path.read_text(errors="replace") if stderr_path.is_file() else "" if required_nonpass.get("exit_code") == 134 and "ONNXDequantizeLinear" in stderr_text: signature = "ONNXDequantizeLinear per-axis operand assertion (signal 6)" elif required_nonpass.get("exit_code") == 134 and "expected only ranked shapes" in stderr_text: signature = "KrnlTypeConverter unranked tensor assertion (signal 6)" elif required_nonpass.get("exit_code") == 134 and "cast() argument of incompatible type" in stderr_text: signature = "Krnl lowering RankedTensorType cast assertion (signal 6)" else: match = re.search(r"failed to legalize operation '([^']+)'", stderr_text) if match: signature = f"unlegalized {match.group(1)}" elif required_nonpass.get("secondary_failure_code") == "MISSING_ONNX_PREREQUISITE": signature = "missing T40 public-quantized ONNX prerequisite" else: signature = str(required_nonpass.get("failure_code") or "UNKNOWN") required_failure_signatures[signature] += 1 if required_nonpass.get("status") == "BLOCKED": blockers.append({ "model_id": model_id, "variant": variant_name, "stage": required_nonpass["stage"], "failure_code": required_nonpass.get("failure_code"), "secondary_failure_code": required_nonpass.get("secondary_failure_code"), "stderr_log": required_nonpass.get("stderr_log"), "resolution_condition": ( "A T40 exporter must produce the public quantized ONNX from the " "unchanged public checkpoint, then checker/runtime/quantization-preservation " "validation and checksum recording must pass. No self-quantization is allowed." ), }) for stage in stages: stage_count += 1 command_count += bool(stage.get("command")) command_argv = stage.get("command_argv", []) executable = Path(command_argv[0]).name if command_argv else "" if executable not in {"onnx-mlir", "onnx-mlir-opt", "mlir-opt", "test"}: add_error(errors, f"{model_id}/{variant_name}/{stage['stage']}: unexpected executable {executable}") else: command_policy_passes += 1 command_text = " ".join(command_argv).lower() forbidden_tokens = ( "representative_dataset", "calibration", "quantize_static", "quantize_dynamic", "fine_tune", "finetune", "optimizer.step", "backward()", ) if any(token in command_text for token in forbidden_tokens): add_error(errors, f"{model_id}/{variant_name}/{stage['stage']}: forbidden command token") if stage.get("failure_code"): failures[stage["failure_code"]] += 1 for key in ("stdout_log", "stderr_log", "resource_log", "command_log"): path_value = stage.get(key) if not path_value or not Path(path_value).is_file(): add_error(errors, f"{model_id}/{variant_name}/{stage['stage']}: missing {key}") else: log_count += 1 for record in stage.get("inputs", []) + stage.get("outputs", []): if not record.get("exists"): continue path = Path(record["path"]) if not path.is_file(): add_error(errors, f"{model_id}/{variant_name}/{stage['stage']}: artifact missing {path}") elif batch.sha256(path) != record.get("sha256") or path.stat().st_size != record.get("bytes"): add_error(errors, f"{model_id}/{variant_name}/{stage['stage']}: artifact checksum/size mismatch {path}") for name, record in variant.get("artifacts", {}).items(): if not record.get("exists"): continue path = Path(record["path"]) actual = batch.file_record(path) if actual != record: add_error(errors, f"{model_id}/{variant_name}/{name}: canonical artifact record mismatch") continue prior = artifact_records.get(str(path)) if prior and prior != record: add_error(errors, f"{model_id}/{variant_name}/{name}: conflicting artifact record") artifact_records[str(path)] = record routes = variant.get("optional_routes", {}) if "tosa" in routes: optional_tosa[routes["tosa"]["status"]] += 1 if "stablehlo" in routes: optional_stablehlo[routes["stablehlo"]["status"]] += 1 if variant["required_path_status"] == "PASS": suffix = "fp32" if variant_name == "fp32" else "quantized" ir_root = batch.resolve_config_path(config["model_dir"]) / "mlir" / suffix for filename in ("onnx.mlir", "krnl.mlir", "affine_scf_memref.mlir", "llvm.mlir"): if not (ir_root / filename).is_file() or (ir_root / filename).stat().st_size == 0: add_error(errors, f"{model_id}/{variant_name}: PASS required route lacks {filename}") mlir_root = batch.resolve_config_path(config["model_dir"]) / "mlir" supporting = { "central_config": config_path, "model_config": ( batch.resolve_config_path(config["model_dir"]) / "config" / "mlir_batch_config.json" ), "result": result_path, "stage_matrix": mlir_root / "mlir_stage_matrix.json", "quantization_preservation": mlir_root / "mlir_quantization_preservation.json", "artifact_manifest": mlir_root / "artifact_manifest.json", "conversion_log": mlir_root / "mlir_conversion.log", } for name, path in supporting.items(): if not path.is_file() or path.stat().st_size == 0: add_error(errors, f"{model_id}: supporting artifact missing/empty {name}: {path}") continue record = batch.file_record(path) artifact_records[str(path)] = record if path.suffix == ".json": try: json.loads(path.read_text()) except (OSError, json.JSONDecodeError) as error: add_error(errors, f"{model_id}: invalid JSON {name}: {error}") local_config = supporting["model_config"] if local_config.is_file() and json.loads(local_config.read_text()) != config: add_error(errors, f"{model_id}: model-local config differs from central config") model_manifest = supporting["artifact_manifest"] if model_manifest.is_file(): document = json.loads(model_manifest.read_text()) if len(document.get("artifacts", [])) != 8: add_error(errors, f"{model_id}: required artifact manifest must contain 8 records") for record in document.get("artifacts", []): if not record.get("exists"): continue actual = batch.file_record(Path(record["path"])) expected = {key: record.get(key) for key in ("path", "exists", "sha256", "bytes")} if actual != expected: add_error(errors, f"{model_id}: per-model manifest mismatch {record['path']}") if expected_ids != {path.stem.split("_", 1)[0] for path in batch.CONFIG_DIR.glob("*_mlir.json")}: add_error(errors, "config model ID set does not exactly match 21 eligible models excluding AD01") if "AD01" in expected_ids: add_error(errors, "AD01 was not excluded") if forbidden_operations: add_error(errors, f"forbidden operation records present: {forbidden_operations}") manifest = { "schema_version": "1.0", "generated_at": batch.utc_now(), "artifact_count": len(artifact_records), "artifacts": [artifact_records[path] for path in sorted(artifact_records)], } batch.write_json(args.artifact_manifest, manifest) report = { "schema_version": "1.0", "generated_at": batch.utc_now(), "status": "PASS" if not errors else "FAIL", "errors": errors, "expected_model_count": 21, "config_schema_pass_count": config_count, "result_schema_pass_count": result_count, "variant_count": variant_count, "variant_status_counts": dict(sorted(variants.items())), "required_route_status_counts": dict(sorted(required.items())), "failure_code_counts": dict(sorted(failures.items())), "optional_tosa_route_status_counts": dict(sorted(optional_tosa.items())), "optional_stablehlo_route_status_counts": dict(sorted(optional_stablehlo.items())), "quant_source_representation_counts": dict(sorted(quant_source_representation.items())), "quant_onnx_dialect_preservation_counts": dict(sorted(quant_onnx_preservation.items())), "quant_low_level_status_counts": dict(sorted(quant_low_level.items())), "required_failure_signature_counts": dict(sorted(required_failure_signatures.items())), "blockers": blockers, "t60_prerequisite_attempt_count": prerequisite_attempts, "actual_onnx_mlir_import_command_count": import_compiler_attempts, "stage_record_count": stage_count, "exact_command_record_count": command_count, "command_policy_pass_count": command_policy_passes, "log_file_validation_count": log_count, "canonical_artifact_checksum_pass_count": len(artifact_records), "artifact_manifest": str(args.artifact_manifest), "artifact_manifest_sha256": batch.sha256(args.artifact_manifest), "forbidden_operations_performed": forbidden_operations, "ad01_excluded": "AD01" not in expected_ids, "registry_sha256": batch.sha256(batch.REGISTRY), "toolchain_lock_sha256": batch.sha256(batch.TOOLCHAIN_LOCK), "command_argv": [sys.executable, *sys.argv], "command": shlex.join([sys.executable, *sys.argv]), "working_directory": str(batch.REPO_ROOT), } batch.write_json(args.report, report) print(json.dumps(report, indent=2, sort_keys=True)) return 0 if not errors else 1 if __name__ == "__main__": raise SystemExit(main())