diff --git a/tinygrad/codegen/gpudims.py b/tinygrad/codegen/gpudims.py index 60b28c8397..7ffd44c2cc 100644 --- a/tinygrad/codegen/gpudims.py +++ b/tinygrad/codegen/gpudims.py @@ -91,7 +91,7 @@ def add_gpudims(ctx:Renderer, s:UOp): subs = {} for r in s_topo: # look for local INDEXes that are not used in the GLOBAL store, then add them as an INVALID - if r.op is Ops.STORE and (idx := r.src[0]).src[0].ptrdtype.addrspace == AddrSpace.GLOBAL: + if r.op is Ops.STORE and (idx := r.src[0]).src[0].addrspace == AddrSpace.GLOBAL: missing_locals = [all_ranges[rng] for rng in local_dims if all_ranges[rng] not in idx.ranges] if len(missing_locals): assert len(idx.src) == 2, "index has 2 sources" diff --git a/tinygrad/codegen/late/devectorizer.py b/tinygrad/codegen/late/devectorizer.py index cde5fda449..8268317392 100644 --- a/tinygrad/codegen/late/devectorizer.py +++ b/tinygrad/codegen/late/devectorizer.py @@ -104,7 +104,7 @@ def fold_expanded_index(midx:UOp): for grp in grouped_offsets: # get the index offset for this element. using [0] is okay, because they are the same lidx = midx.src[offsets[grp[0]][0]] - if len(grp) > 1: lidx = lidx.cast(buf.ptrdtype.base.vec(len(grp)).ptr(size=buf.ptrdtype.size, addrspace=buf.ptrdtype.addrspace)) + if len(grp) > 1: lidx = lidx.cast(buf.ptrdtype.base.vec(len(grp)).ptr(size=buf.max_numel(), addrspace=buf.addrspace)) # set the idxs of the output for i,g in enumerate(grp): for oo in offsets[g]: idxs[oo] = global_offset+i @@ -113,7 +113,7 @@ def fold_expanded_index(midx:UOp): global_offset += len(grp) assert None not in idxs, f"some idxs are missing {idxs}" # this base thing is for image, we want the CAT to be a normal pointer - post_cat = UOp(Ops.PTRCAT, buf.ptrdtype.base.ptr(size=buf.ptrdtype.size, addrspace=buf.ptrdtype.addrspace).vec(global_offset), tuple(ret)) + post_cat = UOp(Ops.PTRCAT, buf.ptrdtype.base.ptr(size=buf.max_numel(), addrspace=buf.addrspace).vec(global_offset), tuple(ret)) return post_cat.gep(tuple(cast(list[int], idxs))) def cat_after_store(cat:UOp, data:UOp): @@ -165,7 +165,7 @@ def split_load_store(ctx:Renderer|None, ls:UOp, idx:UOp): must_divide = False elif buf.dtype.base not in (dtypes.float, dtypes.half, *dtypes.fp8s) and not isinstance(buf.dtype, ImageDType): pass - elif buf.ptrdtype.addrspace == AddrSpace.REG: + elif buf.addrspace == AddrSpace.REG: pass elif isinstance(buf.dtype, ImageDType): lengths = [4] @@ -186,7 +186,7 @@ def split_load_store(ctx:Renderer|None, ls:UOp, idx:UOp): for fold_length in lengths: if global_offset+fold_length > sz: continue lidx = buf.index((offset + global_offset).valid(mask), ptr=True) - if fold_length > 1: lidx = lidx.cast(buf.ptrdtype.base.vec(fold_length).ptr(size=buf.ptrdtype.size, addrspace=buf.ptrdtype.addrspace)) + if fold_length > 1: lidx = lidx.cast(buf.ptrdtype.base.vec(fold_length).ptr(size=buf.max_numel(), addrspace=buf.addrspace)) if ls.op is Ops.STORE: ret.append(ls.replace(src=(lidx,ls.src[1].gep(tuple(range(global_offset, global_offset+fold_length)))))) else: ret.append(ls.replace(src=(lidx,)+ls.src[1:], dtype=ls.dtype.scalar().vec(fold_length))) global_offset += fold_length @@ -243,7 +243,9 @@ def no_vectorized_alu(alu:UOp): return UOp(Ops.STACK, alu.dtype, alus) def no_vectorized_buf(buf:UOp): - return buf.replace(dtype=buf.ptrdtype.base.scalar().ptr(buf.ptrdtype.size*buf.ptrdtype.count, buf.ptrdtype.addrspace)).cast(buf.dtype) + # TODO: this fails on regs + #assert buf.max_numel() == buf.ptrdtype.size + return buf.replace(dtype=buf.ptrdtype.base.scalar().ptr(buf.ptrdtype.size*buf.ptrdtype.count, buf.addrspace)).cast(buf.dtype) def no_vectorized_index(buf:UOp, cast:UOp, idx:UOp, bcast:UOp|None=None): cnt = cast.dtype.count diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index 8ec792dd43..0fd700bf46 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -330,7 +330,7 @@ class Scheduler: def bufs_from_ast(ast:UOp, dname:str) -> list[Buffer]: glbls = sorted([x for x in ast.backward_slice if x.op is Ops.PARAM], key=lambda x: x.arg) - return [Buffer(dname, x.ptrdtype.size, x.dtype.base) for x in glbls] + return [Buffer(dname, x.max_numel(), x.dtype.base) for x in glbls] def apply_opts(ast:UOp, ren:Renderer, beam:int=0) -> UOp: if ast.tag is not None: return ast diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 0ff94de701..c63349500e 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -179,7 +179,7 @@ class CStyleLanguage(Renderer): if u.op in (Ops.PARAM, Ops.DEFINE_VAR): if u.op is not Ops.PARAM: r[u] = u.arg[0] elif isinstance(u.dtype, ImageDType): r[u] = f"data{u.arg}_{u.dtype.shape[0]}x{u.dtype.shape[1]}" - else: r[u] = f"data{u.arg}_{sz}" if (sz:=u.ptrdtype.size) > 0 else f"data{u.arg}" + else: r[u] = f"data{u.arg}_{sz}" if (sz:=u.max_numel()) > 0 else f"data{u.arg}" bufs[u] = (r[u], (u.dtype, u in writable_params)) continue @@ -198,7 +198,7 @@ class CStyleLanguage(Renderer): if u.op in {Ops.ENDIF, Ops.END}: depth -= 1 if (u.op is not Ops.CAST or u.dtype.vcount == 1) and (u.op in {Ops.CONST, Ops.GEP, Ops.INDEX, Ops.CUSTOMI} or \ - (u.op is Ops.LOAD and u.src[0].ptrdtype.addrspace == AddrSpace.REG) or \ + (u.op is Ops.LOAD and u.src[0].addrspace == AddrSpace.REG) or \ (u.op is Ops.CAST and isinstance(u.dtype, PtrDType)) or \ (u.op in {Ops.STACK, *(GroupOp.ALU-{Ops.WHERE}), Ops.CAST, Ops.BITCAST} and child_count[u] == 1 and not getenv("EXPAND_SSA"))): r[u] = l diff --git a/tinygrad/renderer/nir.py b/tinygrad/renderer/nir.py index 32ca66a113..e81713a78c 100644 --- a/tinygrad/renderer/nir.py +++ b/tinygrad/renderer/nir.py @@ -149,12 +149,12 @@ class NIRRenderer(Renderer): (UPat(Ops.DEFINE_VAR, name="x"), lambda ctx,x: ctx.param(ctx.b, x, 4)), (UPat(Ops.SPECIAL, name="x"), lambda ctx,x: nchannel(ctx.b, {'g':ngid, 'l':nlid, 'i': nid}[x.arg[0]](ctx.b), int(x.arg[-1]))), (UPat(Ops.STORE, src=(UPat(Ops.INDEX, src=(UPat.var("buf"),UPat.var("off"))).or_casted(), UPat.var("val"))), - lambda ctx,buf,off,val: nstore(ctx.b, buf.ptrdtype.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype), ctx.r[val], val.dtype)), + lambda ctx,buf,off,val: nstore(ctx.b, buf.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype), ctx.r[val], val.dtype)), (UPat(Ops.LOAD, src=(UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("off"))).or_casted(), UPat.var("alt"), UPat.var("gate")), name="x"), lambda ctx,x,buf,off,alt,gate: if_phi(ctx.b, ctx.r[gate], - lambda: nload(ctx.b, buf.ptrdtype.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype, ctx.r[gate]), x.dtype), lambda: ctx.r[alt])), + lambda: nload(ctx.b, buf.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype, ctx.r[gate]), x.dtype), lambda: ctx.r[alt])), (UPat(Ops.LOAD, src=(UPat(Ops.INDEX, src=(UPat.var("buf"), UPat.var("off"))).or_casted(),), name="x"), - lambda ctx,x,buf,off: nload(ctx.b, buf.ptrdtype.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype), x.dtype)), + lambda ctx,x,buf,off: nload(ctx.b, buf.addrspace, nidx(ctx.b, ctx.r[buf], ctx.r[off], buf.dtype), x.dtype)), (UPat(Ops.STACK, name="x"), lambda ctx,x: nalu(ctx.b, f"vec{x.dtype.count}", *[ctx.r[src] for src in x.src])), (UPat(GroupOp.ALU, name="x"), lambda ctx,x: nalu(ctx.b, aop[x.src[0].dtype.scalar()][x.op], *[ctx.r[src] for src in x.src])), (UPat(Ops.CAST, name="x"), lambda ctx,x: ncast(ctx.b, ctx.r[x.src[0]], x.src[0].dtype, x.dtype)), diff --git a/tinygrad/renderer/ptx.py b/tinygrad/renderer/ptx.py index e2d6b9eadf..c090094dd3 100644 --- a/tinygrad/renderer/ptx.py +++ b/tinygrad/renderer/ptx.py @@ -202,7 +202,7 @@ class PTXRenderer(Renderer): r[u] = r[u.src[0]] continue if u.op is Ops.DEFINE_REG: - r[u] = [ssa("reg", u, self.types[u.dtype.base.scalar()]) for _ in range(u.ptrdtype.size)] + r[u] = [ssa("reg", u, self.types[u.dtype.base.scalar()]) for _ in range(u.max_numel())] continue if u.op in {Ops.INDEX, Ops.LOAD, Ops.STORE} and isinstance(u.src[0].dtype, PtrDType) and u.src[0].dtype.addrspace == AddrSpace.REG: if u.op is Ops.INDEX: diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index fb090bf3f6..cb6e05fb0a 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -489,7 +489,7 @@ def unbind_kernel(ctx:LocalAddBufferContext, b:UOp): return b.src[0] def handle_after(ctx:LocalAddBufferContext, after:UOp): - if isinstance(after.dtype, PtrDType) and after.ptrdtype.addrspace == AddrSpace.LOCAL: return None + if isinstance(after.dtype, PtrDType) and after.addrspace == AddrSpace.LOCAL: return None buf = after.buf_uop # HACK to put the buffer in the MAP instead of MSTACK/MSELECT if buf.op in {Ops.MSTACK, Ops.MSELECT}: buf = buf.src[0] diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 9d81519e1c..cd2aaff9fd 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -368,6 +368,7 @@ class UOp(OpMixin, metaclass=UOpMetaClass): @property def max_shape(self) -> tuple[int, ...]: return to_max_shape(self.shape) + def max_numel(self) -> int: return prod(self.max_shape) @property def shard_shape(self) -> tuple[sint, ...]: @@ -743,6 +744,17 @@ class UOp(OpMixin, metaclass=UOpMetaClass): if x.device is not None: return x.device return None @property + def addrspace(self) -> AddrSpace: + if self.op in {Ops.PARAM, Ops.BUFFER}: return AddrSpace.GLOBAL + if self.op is Ops.DEFINE_LOCAL: return AddrSpace.LOCAL + if self.op is Ops.DEFINE_REG: return AddrSpace.REG + if self.op is Ops.STACK: + assert all_same([x.addrspace for x in self.src]), "addrspace mismatch" + return self.src[0].addrspace + if self.op in {Ops.INDEX, Ops.CAST, Ops.AFTER}: return self.src[0].addrspace + if self.op in GroupOp.Movement: return self.src[0].addrspace + raise Exception(f"{self.op} doesn't have addrspace") + @property def buf_uop(self) -> UOp: if self.op in {Ops.BUFFER, Ops.PARAM}: return self if self.op is Ops.MSELECT: return self.src[0].buf_uop.mselect(self.arg)