Compare commits

...
10 Commits
Author SHA1 Message Date
geohot f082cbcb36 respect the 8x8 tiles 2025-10-07 17:51:24 +08:00
geohot 5ad62f130d split on tile_dim 2025-10-07 17:40:00 +08:00
geohot f129d75ee5 fix on cpu 2025-10-07 16:43:52 +08:00
geohot 51f3a5cbb4 gpu 2025-10-07 16:00:10 +08:00
geohot 1d7a8b33c1 gemm works on pc 2025-10-07 15:52:00 +08:00
George HotzandGitHub 3fae886aa9 Merge branch 'master' into clone_tg 2025-10-07 14:02:36 +08:00
George HotzandGitHub 3f44ef699f Merge branch 'master' into clone_tg 2025-10-07 13:08:35 +08:00
geohot fa23f37e33 clone thunderkittens in uops 2025-10-07 13:06:58 +08:00
geohot 284db26a12 cleanup 2025-10-07 12:13:01 +08:00
geohot 0a0cb0b9e8 merge tagless reshapes 2025-10-07 12:11:11 +08:00
5 changed files with 123 additions and 6 deletions
+109
View File
@@ -0,0 +1,109 @@
from tinygrad import Device, Tensor, Context
from tinygrad.uop.ops import UOp, Ops, graph_rewrite, AxisType, PatternMatcher, UPat, pm_lower_index_dtype, GroupOp, KernelInfo
from tinygrad.dtype import dtypes, AddrSpace
from tinygrad.helpers import prod
from tinygrad.schedule.rangeify import pm_mops
from tinygrad.codegen.simplify import pm_flatten_range
TILE_DIM = 8
N_BLOCK = 4
K_BLOCK = 2
M_BLOCK = 4
#M = N = K = 4096
M = N = K = 1024
range_num = 0
def rng(x, typ=AxisType.LOOP) -> UOp:
global range_num
range_num += 1
return UOp.range(x, range_num-1, typ)
def glbl(nm, dtype, sz): return UOp(Ops.DEFINE_GLOBAL, dtype.ptr(prod(sz), AddrSpace.GLOBAL), arg=nm).reshape(sz)
def rt(nm, dtype, sz): return UOp(Ops.DEFINE_REG, dtype.ptr(prod(sz), AddrSpace.REG), arg=nm).reshape(sz)
def zero(reg:UOp, *endrngs):
rngs = [rng(s//TILE_DIM)*TILE_DIM for s in reg.shape]
rngs = [x+rng(TILE_DIM) for x in rngs]
return reg[*rngs].store(UOp.const(reg.dtype.base, 0.0), *rngs, *endrngs, dtype=reg.dtype).reshape(reg.shape)
def load(reg:UOp, gl:UOp, *idxs):
rngs = [rng(s//TILE_DIM)*TILE_DIM for s in reg.shape]
rngs = [x+rng(TILE_DIM) for x in rngs]
grngs = [i*(r.vmax+1)+r for i,r in zip(idxs,rngs)]
return reg[*rngs].store(gl[*grngs].load(), *rngs, dtype=reg.dtype).reshape(reg.shape)
def store(gl:UOp, reg:UOp, *idxs):
rngs = [rng(s//TILE_DIM)*TILE_DIM for s in reg.shape]
rngs = [x+rng(TILE_DIM) for x in rngs]
# TODO: why does this not have shape?
#rngs = [rng(s) for s in (N_BLOCK*TILE_DIM, M_BLOCK*TILE_DIM)]
grngs = [i*(r.vmax+1)+r for i,r in zip(idxs,rngs)]
return gl[*grngs].store(reg[*rngs].load(), *rngs)
def mma_AB(outacc:UOp, a:UOp, b:UOp, *endrngs):
assert a.shape[1] == b.shape[0]
# meta::unroll_i_j_in_range -- split on TILE_DIM
rngs = [rng(s//TILE_DIM)*TILE_DIM for s in outacc.shape]
red = rng(a.shape[1]//TILE_DIM, AxisType.REDUCE)*TILE_DIM
# meta::unroll_i_in_range -- split reduce on TILE_DIM
rngs = [x+rng(TILE_DIM) for x in rngs]
red = red + rng(TILE_DIM, AxisType.REDUCE)
acc = outacc[*rngs].load(red) + a[rngs[0],red].load() * b[red,rngs[1]].load()
return outacc[*rngs].store(acc, *rngs, red, *endrngs, dtype=outacc.dtype).reshape(outacc.shape)
if __name__ == "__main__":
# TODO: support string ranges
tg_id_y = UOp.range(M // (M_BLOCK * TILE_DIM), -3, AxisType.GLOBAL if Device.DEFAULT != "CPU" else AxisType.LOOP)
tg_id_x = UOp.range(N // (N_BLOCK * TILE_DIM), -2, AxisType.GLOBAL if Device.DEFAULT != "CPU" else AxisType.LOOP)
gl_d = glbl("gl0_d", dtypes.float, (N, M))
gl_a = glbl("gl1_a", dtypes.float, (N, K))
gl_b = glbl("gl2_b", dtypes.float, (K, M))
a_reg = rt("a_reg", dtypes.float, (N_BLOCK*TILE_DIM, K_BLOCK*TILE_DIM))
b_reg = rt("b_reg", dtypes.float, (K_BLOCK*TILE_DIM, M_BLOCK*TILE_DIM))
d_reg = rt("d_reg", dtypes.float, (N_BLOCK*TILE_DIM, M_BLOCK*TILE_DIM))
d_reg = zero(d_reg, UOp(Ops.NOOP, src=(tg_id_y, tg_id_x)))
k = UOp.range(K // (K_BLOCK * TILE_DIM), -1, AxisType.REDUCE)
a_reg = load(a_reg, gl_a, tg_id_y, k)
b_reg = load(b_reg, gl_b, k, tg_id_x)
d_reg = mma_AB(d_reg, a_reg, b_reg, k)
sink = store(gl_d, d_reg, tg_id_y, tg_id_x).sink(arg=KernelInfo())
sink = graph_rewrite(sink, pm_mops+pm_flatten_range, name="pm_mops")
from tinygrad.codegen.gpudims import pm_add_gpudims
sink = graph_rewrite(sink, pm_add_gpudims, ctx=Device.default.renderer, name="gpudims")
pm_lower_index_dtype_simple = PatternMatcher([
(UPat(GroupOp.All, dtype=dtypes.index, name="x"), lambda x: x.replace(dtype=dtypes.int))
])
sink = graph_rewrite(sink, pm_lower_index_dtype_simple, name="index_dtype")
from tinygrad.codegen import rewrites_for_linearizer, apply_rewrites
lin = apply_rewrites(sink, rewrites_for_linearizer)
src = Device.default.renderer.render(lin.arg.lst)
print(src)
#exit(0)
from tinygrad.engine.realize import CompiledRunner, ExecItem
from tinygrad.renderer import ProgramSpec
ps = ProgramSpec("test", src, Device.DEFAULT, sink, lin.arg.lst, [1,1,1], [1,1,1])
run = CompiledRunner(ps)
a = Tensor.randn(N, N)
b = Tensor.randn(N, N)
c = Tensor.empty(N, N)
Tensor.realize(a, b, c)
ei = ExecItem(run, [x.uop.buffer.ensure_allocated() for x in (c,a,b)])
with Context(DEBUG=2):
for i in range(5): ei.run()
for i in range(5): ref = (a@b).realize()
print((ref-c).mean().item())
+3 -3
View File
@@ -5,10 +5,10 @@ from tinygrad.dtype import dtypes
def flatten_range(r:UOp):
off = range_start[r.op]
rngs = r.src[off:]
rngs, noops = partition(r.src[off:], lambda x: x.op is not Ops.NOOP)
if not len(rngs): return None
new_rngs = [x for x in UOp.sink(*rngs).toposort() if x.op is Ops.RANGE]
return r.replace(src=r.src[:off]+tuple(new_rngs))
new_rngs = [x for x in UOp.sink(*rngs).toposort(lambda x: x.op is not Ops.NOOP) if x.op is Ops.RANGE]
return r.replace(src=r.src[:off]+tuple(new_rngs)+tuple(noops))
pm_flatten_range = PatternMatcher([
# real ranges only
+1
View File
@@ -159,6 +159,7 @@ class CStyleLanguage(Renderer):
prefix = None
if u.op is Ops.SPECIAL: r[u] = u.arg
elif u.op is Ops.RANGE: r[u] = "ridx"+range_str(u)
elif u.op is Ops.STORE: r[u] = r[u.src[0].src[0]]
else:
prefix = {Ops.WMMA: "wmma", Ops.DEFINE_LOCAL: "temp", Ops.CONST: "const",
Ops.CAST: "cast", Ops.BITCAST: "cast", Ops.GEP: "gep", Ops.VECTORIZE: "cast", Ops.PRECAST: "precast",
+2 -1
View File
@@ -175,7 +175,8 @@ class RangeifyContext:
def map_reshape(idx:UOp, r:UOp):
acc = 1
to_sum = []
for s,src in list(zip(idx.shape, idx.src[1:]))[::-1]:
idx_shape = [x.vmax+1 for x in idx.src[1:]]
for s,src in list(zip(idx_shape, idx.src[1:]))[::-1]:
to_sum.append(acc*src)
acc *= s
mish = sum(to_sum, start=UOp.const(dtypes.index, 0))
+8 -2
View File
@@ -182,6 +182,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
Ops.MSELECT, Ops.BUFFER, Ops.BUFFERIZE, Ops.VECTORIZE, Ops.STORE}:
return None
if self.op is Ops.INDEX and self.src[0].op is Ops.ASSIGN and self.src[0].src[1].op is Ops.KERNEL: return None
if self.op is Ops.INDEX: return None
if self.op is Ops.BARRIER: return None
if self.op in GroupOp.Block: return None
from tinygrad.shape.shapetracker import ShapeTracker
@@ -194,6 +195,9 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
# CONST with a DEVICE has a shape of ()
if self.op is Ops.CONST and len(self.src) and self.src[0].op is Ops.DEVICE: return ShapeTracker.from_shape(())
if self.op is Ops.STORE and isinstance(self.dtype, PtrDType): return ShapeTracker.from_shape((self.dtype.size,))
#if self.op is Ops.LOAD: return ShapeTracker.from_shape((self.dtype.count,))
# skip the INDEX
if self.op is Ops.STORE and self.dtype is not dtypes.void: return self.src[0].src[0].st
# BufferOps and ASSIGN flow ShapeTracker from a direct edge
if self.op in {Ops.STORE, Ops.ASSIGN, Ops.LOAD}: return self.src[0].st
@@ -308,7 +312,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
def detach(self): return UOp(Ops.DETACH, self.dtype, (self,))
def index(self, *srcs:UOp|None, **kwargs):
return UOp(Ops.INDEX, kwargs.pop("dtype", self.dtype), (self,)+tuple([x for x in srcs if x is not None]), **kwargs)
def __getitem__(self, idx): return self.index(idx)
def __getitem__(self, idx): return self.index(*idx)
def const_like(self, b:ConstLike):
# constants can optionally have a DEVICE source
return UOp.const(self.dtype, b, device=self._device, shape=self.shape if self.st is not None else None)
@@ -332,7 +336,9 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
i = (i,)
return UOp(Ops.GEP, self.dtype.scalar().vec(len(i)) if len(i) > 1 else self.dtype.scalar(), (self,), i)
def load(self, *src:UOp, **kwargs): return UOp(Ops.LOAD, dtype=kwargs.pop("dtype", self.dtype.base), src=(self,)+src, **kwargs)
def store(self, *src:UOp, **kwargs): return UOp(Ops.STORE, kwargs.pop("dtype", dtypes.void), (self,)+src, **kwargs)
def store(self, *src:UOp, **kwargs):
return UOp(Ops.STORE, kwargs.pop("dtype", dtypes.void), (self,)+src, **kwargs)
#return UOp(Ops.STORE, self.dtype, (self,)+src, **kwargs)
def assign(self, x:UOp): return UOp(Ops.ASSIGN, self.dtype, (self, x))
def barrier(self, *src:UOp): return UOp(Ops.BARRIER, src=(self,)+src)
def alu(self, op, *src:UOp, **kwargs):