hexagon: support for multi-NPU devices (IQ9, IQ10) and fully asynchronous backend (#26501)

* hexagon: use non-host bufs by default and make the backend fully async

* hex-hb: remove optional hostbuf support and fix async copy

* hex-unary: relax supported unary check

* hex-bufs: use same get_alignment for host bufs

* snapdragon: bump android_platform to 34

* hex-rows: super hacky get/set rows for q8_0

* hex-get-rows: fix q8_0

* hex-get-rows: supprot for f16 and cleanup for q8_0

* hex-get-rows: generic macros and specialized thread funcs

* hex-get-rows: add DMA pipeline, vtcm_layout and kernel params

* hex-set-rows: fix q8_0 support, add dma and tracing

* hex-tests: override nmse threshold for HTP of Q8_0 quants

* hex-fa: add support for Q8_0 with inplace dequantizers

* hex-get-rows: simplify type dispatch

* hex-rows: simplify GET/SET_ROWS DMA pipeline

* hex-async: add events, set/get-tensor-async and rest of the async api support

* hex-repack: use slice instead of expert in repack functions

* hex-cpy: update event/async-cpy logging

* hex-set-rows: optimize smaller tensors

* hex-geglu: fix perf regression with larger tensors

* hex-get-rows: add missing header

* hex-set-rows: add missing header

* hex-bufs: ressurect GGML_HEXAGON_HOSTBUF but disable it by default

* hexagon: do not reject ops with non-heaxon buffers

* hex-get-rows: apply >=32 restriction only for q8_0

* hex-res: bump vtcm acquire timeout to 10 seconds

* hex-bufs: add support for cloning buffers between sessions to speed up tensor copies

* hex-async: rework event recording and batch flushing and integrate with meta backend

* hex-bufs: improved handling of repacked tensors

* hex-repack: handle get_tensor_2d offsets

* hex-dev: add support for devices with multiple NPUs

* hex-sync: add support for sync tokens to synchronize npu devices for async splits

* hex-mmap: cleanup mmap calls and add a retry for robustness

* hex-sync: add failsafe if sync wait gets stuck

* hex-sync: use sync_seq to check for completed events

* hex-sync: rotate tokens for extra robustness

* hex-devs: add supprot for legacy device names for now

* hex-bufs: add support for auto-cloning buffers from diff sessions

* hex-fusion: simplify and optimize htp-opnode fusion handling

* hex-sync: override opnode name so that it shows up in the profiles

* hex-trace: update scripts to handle multiple devices

* hex-sync: bump the size of the opbatch queue and number of sync tokens

* hex-cpy-sync: do not explicitly flush opbatches in cpy_tensor_async and add support for cpy-dma

* hex-sync: add graph-flush threshold to avoid single op batches

* hex-sync: add sync_peer so that we can flush peers we depend on during cross-device ops

* hex-bufs: introduce tensor->extra and shadow_bufs for repacking

* hex-l2: flush tiny tensors inline

* hex-sync: use explicit l2flush for sync tokens

* hex-extra: track weight flags via tensor extra

* hex-fence: rename sync to fence

* hex-repack: proper handling of set-tensor-2d in the shadow_buf

* hex-trace: remove obsolete opstage mask that we used for profiling

* hex-env: remove obsolete use_hmx variable

* hexagon: new unified run.py and build.py and updated docs

* snapdragon: update run script to auto-escapt test-backend-op -p argument

* hex-scripts: fix trailing spaces

* hex-scripts: fix flake8 warnings

* snapdragon: cleanup dst lib/bin dirs before copying new build

* hex-ops: add support for allreduce

* hex-ar: improved allreduce with dma pipeline

* hex-ar: align macros

* hex-ar: consistent use of fence_seq

* hex-ar: add AR_SELECT env var to select ALLREDUCE kernel or fallback

* hex-ar: add proper synchronize handling for ALLREDUCE

* hex-opbatch: looks like we now just rely on backend.synchronise to flush the batches, no need to flush them by threshold

* hex-ar: bump block size to improve dma efficiency

* hex-ar: fused ALLREDUCE+ADD

* hex-ar: cleaner fence buffer management

* hex-ar: futher allreduce tweaking to remove race conditions

* hex-ar: add simple solver and remove non-dma kernels

* hex-ar: add row-broadcast to fuse with bias ADD

* hex-fence: pass seq numbers via op_params

* hex-ar: allow for both entry/exit seq for completing entry wait

* hex-ar: align macros

* hex-ar: do not refetch broadcast row

* hex-fusion: move all fusion into opbatch::add_op for consistency with ALLREDUCE and things

* hex-fusion: fix incorrect MUL_MAT reordering

* hex-mm: make fused 2x and 3x matmuls more generic

* hex-fusion: move tensor fusion tagging to graph_compute

* hexagon: make sure to copy tensor->extra by value

* hex-get-rows: fix offset calc with row-chunking

* hex-repack: get_tensor_2d fixes for non-zero offsets

* snapdragon: make profile/trace scripts more robust and donot mix stdout/stderr by default

* hex-devices: use legacy device nameing by default to ease the transition

* hex-devices: hardcode CDSP domain IDs for current devices for now

* hex-optrace: improve multi-NPU timestamp alignment and overall handling of cycle values

* hex-optrace: more robust handling of the fence events
This commit is contained in:
Max Krasnyansky
2026-08-26 18:46:50 -07:00
committed by GitHub
parent 925e117994
commit 192067b72d
44 changed files with 5314 additions and 3139 deletions
+193 -62
View File
@@ -34,6 +34,26 @@ trace_pattern = re.compile(
r"trace-evt\s+(?P<event>[A-Z_0-9\-]+):\s+thread\s+(?P<thread>\d+)\s+info\s+(?P<info>\d+)\s+(?P<state>start|stop)\s+(?P<cycles>\d+)"
)
device_pattern = re.compile(r"\b(HTP\d+(?::\d+)?)\s+(?:profile-op|trace-evt)\b")
def extract_device(line):
m = device_pattern.search(line)
if m:
return m.group(1)
return "HTP0"
def device_matches(record_device, target_device):
targets = [t.strip() for t in target_device.split(',')]
for target in targets:
if record_device == target:
return True
if record_device.startswith(target + ":"):
return True
return False
logger = logging.getLogger("ggml-hexagon-profile")
@@ -72,7 +92,7 @@ class CycleUnwrapper:
return raw + self.high_part
def parse_log(file_path, pmu_index=None):
def parse_log(file_path, pmu_index=None, limit=None, device_filter=None, op_filter_re=None):
try:
if file_path != "-":
f = open(file_path, 'r', encoding='utf-8', errors='ignore')
@@ -85,13 +105,22 @@ def parse_log(file_path, pmu_index=None):
all_ops: List[Dict[str, Any]] = []
all_traces: List[Dict[str, Any]] = []
current_op: Optional[Dict[str, Any]] = None
ops_count_per_device = {}
if device_filter is not None:
for target in device_filter.split(','):
ops_count_per_device[target.strip()] = 0
limit_reached = False
timestamp_pattern = re.compile(r"^(?P<min>\d+)\.(?P<sec>\d+)\.(?P<ms>\d+)\.(?P<us>\d+)\s+[A-Z]\s+")
unwrapper = None
trace_unwrapper = None
timestamp_pattern = re.compile(r"(?P<min>\d+)\.(?P<sec>\d+)\.(?P<ms>\d+)\.(?P<us>\d+)\s+[A-Z]\s+")
unwrappers = {}
last_batch_start = {}
trace_unwrappers = {}
for line in f:
ts_match = timestamp_pattern.match(line)
if "profile-op" not in line and "trace-evt" not in line:
continue
ts_match = timestamp_pattern.search(line)
abs_usec = 0
if ts_match:
abs_usec = (
@@ -100,8 +129,11 @@ def parse_log(file_path, pmu_index=None):
+ int(ts_match.group('us'))
)
if "|" in line and "profile-op" in line:
parts = [p.strip() for p in line.split("|")]
device = extract_device(line)
idx = line.find("profile-op")
if idx != -1 and "|" in line[idx:]:
parts = [p.strip() for p in line[idx:].split("|")]
prefix = parts[0]
prefix_match = re.search(r"profile-op\s+(?P<op_name>[A-Z_0-9+]+)", prefix)
if not prefix_match:
@@ -145,7 +177,6 @@ def parse_log(file_path, pmu_index=None):
except (ValueError, IndexError):
pmu_val = None
evt_val = None
evt_val = None
if types.startswith("evt-cnt "):
try:
@@ -158,14 +189,18 @@ def parse_log(file_path, pmu_index=None):
if op_name == "OPBATCH":
if cycles_start_raw:
unwrapped_cycles_start = int(cycles_start_raw)
unwrapper = CycleUnwrapper(unwrapped_cycles_start)
trace_unwrapper = CycleUnwrapper(unwrapped_cycles_start)
unwrappers[device] = CycleUnwrapper(unwrapped_cycles_start)
last_batch_start[device] = unwrapped_cycles_start
for k in list(trace_unwrappers.keys()):
if k[0] == device:
del trace_unwrappers[k]
else:
if cycles_start_raw and unwrapper is not None:
unwrapped_cycles_start = unwrapper.unwrap(int(cycles_start_raw))
if cycles_start_raw:
device_unwrapper = unwrappers.get(device)
if device_unwrapper is not None:
unwrapped_cycles_start = device_unwrapper.unwrap(int(cycles_start_raw))
idx = line.find("profile-op ")
op_text = line[idx + 11:].strip() if idx != -1 else line.strip()
op_text = re.sub(r"^profile-op\s+", "", line[idx:]).strip() if idx != -1 else line.strip()
current_op = {
'name': op_name,
@@ -180,24 +215,58 @@ def parse_log(file_path, pmu_index=None):
'pmu_val': pmu_val,
'evt_val': evt_val,
'abs_usec': abs_usec,
'trace_events': []
'trace_events': [],
'device': device
}
all_ops.append(current_op)
# Check if matching early exit criteria
matched = False
matched_target = None
if device_filter is not None:
targets = [t.strip() for t in device_filter.split(',')]
for target in targets:
if device == target or device.startswith(target + ":"):
matched = True
matched_target = target
break
else:
matched = True
matched_target = device
if op_filter_re is not None and not op_filter_re.search(op_text):
matched = False
if matched:
if matched_target not in ops_count_per_device:
ops_count_per_device[matched_target] = 0
ops_count_per_device[matched_target] += 1
if limit is not None and len(ops_count_per_device) > 0 and all(count >= limit for count in ops_count_per_device.values()):
limit_reached = True
if limit_reached and op_name == "OPBATCH":
break
continue
trace_match = trace_pattern.search(line)
if trace_match:
thread = int(trace_match.group('thread'))
raw_cyc = int(trace_match.group('cycles'))
unwrapped_cyc = None
if trace_unwrapper is not None:
unwrapped_cyc = trace_unwrapper.unwrap(raw_cyc)
th_key = (device, thread)
if th_key not in trace_unwrappers:
batch_start = last_batch_start.get(device)
trace_unwrappers[th_key] = CycleUnwrapper(batch_start)
unwrapped_cyc = trace_unwrappers[th_key].unwrap(raw_cyc)
all_traces.append({
'thread': int(trace_match.group('thread')),
'thread': thread,
'event': trace_match.group('event'),
'info': int(trace_match.group('info')),
'cycles': raw_cyc,
'unwrapped_cycles': unwrapped_cyc,
'state': trace_match.group('state')
'state': trace_match.group('state'),
'device': device
})
f.close()
@@ -207,39 +276,45 @@ def parse_log(file_path, pmu_index=None):
op['start_cycles'] = op['unwrapped_cycles_start']
op['end_cycles'] = op['start_cycles'] + op['cycles'] if op['start_cycles'] is not None else None
# Filter ops with valid start_cycles
valid_ops = [op for op in all_ops if op['start_cycles'] is not None and op['end_cycles'] is not None]
# Group ops by device
valid_ops_by_dev = defaultdict(list)
for op in all_ops:
if op['start_cycles'] is not None and op['end_cycles'] is not None:
valid_ops_by_dev[op['device']].append(op)
# Separate OPBATCH ops from other ops
opbatch_ops = [op for op in valid_ops if op['name'] == "OPBATCH"]
other_ops = [op for op in valid_ops if op['name'] != "OPBATCH"]
# Sort them by start_cycles to enable binary search
opbatch_ops.sort(key=lambda op: op['start_cycles'])
other_ops.sort(key=lambda op: op['start_cycles'])
opbatch_starts = [op['start_cycles'] for op in opbatch_ops]
other_starts = [op['start_cycles'] for op in other_ops]
# Map trace events to any operator whose cycles contain them
# Group trace events by device
traces_by_dev = defaultdict(list)
for e in all_traces:
cyc = e['unwrapped_cycles']
if cyc is None:
continue
if e['unwrapped_cycles'] is not None:
traces_by_dev[e['device']].append(e)
# Map to OPBATCH
idx = bisect.bisect_right(opbatch_starts, cyc) - 1
if idx >= 0:
op = opbatch_ops[idx]
if op['start_cycles'] <= cyc <= op['end_cycles']:
op['trace_events'].append(e)
for device, dev_ops in valid_ops_by_dev.items():
opbatch_ops = [op for op in dev_ops if op['name'] == "OPBATCH"]
other_ops = [op for op in dev_ops if op['name'] != "OPBATCH"]
# Map to other ops
idx = bisect.bisect_right(other_starts, cyc) - 1
if idx >= 0:
op = other_ops[idx]
if op['start_cycles'] <= cyc <= op['end_cycles']:
op['trace_events'].append(e)
opbatch_ops.sort(key=lambda op: op['start_cycles'])
other_ops.sort(key=lambda op: op['start_cycles'])
opbatch_starts = [op['start_cycles'] for op in opbatch_ops]
other_starts = [op['start_cycles'] for op in other_ops]
dev_traces = traces_by_dev.get(device, [])
for e in dev_traces:
cyc = e['unwrapped_cycles']
# Map to OPBATCH
idx = bisect.bisect_right(opbatch_starts, cyc) - 1
if idx >= 0:
op = opbatch_ops[idx]
if op['start_cycles'] <= cyc <= op['end_cycles']:
op['trace_events'].append(e)
# Map to other ops
idx = bisect.bisect_right(other_starts, cyc) - 1
if idx >= 0:
op = other_ops[idx]
if op['start_cycles'] <= cyc <= op['end_cycles']:
op['trace_events'].append(e)
return all_ops
@@ -563,6 +638,7 @@ def main():
parser.add_argument("--timeline", type=str, nargs='?', const='summary', choices=["summary", "bubbles"],
help="Output ASCII art event summary or thread idle bubble analysis (default: summary)")
parser.add_argument("--filter", type=str, help="Regex filter matching against the original profile-op line")
parser.add_argument("--device", type=str, help="Device to filter by (e.g. HTP0, HTP0:0) or 'split' to generate separate reports per device")
group = parser.add_mutually_exclusive_group()
group.add_argument("--head", type=int, help="Limit to first N ops")
@@ -586,29 +662,84 @@ def main():
logger.warning(f"Invalid width format '{w}'")
final_pmu_name = (args.pmu_name or f"#{args.pmu_index}") if args.pmu_index is not None else None
ops = parse_log(args.logfile, pmu_index=args.pmu_index)
op_filter_re = None
if args.filter:
try:
filter_re = re.compile(args.filter)
op_filter_re = re.compile(args.filter)
except re.error as e:
logger.error(f"Invalid regex filter: {e}")
sys.exit(1)
ops = [op for op in ops if filter_re.search(op['op_text'])]
if args.head is not None:
ops = ops[:args.head]
elif args.tail is not None:
ops = ops[-args.tail:]
limit = args.head if args.head is not None else None
device_filter = args.device if (args.device and args.device != "split") else None
ops = parse_log(args.logfile, pmu_index=args.pmu_index, limit=limit, device_filter=device_filter, op_filter_re=op_filter_re)
if args.timeline:
for op in ops:
if args.timeline == "summary":
print_ascii_summary(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events'])
elif args.timeline == "bubbles":
print_bubbles_timeline(op)
if args.device and args.device != "split":
ops = [op for op in ops if device_matches(op['device'], args.device)]
if args.device == "split":
unique_devices = sorted(list(set(op['device'] for op in ops)))
for dev in unique_devices:
dev_ops = [op for op in ops if device_matches(op['device'], dev)]
if args.filter:
try:
filter_re = re.compile(args.filter)
except re.error as e:
logger.error(f"Invalid regex filter: {e}")
sys.exit(1)
dev_ops = [op for op in dev_ops if filter_re.search(op['op_text'])]
if args.head is not None:
dev_ops = dev_ops[:args.head]
elif args.tail is not None:
dev_ops = dev_ops[-args.tail:]
logger.info("\n=========================================")
logger.info(f" Device: {dev}")
logger.info("=========================================")
if args.timeline:
for op in dev_ops:
if args.timeline == "summary":
print_ascii_summary(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events'])
elif args.timeline == "bubbles":
print_bubbles_timeline(op)
else:
generate_report(dev_ops, args.top, overrides, args.sort, pmu_name=final_pmu_name)
else:
generate_report(ops, args.top, overrides, args.sort, pmu_name=final_pmu_name)
if args.filter:
try:
filter_re = re.compile(args.filter)
except re.error as e:
logger.error(f"Invalid regex filter: {e}")
sys.exit(1)
ops = [op for op in ops if filter_re.search(op['op_text'])]
if args.head is not None or args.tail is not None:
ops_by_dev = defaultdict(list)
for op in ops:
ops_by_dev[op['device']].append(op)
filtered_ops = []
for dev in sorted(ops_by_dev.keys()):
dev_ops = ops_by_dev[dev]
if args.head is not None:
dev_ops = dev_ops[:args.head]
elif args.tail is not None:
dev_ops = dev_ops[-args.tail:]
filtered_ops.extend(dev_ops)
ops = filtered_ops
if args.timeline:
for op in ops:
if args.timeline == "summary":
print_ascii_summary(op['name'], op['dims'], op['types'], op['usec'], op['cycles'], op['trace_events'])
elif args.timeline == "bubbles":
print_bubbles_timeline(op)
else:
generate_report(ops, args.top, overrides, args.sort, pmu_name=final_pmu_name)
if __name__ == "__main__":