From 2f0690dcc24371ef0d0dcd063ea60d3b366499b3 Mon Sep 17 00:00:00 2001 From: chenyu Date: Fri, 3 Jul 2026 11:23:23 -0400 Subject: [PATCH] minor gpudims cleanup [PR] (#16843) --- tinygrad/codegen/gpudims.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/tinygrad/codegen/gpudims.py b/tinygrad/codegen/gpudims.py index 4688470e8f..96b6152c8f 100644 --- a/tinygrad/codegen/gpudims.py +++ b/tinygrad/codegen/gpudims.py @@ -1,6 +1,5 @@ import math from tinygrad.uop.ops import UOp, Ops, sint, PatternMatcher, UPat, KernelInfo, ssimplify, AxisType -from tinygrad.helpers import dedup from tinygrad.dtype import dtypes, AddrSpace, Invalid from tinygrad.renderer import Renderer @@ -23,7 +22,7 @@ def _split_dims(dims, max_sizes): 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) + return tuple(_dims[:2] if _dims[2] == 1 else _dims) def get_grouped_dims(prefix, dims:tuple[sint, ...], max_sizes:tuple[int, ...]|None, reverse=False) -> list[UOp]: if reverse: return get_grouped_dims(prefix, dims[::-1], max_sizes)[::-1] @@ -48,14 +47,13 @@ def add_gpudims(ctx:Renderer, s:UOp): all_ranges = {x.arg[0:-1]:x for x in s_topo if x.op is Ops.RANGE} # extract global/local dims - global_dims = sorted(dedup([x.arg[0:-1] for x in all_ranges.values() if x.arg[-1] in (AxisType.GLOBAL, AxisType.THREAD)])) - local_dims = sorted(dedup([x.arg[0:-1] for x in all_ranges.values() if x.arg[-1] in (AxisType.WARP, AxisType.LOCAL, AxisType.GROUP_REDUCE)])) + global_dims = sorted([x.arg[0:-1] for x in all_ranges.values() if x.arg[-1] in (AxisType.GLOBAL, AxisType.THREAD)]) + local_dims = sorted([x.arg[0:-1] for x in all_ranges.values() if x.arg[-1] in (AxisType.WARP, AxisType.LOCAL, AxisType.GROUP_REDUCE)]) if not global_dims and not local_dims: return None # get global and local shape - ranges = [all_ranges[r] for r in global_dims+local_dims if r in all_ranges] - global_shape = tuple([ssimplify(r.src[0]) for r in ranges if r.arg[0:-1] in global_dims]) - local_shape = tuple([ssimplify(r.src[0]) for r in ranges if r.arg[0:-1] in local_dims]) + global_shape = tuple(ssimplify(all_ranges[r].src[0]) for r in global_dims) + local_shape = tuple(ssimplify(all_ranges[r].src[0]) for r in local_dims) # get the idxs ki: KernelInfo = s.arg