From 3688afa51395eb8413e70d706d2ef055fb194348 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Thu, 2 Oct 2025 18:17:03 +0800 Subject: [PATCH] fix swap --- test/test_rangeify.py | 6 ++++-- tinygrad/codegen/opt/postrange.py | 6 +++--- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/test/test_rangeify.py b/test/test_rangeify.py index 84c19e44e1..ab221ab45f 100644 --- a/test/test_rangeify.py +++ b/test/test_rangeify.py @@ -248,8 +248,10 @@ class TestRangeify(unittest.TestCase): args += (Opt(OptOps.DEMOTE, 5, 8),) args += (Opt(OptOps.TC, 0, (0,0,1,3)),) args += (Opt(OptOps.TC, 0, (0,0,1,0)),) - args += (Opt(OptOps.WARP, 1, 32),) - args += (Opt(OptOps.WARP, 2, 32),) + args += (Opt(OptOps.SWAP, 1, 4),) + args += (Opt(OptOps.SWAP, 2, 5),) + args += (Opt(OptOps.WARP, 4, 32),) + args += (Opt(OptOps.WARP, 5, 32),) ret = fa().contiguous(arg=args).realize() with Context(RANGEIFY=0): with Context(DEBUG=2): diff --git a/tinygrad/codegen/opt/postrange.py b/tinygrad/codegen/opt/postrange.py index c06ebd66e1..6b991bf852 100644 --- a/tinygrad/codegen/opt/postrange.py +++ b/tinygrad/codegen/opt/postrange.py @@ -182,13 +182,13 @@ class Scheduler: mr = ctx[0] nr = mr.replace(arg=ctx[0].arg[0:-2]+(mr.arg[-2]+1, mr.arg[-1])) ctx[0] = nr - buf = x.replace(src=x.src+(mr,), tag=1).substitute({mr:nr}) + buf = x.replace(src=(x.src[0], mr)+x.src[1:], tag=1).substitute({mr:nr}) return UOp(Ops.APPENDINDEX, dtypes.void, (buf,mr)) # do the demotion pm_demote = PatternMatcher([ (UPat(Ops.BUFFERIZE, name="x"), do_demote), (UPat(Ops.INDEX, src=(UPat(Ops.APPENDINDEX, name="x"),), name="y", allow_any_len=True), - lambda x,y: y.replace(src=(x.src[0],)+y.src[1:]+x.src[1:])), + lambda x,y: y.replace(src=(x.src[0],)+x.src[1:]+y.src[1:])), ]) self.ast = graph_rewrite(self.ast, pm_demote, ctx=[rr], bottom_up=True, name="demote") elif opt.op is OptOps.TC: @@ -222,7 +222,7 @@ class Scheduler: altrng = self.rngs[opt.arg] except IndexError: raise KernelOptError - check(rng.arg[-1] == AxisType.GLOBAL and altrng.arg[-1] == AxisType.GLOBAL, "swap only for globals") + check(rng.arg[-1] == AxisType.LOOP and altrng.arg[-1] == AxisType.LOOP, "swap only for globals") self.ast = self.ast.substitute({rng:rng.replace(arg=(*altrng.arg[0:-1], rng.arg[-1]), tag=1), altrng:altrng.replace(arg=(*rng.arg[0:-1], altrng.arg[-1]), tag=1)}) self.ast = graph_rewrite(self.ast, remove_tags)