From 4b3fcb4064e516e6661f1641e2db778e39d1987e Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Mon, 18 Aug 2025 13:28:53 -0700 Subject: [PATCH] Revert "REDUCE_AXIS keepdim=False (#11311)" (#11718) This reverts commit b518a7378adb53750e158f21a877af62b598eed3. --- test/test_linearizer.py | 8 ++++---- test/test_linearizer_dumb.py | 22 +++++++++++----------- test/test_quantize_onnx.py | 6 +++--- test/unit/test_uop_spec.py | 2 +- tinygrad/codegen/lowerer.py | 11 ++++------- tinygrad/codegen/opt/kernel.py | 2 -- tinygrad/codegen/opt/swizzler.py | 13 ++++++------- tinygrad/gradient.py | 3 +-- tinygrad/schedule/grouper.py | 2 +- tinygrad/schedule/kernelize.py | 2 +- tinygrad/shape/shapetracker.py | 2 +- tinygrad/tensor.py | 2 +- tinygrad/uop/ops.py | 23 ++++++----------------- 13 files changed, 40 insertions(+), 58 deletions(-) diff --git a/test/test_linearizer.py b/test/test_linearizer.py index c310a5354d..72b5dc0f09 100644 --- a/test/test_linearizer.py +++ b/test/test_linearizer.py @@ -123,7 +123,7 @@ class TestLinearizer(unittest.TestCase): idxs = Tensor([0,3,5,6]).realize() with Context(FUSE_ARANGE=1): sink = dataset[idxs].contiguous().kernelize().uop.base.src[1].arg.ast - real_index = dataset.numpy()[idxs.numpy()].reshape(4, 256) + real_index = dataset.numpy()[idxs.numpy()].reshape(4, 256, 1, 1) helper_linearizer_ast(push_views(sink), [dataset, idxs], wanna_output=[real_index]) def test_two_nested_range(self): @@ -163,7 +163,7 @@ class TestLinearizer(unittest.TestCase): a = Tensor.randn(4, 1).realize() b = Tensor.randn(1, 1).realize() out = (a + b[0]).sum() + b[0] - lin = helper_linearizer_opt(out, wanna_output=[(a.numpy()+b.numpy()[0]).sum()+b.numpy()[0]])[0] + lin = helper_linearizer_opt(out, wanna_output=[(a.numpy()+b.numpy()[0]).sum()+b.numpy()])[0] uops = get_program(lin.get_optimized_ast(), lin.opts).uops ranges = [i for i,u in enumerate(uops) if u.op is Ops.RANGE] # LOAD -> RANGE -> LOAD -> ASSIGN @@ -173,7 +173,7 @@ class TestLinearizer(unittest.TestCase): a = Tensor.randn(2, ).realize() b = Tensor.randn(1, 1).realize() out = (a.reshape(2, 1).expand(2, 3) + b[0]).sum() + b[0] - lin = helper_linearizer_opt(out, wanna_output=[(np.broadcast_to(a.numpy().reshape(2, 1), (2, 3)) + b.numpy()[0]).sum() + b.numpy()[0]])[0] + lin = helper_linearizer_opt(out, wanna_output=[(np.broadcast_to(a.numpy().reshape(2, 1), (2, 3)) + b.numpy()[0]).sum() + b.numpy()])[0] uops = get_program(lin.get_optimized_ast(), lin.opts).uops ranges = [i for i,u in enumerate(uops) if u.op is Ops.RANGE] assert len(ranges) == 1 # NOTE: it collapses now @@ -867,7 +867,7 @@ class TestFloat4(unittest.TestCase): # from llama 7B shard 4 gpus ast = UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.float.ptr(96000), arg=ShapeTracker(views=(View(shape=(1, 3, 32000,), strides=(0, 32000, 1,), offset=0, mask=None, contiguous=True),)), src=( # noqa: E501 + UOp(Ops.VIEW, dtypes.float.ptr(96000), arg=ShapeTracker(views=(View(shape=(1, 3, 32000, 1), strides=(0, 32000, 1, 0), offset=0, mask=None, contiguous=True),)), src=( # noqa: E501 UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(96000), arg=0, src=()),)), UOp(Ops.REDUCE_AXIS, dtypes.float, arg=(Ops.ADD, (3,)), src=( UOp(Ops.CAST, dtypes.float, arg=None, src=( diff --git a/test/test_linearizer_dumb.py b/test/test_linearizer_dumb.py index a3d9de94ae..c8ebbbb365 100644 --- a/test/test_linearizer_dumb.py +++ b/test/test_linearizer_dumb.py @@ -17,7 +17,7 @@ class TestLinearizerDumb(unittest.TestCase): def test_unmerged_ifs(self): ast = UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.half.ptr(1605632), arg=ShapeTracker(views=(View(shape=(64, 1, 512, 7, 7), strides=(25088, 0, 49, 7, 1), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.half.ptr(1605632), arg=ShapeTracker(views=(View(shape=(64, 1, 512, 7, 7, 1, 1, 1), strides=(25088, 0, 49, 7, 1, 0, 0, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.half.ptr(1605632), arg=0, src=()),)), UOp(Ops.MAX, dtypes.half, arg=None, src=( UOp(Ops.MUL, dtypes.half, arg=None, src=( @@ -32,7 +32,7 @@ class TestLinearizerDumb(unittest.TestCase): UOp(Ops.VIEW, dtypes.half.ptr(2359296), arg=ShapeTracker(views=(View(shape=(64, 1, 512, 7, 7, 512, 3, 3), strides=(0, 0, 4608, 0, 0, 9, 3, 1), offset=0, mask=None, contiguous=False),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.half.ptr(2359296), arg=2, src=()),)),)),)),)),)),)), UOp(Ops.CONST, dtypes.half, arg=0.9999950000374996, src=( - x16:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(64, 1, 512, 7, 7), strides=(0, 0, 0, 0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)), + x16:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(64, 1, 512, 7, 7, 1, 1, 1), strides=(0, 0, 0, 0, 0, 0, 0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)), UOp(Ops.CONST, dtypes.half, arg=0.0, src=( x16,)),)),)),)) opts = [Opt(op=OptOps.TC, axis=2, arg=(-1, 2, 1)), Opt(op=OptOps.UPCAST, axis=2, arg=0), Opt(op=OptOps.UNROLL, axis=1, arg=0)] @@ -49,20 +49,20 @@ class TestLinearizerDumb(unittest.TestCase): def test_max_simplify_and_cancel(self): ast = UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.int.ptr(1000), arg=ShapeTracker(views=(View(shape=(1000,), strides=(1,), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.int.ptr(1000), arg=ShapeTracker(views=(View(shape=(1000, 1), strides=(1, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.int.ptr(1000), arg=0, src=()),)), UOp(Ops.MUL, dtypes.int, arg=None, src=( UOp(Ops.CAST, dtypes.int, arg=None, src=( UOp(Ops.CMPNE, dtypes.bool, arg=None, src=( UOp(Ops.CMPNE, dtypes.bool, arg=None, src=( UOp(Ops.LOAD, dtypes.float, arg=None, src=( - UOp(Ops.VIEW, dtypes.float.ptr(1000), arg=ShapeTracker(views=(View(shape=(1000,), strides=(1,), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.float.ptr(1000), arg=ShapeTracker(views=(View(shape=(1000, 1), strides=(1, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(1000), arg=1, src=()),)),)), UOp(Ops.LOAD, dtypes.float, arg=None, src=( - UOp(Ops.VIEW, dtypes.float.ptr(1), arg=ShapeTracker(views=(View(shape=(1000,), strides=(0,), offset=0, mask=None, contiguous=False),)), src=( + UOp(Ops.VIEW, dtypes.float.ptr(1), arg=ShapeTracker(views=(View(shape=(1000, 1), strides=(0, 0), offset=0, mask=None, contiguous=False),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(1), arg=2, src=()),)),)),)), UOp(Ops.CONST, dtypes.bool, arg=True, src=( - x14:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(1000,), strides=(0,), offset=0, mask=None, contiguous=False),)), src=()),)),)),)), + x14:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(1000, 1), strides=(0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)),)), UOp(Ops.ADD, dtypes.int, arg=None, src=( UOp(Ops.REDUCE_AXIS, dtypes.int, arg=(Ops.ADD, (1,)), src=( UOp(Ops.WHERE, dtypes.int, arg=None, src=( @@ -86,7 +86,7 @@ class TestLinearizerDumb(unittest.TestCase): def test_expander_new_srcs(self): ast = UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.float.ptr(25), arg=ShapeTracker(views=(View(shape=(25,), strides=(1,), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.float.ptr(25), arg=ShapeTracker(views=(View(shape=(25, 1), strides=(1, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(25), arg=0, src=()),)), UOp(Ops.REDUCE_AXIS, dtypes.float, arg=(Ops.ADD, (1,)), src=( UOp(Ops.LOAD, dtypes.float, arg=None, src=( @@ -107,7 +107,7 @@ class TestLinearizerDumb(unittest.TestCase): def test_llama_embedding(self): ast = UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.half.ptr(4096), arg=ShapeTracker(views=(View(shape=(4096,), strides=(1,), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.half.ptr(4096), arg=ShapeTracker(views=(View(shape=(4096, 1, 1), strides=(1, 0, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.half.ptr(4096), arg=0, src=()),)), UOp(Ops.CAST, dtypes.half, arg=None, src=( UOp(Ops.REDUCE_AXIS, dtypes.float, arg=(Ops.ADD, (1,)), src=( @@ -126,14 +126,14 @@ class TestLinearizerDumb(unittest.TestCase): UOp(Ops.CONST, dtypes.int, arg=0, src=( x16,)),)),)), UOp(Ops.CONST, dtypes.int, arg=-1, src=( - x19:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(4096, 32000), strides=(0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)), + x19:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(4096, 32000, 1), strides=(0, 0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)), UOp(Ops.LOAD, dtypes.int, arg=None, src=( - UOp(Ops.VIEW, dtypes.int.ptr(1), arg=ShapeTracker(views=(View(shape=(4096, 32000), strides=(0, 0), offset=0, mask=None, contiguous=False),)), src=( + UOp(Ops.VIEW, dtypes.int.ptr(1), arg=ShapeTracker(views=(View(shape=(4096, 32000, 1), strides=(0, 0, 0), offset=0, mask=None, contiguous=False),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.int.ptr(1), arg=1, src=()),)),)),)), UOp(Ops.CONST, dtypes.bool, arg=True, src=( x19,)),)),)), UOp(Ops.LOAD, dtypes.half, arg=None, src=( - UOp(Ops.VIEW, dtypes.half.ptr(131072000), arg=ShapeTracker(views=(View(shape=(4096, 32000), strides=(1, 4096), offset=0, mask=None, contiguous=False),)), src=( + UOp(Ops.VIEW, dtypes.half.ptr(131072000), arg=ShapeTracker(views=(View(shape=(4096, 32000, 1), strides=(1, 4096, 0), offset=0, mask=None, contiguous=False),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.half.ptr(131072000), arg=2, src=()),)),)),)),)),)),)),)),)) k = Kernel(ast, opts=Device[Device.DEFAULT].renderer) prg = get_program(k.get_optimized_ast(), k.opts) diff --git a/test/test_quantize_onnx.py b/test/test_quantize_onnx.py index feb852483c..1555b6befd 100644 --- a/test/test_quantize_onnx.py +++ b/test/test_quantize_onnx.py @@ -242,7 +242,7 @@ class TestDSPCache(unittest.TestCase): # string becuase this breaks Python language server for syntax highlight for some reason ast = eval("""UOp(Ops.SINK, dtypes.void, arg=None, src=( UOp(Ops.STORE, dtypes.void, arg=None, src=( - UOp(Ops.VIEW, dtypes.uchar.ptr(25088), arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32), strides=(0, 896, 32, 1), offset=0, mask=None, contiguous=True),)), src=( + UOp(Ops.VIEW, dtypes.uchar.ptr(25088), arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32, 1), strides=(0, 896, 32, 1, 0), offset=0, mask=None, contiguous=True),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.uchar.ptr(25088), arg=0, src=()),)), UOp(Ops.CAST, dtypes.uchar, arg=None, src=( UOp(Ops.XOR, dtypes.int, arg=None, src=( @@ -275,10 +275,10 @@ class TestDSPCache(unittest.TestCase): UOp(Ops.MUL, dtypes.float, arg=None, src=( UOp(Ops.CAST, dtypes.float, arg=None, src=( UOp(Ops.LOAD, dtypes.int, arg=None, src=( - UOp(Ops.VIEW, dtypes.int.ptr(32), arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32), strides=(0, 0, 0, 1), offset=0, mask=None, contiguous=False),)), src=( + UOp(Ops.VIEW, dtypes.int.ptr(32), arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32, 1), strides=(0, 0, 0, 1, 0), offset=0, mask=None, contiguous=False),)), src=( UOp(Ops.DEFINE_GLOBAL, dtypes.int.ptr(32), arg=3, src=()),)),)),)), UOp(Ops.CONST, dtypes.float, arg=9.203465015161783e-05, src=( - x36:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32), strides=(0, 0, 0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)),)), + x36:=UOp(Ops.VIEW, dtypes.void, arg=ShapeTracker(views=(View(shape=(1, 28, 28, 32, 1), strides=(0, 0, 0, 0, 0), offset=0, mask=None, contiguous=False),)), src=()),)),)),)), UOp(Ops.CONST, dtypes.float, arg=33.812857328652136, src=( x36,)),)), UOp(Ops.CONST, dtypes.float, arg=0.4999999, src=( diff --git a/test/unit/test_uop_spec.py b/test/unit/test_uop_spec.py index a8ed31f4ae..0244ba3531 100644 --- a/test/unit/test_uop_spec.py +++ b/test/unit/test_uop_spec.py @@ -66,7 +66,7 @@ class TestUOpSpec(unittest.TestCase): buf = UOp(Ops.DEFINE_GLOBAL, dtypes.float.ptr(), (), 0) a = UOp(Ops.LOAD, dtypes.float, (buf.view(ShapeTracker.from_shape((32, 1))),)) r = UOp(Ops.REDUCE_AXIS, dtypes.float, (a,), (Ops.ADD, (0,))) - st = UOp.store(buf.view(ShapeTracker.from_shape((32,))), r.view(r.st.expand((32,)))+a) + st = UOp.store(buf.view(ShapeTracker.from_shape((32, 1))), r.view(r.st.expand((32, 1)))+a) with self.assertRaisesRegex(InvalidASTException, "UOp verification failed"): helper_test_verify_ast(st) def test_const_view_always_valid(self): diff --git a/tinygrad/codegen/lowerer.py b/tinygrad/codegen/lowerer.py index ba43fd01d6..06e8cc5c3d 100644 --- a/tinygrad/codegen/lowerer.py +++ b/tinygrad/codegen/lowerer.py @@ -39,8 +39,7 @@ def subblock(ctx: IndexContext, full_new_idx: list[UOp], src: UOp): return graph_rewrite(src, pm_lowerer, lc, name="subblock", bottom_up=True) def lower_reduce_axis(ctx: IndexContext, x: UOp): - src_shape = x.src[0].shape_with_reduced if x.axis_arg and len(x.src[0].shape) < (max(x.axis_arg)+1) else x.src[0].shape - new_idxs = shape_to_idx(src_shape, ctx.axis_types, ctx.start) + new_idxs = shape_to_idx(x.src[0].shape, ctx.axis_types, ctx.start) full_new_idx = list(ctx.idxs) for a in x.axis_arg: full_new_idx[a] = new_idxs[a] @@ -58,12 +57,11 @@ def lower_store(ctx: IndexContext, x: UOp, buf: UOp): # TODO: reenable after REDUCE_AXIS is fixed #assert x.src[1].shape == x.src[0].shape, f"shape mismatch on store {x.src[1].shape} != {x.src[0].shape}" - src_shape = x.src[0].shape + (1,) * (len(x.full_shape)-len(x.src[0].shape)) - new_idxs = shape_to_idx(src_shape, ctx.axis_types, ctx.start) + new_idxs = shape_to_idx(x.src[0].shape, ctx.axis_types, ctx.start) idx, valid = x.st_arg.to_indexed_uops(new_idxs) used_idxs = [x for x in UOp.sink(idx, valid).toposort() if x in new_idxs] real_new_idxs = [] - for i in range(len(src_shape)): + for i in range(len(x.src[0].shape)): if new_idxs[i] in used_idxs or len(ctx.idxs) <= i: real_new_idxs.append(new_idxs[i]) else: real_new_idxs.append(ctx.idxs[i]) @@ -81,8 +79,7 @@ def lower_store(ctx: IndexContext, x: UOp, buf: UOp): def fixup_wmma(ctx:IndexContext, x:UOp): if x.tag is not None: return None - src_shape = x.src[0].shape_with_reduced if x.axis_arg and len(x.src[0].shape) < (max(x.axis_arg)+1) else x.src[0].shape - new_idxs = shape_to_idx(src_shape, ctx.axis_types, ctx.start) + new_idxs = shape_to_idx(x.src[0].shape, ctx.axis_types, ctx.start) full_new_idx = list(ctx.idxs) for a in x.arg[-1]: full_new_idx[a] = new_idxs[a] diff --git a/tinygrad/codegen/opt/kernel.py b/tinygrad/codegen/opt/kernel.py index f0b1c2c946..3df7100e78 100644 --- a/tinygrad/codegen/opt/kernel.py +++ b/tinygrad/codegen/opt/kernel.py @@ -76,8 +76,6 @@ class Kernel: full_shape = ast.full_shape self.sts.append(ShapeTracker.from_shape(full_shape, (0,)*len(full_shape))) - self.sts = [st.reshape(st.shape + (1,) * (len(full_shape) - len(st.shape))) for st in self.sts] - # parameters for optimization self.tensor_core: TensorCore|None = None self.tensor_core_opts: TensorCoreOptions|None = None diff --git a/tinygrad/codegen/opt/swizzler.py b/tinygrad/codegen/opt/swizzler.py index 1fd88cc7d4..4e4b6bce46 100644 --- a/tinygrad/codegen/opt/swizzler.py +++ b/tinygrad/codegen/opt/swizzler.py @@ -22,11 +22,11 @@ merge_views = PatternMatcher([ def reduce_push_add_ones(src:UOp, r:UOp, view:UOp): # contiguous, expand, and the same with ones removed - if unwrap(view.st).contiguous and len(r.shape_with_reduced) < len(view.shape) and \ + if unwrap(view.st).contiguous and len(r.shape) < len(view.shape) and \ tuple(x for x in r.shape if resolve(x != 1)) == tuple(x for x in view.shape if resolve(x != 1)): new_shape: list[sint] = [] new_reduce_axis = [] - if (contraction:=get_contraction_with_reduce(view.shape, r.shape_with_reduced, r.arg[1])) is None: return None + if (contraction:=get_contraction_with_reduce(view.shape, r.shape, r.arg[1])) is None: return None for i,pairs in enumerate(contraction): new_shape_chunk = [view.shape[p] for p in pairs] if i in r.arg[1]: @@ -37,7 +37,7 @@ def reduce_push_add_ones(src:UOp, r:UOp, view:UOp): else: # otherwise, pass through the new_shape_chunk new_shape += new_shape_chunk - ret = r.replace(src=(src.reshape(tuple(new_shape)),), arg=(r.arg[0], tuple(new_reduce_axis))+r.arg[2:]).reshape(view.shape) + ret = r.replace(src=(src.reshape(tuple(new_shape)),), arg=(r.arg[0], tuple(new_reduce_axis))+r.arg[2:]) assert ret.shape == view.shape, f"shape mismatch on reduce_push_add_ones, {ret.shape} != {view.shape}" return ret return None @@ -62,11 +62,10 @@ def apply_swizzle(u:UOp) -> UOp: return graph_rewrite(u, view_left, name="Sub Vi def swizzle_reduceop(r:UOp, src:UOp, view:UOp, fuse=False): # contiguous and same size can push to children # if there's a reduce child, shapes match with ones removed - if unwrap(view.st).contiguous and view.shape == r.shape_with_reduced and \ + if unwrap(view.st).contiguous and view.size == r.size and \ (not (len(r.arg) == 3 and r.arg[2]) or # arg[2] = True is fuse marker tuple((i,x) for i,x in enumerate(r.shape) if resolve(x != 1)) == tuple((i,x) for i,x in enumerate(view.shape) if resolve(x != 1))): return None - if isinstance(r.dtype, ImageDType): return None # swizzle the input input_st = ShapeTracker.from_shape(src.shape) tmp = input_st.permute(tuple(i for i in range(len(input_st.shape)) if i not in r.axis_arg)+r.axis_arg) @@ -84,8 +83,8 @@ def swizzle_reduceop(r:UOp, src:UOp, view:UOp, fuse=False): def reduceop_view_right(src:UOp, v:UOp, r:UOp): assert unwrap(v.st).contiguous and v.size == src.size, f"can't compute new axis for {src.shape} -> {r.shape}" - new_axis = [i for i,(s,u) in enumerate(zip(src.shape, r.shape_with_reduced)) if s != u] - return src.r(r.arg[0], tuple(new_axis)).reshape(r.shape_with_reduced) + new_axis = [i for i,(s,u) in enumerate(zip(src.shape, r.shape)) if s != u] + return src.r(r.arg[0], tuple(new_axis)).reshape(r.shape) def elementwise_view_right(root:UOp): if not (swizzles:=[x for x in root.src if x.op is Ops.VIEW and x.base.op not in ALWAYS_CONTIGUOUS]): return None diff --git a/tinygrad/gradient.py b/tinygrad/gradient.py index ef6d8d4f34..77979148c8 100644 --- a/tinygrad/gradient.py +++ b/tinygrad/gradient.py @@ -37,8 +37,7 @@ pm_gradient = PatternMatcher([ (UPat(Ops.PAD, name="ret"), lambda ctx, ret: (ctx.shrink(tuple([(p[0], s+p[0]) for s,p in zip(ret.src[0].shape, ret.arg)])),)), (UPat(Ops.SHRINK, name="ret"), lambda ctx, ret: (ctx.pad(tuple([(p[0], s-p[1]) for s,p in zip(ret.src[0].shape, ret.arg)])),)), (UPat(Ops.FLIP, name="ret"), lambda ctx, ret: (ctx.flip(ret.arg),)), - (UPat(Ops.EXPAND, name="ret"), lambda ctx, ret: (ctx.r(Ops.ADD, tuple(i for i,(si,so) in enumerate(zip(ret.src[0].shape, ret.arg)) \ - if si!=so)).reshape(ret.src[0].shape),)), + (UPat(Ops.EXPAND, name="ret"), lambda ctx, ret: (ctx.r(Ops.ADD, tuple(i for i,(si,so) in enumerate(zip(ret.src[0].shape, ret.arg)) if si!=so)),)), (UPat(Ops.MULTI, name="ret"), lambda ctx, ret: ctx.shard(ret.device, ret.axis).src), # there's no gradient for bitcast (UPat(Ops.BITCAST), lambda ctx: (None,)), diff --git a/tinygrad/schedule/grouper.py b/tinygrad/schedule/grouper.py index 05f728a4fa..685bc70b7a 100644 --- a/tinygrad/schedule/grouper.py +++ b/tinygrad/schedule/grouper.py @@ -71,7 +71,7 @@ def group_realizes(sink:UOp) -> dict[UOp, None]: for r in toposort: if r.op is not Ops.REDUCE_AXIS: continue if len(r.arg) == 3 and r.arg[2] is True: continue - if FUSE_CONV_BW and r.src[0].base.op is Ops.REDUCE_AXIS: double_reduces.append(r) + if FUSE_CONV_BW and r.src[0].base.op is Ops.REDUCE_AXIS and r.src[0] is not r.src[0].base: double_reduces.append(r) if r in realizes: continue group: dict[UOp, None] = {} recursive_group(r, unwrap(r.st), r, children, realizes, reduce_for_op, group, cache={}) diff --git a/tinygrad/schedule/kernelize.py b/tinygrad/schedule/kernelize.py index a5fdd9c42b..d02ec8fc0e 100644 --- a/tinygrad/schedule/kernelize.py +++ b/tinygrad/schedule/kernelize.py @@ -21,7 +21,7 @@ def simplify_stride0_reduce(reduce:UOp, x:UOp): # must have all stride 0 in the relevant axis (NOTE: can do partial) if not all(unwrap(x.st).views[-1].strides[axis] == 0 for axis in reduce.arg[1]) or not all_int(x.shape): return None prshape = prod(x.shape[i] for i in reduce.arg[1]) - ret = x.shrink(tuple((0,s) if i not in reduce.arg[1] else (0,1) for i,s in enumerate(x.shape))).reshape(reduce.shape) + ret = x.shrink(tuple((0,s) if i not in reduce.arg[1] else (0,1) for i,s in enumerate(x.shape))) match reduce.arg[0]: case Ops.ADD: return ret*prshape case Ops.MUL: return ret.pow(prshape) diff --git a/tinygrad/shape/shapetracker.py b/tinygrad/shape/shapetracker.py index 33e8f91808..3e9b09f7d7 100644 --- a/tinygrad/shape/shapetracker.py +++ b/tinygrad/shape/shapetracker.py @@ -81,7 +81,7 @@ class ShapeTracker: @property def size(self) -> int: return self.views[-1].size() - def reduce(self, axis:tuple[int, ...]) -> tuple[sint, ...]: return tuple(s for i,s in enumerate(self.shape) if i not in axis) + def reduce(self, axis:tuple[int, ...]) -> tuple[sint, ...]: return tuple(1 if i in axis else s for i,s in enumerate(self.shape)) def to_indexed_uops(self, _idxs:list[UOp]|tuple[UOp, ...]|None=None) -> tuple[UOp, UOp]: return views_to_indexed_uops(self.views, tuple(_idxs) if _idxs is not None else None) diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index fd1f8d09d9..0de027819c 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -1691,7 +1691,7 @@ class Tensor(MathTrait): axis = tuple(self._resolve_dim(x) for x in (range(self.ndim) if axis is None else make_tuple(axis, 1))) if self.ndim == 0: axis = () ret = self._apply_uop(UOp.r, op=op, axis=axis) - return ret if not keepdim else ret.reshape(tuple([s if i not in axis else 1 for i,s in enumerate(self.shape)])) + return ret if keepdim else ret.reshape(tuple(s for i,s in enumerate(self.shape) if i not in axis)) def sum(self, axis:int|Sequence[int]|None=None, keepdim=False, dtype:DTypeLike|None=None) -> Tensor: """ diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index dea47c7e28..f26e02518b 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -134,17 +134,6 @@ class UOp(MathTrait, metaclass=UOpMetaClass): # *** uop shape stuff *** - @functools.cached_property - def reduced(self) -> tuple[int, ...]: - return tuple((set(self.axis_arg) if self.op in (Ops.REDUCE_AXIS, Ops.WMMA) else set()).union( - *(x.reduced for x in self.src if x.op not in GroupOp.Movement and x.op != Ops.VIEW))) - - @functools.cached_property - def shape_with_reduced(self) -> tuple[sint, ...]: - shape = self.shape - for i in sorted(self.reduced): shape = shape[:i] + (1,) + shape[i:] - return shape - @functools.cached_property def st(self) -> ShapeTracker|None: if self.op in GroupOp.Block or self.op is Ops.INDEX: return None @@ -170,13 +159,13 @@ class UOp(MathTrait, metaclass=UOpMetaClass): # otherwise we get the shape from sources if not (src_sts := [x.st for x in self.src if x.st is not None]): return None - # assert all_same([x.shape for x in src_sts]), f"UOp sources must have the same shape {self} {[x.shape for x in src_sts]}" + assert all_same([x.shape for x in src_sts]), f"UOp sources must have the same shape {self} {[x.shape for x in src_sts]}" match self.op: case Ops.MULTI: shape = tuple(self.src[0].shape[a]*len(self.device) if a == self.axis else s for a,s in enumerate(self.src[0].shape)) case Ops.BITCAST: shape = src_sts[0].shape if self.dtype.itemsize != (input_sz:=self.src[0].dtype.itemsize): shape = shape[:-1]+((shape[-1]*input_sz) // self.dtype.itemsize,) - case Ops.REDUCE_AXIS | Ops.WMMA: shape = src_sts[0].reduce(self.reduced) + case Ops.REDUCE_AXIS | Ops.WMMA: shape = src_sts[0].reduce(self.axis_arg) case _: shape = src_sts[0].shape return ShapeTracker.from_shape(shape) @@ -291,16 +280,16 @@ class UOp(MathTrait, metaclass=UOpMetaClass): @staticmethod def range(dtype:DType, end:sint, idx:int): return UOp(Ops.RANGE, dtype=dtype, src=(sint_to_uop(end),), arg=idx) def r(self, op:Ops, axis:tuple[int, ...]): - axis = tuple(sorted(x for x in axis)) + axis = tuple(sorted([x for x in axis if resolve(self.shape[x] != 1)])) if len(axis) == 0: return self # move any non reduce axis before the first reduce axis - move_early, rest = partition(range(axis[0], len(self.shape)), lambda i: i not in axis) + move_early, rest = partition(range(axis[0], len(self.shape)), lambda i: i not in axis and resolve(self.shape[i] != 1)) permaxis = tuple(range(axis[0])) + tuple(move_early) + tuple(rest) ret = self.permute(permaxis) - new_axis = tuple(x for x in range(axis[0]+len(move_early), len(self.shape))) + new_axis = tuple([x for x in range(axis[0]+len(move_early), len(self.shape)) if resolve(ret.shape[x] != 1)]) assert len(axis) == len(new_axis) ret = UOp(Ops.REDUCE_AXIS, self.dtype, (ret,), (op, new_axis)) - return ret.reshape(tuple(x for i,x in enumerate(self.shape) if i not in axis)) + return ret.reshape(tuple([x if i not in axis else 1 for i,x in enumerate(self.shape)])) def reduce(self, *src:UOp, **kwargs): return UOp(Ops.REDUCE, kwargs.pop('dtype', self.dtype), src=(self,)+src, **kwargs) def contiguous(self, *args, **kwargs): return UOp(Ops.CONTIGUOUS, dtype=self.dtype, src=(self,)+args, **kwargs) def contiguous_backward(self): return self.alu(Ops.CONTIGUOUS_BACKWARD)