From 5e433685b1f7433d84d3cf2eddaaff0719630fd0 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 27 Jul 2025 19:21:08 -0700 Subject: [PATCH] reorder --- extra/gemm/amd_uop_matmul.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/extra/gemm/amd_uop_matmul.py b/extra/gemm/amd_uop_matmul.py index b54f0e6188..ac827d567a 100644 --- a/extra/gemm/amd_uop_matmul.py +++ b/extra/gemm/amd_uop_matmul.py @@ -152,8 +152,8 @@ def hand_spec_kernel3(): # do the GEMM math iterWaveM = UOp.range(dtypes.int, nbIterWaveM, 8) - iterWaveN = UOp.range(dtypes.int, nbIterWaveN, 9) - yt = UOp.range(dtypes.int, TM, 10) + yt = UOp.range(dtypes.int, TM, 9) + iterWaveN = UOp.range(dtypes.int, nbIterWaveN, 10) xt = UOp.range(dtypes.int, TN, 11) x = iterWaveN * TN + xt y = iterWaveM * TM + yt @@ -163,8 +163,8 @@ def hand_spec_kernel3(): # store c_regs into c iterWaveM = UOp.range(dtypes.int, nbIterWaveM, 12) - iterWaveN = UOp.range(dtypes.int, nbIterWaveN, 13) - yt = UOp.range(dtypes.int, TM, 14) + yt = UOp.range(dtypes.int, TM, 13) + iterWaveN = UOp.range(dtypes.int, nbIterWaveN, 14) xt = UOp.range(dtypes.int, TN, 15) xOut = blockIdx_x * BN + waveIdx * WN + iterWaveN * SUBWN + TN * idxInWave yOut = blockIdx_y * BM + waveIdy * WM + iterWaveM * SUBWM + TM * idyInWave