This commit is contained in:
@@ -40,7 +40,7 @@ def _format_command(cmd):
|
||||
|
||||
|
||||
def compile_with_raptor(network_path, raptor_onnx_path: Path, output_base: Path,
|
||||
crossbar_size, crossbar_count, core_count=None, pim_merge_scheduler="peft",
|
||||
crossbar_size, crossbar_count, core_count=None,
|
||||
pim_memory_report="none", raptor_extra_args=None, cwd=None, verbose=False,
|
||||
reporter=None, timeout_sec=None):
|
||||
# Define the arguments, with the possibility to set crossbar size and count
|
||||
@@ -52,7 +52,6 @@ def compile_with_raptor(network_path, raptor_onnx_path: Path, output_base: Path,
|
||||
"--EmitPimCodegen",
|
||||
f"--crossbar-size={crossbar_size}",
|
||||
f"--crossbar-count={crossbar_count}",
|
||||
f"--pim-merge-scheduler={pim_merge_scheduler}",
|
||||
]
|
||||
if core_count is not None:
|
||||
args.append(f"--core-count={core_count}")
|
||||
|
||||
@@ -76,8 +76,6 @@ def main():
|
||||
ap.add_argument("--crossbar-count", type=int, default=8)
|
||||
ap.add_argument("--core-count", type=int, default=None,
|
||||
help="Core count to pass to Raptor. Required for PIM validation.")
|
||||
ap.add_argument("--pim-merge-scheduler", choices=("peft"), default="peft",
|
||||
help="Scheduler used by the Spatial merge-compute-nodes pass.")
|
||||
ap.add_argument("--pim-memory-report", choices=("none", "summary", "full"), default="none",
|
||||
help="Emit a human-readable PIM memory planning report during codegen.")
|
||||
ap.add_argument("--raptor-extra-arg", action="append", default=[],
|
||||
@@ -149,7 +147,7 @@ def main():
|
||||
result = validate_network(
|
||||
onnx_path, a.raptor_path, a.onnx_include_dir, simulator_dir,
|
||||
crossbar_size=a.crossbar_size, crossbar_count=a.crossbar_count, core_count=a.core_count,
|
||||
pim_merge_scheduler=a.pim_merge_scheduler, pim_memory_report=a.pim_memory_report,
|
||||
pim_memory_report=a.pim_memory_report,
|
||||
raptor_extra_args=a.raptor_extra_arg,
|
||||
command_timeout_seconds=a.command_timeout_seconds,
|
||||
threshold=a.threshold,
|
||||
|
||||
+42
-23
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
import numpy as np
|
||||
import shutil
|
||||
import sys
|
||||
@@ -61,32 +62,37 @@ class ProgressReporter:
|
||||
self.passed_models = 0
|
||||
self.failed_models = 0
|
||||
self.current_label = ""
|
||||
self.enabled = sys.stdout.isatty() if enabled is None else enabled
|
||||
self.enabled = (
|
||||
sys.stdout.isatty() and "CODEX_CI" not in os.environ
|
||||
if enabled is None else enabled
|
||||
)
|
||||
self.verbose = verbose
|
||||
self.columns = shutil.get_terminal_size((100, 20)).columns
|
||||
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:
|
||||
sys.stdout.write("\r" + (" " * self.rendered_width) + "\r")
|
||||
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
|
||||
bar_width = 24
|
||||
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}"
|
||||
if len(prefix_text) > self.columns:
|
||||
prefix_text = f"{self.completed_steps}/{self.total_steps}"
|
||||
|
||||
if prefix_text.startswith("["):
|
||||
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
|
||||
else:
|
||||
prefix = Fore.CYAN + prefix_text + Style.RESET_ALL
|
||||
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 = (
|
||||
" "
|
||||
@@ -109,14 +115,28 @@ class ProgressReporter:
|
||||
elif self.current_label:
|
||||
label = f" {self.current_label}"
|
||||
|
||||
available_label_width = max(0, self.columns - len(prefix_text) - len(model_counter) - len(counts_text) - 3)
|
||||
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_line = prefix_text + model_counter + f" P:{self.passed_models} F:{self.failed_models}" + label
|
||||
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
|
||||
padded_width = max(self.rendered_width, len(plain_line))
|
||||
sys.stdout.write("\r" + rendered_line + (" " * max(0, padded_width - len(plain_line))))
|
||||
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:
|
||||
@@ -158,7 +178,6 @@ class ProgressReporter:
|
||||
if self.enabled:
|
||||
self.suspended = True
|
||||
self._clear()
|
||||
self.rendered_width = 0
|
||||
|
||||
|
||||
def run_command(cmd, cwd=None, reporter=None, timeout_sec=None):
|
||||
@@ -293,7 +312,7 @@ def validate_outputs(sim_arrays, runner_out_dir, outputs_descriptor, threshold=1
|
||||
|
||||
def validate_network(network_onnx_path, raptor_path, onnx_include_dir,
|
||||
simulator_dir, crossbar_size=64, crossbar_count=8, core_count=None,
|
||||
pim_merge_scheduler="peft", pim_memory_report="none", raptor_extra_args=None,
|
||||
pim_memory_report="none", raptor_extra_args=None,
|
||||
threshold=1e-3, rtol=1e-5,
|
||||
seed=0, reporter=None, model_index=1, model_total=1, verbose=False,
|
||||
command_timeout_seconds=60.0, mode=MODE_FULL):
|
||||
@@ -346,8 +365,8 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir,
|
||||
if mode == MODE_COMPILE_ONLY:
|
||||
print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM")
|
||||
pim_pass_timings = compile_with_raptor(
|
||||
network_mlir_path, raptor_path, pim_output_base, crossbar_size, crossbar_count,
|
||||
core_count=core_count, pim_merge_scheduler=pim_merge_scheduler,
|
||||
network_onnx_path, raptor_path, pim_output_base, crossbar_size, crossbar_count,
|
||||
core_count=core_count,
|
||||
pim_memory_report=pim_memory_report, raptor_extra_args=raptor_extra_args,
|
||||
cwd=raptor_dir, verbose=verbose, reporter=reporter, timeout_sec=command_timeout_seconds)
|
||||
print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}")
|
||||
@@ -386,8 +405,8 @@ def validate_network(network_onnx_path, raptor_path, onnx_include_dir,
|
||||
if mode != MODE_RUN_ONLY:
|
||||
print_stage(reporter, model_index, model_total, network_onnx_path.name, "Compile PIM")
|
||||
pim_pass_timings = compile_with_raptor(
|
||||
network_mlir_path, raptor_path, pim_output_base, crossbar_size, crossbar_count,
|
||||
core_count=core_count, pim_merge_scheduler=pim_merge_scheduler,
|
||||
network_onnx_path, raptor_path, pim_output_base, crossbar_size, crossbar_count,
|
||||
core_count=core_count,
|
||||
pim_memory_report=pim_memory_report, raptor_extra_args=raptor_extra_args,
|
||||
cwd=raptor_dir, verbose=verbose, reporter=reporter, timeout_sec=command_timeout_seconds)
|
||||
print_info(reporter, f"PIM artifacts saved to {raptor_dir / 'pim'}")
|
||||
|
||||
Reference in New Issue
Block a user