#!/usr/bin/env python3 import argparse import csv import os import signal import subprocess import sys from concurrent.futures import ProcessPoolExecutor, as_completed from contextlib import nullcontext, redirect_stderr, redirect_stdout from itertools import groupby from pathlib import Path from tempfile import TemporaryDirectory from colorama import Style, Fore from raptor_validation.validate_one import ( MODE_COMPILE_ONLY, MODE_FULL, MODE_RUN_ONLY, MODE_STAGE_TITLES, PIMSIM_DONE, PIMSIM_FAILED, PIMSIM_NOT_RUN, PIMSIM_SKIPPED, PIMSIM_UNSUPPORTED, GENERATED_DIR_NAMES, ProgressReporter, ValidationResult, clean_workspace_artifacts, load_pimcomp_hardware, pimcomp_compatibility_errors, validate_network, ) from raptor_validation.raptor import PIM_PASS_LABELS DEFAULT_PIMCOMP_CONFIG = "arch-a" PIMCOMP_CONFIG_CHOICES = ("arch-a", "arch-b", "arch-c") THROUGHPUT_PIPELINE = 4 THROUGHPUT_BATCH_SIZE = 4 def discover_onnx_files(root): return sorted( path for path in root.rglob("*.onnx") if not any(part in GENERATED_DIR_NAMES for part in path.relative_to(root).parts[:-1]) ) def run_validation_job(job): models, options = job completed = [] for index, onnx_path, rel, log_path in models: reporter = ProgressReporter( options["model_total"], stages_per_model=len(MODE_STAGE_TITLES[options["mode"]]), enabled=False, verbose=options["verbose"], ) def validate(): try: return validate_network( onnx_path, reporter=reporter, model_index=index, model_total=options["model_total"], verbose=options["verbose"], mode=options["mode"], **options["validation_kwargs"], ) except Exception as exc: print_validation_error(reporter, rel, exc) return ValidationResult( False, latency_passed=False, throughput_passed=False, pimsim_status=PIMSIM_NOT_RUN, throughput_pimsim_status=PIMSIM_NOT_RUN, ) sys.stdout.flush() sys.stderr.flush() if log_path is None: result = validate() else: saved_stdout = os.dup(1) saved_stderr = os.dup(2) try: with open(log_path, "w", encoding="utf-8", buffering=1) as log: os.dup2(log.fileno(), 1) os.dup2(log.fileno(), 2) with redirect_stdout(log), redirect_stderr(log): result = validate() finally: os.dup2(saved_stdout, 1) os.dup2(saved_stderr, 2) os.close(saved_stdout) os.close(saved_stderr) reporter.finish() completed.append((str(rel), result, str(log_path) if log_path else None)) return completed def format_return_status(returncode): if returncode < 0: signal_num = -returncode try: signal_name = signal.Signals(signal_num).name except ValueError: signal_name = "UNKNOWN" return f"Program terminated by signal {signal_name} ({signal_num})." return f"Program exited with code {returncode}." def print_validation_error(reporter, rel, exc): reporter.suspend() print(Style.BRIGHT + Fore.RED + f"Exception while validating {rel}" + Style.RESET_ALL, file=sys.stderr, flush=True) if isinstance(exc, subprocess.CalledProcessError): print(format_return_status(exc.returncode), file=sys.stderr, flush=True) if getattr(exc, "output_already_streamed", False): print("Failure log already printed above.", file=sys.stderr, flush=True) elif exc.output: output_text = exc.output.decode("utf-8", errors="replace") if isinstance(exc.output, bytes) else str(exc.output) if output_text: print(output_text, file=sys.stderr, end="" if output_text.endswith("\n") else "\n", flush=True) else: print(f"{type(exc).__name__}: {exc}", file=sys.stderr, flush=True) print("=" * 72, file=sys.stderr, flush=True) reporter.resume() def print_average_pim_pass_timings(pass_timing_sums, pass_timing_counts, total_timing_sum, timed_benchmark_count): if timed_benchmark_count == 0: return print("\n" + Style.BRIGHT + Fore.CYAN + "Average PIM Pass Timings" + Style.RESET_ALL) for _, label in PIM_PASS_LABELS: count = pass_timing_counts[label] if count == 0: continue print(f" {label.ljust(28)} {pass_timing_sums[label] / count:.4f}s") print(f" {'Total'.ljust(28)} {total_timing_sum / timed_benchmark_count:.4f}s") def format_pimsim_metric(status, value, unit): if status == PIMSIM_DONE: return f"{value:.2f} {unit}" return status def format_memory(byte_count): if byte_count is None: return "-" return f"{byte_count / (1 << 20):.2f} MiB" def print_results_table(title, headers, rows): widths = [max(len(header), *(len(row[index]) for row in rows)) for index, header in enumerate(headers)] separator = "+-" + "-+-".join("-" * width for width in widths) + "-+" def table_line(row): return "| " + " | ".join( value.ljust(widths[index]) if index < 3 else value.rjust(widths[index]) for index, value in enumerate(row)) + " |" print("\n" + Style.BRIGHT + Fore.CYAN + title + Style.RESET_ALL) print(separator) print(table_line(headers)) print(separator) for row in rows: line = table_line(row) status = row[2].ljust(widths[2]) color = Fore.GREEN if row[2] == "PASS" else Fore.RED print(line.replace(status, color + status + Style.RESET_ALL, 1)) print(separator) def mode_common_metrics(result, mode): metrics = result.mode_metrics.get(mode, {}) fallback = { "compile_time_s": result.compile_time_s if mode == "latency" else None, "host_memory_bytes": result.host_memory_bytes if mode == "latency" else None, "cores_memory_bytes": result.cores_memory_bytes if mode == "latency" else None, "used_core_count": result.used_core_count if mode == "latency" else None, "used_crossbar_count": result.used_crossbar_count if mode == "latency" else None, } values = {**fallback, **metrics} return ( f"{values['compile_time_s']:.3f} s" if values["compile_time_s"] is not None else "-", format_memory(values["host_memory_bytes"]), format_memory(values["cores_memory_bytes"]), str(values["used_core_count"]) if values["used_core_count"] is not None else "-", str(values["used_crossbar_count"]) if values["used_crossbar_count"] is not None else "-", ) def operation_label(relative_path): path = Path(relative_path) return str(path.parent) if path.parent != Path(".") else path.stem def main(): script_dir = Path(__file__).parent.resolve() pimcomp_configs_dir = script_dir / "pimsim_configs" / "pimcomp" default_pimsim_config = pimcomp_configs_dir / DEFAULT_PIMCOMP_CONFIG / "latency_config.json" default_hardware = load_pimcomp_hardware(default_pimsim_config) ap = argparse.ArgumentParser(description="Validate all ONNX operations under the operations/ directory.") ap.add_argument("--raptor-path", help="Path to the Raptor compiler binary.") ap.add_argument("--onnx-include-dir", help="Path to OnnxMlirRuntime include directory.") ap.add_argument("--operations-dir", default=None, help="Root of the operations tree (default: operations).") ap.add_argument("--simulator-dir", default=None, help="Path to the functional pim-simulator crate root " "(default: auto-detected relative to script).") ap.add_argument("--non-functional-simulator-build-dir", metavar="PATH", default=None, help="Path to the non-functional simulator build directory " "(default: auto-detected relative to script).") ap.add_argument("--pimcomp-config", choices=PIMCOMP_CONFIG_CHOICES, default=DEFAULT_PIMCOMP_CONFIG, help="Hardware/timing profile for non-functional simulation " "(default: arch-a).") ap.add_argument("--skip-non-functional-simulation", action="store_true", help="Skip non-functional simulation.") ap.add_argument("--no-fast", action="store_true", help="Disable fast pimsim-nn throughput convergence for authoritative experiments.") ap.add_argument("--threshold", type=float, default=1e-3, help="Absolute tolerance for per-element output comparison.") ap.add_argument("--relative-threshold", type=float, default=1e-5, help="Relative tolerance for per-element output comparison.") ap.add_argument("--seed", type=int, default=0, help="RNG seed for generated validation inputs.") ap.add_argument("--crossbar-size", type=int, default=default_hardware["crossbar_rows"], help=f"Crossbar size (default: {default_hardware['crossbar_rows']}).") ap.add_argument("--crossbar-count", type=int, default=default_hardware["crossbar_count"], help=f"Crossbars per core (default: {default_hardware['crossbar_count']}).") ap.add_argument("--core-count", type=int, default=default_hardware["core_count"], help=f"Core count (default: {default_hardware['core_count']}).") ap.add_argument("--raptor-extra-arg", action="append", default=[], help="Additional argument to pass through to the Raptor compiler. Repeat as needed.") ap.add_argument("--command-timeout-seconds", type=float, default=1000000.0, help="Per-subprocess timeout in seconds for compiler, runner, and simulation commands.") ap.add_argument("-j", "--jobs", type=int, default=4, help="Number of model validations to run in parallel (default: 4).") ap.add_argument("--clean", action="store_true", help="Remove generated validation artifacts under each model workspace and exit.") mode_group = ap.add_mutually_exclusive_group() mode_group.add_argument("--compile-only", action="store_true", help="Compile reference and PIM artifacts only; do not run reference execution, " "simulations, or comparison.") mode_group.add_argument("--run-only", action="store_true", help="Reuse existing compiled artifacts and only run inputs, reference execution, " "simulations, and comparison.") ap.add_argument("--verbose", action="store_true", help="Print per-stage progress and subprocess logs for passing validations too.") a = ap.parse_args() if a.jobs < 1: ap.error("--jobs must be at least 1") operations_dir = Path(a.operations_dir).resolve() if a.operations_dir else script_dir / "operations" simulator_dir = Path(a.simulator_dir).resolve() if a.simulator_dir else ( script_dir / ".." / "backend-simulators" / "pim" / "pim-simulator" ) pimsim_nn_build_dir = ( Path(a.non_functional_simulator_build_dir).resolve() if a.non_functional_simulator_build_dir else script_dir / ".." / "backend-simulators" / "pim" / "pimsim-nn" / "build" ) pimsim_config_path = pimcomp_configs_dir / a.pimcomp_config / "latency_config.json" throughput_pimsim_config_path = ( pimcomp_configs_dir / a.pimcomp_config / "throughput_config_1000ms.json" ) if not operations_dir.is_dir(): print(Fore.RED + f"Operations directory not found: {operations_dir}" + Style.RESET_ALL) sys.exit(1) onnx_files = discover_onnx_files(operations_dir) if not onnx_files: print(Fore.YELLOW + f"No .onnx files found under {operations_dir}" + Style.RESET_ALL) sys.exit(1) if a.clean: removed_count = 0 for onnx_path in onnx_files: removed_count += len(clean_workspace_artifacts(onnx_path.parent, onnx_path.stem)) (operations_dir / "validation_results.csv").unlink(missing_ok=True) print(Style.BRIGHT + f"Removed {removed_count} generated artifact path(s)." + Style.RESET_ALL) sys.exit(0) missing_args = [] if not a.raptor_path: missing_args.append("--raptor-path") if not a.onnx_include_dir: missing_args.append("--onnx-include-dir") if missing_args: ap.error("the following arguments are required unless --clean is used: " + ", ".join(missing_args)) print(Style.BRIGHT + f"Found {len(onnx_files)} ONNX file(s) to validate." + Style.RESET_ALL) print(f"Operations root: {operations_dir}") print(f"Max parallel jobs: {a.jobs}") print("=" * 72) mode = MODE_FULL if a.compile_only: mode = MODE_COMPILE_ONLY elif a.run_only: mode = MODE_RUN_ONLY selected_pimsim_config = None selected_throughput_pimsim_config = None if not a.skip_non_functional_simulation: compatibility_errors = pimcomp_compatibility_errors( pimsim_config_path, core_count=a.core_count, crossbar_count=a.crossbar_count, crossbar_size=a.crossbar_size, ) if compatibility_errors: print( Fore.RED + "Non-functional simulation disabled: " + "; ".join(compatibility_errors) + Style.RESET_ALL, file=sys.stderr, ) else: selected_pimsim_config = pimsim_config_path selected_throughput_pimsim_config = throughput_pimsim_config_path results = {} # relative_path -> ValidationResult pass_timing_sums = {label: 0.0 for _, label in PIM_PASS_LABELS} pass_timing_counts = {label: 0 for _, label in PIM_PASS_LABELS} total_timing_sum = 0.0 timed_benchmark_count = 0 reporter = ProgressReporter(len(onnx_files), stages_per_model=1, verbose=a.verbose) raptor_extra_args = list(a.raptor_extra_arg) if not any(str(arg).startswith("--pim-target-config=") for arg in raptor_extra_args): raptor_extra_args.append(f"--pim-target-config={pimsim_config_path}") validation_kwargs = { "raptor_path": a.raptor_path, "onnx_include_dir": a.onnx_include_dir, "simulator_dir": simulator_dir, "crossbar_size": a.crossbar_size, "crossbar_count": a.crossbar_count, "core_count": a.core_count, "raptor_extra_args": raptor_extra_args, "pimsim_nn_build_dir": pimsim_nn_build_dir, "pimsim_config_path": selected_pimsim_config, "throughput_pipeline": THROUGHPUT_PIPELINE, "throughput_batch_size": THROUGHPUT_BATCH_SIZE, "throughput_pimsim_config_path": selected_throughput_pimsim_config, "pimsim_fast": not a.no_fast, "command_timeout_seconds": a.command_timeout_seconds, "threshold": a.threshold, "rtol": a.relative_threshold, "seed": a.seed, } indexed_files = list(enumerate(onnx_files, start=1)) workspace_groups = [ list(group) for _, group in groupby(indexed_files, key=lambda indexed_path: indexed_path[1].parent) ] print_directly = min(a.jobs, len(workspace_groups)) == 1 with (nullcontext(None) if print_directly else TemporaryDirectory(prefix="raptor-validation-")) as log_dir: jobs = [ ( [ ( index, onnx_path, onnx_path.relative_to(operations_dir), Path(log_dir) / f"{index}.log" if log_dir else None, ) for index, onnx_path in workspace_group ], { "model_total": len(onnx_files), "mode": mode, "verbose": a.verbose, "validation_kwargs": validation_kwargs, }, ) for workspace_group in workspace_groups ] with (nullcontext(None) if print_directly else ProcessPoolExecutor(max_workers=a.jobs)) as executor: completed_groups = ( map(run_validation_job, jobs) if print_directly else (future.result() for future in as_completed( executor.submit(run_validation_job, job) for job in jobs )) ) for completed_group in completed_groups: for rel, result, log_path in completed_group: if log_path: reporter.suspend() output = Path(log_path).read_text(encoding="utf-8", errors="replace") if output: print(output, end="" if output.endswith("\n") else "\n") reporter.resume() reporter.advance() reporter.record_result(result.passed) results[rel] = result if result.pim_pass_timings: benchmark_total = 0.0 for label, duration in result.pim_pass_timings.items(): pass_timing_sums[label] += duration pass_timing_counts[label] += 1 benchmark_total += duration total_timing_sum += benchmark_total timed_benchmark_count += 1 reporter.finish() results = dict(sorted(results.items())) # Summary n_passed = sum(1 for result in results.values() if result.passed) n_total = len(results) latency_headers = ("Operation", "Arch", "Result", "Compile", "Host mem", "Cores mem", "Cores", "Xbars", "Latency", "Power", "Energy") throughput_headers = ("Operation", "Arch", "Result", "Compile", "Host mem", "Cores mem", "Cores", "Xbars", "Throughput", "Avg latency", "Avg power", "Avg energy") csv_headers = ( "Operation", "Arch", "Result (l)", "Result (t)", "Compile (l)", "Host mem (l)", "Cores mem (l)", "Cores (l)", "Xbars (l)", "Latency (l)", "Power (l)", "Energy (l)", "Compile (t)", "Host mem (t)", "Cores mem (t)", "Cores (t)", "Xbars (t)", "Avg latency (t)", "Throughput (t)", "Avg power (t)", "Avg energy (t)", ) latency_rows = [] throughput_rows = [] csv_rows = [] for rel, result in results.items(): operation = operation_label(rel) latency_status = "PASS" if result.latency_passed else "FAIL" throughput_status = "PASS" if result.throughput_passed else "FAIL" latency_common = mode_common_metrics(result, "latency") throughput_common = mode_common_metrics(result, "throughput") latency_metrics = ( format_pimsim_metric(result.pimsim_status, result.pimsim_latency_ms, "ms"), format_pimsim_metric(result.pimsim_status, result.pimsim_power_mw, "mW"), format_pimsim_metric(result.pimsim_status, result.pimsim_energy_pj, "pJ"), ) throughput_metrics = ( format_pimsim_metric( result.throughput_pimsim_status, result.pimsim_throughput_samples_s, "samples/s", ), format_pimsim_metric( result.throughput_pimsim_status, result.pimsim_throughput_average_latency_ms, "ms", ), format_pimsim_metric( result.throughput_pimsim_status, result.pimsim_throughput_average_power_mw, "mW", ), format_pimsim_metric( result.throughput_pimsim_status, result.pimsim_throughput_average_energy_pj, "pJ/it", ), ) latency_rows.append((operation, a.pimcomp_config, latency_status, *latency_common, *latency_metrics)) throughput_rows.append((operation, a.pimcomp_config, throughput_status, *throughput_common, *throughput_metrics)) csv_rows.append(( operation, a.pimcomp_config, latency_status, throughput_status, *latency_common, *latency_metrics, *throughput_common, *throughput_metrics, )) print_results_table("Latency", latency_headers, latency_rows) print_results_table("Throughput", throughput_headers, throughput_rows) with (operations_dir / "validation_results.csv").open( "w", encoding="utf-8", newline="" ) as results_file: csv.writer(results_file, lineterminator="\n").writerows((csv_headers, *csv_rows)) print("\n" + Style.BRIGHT + Fore.CYAN + "Summary" + Style.RESET_ALL) print(Style.BRIGHT + f"Passed: {n_passed}" + Style.RESET_ALL) print(Style.BRIGHT + f"Failed: {n_total - n_passed}" + Style.RESET_ALL) measured_latencies = [ result.pimsim_latency_ms for result in results.values() if result.pimsim_status == PIMSIM_DONE ] pimsim_statuses = [ status for result in results.values() for status in (result.pimsim_status, result.throughput_pimsim_status) ] pimsim_failed = pimsim_statuses.count(PIMSIM_FAILED) pimsim_skipped = sum(status in (PIMSIM_SKIPPED, PIMSIM_NOT_RUN) for status in pimsim_statuses) pimsim_unsupported = pimsim_statuses.count(PIMSIM_UNSUPPORTED) print( Style.BRIGHT + f"pimsim-nn: {pimsim_statuses.count(PIMSIM_DONE)} measured, " f"{pimsim_failed} failed, {pimsim_unsupported} unsupported, " f"{pimsim_skipped} skipped" + Style.RESET_ALL ) if measured_latencies: print( Style.BRIGHT + f"Total latency: {sum(measured_latencies):.2f} ms" + Style.RESET_ALL ) if a.verbose: print_average_pim_pass_timings( pass_timing_sums, pass_timing_counts, total_timing_sum, timed_benchmark_count, ) sys.exit(0 if n_passed == n_total else 1) if __name__ == "__main__": main()