From fa61b692fc987fb1e8045335124e4f7a238a85f3 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 24 Aug 2025 12:44:49 -0700 Subject: [PATCH] use simple gidxs on GPU --- test/test_linearizer.py | 73 +------------------------------------ tinygrad/codegen/gpudims.py | 56 +++++----------------------- 2 files changed, 10 insertions(+), 119 deletions(-) diff --git a/test/test_linearizer.py b/test/test_linearizer.py index bdcd2c6c9d..30f7211211 100644 --- a/test/test_linearizer.py +++ b/test/test_linearizer.py @@ -3,7 +3,6 @@ import unittest from dataclasses import replace from tinygrad.codegen.opt.kernel import Opt, OptOps, KernelOptError, Kernel, AxisType -from tinygrad.codegen.gpudims import get_grouped_dims from tinygrad.uop.ops import UOp, Ops, GroupOp, KernelInfo from tinygrad.device import Device, Buffer, is_dtype_supported from tinygrad.shape.shapetracker import ShapeTracker @@ -467,77 +466,7 @@ class TestLinearizer(unittest.TestCase): end_range = [i for i, x in enumerate(uops) if x.op is Ops.ENDRANGE][0] assert end_range < uops.index(u) - def test_grouped_dims(self): - def _assert_grouped_dims(prefix, dims, max_sizes, reverse_dims, expected_sizes, assert_same_length = True): - idxs = get_grouped_dims(prefix, dims, max_sizes, reverse_dims) - loop_idxs = dedup(flatten([[y for y in x.toposort() if y.op is Ops.SPECIAL] for x in idxs])) - loop_idxs = sorted(loop_idxs, key=lambda uop: uop.arg[0]) - sizes = [x.arg[1] for x in loop_idxs] - assert len(idxs) == len(dims), f"expected idxs to have same length as dims {len(dims)}, got {len(idxs)}" - if assert_same_length: - assert len(loop_idxs) == min(len(sizes), len(dims)), f"expected idxs to have length {min(len(sizes), len(dims))}, got {len(loop_idxs)}" - assert sizes == expected_sizes, f"expected sizes={expected_sizes}, got {sizes=}" - # TODO: add these back after uop symbolic - # for i in range(len(dims)): - # assert idxs[i].max+1 == dims[i], f"idxs[{i}] should have max {dims[i]-1}" - # for i in range(len(loop_idxs)): - # assert loop_idxs[i].expr.startswith(prefix), f"loop_idxs[{i}] must start with {prefix}" - # assert loop_idxs[i].max+1 == sizes[i], f"loop_idxs[{i}] should have max {sizes[i]-1}" - - # no-op - _assert_grouped_dims("gidx", (2,), (16,16,16), False, [2]) - _assert_grouped_dims("gidx", (2,3), (16,16,16), False, [2,3]) - - # check reverse dims - _assert_grouped_dims("gidx", (2,3), (16,16,16), True, [3,2]) - _assert_grouped_dims("gidx", (2,3,4), (16,16,16), False, [2,3,4]) - - # test splitting globals: len(dims) == len(max) - _assert_grouped_dims("gidx", (64,3,4), (16,16,16), False, [16,12,4]) - _assert_grouped_dims("gidx", (64,3,4), (16,4,16), False, [16,3,16]) - _assert_grouped_dims("gidx", (64,3,4), (16,16,16), True, [16,3,16]) - _assert_grouped_dims("gidx", (128,3,4), (16,4,256), False, [16,3,32]) - _assert_grouped_dims("gidx", (4,4,512), (16,4,256), False, [8,4,256]) - - # prefer group_dim strategy when possible - _assert_grouped_dims("gidx", (512,4,2), (8192,2,2), False, [2048,2]) - - # test splitting globals: len(dims) < len(max) - # len(dim) -> len(limited) - # 1 -> 2 - _assert_grouped_dims("gidx", (128,), (16,16,256), False, [16,8], False) - # 1 -> 3 - _assert_grouped_dims("gidx", (65536,), (16,16,256), False, [16,16,256], False) - # 2 -> 3 - _assert_grouped_dims("gidx", (128,128), (16,16,256), False, [16,16,64], False) - # test when the only divisor is the square root of dim - _assert_grouped_dims("gidx", (121,), (12,12,12), False, [11,11], False) - - # collapse on onto the left most axis - _assert_grouped_dims("gidx", (2,3,4,5), (16,16,16), False, [6,4,5]) - _assert_grouped_dims("gidx", (2,3,4,5), (32,16,16), True, [20,3,2]) - # _assert_grouped_dims("gidx", (Variable("start_pos",1,2),3,4,5), (32,16,16), True, [20,3,Variable("start_pos",1,2)]) - - # collapse on left-most available axis (the left most is too small) - _assert_grouped_dims("gidx", (2,3,4,5), (4,16,16), False, [2,12,5]) - _assert_grouped_dims("gidx", (2,3,4,5), (16,16,16), True, [5,12,2]) - - # _assert_grouped_dims("gidx", (Variable("start_pos",1,2),3,4,5), (16,16,16), False, [Variable("start_pos",1,2)*3,4,5]) - - # dim too large and not factorable - with self.assertRaises(RuntimeError): - get_grouped_dims("gidx", (23,), (16,16,16), False,) - with self.assertRaises(RuntimeError): - get_grouped_dims("gidx", (128,3,4), (16,2,2), False,) - - # too large for sizes - with self.assertRaises(RuntimeError): - get_grouped_dims("gidx", (2,3,4,5,6), (16,16,16)) - - # # variable too large - # with self.assertRaises(AssertionError): - # get_grouped_dims("gidx", (Variable("start_pos",0,16),3,4), (16,16,16), False,) - + @unittest.skip("only one global now") @unittest.skipUnless(Device[Device.DEFAULT].renderer.has_local, "test requires locals") def test_default_global_reversed(self): # shrink so that the dims do not collapse diff --git a/tinygrad/codegen/gpudims.py b/tinygrad/codegen/gpudims.py index 8be324a6e6..280e0a3ab5 100644 --- a/tinygrad/codegen/gpudims.py +++ b/tinygrad/codegen/gpudims.py @@ -1,53 +1,15 @@ -import math from tinygrad.uop.ops import UOp, Ops, sint, PatternMatcher, UPat, KernelInfo, ssimplify, AxisType -from tinygrad.helpers import all_int, partition, flatten, prod, dedup +from tinygrad.helpers import partition, flatten, prod, dedup from tinygrad.dtype import dtypes -from tinygrad.shape.view import get_contraction from tinygrad.renderer import Renderer -def _group_dims(dims:tuple[sint, ...], max_sizes:tuple[int, ...]): - # TODO: symbolic shape - if not all_int(dims): return dims - while len(dims) > len(max_sizes) or any(d > m for d,m in zip(dims, max_sizes)): - for i,m in enumerate(max_sizes): - if i < (len(dims)-1) and dims[i] * dims[i+1] <= m: - dims = dims[:i] + (dims[i]*dims[i+1],) + dims[i+2:] - break - else: return None - return dims - -def _split_dims(dims, max_sizes): - if all(d <= m for d,m in zip(dims, max_sizes)): return dims - _dims = list(dims) + [1]*(3-len(dims)) - for i in range(len(_dims)): - while _dims[i] > max_sizes[i]: - div = next((d for d in range(2, math.ceil(math.sqrt(_dims[i])) + 1) if (_dims[i] % d) == 0), 1) - if div == 1: raise RuntimeError(f"cannot limit dim {dims=}, {max_sizes=}") - _dims[i], _dims[(i+1)%len(_dims)] = _dims[i]//div, _dims[(i+1)%len(_dims)]*div - return tuple(_dims[:2] if _dims[2] == 1 else _dims[0] if _dims[1:3] == [1,1] else _dims) - -def get_grouped_dims(prefix, dims:tuple[sint, ...], max_sizes:tuple[int, ...]|None, reverse=False) -> list[UOp]: +def get_grouped_dims(prefix, dims:tuple[sint, ...], reverse=False) -> list[UOp]: if reverse: dims = dims[::-1] - # try to group first: (a, b, c, d) -> (ab, c, d) - limited = (grouped if (grouped := _group_dims(dims, max_sizes)) else dims) if max_sizes is not None else dims - # check if grouping failed - if max_sizes is not None and len(limited) > len(max_sizes): raise RuntimeError(f"cannot limit dim {dims=}, {max_sizes=}") - # try to split up dims: (a,) -> (b, c) - if limited == dims: limited = _split_dims(dims, max_sizes) if max_sizes is not None else dims - ret = raw_idxs = [UOp(Ops.SPECIAL, dtypes.int, (), (f"{prefix}{i}", s)) for i,s in enumerate(limited)] - if len(limited) < len(dims): - ret = [] - if (contraction:=get_contraction(dims, limited)) is None: raise AssertionError(f"get_contraction should not be None {dims=} {limited=}") - for idx, contraction_group in zip(raw_idxs, contraction): - for c in contraction_group[:-1]: - ret.append(idx % dims[c]) - idx //= dims[c] - ret.append(idx) - elif len(limited) > len(dims): - a, b = len(limited), len(dims) - if a == 2 and b == 1: ret = [raw_idxs[0] * limited[1] + raw_idxs[1]] - if a == 3 and b == 1: ret = [raw_idxs[0] * (limited[1] * limited[2]) + raw_idxs[1] * limited[2] + raw_idxs[2]] - if a == 3 and b == 2: ret = [raw_idxs[0] * limited[1] + raw_idxs[1], raw_idxs[2]] + spec = UOp(Ops.SPECIAL, dtypes.int, (), (f"{prefix}0", ssimplify(prod(dims)))) + ret = [] + for d in dims: + ret.append(spec % d) + spec //= d return ret[::-1] if reverse else ret def add_gpudims(ctx:Renderer, s:UOp): @@ -72,10 +34,10 @@ def add_gpudims(ctx:Renderer, s:UOp): ki: KernelInfo = s.arg if ki.dont_use_locals: assert not local_dims, "can't use locals if there's no local dims" - idxs = get_grouped_dims("idx", global_shape, ctx.global_max, reverse=True) + idxs = get_grouped_dims("idx", global_shape, reverse=True) else: # define indexes for GPU-like execution - idxs = get_grouped_dims("gidx", global_shape, ctx.global_max, reverse=True) + get_grouped_dims("lidx", local_shape, ctx.local_max) + idxs = get_grouped_dims("gidx", global_shape, reverse=True) + get_grouped_dims("lidx", local_shape) # apply to multiple ranges subs = {}