diff --git a/extra/hcq2/hcq2.py b/extra/hcq2/hcq2.py index 7021dd8e0b..7fccc0cd39 100644 --- a/extra/hcq2/hcq2.py +++ b/extra/hcq2/hcq2.py @@ -138,7 +138,7 @@ def unwrap_after(uop): def make_getaddr(u, device=None): if unwrap_after(u).op not in (Ops.BUFFER, Ops.SLICE, Ops.BINARY, Ops.MSTACK, Ops.MSELECT, Ops.PARAM): return u - return UOp(Ops.GETADDR, dtypes.uint64, src=(u, UOp(Ops.DEVICE, arg=device or to_tuple(u.device)[0]))) + return UOp(Ops.GETADDR, dtypes.uint64, src=(u,), arg=device or to_tuple(u.device)[0]) def make_ins(op, *srcs): return UOp(Ops.INS, dtypes.void, tuple(UOp.const(dtypes.uint32, s) if isinstance(s, int) else s.cast(dtypes.uint32) for s in srcs), op) @@ -214,9 +214,9 @@ def prep_program(call:UOp, prg:UOp) -> UOp|None: return prg.replace(src=(buf.after(buf.store(blob)),), arg=(data, prg.arg)).call(*call.src[1:], aux=HCQInfo.from_call(call)) def prep_kernargs(call:UOp, prg:UOp) -> UOp: - (data, info), dev_uop = prg.arg, UOp(Ops.DEVICE, arg=call.src[1].device) - buf = UOp.new_buffer(dev_uop.arg, data.kernargs_alloc_size, dtypes.uint8).rtag("kernargs") - patches = [make_patch(buf, i*8, make_getaddr(call.src[1+gi], dev_uop.arg)) for i,gi in enumerate(info.globals)] \ + (data, info), dev = prg.arg, call.src[1].device + buf = UOp.new_buffer(dev, data.kernargs_alloc_size, dtypes.uint8).rtag("kernargs") + patches = [make_patch(buf, i*8, make_getaddr(call.src[1+gi], dev)) for i,gi in enumerate(info.globals)] \ + [make_patch(buf, len(info.globals)*8 + i*4, v, dtypes.uint32) for i,v in enumerate(info.vars)] return call.replace(src=(prg.replace(src=prg.src + (buf.after(*patches),), arg=(data, info)),) + call.src[1:]) @@ -517,15 +517,15 @@ def fold_const_store(buf:UOp, off:UOp, val:UOp) -> UOp: def resolve_getaddr(buf:UOp, g:UOp) -> UOp: if buf.op not in (Ops.BUFFER, Ops.MSTACK, Ops.MSELECT): return buf - devs, b = to_tuple(g.src[1].arg), buf.buffer + devs, b = to_tuple(g.arg), buf.buffer bufs = tuple(cast(Buffer, x.buffer) for x in buf.src) if buf.op is Ops.MSTACK else tuple(b.bufs if isinstance(b, MultiBuffer) else (b,)*len(devs)) assert len(bufs) == len(devs), f"can't resolve {len(bufs)} buffers on {len(devs)} devices" addrs = tuple(UOp.const(dtypes.uint64, x.get_buf(d).va_addr) for x, d in zip(bufs, devs)) return addrs[0] if len(addrs) == 1 else UOp(Ops.STACK, dtypes.uint64.vec(len(addrs)), addrs) -def resolve_getaddr_slice(bv:UOp, dev:UOp) -> UOp: +def resolve_getaddr_slice(bv:UOp, g:UOp) -> UOp: itemsize = bv.src[0].dtype.itemsize if unwrap_after(bv.src[0]).op in (Ops.BUFFER, Ops.SLICE, Ops.MSTACK, Ops.MSELECT) else bv.dtype.itemsize - return UOp(Ops.GETADDR, dtypes.uint64, src=(bv.src[0], dev)) + UOp.const(dtypes.uint64, bv.src[1].arg * itemsize) + return UOp(Ops.GETADDR, dtypes.uint64, src=(bv.src[0],), arg=g.arg) + UOp.const(dtypes.uint64, bv.src[1].arg * itemsize) pm_resolve_patches = PatternMatcher([ # multi @@ -537,8 +537,8 @@ pm_resolve_patches = PatternMatcher([ lambda shr, bv: shr.replace(src=(bv.src[0], shr.src[1] + bv.src[1].cast(shr.src[1].dtype), shr.src[2]))), # getaddr - (UPat(Ops.GETADDR, src=(UPat(Ops.SLICE, name="bv"), UPat(Ops.DEVICE, name="dev"))), resolve_getaddr_slice), # getaddr(slice(x)) -> offset+getaddr(x) - (UPat(Ops.GETADDR, src=(UPat(name="buf"), UPat(Ops.DEVICE)), name="g"), resolve_getaddr), + (UPat(Ops.GETADDR, src=(UPat(Ops.SLICE, name="bv"),), name="g"), resolve_getaddr_slice), # getaddr(slice(x)) -> offset+getaddr(x) + (UPat(Ops.GETADDR, src=(UPat(name="buf"),), name="g"), resolve_getaddr), # folders (UPat({Ops.BUFFER, Ops.SLICE, Ops.MSTACK}, name="buf").store(UPat(Ops.BINARY, name="blob")), fold_blob_store), diff --git a/spec/tinyspec.pdf b/spec/tinyspec.pdf index 7ff340de64..04d4c67439 100644 Binary files a/spec/tinyspec.pdf and b/spec/tinyspec.pdf differ diff --git a/spec/tinyspec.tex b/spec/tinyspec.tex index 5c87166a87..df371b60b2 100644 --- a/spec/tinyspec.tex +++ b/spec/tinyspec.tex @@ -227,7 +227,7 @@ Ternary & $(P, A, B)$ \midrule \op{Barrier} & (deps\ldots) & --- & Synchronize threads within a workgroup. \\ \op{Ins} & \ldots & \ldots & A single machine instruction (e.g.\ AMD ISA). \\ -\op{GetAddr} & (buf, dev) & --- & Lower buf to its address on device dev. \\ +\op{GetAddr} & (buf,) & dev & Lower buf to its address on device dev. \\ \op{Special} & (bound,) & name & GPU thread/workgroup index (e.g.\ \texttt{gidx0}, \texttt{lidx1}). \\ \op{If} & (gate,) & --- & Begin conditional execution block. \\ \op{Endif} & (if,) & --- & End conditional execution block. \\