import json import os import re import shutil import subprocess import sys import time import numpy as np from dataclasses import dataclass, field from pathlib import Path from colorama import Style, Fore from .gen_network_runner import gen_network_runner from .onnx_utils import gen_random_inputs, save_inputs_to_files, onnx_io, write_inputs_to_memory_bin, _ONNX_TO_NP from .raptor import compile_with_raptor from .pimsim_nn import export_raptor_latency_artifact, parse_pimsim_nn_metrics, read_raptor_instruction_count from .subprocess_utils import run_command_with_reporter STAGE_TITLES = ( "Compile ONNX", "Build Runner", "Generate Inputs", "Run Reference", "Compile PIM", "Run Functional Simulation", "Compare Outputs", "Run Non-functional Simulation", ) STAGE_COLORS = { STAGE_TITLES[0]: Fore.BLUE, STAGE_TITLES[1]: Fore.MAGENTA, STAGE_TITLES[2]: Fore.YELLOW, STAGE_TITLES[3]: Fore.GREEN, STAGE_TITLES[4]: Fore.CYAN, STAGE_TITLES[5]: Fore.MAGENTA, STAGE_TITLES[6]: Fore.YELLOW, STAGE_TITLES[7]: Fore.BLUE, } STAGE_COUNT = len(STAGE_TITLES) GENERATED_DIR_NAMES = ("inputs", "outputs", "pimcomp", "raptor", "runner", "simulation") MODE_FULL = "full" MODE_COMPILE_ONLY = "compile_only" MODE_RUN_ONLY = "run_only" MODE_STAGE_TITLES = { MODE_FULL: STAGE_TITLES, MODE_COMPILE_ONLY: ( "Compile ONNX", "Build Runner", "Compile PIM", ), MODE_RUN_ONLY: ( "Generate Inputs", "Run Reference", "Run Functional Simulation", "Compare Outputs", "Run Non-functional Simulation", ), } PIMSIM_DONE = "DONE" PIMSIM_FAILED = "ERROR" PIMSIM_UNSUPPORTED = "UNSUPPORTED" PIMSIM_SKIPPED = "SKIP" PIMSIM_NOT_RUN = "-" PIMSIM_UNSUPPORTED_VSOFTMAX = "pimsim-nn does not support opcode vsoftmax" class PimSimUnsupportedError(RuntimeError): pass def sanitize_output_name(name): return "".join(ch if ch.isalnum() or ch in "_.-" else "_" for ch in name[:255]) @dataclass class ValidationResult: passed: bool pim_pass_timings: dict[str, float] = field(default_factory=dict) pimsim_latency_ms: float | None = None pimsim_power_mw: float | None = None pimsim_energy_pj: float | None = None pimsim_status: str = PIMSIM_SKIPPED compile_time_s: float | None = None host_memory_bytes: int | None = None cores_memory_bytes: int | None = None used_core_count: int | None = None used_crossbar_count: int | None = None _MEMORY_UNITS = {"B": 1, "KB": 1 << 10, "MB": 1 << 20, "GB": 1 << 30} def collect_pim_resource_metrics(pim_dir): pim_dir = Path(pim_dir) report_path = pim_dir.parent / "reports" / "memory_report.txt" report = report_path.read_text(encoding="utf-8") if report_path.exists() else "" def memory_bytes(label): match = re.search(rf"^\s*{re.escape(label)}:\s+([0-9.]+)\s+(B|KB|MB|GB)$", report, re.MULTILINE) return round(float(match.group(1)) * _MEMORY_UNITS[match.group(2)]) if match else None with open(pim_dir / "config.json", encoding="utf-8") as f: config = json.load(f) used_cores = sum(read_raptor_instruction_count(path) > 0 for path in pim_dir.glob("core_*.pim")) used_crossbars = sum(sum(groups) for groups in config.get("array_group_map", {}).values()) return { "host_memory_bytes": memory_bytes("Host memory"), "cores_memory_bytes": memory_bytes("Local memory after reuse"), "used_core_count": used_cores, "used_crossbar_count": used_crossbars, } class ProgressReporter: def __init__(self, total_models, stages_per_model=STAGE_COUNT, enabled=None, verbose=False): self.total_models = total_models self.stages_per_model = stages_per_model self.total_steps = max(1, total_models * stages_per_model) self.completed_steps = 0 self.passed_models = 0 self.failed_models = 0 self.current_label = "" self.enabled = ( sys.stdout.isatty() and "CODEX_CI" not in os.environ if enabled is None else enabled ) self.verbose = verbose self.columns = max(1, shutil.get_terminal_size((100, 20)).columns) self.suspended = False self.rendered_width = 0 self.rendered_rows = 0 def _clear(self): if self.enabled and self.rendered_rows: columns = max(1, shutil.get_terminal_size((100, 20)).columns) rows = max(self.rendered_rows, (self.rendered_width + columns - 1) // columns) sys.stdout.write("\r\033[2K") for _ in range(rows - 1): sys.stdout.write("\033[1A\r\033[2K") sys.stdout.flush() self.rendered_width = 0 self.rendered_rows = 0 def _render(self): if not self.enabled or self.suspended: return self.columns = max(1, shutil.get_terminal_size((100, 20)).columns) bar_width = min(24, max(4, self.columns - 24)) filled = int(bar_width * self.completed_steps / self.total_steps) counts_text = f"P:{self.passed_models} F:{self.failed_models}" prefix_text = f"[{'#' * filled}{'-' * (bar_width - filled)}] {self.completed_steps}/{self.total_steps}" bar = Fore.GREEN + ("#" * filled) + Fore.CYAN + ("-" * (bar_width - filled)) prefix = Fore.CYAN + f"[{bar}{Fore.CYAN}] {self.completed_steps}/{self.total_steps}" + Style.RESET_ALL counts = ( " " + Style.BRIGHT + Fore.GREEN + f"P:{self.passed_models}" + Style.RESET_ALL + " " + Style.BRIGHT + Fore.RED + f"F:{self.failed_models}" + Style.RESET_ALL ) model_counter = "" label = "" if self.current_label.startswith("[") and "] " in self.current_label: model_counter, label = self.current_label.split("] ", 1) model_counter = f" {model_counter}]" label = f" {label}" elif self.current_label: label = f" {self.current_label}" fixed_width = len(prefix_text) + len(model_counter) + len(counts_text) + 2 if fixed_width > self.columns: model_counter = "" fixed_width = len(prefix_text) + len(counts_text) + 2 if fixed_width > self.columns: prefix_text = f"{self.completed_steps}/{self.total_steps}" prefix = Fore.CYAN + prefix_text + Style.RESET_ALL fixed_width = len(prefix_text) + len(counts_text) + 2 if fixed_width > self.columns: counts = "" counts_text = "" fixed_width = len(prefix_text) + 1 available_label_width = max(0, self.columns - fixed_width) label = label[:available_label_width] plain_counts = f" {counts_text}" if counts_text else "" plain_line = prefix_text + model_counter + plain_counts + label rendered_line = prefix + model_counter + counts + label + Style.RESET_ALL self._clear() sys.stdout.write(rendered_line) sys.stdout.flush() self.rendered_width = len(plain_line) self.rendered_rows = max(1, (self.rendered_width + self.columns - 1) // self.columns) def log(self, message="", color=None): if not self.verbose: self._render() return if self.enabled: self._clear() if color: print(color + message + Style.RESET_ALL, flush=True) else: print(message, flush=True) self._render() def set_stage(self, model_index, model_total, model_name, stage_name): self.current_label = f"[{model_index}/{model_total}] {model_name} ยท {stage_name}" self._render() def advance(self): self.completed_steps = min(self.total_steps, self.completed_steps + 1) self._render() def record_result(self, passed): if passed: self.passed_models += 1 else: self.failed_models += 1 self._render() def suspend(self): if self.enabled: self._clear() self.suspended = True def resume(self): self.suspended = False self._render() def finish(self): if self.enabled: self.suspended = True self._clear() def run_command(cmd, cwd=None, reporter=None, timeout_sec=None, capture_output=False): return run_command_with_reporter( cmd, cwd=cwd, reporter=reporter, timeout_sec=timeout_sec, capture_output=capture_output, ) def load_pimcomp_hardware(config_path): with open(config_path, encoding="utf-8") as f: config = json.load(f) matrix = config["chip_config"]["core_config"]["matrix_config"] rows, cols = config["chip_config"]["network_config"]["layout"] xbar_rows, xbar_cols = matrix["xbar_size"] return { "core_count": config["chip_config"]["core_cnt"], "crossbar_count": matrix["xbar_array_count"], "crossbar_rows": xbar_rows, "crossbar_cols": xbar_cols, "mesh_rows": rows, "mesh_cols": cols, } def pimcomp_compatibility_errors(config_path, *, core_count, crossbar_count, crossbar_size): hardware = load_pimcomp_hardware(config_path) errors = [] if hardware["mesh_rows"] * hardware["mesh_cols"] != hardware["core_count"]: errors.append( f"config layout {hardware['mesh_rows']}x{hardware['mesh_cols']} does not match " f"{hardware['core_count']} cores" ) if core_count != hardware["core_count"]: errors.append(f"--core-count={core_count}, config requires {hardware['core_count']}") if crossbar_count != hardware["crossbar_count"]: errors.append( f"--crossbar-count={crossbar_count}, config requires {hardware['crossbar_count']}" ) if ( hardware["crossbar_rows"] != hardware["crossbar_cols"] or crossbar_size != hardware["crossbar_rows"] ): errors.append( f"--crossbar-size={crossbar_size}, config requires " f"{hardware['crossbar_rows']}x{hardware['crossbar_cols']}" ) return errors def run_pimsim_nn(pimsim_nn_build_dir, pim_dir, config_path, reporter=None, timeout_sec=None): latency_artifact = export_raptor_latency_artifact(pim_dir, Path(pim_dir).parent / "pimsim_nn") try: output = run_command( [pimsim_nn_build_dir / "ChipTest", latency_artifact, config_path, "--gui=false"], cwd=pimsim_nn_build_dir, reporter=reporter, timeout_sec=timeout_sec, capture_output=True, ) except subprocess.CalledProcessError as exc: error_output = exc.output.decode("utf-8", errors="replace") if isinstance(exc.output, bytes) else str(exc.output) if PIMSIM_UNSUPPORTED_VSOFTMAX in error_output: raise PimSimUnsupportedError(PIMSIM_UNSUPPORTED_VSOFTMAX) from exc raise metrics = parse_pimsim_nn_metrics(output) required = ("latency_ms", "average_power_mw", "average_energy_pj") if any(name not in metrics for name in required): raise RuntimeError("pimsim-nn output did not contain latency, average power, and average energy") return tuple(metrics[name] for name in required) def clean_workspace_artifacts(workspace_dir, model_stem): workspace_dir = Path(workspace_dir) removed_paths = [] def remove_path(path): if path.is_symlink() or path.is_file(): path.unlink(missing_ok=True) removed_paths.append(path) elif path.is_dir(): shutil.rmtree(path) removed_paths.append(path) for name in GENERATED_DIR_NAMES: remove_path(workspace_dir / name) for suffix in (".onnx.mlir", ".so", ".tmp"): remove_path(workspace_dir / f"{model_stem}{suffix}") return removed_paths def print_stage(reporter, model_index, model_total, model_name, title): color = STAGE_COLORS.get(title, Fore.WHITE) reporter.log(Style.BRIGHT + color + f"[{title}]" + Style.RESET_ALL) reporter.set_stage(model_index, model_total, model_name, title) def print_info(reporter, message): reporter.log(f" {message}") def compile_onnx_network(network_onnx_path, raptor_path, raptor_dir, runner_dir, reporter=None, timeout_sec=None): stem = network_onnx_path.stem onnx_ir_base = raptor_dir / stem runner_base = runner_dir / stem run_command([raptor_path, network_onnx_path, "-o", onnx_ir_base, "--EmitONNXIR", "--mlir-elide-elementsattrs-if-larger=16"], reporter=reporter, timeout_sec=timeout_sec) run_command([raptor_path, network_onnx_path, "-o", runner_base], reporter=reporter, timeout_sec=timeout_sec) network_so_path = runner_base.with_suffix(".so") network_mlir_path = onnx_ir_base.with_suffix(".onnx.mlir") onnx_ir_base.with_suffix(".tmp").unlink(missing_ok=True) return network_so_path, network_mlir_path def build_onnx_runner(source_dir, build_dir, reporter=None, timeout_sec=None): run_command(["cmake", source_dir], cwd=build_dir, reporter=reporter, timeout_sec=timeout_sec) run_command(["cmake", "--build", ".", "-j"], cwd=build_dir, reporter=reporter, timeout_sec=timeout_sec) return build_dir / "runner" def build_dump_ranges(config_path, outputs_descriptor): with open(config_path) as f: output_addresses = json.load(f)["outputs_addresses"] ranges = [] for addr, (_, _, dtype_code, shape) in zip(output_addresses, outputs_descriptor): byte_size = int(np.prod(shape)) * np.dtype(_ONNX_TO_NP[dtype_code]).itemsize ranges.append(f"{addr},{byte_size}") return ",".join(ranges) def run_pim_simulator(simulator_dir, pim_dir, output_bin_path, dump_ranges, reporter=None, timeout_sec=None): run_command( ["cargo", "run", "--no-default-features", "--release", "--package", "pim-simulator", "--bin", "pim-simulator", "--", "-f", str(pim_dir), "-o", str(output_bin_path), "-d", dump_ranges], cwd=simulator_dir, reporter=reporter, timeout_sec=timeout_sec, ) def parse_pim_simulator_outputs(output_bin_path, outputs_descriptor): raw = output_bin_path.read_bytes() arrays = [] offset = 0 for _, _, dtype_code, shape in outputs_descriptor: dtype = np.dtype(_ONNX_TO_NP[dtype_code]) count = int(np.prod(shape)) array = np.frombuffer(raw, dtype=dtype, count=count, offset=offset).reshape(shape) offset += count * dtype.itemsize arrays.append(array) return arrays def validate_outputs(sim_arrays, runner_out_dir, outputs_descriptor, threshold, rtol, verbose): all_passed = True rows = [] for sim_array, (oi, name, _, shape) in zip(sim_arrays, outputs_descriptor): csv_name = f"output{oi}_{sanitize_output_name(name)}.csv" runner_array = np.loadtxt(runner_out_dir / csv_name, delimiter=',', dtype=np.float32).reshape(shape) sim_array64 = sim_array.astype(np.float64) runner_array64 = runner_array.astype(np.float64) abs_diff = np.abs(sim_array64 - runner_array64) allowed_diff = threshold + rtol * np.abs(runner_array64) max_diff = float(np.max(abs_diff)) passed = bool(np.all(abs_diff <= allowed_diff)) rows.append((name, f"{max_diff:.6e}", passed)) if not passed: all_passed = False name_width = max(len("Output"), *(len(name) for name, _, _ in rows)) diff_width = max(len("Max diff"), *(len(diff) for _, diff, _ in rows)) result_width = len("Result") separator = f" +-{'-' * name_width}-+-{'-' * diff_width}-+-{'-' * result_width}-+" if verbose or not all_passed: print(separator) print(f" | {'Output'.ljust(name_width)} | {'Max diff'.ljust(diff_width)} | {'Result'} |") print(separator) for name, diff_text, passed in rows: status_text = ("PASS" if passed else "FAIL").ljust(result_width) status = Fore.GREEN + status_text + Style.RESET_ALL if passed else Fore.RED + status_text + Style.RESET_ALL print(f" | {name.ljust(name_width)} | {diff_text.ljust(diff_width)} | {status} |") print(separator) return all_passed def validate_network(network_onnx_path, raptor_path, onnx_include_dir, simulator_dir, crossbar_size, crossbar_count, core_count, raptor_extra_args, pimsim_nn_build_dir, pimsim_config_path, threshold, rtol, seed, reporter, model_index, model_total, verbose, command_timeout_seconds, mode): network_onnx_path = Path(network_onnx_path).resolve() raptor_path = Path(raptor_path).resolve() onnx_include_dir = Path(onnx_include_dir).resolve() simulator_dir = Path(simulator_dir).resolve() pimsim_enabled = pimsim_nn_build_dir is not None and pimsim_config_path is not None if pimsim_enabled: pimsim_nn_build_dir = Path(pimsim_nn_build_dir).resolve() pimsim_config_path = Path(pimsim_config_path).resolve() compile_extra_args = list(raptor_extra_args or []) owns_reporter = reporter is None reporter = reporter or ProgressReporter(model_total, stages_per_model=len(MODE_STAGE_TITLES[mode]), verbose=verbose) workspace_dir = network_onnx_path.parent raptor_dir = workspace_dir / "raptor" runner_dir = workspace_dir / "runner" runner_build_dir = runner_dir / "build" if mode != MODE_RUN_ONLY: clean_workspace_artifacts(workspace_dir, network_onnx_path.stem) Path.mkdir(raptor_dir, exist_ok=True) Path.mkdir(runner_build_dir, parents=True, exist_ok=True) reporter.log(Fore.CYAN + f"[{model_index}/{model_total}]" + Style.RESET_ALL + f" {Style.BRIGHT}Validating {network_onnx_path.name}{Style.RESET_ALL}") failed_with_exception = False pim_pass_timings = {} compile_time_s = None resource_metrics = {} try: stem = network_onnx_path.stem network_so_path = runner_dir / f"{stem}.so" network_mlir_path = raptor_dir / f"{stem}.onnx.mlir" runner_path = runner_build_dir / "runner" pim_output_base = raptor_dir / stem def compile_pim(): nonlocal compile_time_s, resource_metrics started = time.perf_counter() timings = compile_with_raptor( network_onnx_path, raptor_path, pim_output_base, crossbar_size, crossbar_count, core_count=core_count, raptor_extra_args=compile_extra_args, cwd=raptor_dir, verbose=verbose, reporter=reporter, timeout_sec=command_timeout_seconds) compile_time_s = time.perf_counter() - started resource_metrics = collect_pim_resource_metrics(raptor_dir / "pim") return timings if mode != MODE_RUN_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile ONNX") network_so_path, network_mlir_path = compile_onnx_network( network_onnx_path, raptor_path, raptor_dir, runner_dir, reporter=reporter, timeout_sec=command_timeout_seconds) print_info(reporter, f"MLIR saved to {network_mlir_path}") print_info(reporter, f"Shared library saved to {network_so_path}") reporter.advance() print_stage(reporter, model_index, model_total, network_onnx_path.name, "Build Runner") gen_network_runner( network_onnx_path, network_so_path, onnx_include_dir, entry="run_main_graph", out=runner_dir / "runner.c", verbose=False, ) runner_path = build_onnx_runner(runner_dir, runner_build_dir, reporter=reporter, timeout_sec=command_timeout_seconds) print_info(reporter, f"Runner built at {runner_path}") reporter.advance() if mode == MODE_COMPILE_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM") pim_pass_timings = compile_pim() print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}") reporter.advance() reporter.record_result(True) reporter.log(Style.BRIGHT + f"Result: {Fore.GREEN}PASS{Style.RESET_ALL}" + Style.RESET_ALL) return ValidationResult( passed=True, pim_pass_timings=pim_pass_timings, compile_time_s=compile_time_s, **resource_metrics) if mode == MODE_RUN_ONLY: required_paths = [ (network_so_path, "compiled reference shared library"), (network_mlir_path, "exported ONNX MLIR"), (runner_path, "built reference runner"), (raptor_dir / "pim" / "config.json", "compiled PIM artifacts"), ] missing = [f"{description} at {path}" for path, description in required_paths if not path.exists()] if missing: raise FileNotFoundError("run-only mode requires existing artifacts:\n " + "\n ".join(missing)) resource_metrics = collect_pim_resource_metrics(raptor_dir / "pim") print_stage(reporter, model_index, model_total, network_onnx_path.name, "Generate Inputs") inputs_descriptor, outputs_descriptor = onnx_io(network_onnx_path) inputs_list, _inputs_dict = gen_random_inputs(inputs_descriptor, seed=seed) flags, _files = save_inputs_to_files(network_onnx_path, inputs_list, out_dir=workspace_dir / "inputs") print_info(reporter, f"Saved {len(inputs_list)} input file(s) to {workspace_dir / 'inputs'}") reporter.advance() print_stage(reporter, model_index, model_total, network_onnx_path.name, "Run Reference") out_dir = workspace_dir / "outputs" Path.mkdir(out_dir, exist_ok=True) run_cmd = [runner_path, *flags] run_cmd += ["--save-csv-dir", f"{out_dir}"] run_command(run_cmd, cwd=runner_build_dir, reporter=reporter, timeout_sec=command_timeout_seconds) print_info(reporter, f"Reference outputs saved to {out_dir}") reporter.advance() if mode != MODE_RUN_ONLY: print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM") pim_pass_timings = compile_pim() print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}") reporter.advance() print_stage( reporter, model_index, model_total, network_onnx_path.name, "Run Functional Simulation") pim_dir = raptor_dir / "pim" write_inputs_to_memory_bin(pim_dir / "memory.bin", pim_dir / "config.json", inputs_list) simulation_dir = workspace_dir / "simulation" Path.mkdir(simulation_dir, exist_ok=True) dump_ranges = build_dump_ranges(pim_dir / "config.json", outputs_descriptor) output_bin_path = simulation_dir / "out.bin" run_pim_simulator(simulator_dir, pim_dir, output_bin_path, dump_ranges, reporter=reporter, timeout_sec=command_timeout_seconds) print_info(reporter, f"Functional simulation output saved to {output_bin_path}") reporter.advance() print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compare Outputs") sim_arrays = parse_pim_simulator_outputs(output_bin_path, outputs_descriptor) reporter.suspend() passed = validate_outputs(sim_arrays, out_dir, outputs_descriptor, threshold, rtol=rtol, verbose=verbose) reporter.resume() reporter.advance() print_stage( reporter, model_index, model_total, network_onnx_path.name, "Run Non-functional Simulation") pimsim_latency_ms = None pimsim_power_mw = None pimsim_energy_pj = None pimsim_status = PIMSIM_SKIPPED if pimsim_enabled: try: pimsim_latency_ms, pimsim_power_mw, pimsim_energy_pj = run_pimsim_nn( pimsim_nn_build_dir, pim_dir, pimsim_config_path, reporter=reporter, timeout_sec=command_timeout_seconds, ) pimsim_status = PIMSIM_DONE print_info( reporter, f"Latency: {pimsim_latency_ms:.6f} ms, " f"Power: {pimsim_power_mw:.6f} mW, " f"Energy: {pimsim_energy_pj:.6f} pJ") except PimSimUnsupportedError as exc: pimsim_status = PIMSIM_UNSUPPORTED print_info(reporter, str(exc)) except Exception as exc: pimsim_status = PIMSIM_FAILED reporter.suspend() print( Fore.RED + f"pimsim-nn non-functional simulation failed: {type(exc).__name__}: {exc}" + Style.RESET_ALL, file=sys.stderr, flush=True, ) reporter.resume() else: print_info(reporter, "pimsim-nn non-functional simulation skipped") reporter.advance() reporter.record_result(passed) status = Fore.GREEN + "PASS" + Style.RESET_ALL if passed else Fore.RED + "FAIL" + Style.RESET_ALL reporter.log(Style.BRIGHT + f"Result: {status}" + Style.RESET_ALL) return ValidationResult( passed=passed, pim_pass_timings=pim_pass_timings, pimsim_latency_ms=pimsim_latency_ms, pimsim_power_mw=pimsim_power_mw, pimsim_energy_pj=pimsim_energy_pj, pimsim_status=pimsim_status, compile_time_s=compile_time_s, **resource_metrics, ) except Exception: failed_with_exception = True reporter.record_result(False) reporter.log(Style.BRIGHT + Fore.RED + "Result: FAIL" + Style.RESET_ALL) reporter.suspend() raise finally: if not failed_with_exception: reporter.log("=" * 72) if owns_reporter: reporter.finish()