From a2e64e16aa1845b06991d194c81a67f211b89790 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Sun, 23 Aug 2026 00:22:39 +0300 Subject: [PATCH] hcq2: early usb (#17683) * hcq2: usb interface and submit * x * x * x * x * r * x --- extra/hcq2/ops_amd2.py | 85 ++++++++++++++++++----- tinygrad/runtime/ops_cpu.py | 8 +-- tinygrad/runtime/support/hcq2.py | 75 +++++++++++++-------- tinygrad/runtime/support/usb.py | 112 ++++++++++++++++++++++++++++++- tinygrad/uop/ops.py | 4 +- 5 files changed, 231 insertions(+), 53 deletions(-) diff --git a/extra/hcq2/ops_amd2.py b/extra/hcq2/ops_amd2.py index e603df20a4..c4f7b6bc3b 100644 --- a/extra/hcq2/ops_amd2.py +++ b/extra/hcq2/ops_amd2.py @@ -19,7 +19,7 @@ from tinygrad.runtime.support.hcq import FileIOInterface, HCQBuffer, MMIOInterfa from tinygrad.runtime.support.am.amdev import AMDev, AMMemoryManager from tinygrad.runtime.support.amd import AMDReg, AMDIP, import_module, import_soc, import_pmc from tinygrad.runtime.support.system import PCIIfaceBase, PCIAllocationMeta, USBPCIDevice, MAP_FIXED, MAP_NORESERVE -from tinygrad.runtime.support.usb import USB3 +from tinygrad.runtime.support.usb import USB3, usb_ib, usb_push, usb_arm_bytes, pm_usb_stage, pm_usb_hostio, pm_usb_bufferize from tinygrad.runtime.support.memory import AddrSpace, BumpAllocator from tinygrad.runtime.ops_amd import SQTT, SQTT_ITRACE_SE_MASK, SQTT_LIMIT_SE, SQTT_SIMD_SEL, SQTT_TOKEN_EXCLUDE, PMC from tinygrad.runtime.ops_amd import EVENT_INDEX_PARTIAL_FLUSH, WAIT_REG_MEM_FUNCTION_EQ, WAIT_REG_MEM_FUNCTION_NEQ, WAIT_REG_MEM_FUNCTION_GEQ @@ -146,11 +146,14 @@ pm_pm4_opsel = PatternMatcher([ (UPat(Ops.INS, arg="store", src=(UPat((Ops.BUFFER, Ops.PARAM), name="dst"), UPat(name="val"))), pm4_store), ]) +def queue_ptrs(devs, qname:str, q:AMDQueueDesc) -> tuple[UOp, ...]: + return tuple(UOp.placeholder((b.size,), b.dtype, 0, device=devs).rtag(f"{qname}_{n}") + for n, b in (("ring", q.ring), ("write_ptr", q.write_ptr), ("doorbell", q.doorbell), ("put_value", q.put_value))) + def pm4_submit(ctx, lin): # ensure compute queues are allocated for d in (devs:=ctx.devs): q = Device[d].compute_queue - ring, wptr, doorbell, put_ptr = (UOp.placeholder((b.size,), b.dtype, 0, device=devs).rtag(f"COMPUTE:0_{name}") - for name, b in (("ring", q.ring), ("write_ptr", q.write_ptr), ("doorbell", q.doorbell), ("put_value", q.put_value))) + ring, wptr, doorbell, put_ptr = queue_ptrs(devs, "COMPUTE:0", q) # the host fence at the start of the batch guarantees the ib is free to reuse size_dw = sum(len(ins.src) for ins in lin.src) @@ -216,8 +219,7 @@ def sdma_submit(cmdbuf, devs): # the sdma queue's ring and its host-side ring/write/put pointers for d in devs: q = Device[d].sdma_queue(0) - ring, wptr, doorbell, put_ptr = (UOp.placeholder((b.size,), b.dtype, 0, device=devs).rtag(f"COPY:0_{name}") - for name, b in (("ring", q.ring), ("write_ptr", q.write_ptr), ("doorbell", q.doorbell), ("put_value", q.put_value))) + ring, wptr, doorbell, put_ptr = queue_ptrs(devs, "COPY:0", q) # sdma needs the cmdbuf contiguous: if it won't fit before the ring end, restart at 0 and zero the tail put_b = put_ptr.index(zero) @@ -244,15 +246,32 @@ def sdma_submit(cmdbuf, devs): pm_sdma_submit = PatternMatcher([(UPat(Ops.LINEAR, name="lin"), lambda ctx, lin: sdma_submit(make_cmdbuf(lin, ctx.devs), ctx.devs))]) +# ***************** +# USB submit + +def amd_usb_submit(ctx, lin): + for d in ctx.devs: q = Device[d].compute_queue if (comp:=ctx.qname.startswith("COMPUTE")) else Device[d].sdma_queue(0) + + if nb:=usb_arm_bytes(ctx.pre, Device[ctx.devs[0]].iface.usb_sram): + poke = (ctx.sdma.SDMA_OP_WRITE, *data64_le(Device[ctx.devs[0]].iface.cq_buf.va_addr + 12), 0, 0) + lin = lin.replace(src=lin.src + (UOp(Ops.INS, arg="poke", src=tuple(UOp.const(x, dtypes.uint32) for x in poke)),)) + + ib_host, ib_gpu, pkt_dw = usb_ib(ctx.devs, lin, 32 if comp else 0x100, nb) + pkt = (ctx.pm4.PACKET3(ctx.pm4.PACKET3_INDIRECT_BUFFER,2),*data64_le(ib_gpu.getaddr(ctx.devs)),pkt_dw|ctx.pm4.INDIRECT_BUFFER_VALID) if comp else () + return usb_push(ctx.devs, *queue_ptrs(ctx.devs, ctx.qname, q), ib_host, ib_gpu, pkt, 4 if comp else 1) + +pm_usb_submit = PatternMatcher([(UPat(Ops.LINEAR, name="lin"), amd_usb_submit)]) + @dataclass(frozen=True) class AMDEncodeCtx: # encode-time constants for one queue: devs (every cmdbuf address resolves into these) + gfx version + packet/ip modules devs: tuple[str, ...]; target: tuple[int, ...]; pm4: Any; sdma: Any; soc: Any # noqa: E702 - gc: AMDIP; nbio: AMDIP; xccs: int; max_copy_size: int; tmpring_size: Callable # noqa: E702 + gc: AMDIP; nbio: AMDIP; xccs: int; max_copy_size: int; tmpring_size: Callable; qname: str; pre: UOp # pre: the queue before opsel def encode_queue(q:UOp) -> UOp|None: d = Device[(devs:=to_tuple(q.arg[0]))[0]] - ctx = AMDEncodeCtx(devs, d.target, d.pm4, d.sdma, d.soc, d.gc, d.nbio, d.xccs, d.max_copy_size, d.tmpring_size) - opsel, submit = (pm_pm4_opsel, pm_pm4_submit) if q.arg[1].startswith("COMPUTE") else (pm_sdma_opsel, pm_sdma_submit) + ctx = AMDEncodeCtx(devs, d.target, d.pm4, d.sdma, d.soc, d.gc, d.nbio, d.xccs, d.max_copy_size, d.tmpring_size, q.arg[1], q) + opsel = pm_pm4_opsel if (comp:=q.arg[1].startswith("COMPUTE")) else pm_sdma_opsel + submit = d.pm_submit if d.pm_submit is not None else (pm_pm4_submit if comp else pm_sdma_submit) return submit.rewrite(graph_rewrite(q, opsel + pm_flatten_linear, walk=True, ctx=ctx, name=f"{q.arg[1]} opsel"), ctx) @dataclass(frozen=True) @@ -282,13 +301,14 @@ def amd_build_program(prg:UOp) -> UOp: wave32=bool(desc.kernel_code_properties & 0x400), private_segment_size=desc.private_segment_fixed_size, kernargs_segment_size=desc.kernarg_size, kernargs_alloc_size=desc.kernarg_size + (ctypes.sizeof(hsa.hsa_kernel_dispatch_packet_t) if edp else 0), enable_dispatch_ptr=edp, enable_private_segment_sgpr=desc.kernel_code_properties & hsa.AMD_KERNEL_CODE_PROPERTIES_ENABLE_SGPR_PRIVATE_SEGMENT_BUFFER) + image = bytes(image).ljust(round_up(len(image), 4), b"\x00") # the program is uploaded as whole dwords buf = UOp.placeholder((len(image),), dtypes.uint8, next(UOp.unique_num), device=prg.device).rtag("program") - cached = _amd_program_cache[key] = prg.replace(src=(buf.after(make_binary_patch(buf, bytes(image))),), arg=(data, prg.arg)) + cached = _amd_program_cache[key] = prg.replace(src=(buf.after(make_binary_patch(buf, image)),), arg=(data, prg.arg)) return cached class AMDAllocator(HCQAllocator['AMDDevice']): def __init__(self, dev:AMDDevice): - super().__init__(dev, supports_copy_from_disk=dev.has_copy_queue, supports_transfer=dev.has_copy_queue and not dev.is_usb()) + super().__init__(dev, supports_copy_from_disk=dev.has_copy_queue, supports_transfer=dev.has_copy_queue and not dev.is_usb) def _alloc(self, size:int, options:BufferSpec) -> HCQBuffer: return self.dev.iface.alloc(size, host=options.host, uncached=options.uncached, cpu_access=options.cpu_access or not self.dev.has_copy_queue) @@ -524,8 +544,7 @@ class PCIIface(PCIIfaceBase): 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, no_sync=True).cast('Q')[0] - 1 + d.signal('timeline')._buf.cpu_view().view(fmt='Q')[0] = d.signal('value', 1, device="CPU")._buf.cpu_view().view(fmt='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))): @@ -539,6 +558,32 @@ class PCIIface(PCIIfaceBase): def device_fini(self): self.dev_impl.fini() +class USBIface(PCIIface): + def __init__(self, dev, dev_id): # pylint: disable=super-init-not-called + if dev_id >= len(visible:=hcq_filter_visible_devices(USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001), "AMD")): + raise RuntimeError(f"AMD:{dev_id} does not exist ({pluralize('device', len(visible))} available)") + self.dev, self.pci_dev, self.vram_bar, self.count = dev, USBPCIDevice("AM", *visible[dev_id]), 0, len(visible) + self.dev_impl = AMDev(self.pci_dev) + self._compute_props() + self.sram = self._dma_region(ctrl_addr=0xf000, sys_addr=0x200000, size=0x80000) + self.cq_buf = self._dma_region(ctrl_addr=0xb800, sys_addr=0x822000, size=0x1000) # +12 is the dword that releases an armed read + self.usb_handle = unwrap(ctypes.cast(self.pci_dev.usb.usb.handle, ctypes.c_void_p).value) + + def _dma_region(self, ctrl_addr, sys_addr, size): + region = self.dev_impl.mm.map_range(vaddr:=self.dev_impl.mm.alloc_vaddr(size=size), size, [(sys_addr, size)], aspace=AddrSpace.SYS, uncached=True) + return HCQBuffer(vaddr, size, meta=PCIAllocationMeta(region, has_cpu_mapping=False), view=self.pci_dev.dma_view(ctrl_addr, size), owner=self.dev) + + def alloc(self, size:int, host=False, uncached=False, cpu_access=False, contiguous=False, force_devmem=False, **kwargs) -> HCQBuffer: + # everything, even host-style signals, lives in vram: gpu writes into the bridge's own memory collide with an armed 0xF2 read stream + return super().alloc(size, host=False, uncached=uncached, cpu_access=cpu_access or host, contiguous=contiguous, force_devmem=True, **kwargs) + + def sleep(self, timeout): pass + + # we don't own the sram region, so the buffer never frees it + @functools.cached_property + def usb_sram(self) -> Buffer: + return Buffer(self.dev.device, (b:=self.sram).size, dtypes.uint8, options=BufferSpec(external_ptr=b.va_addr, nolru=True)).allocate(opaque=b) + def _mock(iface, name=None): return type(name or f"MOCK{iface.__name__}", (iface,), {}) class AMDDevice(HCQ2Compiled): @@ -549,19 +594,21 @@ class AMDDevice(HCQ2Compiled): # encoding of cmdbuf (UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="q"),)), encode_queue), ]) + pm_submit: PatternMatcher|None = None timestamp_divider = 100.0 # AMD GPU clock: ticks/us max_scratch_psize = 0 - ifaces = [KFDIface, PCIIface, _mock(KFDIface, "MOCKIface"), _mock(KFDIface), _mock(PCIIface)] + ifaces = [KFDIface, PCIIface, USBIface, _mock(KFDIface, "MOCKIface"), _mock(KFDIface), _mock(PCIIface), _mock(USBIface)] def device_props(self): return self.iface.props def is_am(self) -> bool: return isinstance(self.iface, (PCIIface,)) - def is_usb(self) -> bool: return False def __init__(self, device:str=""): self.iface = self._select_iface(device) + self.is_usb = isinstance(self.iface, USBIface) + if self.is_usb: self.rt_nbytes = 4 << 20 self.target:tuple[int, ...] = ((trgt:=self.iface.props['gfx_target_version']) // 10000, (trgt // 100) % 100, trgt % 100) self.arch = "gfx%d%x%x" % self.target @@ -586,7 +633,7 @@ class AMDDevice(HCQ2Compiled): self.is_aql = getenv("AMD_AQL", int(self.xccs > 1)) if self.is_aql: - self.pm4_ibs = self.iface.alloc(0x2000 if self.is_usb() else (16 << 20), uncached=True, cpu_access=True) + self.pm4_ibs = self.iface.alloc(0x2000 if self.is_usb else (16 << 20), uncached=True, cpu_access=True) self.pm4_ib_alloc = BumpAllocator(self.pm4_ibs.size, wrap=True) self.max_copy_size = 0x40000000 if self.iface.ip_versions[am.SDMA0_HWIP][0] >= 5 else 0x400000 @@ -599,6 +646,10 @@ class AMDDevice(HCQ2Compiled): self.max_private_segment_size = 0 self.pm_bufferize = PatternMatcher([(UPat(Ops.PARAM, tag="scratch", name="b"), lambda ctx, b: ctx[0].scratch_buffer(b.max_numel()))]) + self.pm_bufferize + if self.is_usb: + self.pm_bufferize = pm_usb_bufferize + self.pm_bufferize + self.pm_stage_copy, self.pm_host_lower, self.pm_submit = pm_usb_stage, pm_usb_hostio, pm_usb_submit + self.pmc_enabled:bool = PROFILE > 0 and PMC > 0 if self.pmc_enabled: self.iface.require_profile_mode() @@ -659,7 +710,7 @@ class AMDDevice(HCQ2Compiled): wg_data_size = round_up((vgpr_size_per_cu + sgrp_size_per_cu + lds_size_per_cu + hwreg_size_per_cu) * self.cu_cnt, mmap.PAGESIZE) ctl_stack_size = round_up((12 if self.target[0] != 9 else 8) * self.wave_cnt + 8 + 40, mmap.PAGESIZE) return self.create_queue(kfd.KFD_IOC_QUEUE_TYPE_COMPUTE_AQL if self.is_aql else kfd.KFD_IOC_QUEUE_TYPE_COMPUTE, - 0x2000 if self.is_usb() else (16 << 20), eop_buffer_size=0x1000, + 0x2000 if self.is_usb else (16 << 20), eop_buffer_size=0x1000, ctx_save_restore_size=0 if self.is_am() else wg_data_size + ctl_stack_size, ctl_stack_size=ctl_stack_size, debug_memory_size=round_up(self.wave_cnt * 32, 64)) @@ -667,7 +718,7 @@ class AMDDevice(HCQ2Compiled): if getenv("AMD_DISABLE_SDMA"): return None if idx in self.sdma_queues: return self.sdma_queues[idx] with contextlib.suppress(OSError): - self.sdma_queues[idx] = self.create_queue(kfd.KFD_IOC_QUEUE_TYPE_SDMA, 0x200 if self.is_usb() else (16 << 20), idx=idx) + self.sdma_queues[idx] = self.create_queue(kfd.KFD_IOC_QUEUE_TYPE_SDMA, 0x2000 if self.is_usb else (16 << 20), idx=idx) return self.sdma_queues.get(idx, None) def tmpring_size(self, private_segment_size): diff --git a/tinygrad/runtime/ops_cpu.py b/tinygrad/runtime/ops_cpu.py index 28f2165571..1e1293733f 100644 --- a/tinygrad/runtime/ops_cpu.py +++ b/tinygrad/runtime/ops_cpu.py @@ -5,7 +5,7 @@ from typing import cast, Callable from tinygrad.helpers import to_mv, from_mv, OSX, WIN, Context, mv_address, suppress_finalizing, unwrap, data64_le, to_tuple from tinygrad.device import Buffer, BufferSpec, TinyELF, Program, Device from tinygrad.runtime.support.hcq import HCQBuffer, MMIOInterface -from tinygrad.runtime.support.hcq2 import HCQ2Compiled, HCQAllocator, make_cmdbuf, make_signal +from tinygrad.runtime.support.hcq2 import HCQ2Compiled, HCQAllocator, make_cmdbuf, make_buf from tinygrad.runtime.support.c import DLL from tinygrad.renderer.cstyle import ClangRenderer from tinygrad.renderer.llvmir import CPULLVMRenderer @@ -80,7 +80,7 @@ pm_cpu_opsel = PatternMatcher([ (UPat(Ops.INS, arg="store", src=(UPat((Ops.BUFFER, Ops.PARAM), name="dst"), UPat(name="val"))), lambda ctx, dst, val: cpu_cmd(ctx, signal_prog, dst.getaddr(ctx), val.cast(dtypes.uint64))), (UPat(Ops.INS, arg="timestamp", src=(UPat(name="dst"),)), - lambda ctx, dst: cpu_cmd(ctx, timestamp_prog, dst.getaddr(ctx), *(() if WIN else (make_signal(ctx, tag="func:clock_gettime").getaddr(ctx),)))), + lambda ctx, dst: cpu_cmd(ctx, timestamp_prog, dst.getaddr(ctx), *(() if WIN else (make_buf(ctx, tag="func:clock_gettime").getaddr(ctx),)))), ]) def encode_queue(q:UOp) -> UOp: @@ -91,7 +91,7 @@ def encode_queue(q:UOp) -> UOp: assert cnt < RING_SLOTS, f"submit of {cnt} entries doesn't fit the ring" cmdbuf = make_cmdbuf(lin, devs, buf=UOp.placeholder((cnt*CMD_SIZE,), dtypes.uint64, next(UOp.unique_num), device=devs).rtag("cmdbuf")) ring = UOp.placeholder((ring_words:=RING_SLOTS*CMD_SIZE,), dtypes.uint64, 0, device=devs, volatile=True).rtag(f"{queue}_ring") - put, done, sem, sysbuf = (make_signal(devs, tag=f"{queue}_{name}") for name in ("put", "done", "sem", "sys")) + put, done, sem, sysbuf = (make_buf(devs, tag=f"{queue}_{name}") for name in ("put", "done", "sem", "sys")) # submits are serialized on the submitter, so they can bump put without atomics ran = done.after(l:=UOp.loop(next(UOp.unique_num))).index(0).load() @@ -104,7 +104,7 @@ def encode_queue(q:UOp) -> UOp: if WIN: return sysbuf.after(bumped).index(0).store(put.after(bumped).index(0).load()) e = UOp.range(cnt, next(UOp.unique_num), dtype=dtypes.int, src=(bumped,)) - return make_signal(devs, tag="func:sem_post").after(e).index(0).load().call(sem.after(e).index(0), ret_dtype=dtypes.void).end(e) + return make_buf(devs, tag="func:sem_post").after(e).index(0).load().call(sem.after(e).index(0), ret_dtype=dtypes.void).end(e) # ***************** diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index d985becfc6..4decf65da5 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -68,16 +68,15 @@ def make_binary_patch(buf:UOp, blob:bytes) -> UOp: r = UOp.range(len(blob) // buf.dtype.itemsize, 0, dtype=dtypes.int, src=(buf, data)) return buf.index(r).store(data.index(r).load()).end(r).rtag("link") -def make_cmdbuf(lin, devs, buf:UOp|None=None): +def make_buf(devs, slot:int=0, tag:str="signal") -> UOp: return UOp.placeholder((1,), dtypes.uint64, slot, device=devs, volatile=True, tag=tag) + +def make_cmdbuf(lin, devs, buf:UOp|None=None, dep:tuple[UOp, ...]=()): blob, patches = bytearray(), [] for s in (s for ins in lin.src for s in ins.src): if s.op is not Ops.CONST: patches.append((len(blob), s)) blob.extend(struct.pack(f'<{s.dtype.fmt}', s.val if s.op is Ops.CONST else 0x0)) cmdbuf = buf if buf is not None else UOp.placeholder((len(blob) // 4,), dtypes.uint32, next(UOp.unique_num), device=devs).rtag("cmdbuf") - return cmdbuf.after(make_binary_patch(cmdbuf, bytes(blob)), *make_patches(cmdbuf, patches)) - -def make_signal(devs, slot:int=0, tag:str="signal") -> UOp: - return UOp.placeholder((1,), dtypes.uint64, slot, device=devs, volatile=True).rtag(tag) + return cmdbuf.after(*dep, make_binary_patch(cmdbuf, bytes(blob)), *make_patches(cmdbuf, patches)) def make_submit(*cmds, devs:str|tuple[str, ...], queue:str) -> UOp: return UOp.custom_function("submit_cmdbuf", UOp(Ops.LINEAR, src=tuple(cmds), arg=(to_tuple(devs), queue))) @@ -101,6 +100,17 @@ def replace_call_buffers(ctx:tuple[list[UOp], dict[UOp, int]], call:UOp) -> UOp| return call.replace(src=call.src[:1] + tuple(s if s.op is Ops.PARAM or s.is_bound_var else s.param_like(slots[s]) for s in call.src[1:])) pm_replace_buffers = PatternMatcher([(UPat(Ops.CALL, name="call"), replace_call_buffers)]) +# ***************** + +def stage_copy_ext(call:UOp) -> UOp|None: + if (d:=next((d for b in call.src[1:] for d in to_tuple(b.device) if not d.startswith("CPU")), None)) is None: return None + return pm.rewrite(call) if (pm:=getattr(Device[d], "pm_stage_copy", None)) is not None else None + +def encode_host_call(call:UOp) -> UOp|None: + if (pm:=getattr(Device[call.arg.aux.device[0]], "pm_host_lower", None)) is None: return None + body = graph_rewrite(call.src[0], pm, name="lower host access", enter_calls=True) + return None if body is call.src[0] else call.replace(src=(body, *call.src[1:])) + # ***************** # 1.1. prep: staging copies @@ -126,6 +136,8 @@ def stage_copy(dst:UOp, src:UOp) -> UOp|None: # 1.2. prep: kernel copies def _get_enqueue_devs(call:UOp) -> Any|None: + if (call.arg.name or "").startswith("hcq_"): return None # host exec is not any device + if not (bufs:=call.src[1:]) or not all(all_devices_in(b.device, HCQ_DEVS) for b in bufs): return None if call.src[0].op is Ops.COPY: bufs = bufs[::-1] # copies push from the src device: p2p writes are faster than reads devs = min(bufs, key=lambda b: to_tuple(b.device)[0].startswith("CPU")).device # prio to enqueue on not CPU device @@ -138,6 +150,7 @@ def kernel_copy(call:UOp, dst:UOp, src:UOp) -> UOp|None: return call.replace(src=(to_program(ast, Device[dev].renderer), dst, src)) pm_insert_copy_staging = PatternMatcher([ + (UPat(Ops.CALL, src=(UPat(Ops.COPY),), name="call", allow_any_len=True), stage_copy_ext), (UPat(Ops.CALL, src=(UPat(Ops.COPY), UPat(name="dst"), UPat(name="src"))), stage_copy), (UPat(Ops.CALL, src=(UPat(Ops.COPY), UPat(name="dst"), UPat(name="src")), name="call"), kernel_copy) ]) @@ -174,7 +187,7 @@ def _build_wait_cmds(slots:dict[str, int], dep_lanes:list[tuple[tuple, int, int] waits = [] for (ddevs, dqueue, dtag), by_lane in deps.items(): for ls in itertools.zip_longest(*(by_lane[lane] for lane in range(len(devices)))): - s = UOp.mstack(*[make_signal(d, tag="sentinel_signal") if dl is None else make_signal(ddevs[dl], slots[dqueue]) for dl, d in zip(ls, devices)]) + s = UOp.mstack(*[make_buf(d, tag="sentinel_signal") if dl is None else make_buf(ddevs[dl], slots[dqueue]) for dl, d in zip(ls, devices)]) waits.append(UOp(Ops.INS, arg="wait", src=(s, UOp.const(dtag + 1, dtypes.uint64)))) return waits, {dtag for _, _, dtag in deps} @@ -196,19 +209,20 @@ def _build_finalizers(batch:list[tuple[UOp, tuple[str, ...]]], batch_info:list[t signal_tags |= cur_signal_tags # wait the syncs and signal the device epoch, then bump the timeline on the host - tl_signal, tl_value = make_signal(devs, tag="timeline_signal"), make_signal(devs, tag="timeline_value") + tl_signal, tl_value = make_buf(devs, tag="timeline_signal"), make_buf(devs, tag="timeline_value") fin_submit = make_submit(*waits, UOp(Ops.INS, arg="store", src=(tl_signal, tl_value.index(0))), devs=devs, queue="COMPUTE:0") epoch = (epoch_slot:=tl_value.after(fin_submit).index(0)).load() # fence once per device group on this schedule's previous epoch qs = dedup([qn for bdevs, qn in batch_info if set(bdevs) & set(devs)]) - sched_epoch = make_signal(devs, next(UOp.unique_num)) + sched_epoch = make_buf(devs, next(UOp.unique_num), tag="epoch") wait_device_epoch = (done:=tl_signal.after(loop:=UOp.loop(0)).index(0).load()).end(loop, done < sched_epoch.index(0).load()) fences.append(make_call("hcq_fence", UOp.sink(wait_device_epoch), HCQInfo(devs))) - # queues of other groups wait on these signals, so reset them only after every group reached its epoch - if qs: resets.append(make_call("hcq_reset", UOp.sink(*[make_signal(devs, slots[q]).index(0).store(0) for q in qs]), HCQInfo(devs))) + # queues of other groups wait on these signals, reset them after every group reached its epoch + rst = functools.reduce(lambda a,q: a+(make_buf(devs, slots[q]).after(*a[-1:]).index(0).store(0),), qs, cast(tuple[UOp, ...], ())) + if rst: resets.append(make_call("hcq_reset", UOp.sink(*rst), HCQInfo(devs))) fins.append(make_call("hcq_finalizer", UOp.sink(epoch_slot.store(epoch + 1), sched_epoch.after(fin_submit).index(0).store(epoch)), HCQInfo(devs))) return fences + resets, fins, signal_tags @@ -234,19 +248,19 @@ def _finalize_batch(batch:list[tuple[UOp, tuple[str, ...]]], profile:bool) -> li for tag, ((call, _), (devices, queue), q) in enumerate(zip(batch, batch_info, call_waits)): # first queue use, sync prior device work with the device timeline if batch_info.index((devices, queue)) == tag: - epoch = make_signal(devices, tag="timeline_value").index(0) - 1 - q = [UOp(Ops.INS, arg="barrier", src=()), UOp(Ops.INS, arg="wait", src=(make_signal(devices, tag="timeline_signal"), epoch))] + q + epoch = make_buf(devices, tag="timeline_value").index(0) - 1 + q = [UOp(Ops.INS, arg="barrier", src=()), UOp(Ops.INS, arg="wait", src=(make_buf(devices, tag="timeline_signal"), epoch))] + q # and make hcq call name, info = get_call_name(call, get_call_arg_uops(call)), HCQInfo(devices, estimate_uop(call)) ts_ids = [next(UOp.unique_num) for _ in range(2)] if profile else [] kerns.append((devices, make_call(name, call.src[0], info), tuple(ts_ids))) - ts_ins = [UOp(Ops.INS, arg="timestamp", src=(make_signal(devices, s),)) for s in ts_ids] + ts_ins = [UOp(Ops.INS, arg="timestamp", src=(make_buf(devices, s),)) for s in ts_ids] q += ts_ins[:1] + [call.replace(arg=replace(call.arg, aux=info))] + ts_ins[1:] # signal the queue if someone waits for us - if tag in signal_tags: q += [UOp(Ops.INS, arg="store", src=(make_signal(devices, slots[queue]), UOp.const(tag + 1, dtypes.uint64)))] + if tag in signal_tags: q += [UOp(Ops.INS, arg="store", src=(make_buf(devices, slots[queue]), UOp.const(tag + 1, dtypes.uint64)))] src.append(make_call(f"submit {name}", make_submit(*q, devs=devices, queue=queue).sink(), info)) # append batch timestamps to finalizers @@ -304,7 +318,8 @@ def encode_cmdbuf(submit:UOp, lin:UOp) -> UOp|None: if (pm:=Device.get_class(lin.arg[0][0]).pm_lower) is None: return None return graph_rewrite(submit, pm, name=f"encode {lin.arg[0]}", enter_calls=True) pm_encode_cmdbufs = PatternMatcher([ - (UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="lin"),), name="submit"), encode_cmdbuf)]) + (UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="lin"),), name="submit"), encode_cmdbuf), + (UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), encode_host_call)]) # ***************** @@ -350,20 +365,20 @@ def split_patches(call:UOp) -> UOp|None: lt_patches:list[UOp] = [] body = graph_rewrite(call.src[0], pm_trim_link_patches, ctx=(rt_patches, lt_patches), name=f"trim link-time patches ({call.arg.name})") - # split patches - inputs, internals = partition(dedup(g for p in rt_patches for g in get_getaddrs(p)), is_input_addr) + # split patches. addresses read in the body go through the tables too + inputs, internals = partition(dedup([g for p in rt_patches for g in get_getaddrs(p)] + get_getaddrs(body)), is_input_addr) 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 ipatches = [p for p in rt_patches if p.tag == "inputs" and all(v in tables[0][3] for v in p.src[1].src)] # only getaddrs go to the table gathers = make_gather_loop(ipatches, tables[0][0], tables[0][3], lt_patches) if ipatches else {} - body = body.substitute({p:p.substitute(gathers | reads) for p in rt_patches}) + body = body.substitute({p:p.substitute(gathers | reads) for p in rt_patches}).substitute(reads) lt_srcs = collections.defaultdict(list) for p in lt_patches: lt_srcs[p.buf_uop].append(p) return call.replace(src=(body, *call.src[1:], *[b.after(*ps) for b,ps in lt_srcs.items()], *fills), - arg=replace(call.arg, aux=replace(call.arg.aux, input_idxs=((call.arg.aux.device, + arg=replace(call.arg, aux=replace(call.arg.aux, input_idxs=((to_tuple(inputs[0].arg), tuple(sorted(dedup(b.arg.slot for g in inputs for b in unwrap_mstack(g.buf_uop))))),) if inputs else call.arg.aux.input_idxs))) pm_split_patches = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), split_patches)]) @@ -500,7 +515,7 @@ def push_stack(op, s): return UOp(Ops.STACK, def fold_binary(buf:UOp, blob:UOp) -> UOp: for b in (m.bufs if isinstance(m:=buf.buffer, MultiBuffer) else (m,)): - b.ensure_allocated().as_memoryview(force_zero_copy=True, no_sync=True).cast('B')[:len(blob.arg)] = blob.arg + b.ensure_allocated()._buf.cpu_view().view(fmt='B')[:len(blob.arg)] = blob.arg return UOp(Ops.NOOP) def fold_const_store(view:UOp, off:UOp, val:UOp) -> UOp: @@ -509,7 +524,7 @@ def fold_const_store(view:UOp, off:UOp, val:UOp) -> UOp: for b,v in zip((bs:=mb.bufs if isinstance((mb:=buf.buffer), MultiBuffer) else (mb,)), val.src if val.op is Ops.STACK else (val,)*len(bs)): data = struct.pack(f'<{v.dtype.fmt}', truncate[v.dtype]((v.src[0] if v.op is Ops.CAST else v).val)) bo = start*buf.dtype.itemsize + off.val*val.dtype.itemsize - b.ensure_allocated().as_memoryview(force_zero_copy=True, no_sync=True).cast('B')[bo:bo+len(data)] = data + b.ensure_allocated()._buf.cpu_view().view(fmt='B')[bo:bo+len(data)] = data return UOp(Ops.NOOP) def resolve_getaddr(buf:UOp, g:UOp) -> UOp: @@ -561,6 +576,7 @@ def hcq_link(linear:UOp, cache=True) -> UOp: class HCQ2Compiled(Compiled): timestamp_divider: float = 1000.0 wait_timeout_ms: float = 30000.0 + rt_nbytes: int = 64 << 20 # scratch that single-run placeholders are carved out of def __init__(self, device:str, allocator:HCQAllocator, compilers:list[type[Renderer]], runtime, can_recover:bool=False, arch=None): self.can_recover = can_recover @@ -568,14 +584,15 @@ class HCQ2Compiled(Compiled): self.pm_bufferize = PatternMatcher([ (UPat(Ops.PARAM, tag="sentinel_signal"), lambda ctx: ctx[0].signal("sentinel", (1 << 64) - 1)), (UPat(Ops.PARAM, tag="timeline_signal"), lambda ctx: ctx[0].signal("timeline")), - (UPat(Ops.PARAM, tag="timeline_value"), lambda ctx: ctx[0].signal("value", 1)), + (UPat(Ops.PARAM, tag="timeline_value"), lambda ctx: ctx[0].signal("value", 1, device="CPU")), + (UPat(Ops.PARAM, tag="epoch", name="b"), lambda ctx, b: ctx[0].signal(b.arg.slot, device="CPU")), (UPat(Ops.PARAM, tag="signal", name="b"), lambda ctx, b: ctx[0].signal(b.arg.slot)), (UPat(Ops.PARAM, name="b"), lambda ctx, b: None if b.tag is None else ctx[0].new_buffer(b, cache=ctx[1])) ]) super().__init__(device, allocator, compilers, runtime, None, arch=arch) - self.rt_allocator = BumpAllocator(64 << 20) + self.rt_allocator = BumpAllocator(self.rt_nbytes) self.prof_ents:dict[int, ProfileGraphEntry] = {} def collect_prof(self): @@ -609,12 +626,12 @@ class HCQ2Compiled(Compiled): self.rt_allocator.alloc(b.max_numel() * b.dtype.itemsize, alignment=128)) @functools.cache - def signal(self, name:str|int, init_value:int=0) -> Buffer: - buf = Buffer(self.device, 1, dtypes.uint64, options=BufferSpec(host=True, uncached=True, cpu_access=True), preallocate=True) - buf.as_memoryview(force_zero_copy=True, no_sync=True).cast('Q')[0] = init_value + def signal(self, name:str|int, init_value:int=0, device:str|None=None) -> Buffer: + buf = Buffer(device or self.device, 1, dtypes.uint64, options=BufferSpec(host=True, uncached=True, cpu_access=True), preallocate=True) + buf._buf.cpu_view().view(fmt='Q')[0] = init_value return buf - def _wait_signal(self, sig:memoryview, value:int, timeout:int|None=None): + def _wait_signal(self, sig:MMIOInterface|memoryview, value:int, timeout:int|None=None): timeout = timeout if timeout is not None and self.can_recover else None st, done = time.perf_counter(), sig[0] while done < value: @@ -624,8 +641,8 @@ class HCQ2Compiled(Compiled): def synchronize(self, timeout:int|None=None): if HCQ_RUNTIME_DEV.value != self.device: Device[HCQ_RUNTIME_DEV.value].synchronize() - 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') + sig = self.signal("timeline")._buf.cpu_view().view(fmt='Q') + tl = self.signal("value", 1, device="CPU")._buf.cpu_view().view(fmt='Q') self._wait_signal(sig, tl[0] - 1, timeout) if self.prof_ents: self.collect_prof() diff --git a/tinygrad/runtime/support/usb.py b/tinygrad/runtime/support/usb.py index 2229f4a5b8..18ce790309 100644 --- a/tinygrad/runtime/support/usb.py +++ b/tinygrad/runtime/support/usb.py @@ -1,6 +1,11 @@ import ctypes, struct, time, functools, itertools +from typing import Any, cast from tinygrad.runtime.autogen import libusb -from tinygrad.helpers import DEBUG, DEV, to_mv, from_mv, round_up, ceildiv +from tinygrad.helpers import DEBUG, DEV, to_mv, from_mv, round_up, ceildiv, unwrap, dedup, to_tuple +from tinygrad.dtype import dtypes +from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher +from tinygrad.device import Buffer, BufferSpec, Device +from tinygrad.runtime.support.hcq2 import HCQInfo, make_buf, make_cmdbuf, make_submit, HCQ_RUNTIME_DEV from tinygrad.runtime.support.hcq import MMIOInterface from tinygrad.runtime.support import c @@ -220,6 +225,7 @@ class USBMMIOInterface(MMIOInterface): return (index * self.el_sz, self.el_sz) def __getitem__(self, index): + Device[HCQ_RUNTIME_DEV.value].synchronize() # one driver on the link: drain the compiled submits before python touches it off, sz = self._off_from_index(index) if self.pcimem: assert sz % 4 == 0 and off % 4 == 0, f"pcie_mem_read requires 4-byte aligned access, got off={off}, sz={sz}" @@ -228,12 +234,114 @@ class USBMMIOInterface(MMIOInterface): return data if isinstance(index, slice) else int.from_bytes(data, "little") def __setitem__(self, index, data): + Device[HCQ_RUNTIME_DEV.value].synchronize() off, _ = self._off_from_index(index) data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data) if not self.pcimem: self.usb.scsi_write(data) if self.addr == 0xf000 else self.usb.write(self.addr + off, data) - else: self.usb.pcie_mem_write(self.addr+off, data) + else: + # writes are whole dwords + assert len(data) % 4 == 0 and off % 4 == 0, f"pcie_mem_write requires 4-byte aligned access, got off={off}, sz={len(data)}" + self.usb.pcie_mem_write(self.addr+off, data) def view(self, offset:int=0, size:int|None=None, fmt=None): return USBMMIOInterface(self.usb, self.addr+offset, self.nbytes-offset if size is None else size, fmt=fmt or self.fmt, pcimem=self.pcimem) +# ***************** + +def _libusb(devs, dep:tuple[UOp, ...], fn:str, *args) -> UOp: + return make_buf(devs, tag=f"func:{fn}").after(*dep).index(0).load().call(make_buf(devs, tag="usb_handle").index(0).load(), + *[UOp.const(a, dtypes.int) if isinstance(a, int) else a for a in args], ret_dtype=dtypes.void) + +def usb_bulk(devs, dep, endpoint:int, data:UOp, length, timeout:int=1000) -> UOp: # NULL actual_length out param + return _libusb(devs, dep, "libusb_bulk_transfer", endpoint, data, length, UOp.const(0, dtypes.uint64), timeout) + +def usb_stream(devs, dep:tuple[UOp, ...], addr:UOp, data:UOp, nbytes:int, write:bool) -> UOp: + hdr = UOp.placeholder((2,), dtypes.uint64, device=devs, tag="usb_scratch").after(*dep) + arm = _libusb(devs, (hdr.index(0).store(addr), hdr.index(1).store(UOp.const(nbytes // 4, dtypes.uint64))), "libusb_control_transfer", + 0x40, 0xF0, (0x60 if write else 0x20) | (0x0F << 8), 1 if write else 2, hdr.index(0), 12, 5000) + return usb_bulk(devs, (arm,), 0x02 if write else 0x81, data, nbytes) + +def usb_writes(devs, ws:list[tuple[UOp, UOp, int]]) -> tuple[UOp, ...]: + return functools.reduce(lambda dep, w: (usb_stream(devs, dep, w[0], w[1], w[2], True),), ws, ()) + +def usb_load(b:UOp, idx:UOp, dt) -> UOp: + got = UOp.placeholder((1,), dt, device=(devs:=to_tuple(b.device)), tag="usb_scratch") + addr = b.getaddr((HCQ_RUNTIME_DEV.value,)) + (idx*dt.itemsize).cast(dtypes.uint64) + return got.after(usb_stream(devs, b.src[1:] if b.op is Ops.AFTER else (), addr, got.index(0), dt.itemsize, False)).index(0).load() + +def usb_write(b:UOp, idx:UOp, v:UOp) -> UOp: + val = (s:=UOp.placeholder((1,), v.dtype, device=(devs:=to_tuple(b.device)), tag="usb_scratch")).after(s.index(0).store(v)) + addr = b.getaddr((HCQ_RUNTIME_DEV.value,)) + (idx*v.dtype.itemsize).cast(dtypes.uint64) + return usb_stream(devs, b.src[1:] if b.op is Ops.AFTER else (), addr, val.index(0), v.dtype.itemsize, True) + +def usb_idle(devs) -> UOp: + v = usb_load(make_buf(devs, tag="timeline_signal").after(loop:=UOp.loop(0)), UOp.const(0, dtypes.int), dtypes.uint64) + return v.end(loop, v + 1 < make_buf(devs, tag="timeline_value").index(0).load()) + +def usb_scsi(devs, read:bool, nbytes:int) -> UOp: + return _libusb(devs, (usb_idle(devs),), "libusb_control_transfer", 0x40, 0xF2, ceildiv(nbytes, 512) | (0x8000 if read else 0), + (ceildiv(nbytes, 0x4000) & 0xFF) << 8, UOp.const(0, dtypes.uint64), 0, 1000) + +def usb_stage_copy(dst:UOp, src:UOp) -> UOp|None: + if (cin:=to_tuple(src.device)[0].startswith("CPU")) == to_tuple(dst.device)[0].startswith("CPU"): return None + + total, ops, win = dst.nbytes(), [], cast(Any, Device[(devs:=to_tuple((dst if cin else src).device))[0]]).iface.usb_sram + for off in range(0, total, win.size): # off and nb are bytes, the two ends of the copy can have different dtypes + sram = UOp.from_buffer(win)[0:(nb:=min(win.size, total - off))] + s, d = src[off // src.dtype.itemsize:(off + nb) // src.dtype.itemsize], dst[off // dst.dtype.itemsize:(off + nb) // dst.dtype.itemsize] + if cin: + push = usb_bulk(devs, (usb_scsi(devs, False, nb),), 0x02, s.getaddr((HCQ_RUNTIME_DEV.value,)), round_up(nb, 512), 10000) + ops += [UOp.custom_function("hcq", push.sink()).call(sram, s, name="hcq_copyin", aux=HCQInfo(devs)), + sram.copy_to_device(d.device).call(d, sram)] + else: + pad = UOp.new_buffer("CPU", round_up(nb, 512), dtypes.uint8)[0:nb] + submit = make_submit(UOp(Ops.CALL, dtypes.void, (UOp(Ops.COPY, dtypes.void, ()), sram, s)), devs=devs, queue="COPY:0") + pull = usb_bulk(devs, (submit,), 0x81, pad.getaddr((HCQ_RUNTIME_DEV.value,)), round_up(nb, 512), 10000) + ops += [UOp.custom_function("hcq", pull.sink()).call(pad, sram, s, name="hcq_copyout", aux=HCQInfo(devs)), + pad.copy_to_device("CPU").call(d, pad)] + return UOp(Ops.LINEAR, src=tuple(ops)) +pm_usb_stage = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.COPY), UPat(name="dst"), UPat(name="src"))), usb_stage_copy)]) + +def usb_arm_bytes(lin:UOp, sram:Buffer) -> int: + dsts = [c.src[1] for c in lin.src if c.op is Ops.CALL and c.src[0].op is Ops.COPY] # the rest of a linear is INS, some with no srcs + return next((d.nbytes() for d in dsts if d.base.op is Ops.BUFFER and d.base.buffer is sram), 0) + +def usb_ib(devs, lin:UOp, align:int, arm:int=0) -> tuple[UOp, UOp, int]: + pkt_dw = sum(s.dtype.itemsize for ins in lin.src for s in ins.src) // 4 # by bytes: sdma packs 64-bit addresses as single srcs + kargs = dedup([b for b in lin.toposort() if b.op is Ops.PARAM and b.tag == "kernargs"]) + offs, up_dw = {}, round_up(pkt_dw, align) + for k in kargs: offs[k], up_dw = up_dw, round_up(up_dw + k.max_numel(), 32) + ib_gpu = UOp.placeholder((up_dw,), dtypes.uint32, device=devs, tag="cmdbuf") + ib_host = UOp.placeholder((up_dw,), dtypes.uint32, device=devs, tag="usb_scratch") + gsubs = {g: g.replace(src=(d if a.op is not Ops.AFTER else d.after(*a.src[1:]),)) for g in lin.toposort() if g.op is Ops.GETADDR + for a in [g.src[0]] if (k:=a.src[0] if a.op is Ops.AFTER else a) in offs for d in [ib_gpu[offs[k]:offs[k] + k.max_numel()]]} + lin = lin.substitute(gsubs, walk=True).substitute({k: ib_host[offs[k]:offs[k] + k.max_numel()] for k in kargs}, walk=True) + return make_cmdbuf(lin, devs, buf=ib_host, dep=(usb_scsi(devs, True, arm),) if arm else ()), ib_gpu, pkt_dw + +def usb_push(devs, ring:UOp, wptr:UOp, doorbell:UOp, put_ptr:UOp, ib_host:UOp, ib_gpu:UOp, pkt:tuple, unit:int) -> UOp: + stage = UOp.placeholder(((n:=round_up(len(pkt), 4)) + 2,), dtypes.uint32, device=devs, tag="usb_scratch") + put, step = put_ptr.index(zero:=UOp.const(0, dtypes.int)), (n * 4 if pkt else ib_host.nbytes()) // unit + st = stage.after(*[stage.index(i).store(UOp.const(v, dtypes.uint32)) for i, v in enumerate(pkt)], + *[stage.index(n + i).store((((put + step) >> (32 * i)) & 0xffffffff).cast(dtypes.uint32)) for i in (0, 1)]) + + writes = [(ib_gpu.getaddr((HCQ_RUNTIME_DEV.value,)), ib_host.index(zero), ib_gpu.nbytes())] if pkt else [] + writes += [(ring.getaddr((HCQ_RUNTIME_DEV.value,)) + ((put % (ring.nbytes() // unit)) * unit).cast(dtypes.uint64), + (st if pkt else ib_host).index(zero), step * unit)] + writes += [(p.getaddr((HCQ_RUNTIME_DEV.value,)), st.index(n), 8) for p in (wptr, doorbell)] + return put_ptr.after(*usb_writes(devs, writes)).index(zero).store(put + step) + +USB_HOST_TAGS = {"signal", "timeline_signal"} +pm_usb_hostio = PatternMatcher([ + (UPat(Ops.LOAD, src=(UPat(Ops.INDEX, src=(UPat(Ops.PARAM, tag=USB_HOST_TAGS).or_after(name="b"), UPat(name="idx"))),), + name="ld"), lambda b, idx, ld: usb_load(b, idx, ld.dtype)), + (UPat(Ops.STORE, src=(UPat(Ops.INDEX, src=(UPat(Ops.PARAM, tag=USB_HOST_TAGS).or_after(name="b"), UPat(name="idx"))), UPat(name="v"))), usb_write)]) + +pm_usb_bufferize = PatternMatcher([ + (UPat(Ops.PARAM, tag={"systems", "runtime", "inputs", "usb_scratch"}, name="b"), + lambda ctx, b: Buffer("CPU", b.max_numel(), b.dtype, options=BufferSpec(nolru=True), preallocate=True)), + (UPat(Ops.PARAM, tag="usb_handle", name="b"), lambda ctx, b: ctx[0].signal(b.tag, ctx[0].iface.usb_handle, device="CPU")), + (UPat(Ops.PARAM, name="b"), lambda ctx, b: None if not isinstance(b.tag, str) or not b.tag.startswith("func:") else + ctx[0].signal(b.tag, unwrap(ctypes.cast(getattr(libusb.dll, b.tag[5:]), ctypes.c_void_p).value), device="CPU")), +]) + if DEV.interface.startswith("MOCK"): from test.mockgpu.usb import MockUSB3 as USB3 # type: ignore # noqa: F811 diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 9102f2af5b..ae40f94eb2 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -1135,14 +1135,16 @@ class UOp(RandMixin, metaclass=UOpMetaClass): # *** uop high level syntactic sugar *** @staticmethod - def placeholder(shape:tuple[int, ...], dtype:DType, slot:int, addrspace=AddrSpace.GLOBAL, device=None, volatile=False): + def placeholder(shape:tuple[int, ...], dtype:DType, slot:int|None=None, addrspace=AddrSpace.GLOBAL, device=None, volatile=False, tag=None): dtype = strong_dtype(dtype) # storage is never weak: a placeholder commits the width of what's put in it + if slot is None: slot = next(UOp.unique_num) if addrspace is AddrSpace.GLOBAL: ret = UOp(Ops.PARAM, src=(shape_to_shape_arg((prod(shape),)),), arg=ParamArg(slot, dtype, addrspace=addrspace, device=device,volatile=volatile)) else: assert addrspace in (AddrSpace.LOCAL, AddrSpace.REG) assert device is None, "LOCAL and REG placeholders cannot have a device" ret = UOp(Ops.BUFFER, src=(shape_to_shape_arg((prod(shape),)),), arg=ParamArg(slot, dtype, addrspace=addrspace)) + if tag is not None: ret = ret.rtag(tag) if len(shape) > 1: ret = ret.reshape(shape) return ret def placeholder_like(self, slot:int, addrspace=AddrSpace.GLOBAL):