From 874d33128b4e4785beea736d97df6716e0321717 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Wed, 5 Aug 2026 10:00:42 +0300 Subject: [PATCH] hcq2 benchmark (#17235) * hcq2 in ci? * fix * traning * x * x * x * recover * debug * impler * x * x * x * hcq2: group input scatter plans by destination * hcq2: simplify input scatter tables * x --- .github/workflows/benchmark.yml | 4 ++++ .github/workflows/test.yml | 10 +++++++- extra/hcq2/ops_amd2.py | 8 ++++--- tinygrad/engine/realize.py | 2 +- tinygrad/runtime/support/hcq2.py | 40 +++++++++++++++++++------------- 5 files changed, 43 insertions(+), 21 deletions(-) diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index b805c99694..739f7242c0 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -94,6 +94,7 @@ jobs: shell: bash -e -o pipefail {0} env: DEV: ${{ matrix.dev }} + HCQ2: ${{ matrix.dev == 'AMD' && '1' || '0' }} if: github.repository_owner == 'tinygrad' steps: - name: Checkout Code @@ -148,6 +149,7 @@ jobs: shell: bash -e -o pipefail {0} env: DEV: ${{ matrix.dev }} + HCQ2: ${{ matrix.dev == 'AMD' && '1' || '0' }} if: github.repository_owner == 'tinygrad' steps: - name: Checkout Code @@ -200,6 +202,7 @@ jobs: shell: bash -e -o pipefail {0} env: DEV: ${{ matrix.dev }} + HCQ2: ${{ matrix.dev == 'AMD' && '1' || '0' }} if: github.repository_owner == 'tinygrad' steps: - name: Checkout Code @@ -249,6 +252,7 @@ jobs: shell: bash -e -o pipefail {0} env: DEV: ${{ matrix.dev }} + HCQ2: ${{ matrix.dev == 'AMD' && '1' || '0' }} if: github.repository_owner == 'tinygrad' steps: - name: Checkout Code diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 54fc715189..605a8e6d5b 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -519,7 +519,15 @@ jobs: deps: testing_unit amd: 'true' - name: Run HCQ2 tests - run: HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/test_tiny.py + run: | + command -v gdb >/dev/null || sudo apt-get install -y -qq gdb || true + ulimit -c unlimited || true + echo 'core.%p' | sudo tee /proc/sys/kernel/core_pattern || true + # -u + -v so the last test name reaches the log before a crash, faulthandler for the python stack on SIGSEGV + HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONFAULTHANDLER=1 PYTHONPATH=. python -u test/test_tiny.py -v || { + st=$?; echo "::error::test_tiny exited $st" + for c in core.*; do gdb -q -batch -ex 'thread apply all bt' "$(command -v python)" "$c" || true; done + exit $st; } - name: Run HCQ2 multi-device tests run: | HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/unit/test_multitensor.py \ diff --git a/extra/hcq2/ops_amd2.py b/extra/hcq2/ops_amd2.py index 788d40f495..4376e1083d 100644 --- a/extra/hcq2/ops_amd2.py +++ b/extra/hcq2/ops_amd2.py @@ -182,7 +182,8 @@ def sdma_copy(ctx, call): src_addr, dst_addr = call.src[2].getaddr(ctx.devs), call.src[1].getaddr(ctx.devs) return call.ins(SDMAOps.COPY, src=tuple(UOp.const(x, dtypes.uint32) for off in range(0, sz, ctx.max_copy_size) for x in ( ctx.sdma.SDMA_OP_COPY | ctx.sdma.SDMA_PKT_COPY_LINEAR_HEADER_SUB_OP(ctx.sdma.SDMA_SUBOP_COPY_LINEAR), - ctx.sdma.SDMA_PKT_COPY_LINEAR_COUNT_COUNT(min(sz-off, ctx.max_copy_size)-1), 0, *data64_le(src_addr+off), *data64_le(dst_addr+off)))) + ctx.sdma.SDMA_PKT_COPY_LINEAR_COUNT_COUNT(min(sz-off, ctx.max_copy_size)-1), 0, + *data64_le(src_addr+UOp.const(off, dtypes.uint64)), *data64_le(dst_addr+UOp.const(off, dtypes.uint64))))) def sdma_wait(ctx, ins, dst, val): op = ctx.sdma.SDMA_OP_POLL_REGMEM | ctx.sdma.SDMA_PKT_POLL_REGMEM_HEADER_FUNC(WAIT_REG_MEM_FUNCTION_GEQ) \ @@ -507,11 +508,12 @@ class PCIIface(PCIIfaceBase): if drain_only: d.iface.dev_impl.ih.drain() else: d.iface.dev_impl.ih.interrupt_handler() - if reset and d.iface.dev_impl.recover(): + if reset and d.iface.dev_impl.recover(force=True): cq = d.compute_queue for b in (cq.put_value, cq.read_ptr, cq.write_ptr): b._buf.view.view(fmt='Q')[0] = 0 d.iface.dev_impl.gfx.setup_ring(*cq.params) - d.signal('timeline')._buf.cpu_view().mv.cast('Q')[0] = d.signal('value', 1).as_memoryview(force_zero_copy=True).cast('Q')[0] - 1 + d.signal('timeline')._buf.cpu_view().mv.cast('Q')[0] = \ + d.signal('value', 1).as_memoryview(force_zero_copy=True, no_sync=True).cast('Q')[0] - 1 def sleep(self, timeout): if hasattr(self.pci_dev, 'irq_poller') and self.pci_dev.irq_poller is not None and (events_cnt:=len(self.pci_dev.irq_poller.poll(timeout))): diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index 38d5c2831d..ed27debcd9 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -222,7 +222,7 @@ def exec_hcq(ctx:ExecContext, call:UOp, ast:UOp) -> float|None: st = time.perf_counter() for d in call.arg.aux.device: with track_stats(ctx, call, d, [], ctx.var_vals): - if ctx.wait: Device[d].synchronize() + if ctx.wait: cast(Any, Device[d]).synchronize(timeout=ctx.timeout) return time.perf_counter() - st # flatten LINEAR-in-LINEAR: any nested LINEAR child gets inlined into its parent's src diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index e01dab947f..2b55cc07f6 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -285,21 +285,26 @@ def make_addr_table(call:UOp, gaddrs:list[UOp], name:str) -> tuple[UOp, dict[UOp fills = (table.after(*make_patches(table, [(i*table.dtype.itemsize, addr) for addr, i in slots.items()])),) if slots else () return table, reads, fills, {g:slots[bare[g]] for g in gaddrs} -def make_scatter_loop(patches:list[UOp], inputs_table:tuple, lt_patches:list[UOp]) -> dict[UOp, UOp]: - (table, _, _, slots), dst, data, subs = inputs_table, patches[0].buf_uop, [], {} - for p in patches: - words = [(off, val, get_getaddrs(val)) for off,val in zip(p.src[0].src[1].src, p.src[1].src)] - data += [off.val << 32 | slots[gaddrs[0]] for off,_,gaddrs in words if gaddrs][::2] - scalars = [(off.val*dst.dtype.itemsize, val) for off,val,gaddrs in words if not gaddrs] - subs[p] = UOp.group(*make_patches(dst, scalars)) if scalars else UOp(Ops.NOOP) +def is_bare_addr(val:UOp) -> bool: return val.op is Ops.CAST and val.src[0].op in (Ops.AND, Ops.SHR) and val.src[0].src[0].op is Ops.GETADDR - # plan entry: dst word offset << 32 | addr table slot - plan = UOp.placeholder((len(data),), dtypes.uint64, next(UOp.unique_num), device=dst.device).rtag("systems") - entry = plan.index(ridx:=UOp.range(len(data), next(UOp.unique_num), dtype=dtypes.int, src=(plan, dst))).load() - slot, widx = ((entry & 0xffffffff) % table.max_numel()).cast(dtypes.int), ((entry >> 32) % (dst.max_numel()-1)).cast(dtypes.int) # CHECK_OOB bounds - loop = UOp.group(*[dst.index(widx+i).store((table.index(slot).load() >> 32*i).cast(dtypes.uint32)) for i in range(2)]).end(ridx) - lt_patches.append(make_binary_patch(plan, struct.pack(f'<{len(data)}Q', *data))) - subs[patches[0]] = UOp.group(loop, subs[patches[0]]) +def make_scatter_loops(patches:list[UOp], inputs_table:tuple, lt_patches:list[UOp]) -> dict[UOp, UOp]: + table, _, _, slots = inputs_table + subs, by_dst = {}, collections.defaultdict(list) + for p in patches: by_dst[p.buf_uop].append(p) + for dst, patches in by_dst.items(): + data = [] + for p in patches: + words = [(off, val, get_getaddrs(val)) for off,val in zip(p.src[0].src[1].src, p.src[1].src)] + data += [(off.val, slots[gaddrs[0]]) for off,_,gaddrs in words if gaddrs][::2] + scalars = [(off.val*dst.dtype.itemsize, val) for off,val,gaddrs in words if not gaddrs] + subs[p] = UOp.group(*make_patches(dst, scalars)) if scalars else UOp(Ops.NOOP) + + word_table, slot_table = (UOp.placeholder((len(data),), dtypes.uint32, next(UOp.unique_num), device=dst.device).rtag("systems") for _ in range(2)) + ridx = UOp.range(len(data), next(UOp.unique_num), dtype=dtypes.int, src=(word_table, slot_table, dst)) + widx, slot = ((p.index(ridx).load() % bound).cast(dtypes.int) for p,bound in ((word_table, dst.max_numel()-1), (slot_table, table.max_numel()))) + loop = UOp.group(*[dst.index(widx+i).store((table.index(slot).load() >> 32*i).cast(dtypes.uint32)) for i in range(2)]).end(ridx) + lt_patches += [make_binary_patch(buf, struct.pack(f'<{len(data)}I', *vals)) for buf,vals in zip((word_table, slot_table), zip(*data))] + subs[patches[0]] = UOp.group(loop, subs[patches[0]]) return subs def is_input_addr(g:UOp) -> bool: return all(x.op is Ops.PARAM and x.tag is None for x in unwrap_mstack(g.buf_uop)) @@ -314,8 +319,9 @@ def split_patches(call:UOp) -> UOp|None: runtimes, systems = partition(internals, lambda g: any(x.tag in {"program", "kernargs", "cmdbuf"} for x in unwrap_mstack(g.buf_uop))) tables = [make_addr_table(call, gs, n) for gs,n in ((inputs, "inputs"), (runtimes, "runtime"), (systems, "systems"))] reads, fills = {k:v for _,r,_,_ in tables for k,v in r.items()}, [f for t in tables[1:] for f in t[2]] # inputs table is filled by exec - input_patches = [p for p in rt_patches if (gs:=get_getaddrs(p)) and all(map(is_input_addr, gs))] - scatter = make_scatter_loop(input_patches, tables[0], lt_patches) if input_patches else {} + input_patches = [p for p in rt_patches if (gs:=get_getaddrs(p)) and all(map(is_input_addr, gs)) + and all(is_bare_addr(v) for v in p.src[1].src if get_getaddrs(v))] + scatter = make_scatter_loops(input_patches, tables[0], lt_patches) body = body.substitute({p:p.substitute(scatter | reads) for p in rt_patches}) lt_srcs = collections.defaultdict(list) @@ -495,6 +501,7 @@ class HCQ2Compiled(Compiled): def __init__(self, device:str, allocator:HCQAllocator, compilers:list[type[Renderer]], runtime, can_recover:bool=False, arch=None): self.device_id:int = int(device.split(":")[1]) if ":" in device else 0 + self.can_recover = can_recover self.pm_bufferize = PatternMatcher([ (UPat(Ops.PARAM, tag="sentinel_signal"), lambda ctx: ctx[0].signal("sentinel", (1 << 64) - 1)), @@ -524,6 +531,7 @@ class HCQ2Compiled(Compiled): if not hasattr(self, 'iface'): return sig = self.signal("timeline").as_memoryview(force_zero_copy=True, no_sync=True).cast('Q') tl = self.signal("value", 1).as_memoryview(force_zero_copy=True, no_sync=True).cast('Q') + timeout = timeout if timeout is not None and self.can_recover else None st = time.perf_counter() while sig[0] < tl[0] - 1: if time.perf_counter() - st > (timeout or 3000) / 1000: self.on_device_hang()