From d386df3ea5697a7b0e70194c205ad830ce6e90bb Mon Sep 17 00:00:00 2001 From: hw-native-sys-bot Date: Tue, 7 Apr 2026 15:11:01 +0800 Subject: [PATCH 1/3] Fix(ci): move retry and pin-commit logic to parent process (#456) The pin-commit retry was inside device_worker_main (subprocess), so a segfault in ChipWorker killed the subprocess before the retry could run. - Split run_hw_tasks_subprocess into sim (no retry) and hw (3 retries) - Move PTO-ISA pin retry to main() where the parent process controls it - Subprocess runs once and exits; all retry decisions made by parent - Remove --max-attempts CLI arg (subprocess always single-shot) Co-authored-by: Chao Wang <26245345+ChaoWao@users.noreply.github.com> --- ci.py | 203 +++++++++++++++++++++++++++++++++------------------------- 1 file changed, 115 insertions(+), 88 deletions(-) diff --git a/ci.py b/ci.py index 9dea2f71a3..1a33e72e0b 100644 --- a/ci.py +++ b/ci.py @@ -722,8 +722,6 @@ def _build_device_worker_base_args(args: argparse.Namespace) -> list[str]: base_args += ["-r", args.runtime] if args.build_runtime: base_args.append("--build-runtime") - if args.pto_isa_commit: - base_args += ["-c", args.pto_isa_commit] if args.run_all_cases: base_args.append("--all") return base_args @@ -734,11 +732,12 @@ def _run_device_worker_subprocess( device_id: int, args: argparse.Namespace, tag: str, - max_attempts: int = MAX_RETRIES, + pto_isa_commit: str | None = None, ) -> list[TaskResult]: """Run a task batch in one device-worker subprocess and return its reported results.""" base_args = _build_device_worker_base_args(args) - base_args += ["--max-attempts", str(max_attempts)] + if pto_isa_commit: + base_args += ["-c", pto_isa_commit] with tempfile.NamedTemporaryFile( prefix=f"ci_{tag}_tasks_dev{device_id}_", @@ -822,12 +821,33 @@ def _normalize_task_result( ) +def run_sim_tasks_subprocess( + tasks: list[TaskSpec], + args: argparse.Namespace, + pto_isa_commit: str | None = None, +) -> list[TaskResult]: + """Run simulation tasks: one subprocess per task, no retry.""" + results: list[TaskResult] = [] + for task in tasks: + task_results = _run_device_worker_subprocess( + [task], + 0, + args, + tag="sim", + pto_isa_commit=pto_isa_commit, + ) + normalized = _normalize_task_result(task, 0, 0, task_results) + results.append(normalized) + return results + + def run_hw_tasks_subprocess( tasks: list[TaskSpec], devices: list[int], args: argparse.Namespace, + pto_isa_commit: str | None = None, ) -> list[TaskResult]: - """Run hardware tasks via a shared parent queue, mirroring ci.sh parallel scheduling.""" + """Run hardware tasks: one subprocess per task, retry up to MAX_RETRIES in parent.""" task_queue: Queue[tuple[TaskSpec, int] | None] = Queue() for task in tasks: task_queue.put((task, 0)) @@ -850,7 +870,7 @@ def _run_device(dev_id: int): dev_id, args, tag=tag, - max_attempts=1, + pto_isa_commit=pto_isa_commit, ) normalized = _normalize_task_result(task, dev_id, attempt, task_results) with results_lock: @@ -974,88 +994,101 @@ def device_worker_main(args: argparse.Namespace) -> int: logger.info("No tasks found") return 0 + all_results = _run_tasks_on_device(tasks, device_id, platform, pto_isa_root, args) + _write_results_json(all_results, args.result_json) + return print_summary(all_results) + + +def _run_tasks_on_device( + tasks: list[TaskSpec], + device_id: int, + platform: str, + pto_isa_root: str, + args: argparse.Namespace, +) -> list[TaskResult]: + """Compile and run all tasks on a single device. Returns all TaskResults.""" logger.info(f"Compiling {len(tasks)} tasks...") - compiled = compile_all_tasks( - tasks, pto_isa_root, build_runtime=args.build_runtime, run_all_cases=args.run_all_cases - ) + try: + compiled = compile_all_tasks( + tasks, pto_isa_root, build_runtime=args.build_runtime, run_all_cases=args.run_all_cases + ) + except RuntimeError: + return [ + TaskResult( + name=t.name, + platform=platform, + passed=False, + device=str(device_id), + attempt=0, + elapsed_s=0, + error="compile failed", + ) + for t in tasks + ] groups = group_by_runtime(compiled) all_results: list[TaskResult] = [] for rt_name, group_tasks in groups.items(): - remaining = list(group_tasks) - - for attempt in range(args.max_attempts): - if not remaining: - break + rt_bins = cast(RuntimeBinariesLike, group_tasks[0].runtime_bins) + worker = ChipWorker() + try: + worker.init( + device_id, + str(rt_bins.host_path), + rt_bins.aicpu_path.read_bytes(), + rt_bins.aicore_path.read_bytes(), + ) + except Exception as e: + logger.error(f"[dev{device_id}] Failed to init ChipWorker for {rt_name}: {e}") + all_results.extend( + TaskResult( + name=ct.spec.name, + platform=platform, + passed=False, + device=str(device_id), + attempt=0, + elapsed_s=0, + error=str(e), + ) + for ct in group_tasks + ) + continue - rt_bins = cast(RuntimeBinariesLike, remaining[0].runtime_bins) - worker = ChipWorker() + for ct in group_tasks: + start = time.monotonic() try: - worker.init( - device_id, - str(rt_bins.host_path), - rt_bins.aicpu_path.read_bytes(), - rt_bins.aicore_path.read_bytes(), + run_single_task(ct, worker, device_id) + elapsed = time.monotonic() - start + logger.info(f"[dev{device_id}] PASS: {ct.spec.name} ({elapsed:.1f}s)") + all_results.append( + TaskResult( + name=ct.spec.name, + platform=platform, + passed=True, + device=str(device_id), + attempt=0, + elapsed_s=elapsed, + ) ) except Exception as e: - logger.error(f"[dev{device_id}] Failed to init ChipWorker for {rt_name}: {e}") - all_results.extend( + elapsed = time.monotonic() - start + logger.error(f"[dev{device_id}] FAIL: {ct.spec.name} ({elapsed:.1f}s): {e}") + all_results.append( TaskResult( name=ct.spec.name, platform=platform, passed=False, device=str(device_id), - attempt=attempt, - elapsed_s=0, + attempt=0, + elapsed_s=elapsed, error=str(e), ) - for ct in remaining ) - remaining = [] - break - - failed_tasks = [] - for ct in remaining: - start = time.monotonic() - try: - run_single_task(ct, worker, device_id) - elapsed = time.monotonic() - start - logger.info(f"[dev{device_id}] PASS: {ct.spec.name} ({elapsed:.1f}s)") - all_results.append( - TaskResult( - name=ct.spec.name, - platform=platform, - passed=True, - device=str(device_id), - attempt=attempt, - elapsed_s=elapsed, - ) - ) - except Exception as e: - elapsed = time.monotonic() - start - logger.error(f"[dev{device_id}] FAIL: {ct.spec.name} ({elapsed:.1f}s): {e}") - all_results.append( - TaskResult( - name=ct.spec.name, - platform=platform, - passed=False, - device=str(device_id), - attempt=attempt, - elapsed_s=elapsed, - error=str(e), - ) - ) - failed_tasks.append(ct) - worker.reset() - remaining = failed_tasks - - if remaining and attempt + 1 >= MAX_RETRIES: - logger.warning(f"[dev{device_id}] Quarantined after exhausting retries") + worker.reset() - _write_results_json(all_results, args.result_json) - return print_summary(all_results) + return all_results # --------------------------------------------------------------------------- @@ -1097,7 +1130,6 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--parallel", action="store_true") parser.add_argument("--all", dest="run_all_cases", action="store_true", help="Run all cases, not just DEFAULT_CASE") parser.add_argument("--device-worker", action="store_true", help=argparse.SUPPRESS) - parser.add_argument("--max-attempts", type=int, default=MAX_RETRIES, help=argparse.SUPPRESS) parser.add_argument("--result-json", default=None, help=argparse.SUPPRESS) parser.add_argument("--task-list-json", default=None, help=argparse.SUPPRESS) return parser.parse_args() @@ -1150,35 +1182,30 @@ def _watchdog_handler(signum, frame): return 0 logger.info(f"Discovered {len(tasks)} tasks") - # Step 2 & 3: Compile and run via subprocess-per-runtime-group - # Each subprocess loads exactly one host .so, avoiding RTLD_GLOBAL symbol conflicts. + # Step 2: Compile and run — each task in its own subprocess. + # sim: no retry; hw: retry up to MAX_RETRIES in parent. if is_sim: - all_results = run_hw_tasks_subprocess(tasks, [0], args) + all_results = run_sim_tasks_subprocess(tasks, args) else: all_results = run_hw_tasks_subprocess(tasks, args.devices, args) - # Step 5: PTO-ISA pinned retry for failures - # Deduplicate results by task name (last result wins, same as print_summary) - # then only retry tasks that did NOT exhaust all retries — tasks that failed - # on every attempt are deterministic and won't benefit from a PTO-ISA pin. - final_by_name: dict[str, TaskResult] = {} + # Step 3: Pin retry — re-run failed tasks with pinned PTO-ISA commit. + final: dict[str, TaskResult] = {} for r in all_results: - final_by_name[r.name] = r - max_attempt = args.max_attempts - 1 - failures = [r for r in final_by_name.values() if not r.passed and r.attempt < max_attempt] + final[r.name] = r + failures = [r for r in final.values() if not r.passed] + if failures and args.pto_isa_commit: failed_names = {r.name for r in failures} - logger.info(f"[CI] {len(failures)} failure(s), retrying with pinned PTO-ISA {args.pto_isa_commit}") - reset_pto_isa(args.pto_isa_commit, args.clone_protocol) - retry_tasks = [task for task in tasks if task.name in failed_names] + failed_tasks = [t for t in tasks if t.name in failed_names] + logger.info(f"[CI] {len(failed_tasks)} failure(s), retrying with pinned PTO-ISA {args.pto_isa_commit}") if is_sim: - retry_results = run_hw_tasks_subprocess(retry_tasks, [0], args) + pin_results = run_sim_tasks_subprocess(failed_tasks, args, pto_isa_commit=args.pto_isa_commit) else: - retry_results = run_hw_tasks_subprocess(retry_tasks, args.devices, args) - - all_results.extend(retry_results) + pin_results = run_hw_tasks_subprocess(failed_tasks, args.devices, args, pto_isa_commit=args.pto_isa_commit) + all_results.extend(pin_results) - # Step 6: Summary + # Step 4: Summary signal.alarm(0) return print_summary(all_results) From c2b2c7f04220c8f39beff9dbd87a0571a36548db Mon Sep 17 00:00:00 2001 From: hw-native-sys-bot Date: Tue, 7 Apr 2026 16:40:34 +0800 Subject: [PATCH 2/3] Refactor: add device quarantine and improve failure logging (#462) - Quarantine device on first failure; re-enqueue task for healthy devices (up to MAX_RETRIES across devices), matching ci.sh semantics - Print subprocess logs on failure: sim during pin-commit retry, hw on last retry during pin-commit - Add progress logging ([n/total] PASS/FAIL) for both sim and hw paths - Remove unused `--parallel` flag (parallelism determined by `-d` device count) - Remove unused `PYTHONDONTWRITEBYTECODE` env setting Co-authored-by: Chao Wang <26245345+ChaoWao@users.noreply.github.com> --- .github/workflows/ci.yml | 4 +- ci.py | 195 +++++++++++++++++++++++++-------------- 2 files changed, 127 insertions(+), 72 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 229c4c58ea..73bd34edce 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -215,7 +215,7 @@ jobs: - name: Run on-device examples (a2a3) run: | export PATH="$HOME/.local/bin:$PATH" - source ${ASCEND_HOME_PATH}/bin/setenv.bash && python ci.py -p a2a3 -d ${DEVICE_RANGE} --parallel -c 882c4db -t 600 --clone-protocol https + source ${ASCEND_HOME_PATH}/bin/setenv.bash && python ci.py -p a2a3 -d ${DEVICE_RANGE} -c 882c4db -t 600 --clone-protocol https # ---------- Detect A5 changes (runs on GitHub server, not A5 machine) ---------- @@ -290,4 +290,4 @@ jobs: export PATH="$HOME/.local/bin:$PATH" source ${ASCEND_HOME_PATH}/bin/setenv.bash DEVICE_LIST=$(python -c "s,e='${DEVICE_RANGE}'.split('-'); print(','.join(str(i) for i in range(int(s),int(e)+1)))") - task-submit --device "$DEVICE_LIST" --run "python ci.py -p a5 -d ${DEVICE_RANGE} --parallel -c 882c4db -t 600 --clone-protocol https" + task-submit --device "$DEVICE_LIST" --run "python ci.py -p a5 -d ${DEVICE_RANGE} -c 882c4db -t 1200 --clone-protocol https" diff --git a/ci.py b/ci.py index 1a33e72e0b..c62b78f0bf 100644 --- a/ci.py +++ b/ci.py @@ -14,7 +14,7 @@ per device, reusing ChipWorker across tasks that share the same runtime. Usage: - python tools/ci.py -p a2a3 -d 5-8 --parallel -c 6622890 -t 600 + python tools/ci.py -p a2a3 -d 5-8 -c 6622890 -t 600 python tools/ci.py -p a2a3sim -r tensormap_and_ringbuffer -c 6622890 -t 600 """ @@ -35,7 +35,7 @@ from pathlib import Path from queue import Empty, Queue from threading import Lock, Thread -from typing import Any, Protocol, cast +from typing import Any, Callable, Protocol, cast # --------------------------------------------------------------------------- # Path setup — mirrors run_example.py @@ -733,6 +733,7 @@ def _run_device_worker_subprocess( args: argparse.Namespace, tag: str, pto_isa_commit: str | None = None, + print_log_on_fail: bool = False, ) -> list[TaskResult]: """Run a task batch in one device-worker subprocess and return its reported results.""" base_args = _build_device_worker_base_args(args) @@ -765,10 +766,11 @@ def _run_device_worker_subprocess( logger.info(f"[{tag}:dev{device_id}] Launching: {' '.join(full_cmd)}") try: - proc = subprocess.run(full_cmd, check=False, capture_output=True, text=True, timeout=args.timeout) + proc = subprocess.run(full_cmd, check=False, capture_output=True, text=True) device_results = _read_results_json(result_path) if proc.returncode != 0: - logger.error(f"[{tag}:dev{device_id}] Failed:\n{proc.stdout}\n{proc.stderr}") + if print_log_on_fail: + logger.error(f"[{tag}:dev{device_id}] Failed:\n{proc.stdout}\n{proc.stderr}") fallback_needed = proc.returncode != 0 and not any(not result.passed for result in device_results) if fallback_needed: device_results.append( @@ -783,20 +785,6 @@ def _run_device_worker_subprocess( ) ) return device_results - except subprocess.TimeoutExpired: - error_msg = f"Timed out after {args.timeout}s" - logger.error(f"[{tag}:dev{device_id}] {error_msg}") - return [ - TaskResult( - name=f"{tag}-device-{device_id}", - platform=args.platform, - passed=False, - device=str(device_id), - attempt=0, - elapsed_s=args.timeout, - error=error_msg, - ) - ] finally: task_list_path.unlink(missing_ok=True) result_path.unlink(missing_ok=True) @@ -827,17 +815,22 @@ def run_sim_tasks_subprocess( pto_isa_commit: str | None = None, ) -> list[TaskResult]: """Run simulation tasks: one subprocess per task, no retry.""" + is_pin_retry = pto_isa_commit is not None results: list[TaskResult] = [] - for task in tasks: + total = len(tasks) + for i, task in enumerate(tasks, 1): task_results = _run_device_worker_subprocess( [task], 0, args, tag="sim", pto_isa_commit=pto_isa_commit, + print_log_on_fail=is_pin_retry, ) normalized = _normalize_task_result(task, 0, 0, task_results) results.append(normalized) + status = "PASS" if normalized.passed else "FAIL" + logger.info(f"[sim] [{i}/{total}] {status}: {task.name} ({normalized.elapsed_s:.1f}s)") return results @@ -847,55 +840,93 @@ def run_hw_tasks_subprocess( args: argparse.Namespace, pto_isa_commit: str | None = None, ) -> list[TaskResult]: - """Run hardware tasks: one subprocess per task, retry up to MAX_RETRIES in parent.""" - task_queue: Queue[tuple[TaskSpec, int] | None] = Queue() + """Run hardware tasks: one subprocess per task. + + On any failure the device is immediately quarantined (worker exits). Healthy + devices keep pulling from the shared queue. Tasks that were never run or failed + are collected so the caller can re-run them in a pin-commit pass with all devices + refreshed. + """ + task_queue: Queue[tuple[TaskSpec, int]] = Queue() + total = len(tasks) for task in tasks: task_queue.put((task, 0)) results: list[TaskResult] = [] results_lock = Lock() + completed = [0] # mutable counter for thread-safe increment + quarantined: set[int] = set() + quarantine_lock = Lock() tag = "hw" + is_pin_retry = pto_isa_commit is not None + def _run_device(dev_id: int): while True: - item = task_queue.get() - if item is None: - task_queue.task_done() + try: + task, attempt = task_queue.get_nowait() + except Empty: return - task, attempt = item - try: - task_results = _run_device_worker_subprocess( - [task], - dev_id, - args, - tag=tag, - pto_isa_commit=pto_isa_commit, - ) - normalized = _normalize_task_result(task, dev_id, attempt, task_results) - with results_lock: - results.append(normalized) + is_last_attempt = attempt + 1 >= MAX_RETRIES + task_results = _run_device_worker_subprocess( + [task], + dev_id, + args, + tag=tag, + pto_isa_commit=pto_isa_commit, + print_log_on_fail=is_pin_retry and is_last_attempt, + ) + normalized = _normalize_task_result(task, dev_id, attempt, task_results) + with results_lock: + results.append(normalized) + if normalized.passed or is_last_attempt: + completed[0] += 1 + n = completed[0] + status = "PASS" if normalized.passed else "FAIL" + attempt_info = f" attempt {attempt + 1}" if attempt > 0 else "" + logger.info( + f"[{tag}:dev{dev_id}] [{n}/{total}] {status}: {task.name}{attempt_info} ({normalized.elapsed_s:.1f}s)" + ) - if normalized.passed: - continue + if normalized.passed: + continue - next_attempt = attempt + 1 - if next_attempt < MAX_RETRIES: - task_queue.put((task, next_attempt)) - else: - logger.warning(f"[{tag}:dev{dev_id}] Exhausted retries on {task.name}") - finally: - task_queue.task_done() + # Failure: re-enqueue with attempt+1 if under limit, quarantine this device + if not is_last_attempt: + task_queue.put((task, attempt + 1)) + logger.warning(f"[{tag}:dev{dev_id}] Quarantined after failure on {task.name}") + with quarantine_lock: + quarantined.add(dev_id) + return threads = [Thread(target=_run_device, args=(device_id,)) for device_id in devices] for t in threads: t.start() - task_queue.join() - for _ in threads: - task_queue.put(None) for t in threads: t.join() + # Tasks stranded in queue — all devices quarantined before queue emptied + while True: + try: + task, attempt = task_queue.get_nowait() + except Empty: + break + results.append( + TaskResult( + name=task.name, + platform=task.platform, + passed=False, + device="N/A", + attempt=attempt, + elapsed_s=0, + error="All devices quarantined", + ) + ) + + if quarantined: + logger.warning(f"[{tag}] Quarantined devices: {sorted(quarantined)}") + return results @@ -1127,7 +1158,6 @@ def parse_args() -> argparse.Namespace: parser.add_argument("-c", "--pto-isa-commit", default=None) parser.add_argument("-t", "--timeout", type=int, default=600) parser.add_argument("--clone-protocol", choices=["ssh", "https"], default="ssh") - parser.add_argument("--parallel", action="store_true") parser.add_argument("--all", dest="run_all_cases", action="store_true", help="Run all cases, not just DEFAULT_CASE") parser.add_argument("--device-worker", action="store_true", help=argparse.SUPPRESS) parser.add_argument("--result-json", default=None, help=argparse.SUPPRESS) @@ -1142,9 +1172,32 @@ def parse_device_range(device_range: str) -> list[int]: return [int(device_range)] +def _run_with_timeout( + phase_name: str, + timeout_s: int, + runner: Callable[[], list[TaskResult]], +) -> list[TaskResult]: + def _watchdog_handler(signum, frame): + print(f"\n{'=' * 40}", flush=True) + print( + f"[CI] TIMEOUT: {phase_name} exceeded {timeout_s}s ({timeout_s // 60}min) limit, aborting", + flush=True, + ) + print(f"{'=' * 40}", flush=True) + os._exit(1) + + previous_handler = signal.getsignal(signal.SIGALRM) + signal.signal(signal.SIGALRM, _watchdog_handler) + signal.alarm(timeout_s) + try: + return runner() + finally: + signal.alarm(0) + signal.signal(signal.SIGALRM, previous_handler) + + def main() -> int: logging.basicConfig(level=logging.INFO, format="[%(levelname)s] %(message)s", force=True) - os.environ["PYTHONDONTWRITEBYTECODE"] = "1" args = parse_args() args.devices = parse_device_range(args.device_range) @@ -1161,20 +1214,6 @@ def main() -> int: if args.device_worker: return device_worker_main(args) - # Watchdog timer - watchdog_fired = False - - def _watchdog_handler(signum, frame): - nonlocal watchdog_fired - watchdog_fired = True - print(f"\n{'=' * 40}", flush=True) - print(f"[CI] TIMEOUT: exceeded {args.timeout}s ({args.timeout // 60}min) limit, aborting", flush=True) - print(f"{'=' * 40}", flush=True) - os._exit(1) - - signal.signal(signal.SIGALRM, _watchdog_handler) - signal.alarm(args.timeout) - # Step 1: Discover tasks tasks = discover_tasks(args.platform, runtime_filter=args.runtime) if not tasks: @@ -1183,11 +1222,15 @@ def _watchdog_handler(signum, frame): logger.info(f"Discovered {len(tasks)} tasks") # Step 2: Compile and run — each task in its own subprocess. - # sim: no retry; hw: retry up to MAX_RETRIES in parent. + # hw: failed device is quarantined; healthy devices keep running remaining tasks. if is_sim: - all_results = run_sim_tasks_subprocess(tasks, args) + all_results = _run_with_timeout("initial pass", args.timeout, lambda: run_sim_tasks_subprocess(tasks, args)) else: - all_results = run_hw_tasks_subprocess(tasks, args.devices, args) + all_results = _run_with_timeout( + "initial pass", + args.timeout, + lambda: run_hw_tasks_subprocess(tasks, args.devices, args), + ) # Step 3: Pin retry — re-run failed tasks with pinned PTO-ISA commit. final: dict[str, TaskResult] = {} @@ -1200,13 +1243,25 @@ def _watchdog_handler(signum, frame): failed_tasks = [t for t in tasks if t.name in failed_names] logger.info(f"[CI] {len(failed_tasks)} failure(s), retrying with pinned PTO-ISA {args.pto_isa_commit}") if is_sim: - pin_results = run_sim_tasks_subprocess(failed_tasks, args, pto_isa_commit=args.pto_isa_commit) + pin_results = _run_with_timeout( + "pin retry", + args.timeout, + lambda: run_sim_tasks_subprocess(failed_tasks, args, pto_isa_commit=args.pto_isa_commit), + ) else: - pin_results = run_hw_tasks_subprocess(failed_tasks, args.devices, args, pto_isa_commit=args.pto_isa_commit) + pin_results = _run_with_timeout( + "pin retry", + args.timeout, + lambda: run_hw_tasks_subprocess( + failed_tasks, + args.devices, + args, + pto_isa_commit=args.pto_isa_commit, + ), + ) all_results.extend(pin_results) # Step 4: Summary - signal.alarm(0) return print_summary(all_results) From a0eff4b0c2052fa631f097c424c07c18f20be0c2 Mon Sep 17 00:00:00 2001 From: chenshengxin Date: Tue, 7 Apr 2026 18:58:41 +0800 Subject: [PATCH 3/3] Perf: cache hash and prefetch chain in TensorMap lookup/insert - Cache the hash(addr) result from lookup() and reuse it in the subsequent insert() call for INOUT tensors, eliminating a redundant 64-bit multiply per tensor - Add software prefetch of next_in_bucket during chain traversal to hide memory latency on chains longer than one entry - Add lookup/insert/link_entry overloads that accept precomputed hash Benchmarked on Ascend910 (device 11, 100 rounds, 3 runs averaged): benchmark_bgemm -3.8%, other workloads -0.2% to -0.7%. --- .../runtime/pto_orchestrator.cpp | 10 ++++-- .../runtime/pto_tensormap.h | 31 ++++++++++++++++++- 2 files changed, 38 insertions(+), 3 deletions(-) diff --git a/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_orchestrator.cpp b/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_orchestrator.cpp index 5e754f9e0e..e05e67b702 100644 --- a/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_orchestrator.cpp +++ b/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_orchestrator.cpp @@ -557,6 +557,7 @@ pto2_submit_mixed_task(PTO2OrchestratorState *orch, const MixedKernels &mixed_ke CYCLE_COUNT_LAP_RECORD(g_orch_sync_cycle, AicpuPhaseId::ORCH_SYNC, task_id.raw); // === STEP 3: Lookup inputs + materialize runtime-created outputs === + uint32_t cached_hashes[MAX_TENSOR_ARGS] = {}; for (int i = 0; i < args.tensor_count(); i++) { TensorArgType ptype = args.tag(i); if (ptype == TensorArgType::OUTPUT) { @@ -587,7 +588,7 @@ pto2_submit_mixed_task(PTO2OrchestratorState *orch, const MixedKernels &mixed_ke } PTO2LookupResult lookup_result; - orch->tensor_map.lookup(*tensor, lookup_result); + orch->tensor_map.lookup(*tensor, lookup_result, cached_hashes[i]); for (int r = 0; r < lookup_result.count; r++) { PTO2TensorMapEntry &entry = *lookup_result.entries[r].entry; @@ -614,7 +615,12 @@ pto2_submit_mixed_task(PTO2OrchestratorState *orch, const MixedKernels &mixed_ke TensorArgType ptype = args.tag(i); if (ptype == TensorArgType::INOUT || ptype == TensorArgType::OUTPUT_EXISTING) { if (!args.tensor(i).ptr->manual_dep) { - orch->tensor_map.insert(*args.tensor(i).ptr, task_id); + if (ptype == TensorArgType::INOUT) { + // Reuse hash cached during lookup (STEP 3) + orch->tensor_map.insert(*args.tensor(i).ptr, task_id, cached_hashes[i]); + } else { + orch->tensor_map.insert(*args.tensor(i).ptr, task_id); + } } } } diff --git a/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_tensormap.h b/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_tensormap.h index 98ff5211e7..e847a86d35 100644 --- a/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_tensormap.h +++ b/src/a2a3/runtime/tensormap_and_ringbuffer/runtime/pto_tensormap.h @@ -301,7 +301,17 @@ struct PTO2TensorMap { * @param result Output: stack-allocated result buffer */ void lookup(const Tensor &tensor, PTO2LookupResult &result) { + uint32_t unused; + lookup(tensor, result, unused); + } + + /** + * Lookup with hash output — returns the computed bucket index for reuse + * in a subsequent insert() call on the same tensor address. + */ + void lookup(const Tensor &tensor, PTO2LookupResult &result, uint32_t &out_hash) { uint32_t bucket_index = hash(tensor.buffer.addr); + out_hash = bucket_index; PTO2TensorMapEntry *cur_entry = buckets[bucket_index]; result.count = 0; @@ -312,6 +322,9 @@ struct PTO2TensorMap { while (cur_entry != nullptr) { PTO2TensorMapEntry *next_entry = cur_entry->next_in_bucket; + if (next_entry != nullptr) { + __builtin_prefetch(next_entry, 0, 1); // Prefetch next entry's cache line 1 (read, moderate locality) + } #if PTO2_TENSORMAP_PROFILING chain_len++; @@ -364,6 +377,16 @@ struct PTO2TensorMap { link_entry(entry, tensor.buffer.addr, producer_task_id); } + /** + * Insert with precomputed hash — avoids recomputing hash(addr) when + * the caller already has it from a prior lookup() on the same address. + */ + void insert(const Tensor &tensor, PTO2TaskId producer_task_id, uint32_t precomputed_hash) { + PTO2TensorMapEntry *entry = new_entry(); + entry->copy_from_tensor(tensor); + link_entry(entry, producer_task_id, precomputed_hash); + } + /** * Cleanup stale entries for retired tasks * @@ -417,10 +440,16 @@ struct PTO2TensorMap { * Link an initialized entry into bucket and task chains. */ void link_entry(PTO2TensorMapEntry *entry, uint64_t addr, PTO2TaskId producer_task_id) { + link_entry(entry, producer_task_id, hash(addr)); + } + + /** + * Link an initialized entry into bucket and task chains (with precomputed hash). + */ + void link_entry(PTO2TensorMapEntry *entry, PTO2TaskId producer_task_id, uint32_t bucket_index) { #if PTO2_TENSORMAP_PROFILING g_insert_count++; #endif - uint32_t bucket_index = hash(addr); auto ring_id = producer_task_id.ring(); auto local_id = producer_task_id.local(); int32_t task_slot = local_id & (task_window_sizes[ring_id] - 1);