From f46c34db6857b72ffe8b066d8064f00ddf8017a8 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 24 Jul 2026 23:32:00 +0300 Subject: [PATCH] add num_stages 2 to amd --- modules/sdnq/kernels/triton_atten.py | 2 +- modules/sdnq/kernels/triton_mm.py | 2 +- modules/sdnq/kernels/triton_scaled_mm.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/modules/sdnq/kernels/triton_atten.py b/modules/sdnq/kernels/triton_atten.py index e23b994c8..8c4c60444 100644 --- a/modules/sdnq/kernels/triton_atten.py +++ b/modules/sdnq/kernels/triton_atten.py @@ -15,7 +15,7 @@ matmul_configs = [ for BM in [int(BM) for BM in os.environ.get("SDNQ_TRITON_ATTEN_BLOCK_SIZE_M_LIST", "64,128").replace(" ","").split(",")] for BN in [int(BN) for BN in os.environ.get("SDNQ_TRITON_ATTEN_BLOCK_SIZE_N_LIST", "32,64").replace(" ","").split(",")] for w in [int(w) for w in os.environ.get("SDNQ_TRITON_ATTEN_NUM_WARPS_LIST", "8,16" if torch.xpu.is_available() else "4,8").replace(" ","").split(",")] - for s in [int(s) for s in os.environ.get("SDNQ_TRITON_ATTEN_NUM_STAGES_LIST", "1" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] + for s in [int(s) for s in os.environ.get("SDNQ_TRITON_ATTEN_NUM_STAGES_LIST", "1,2" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] ] diff --git a/modules/sdnq/kernels/triton_mm.py b/modules/sdnq/kernels/triton_mm.py index be49a94fe..7479c14b5 100644 --- a/modules/sdnq/kernels/triton_mm.py +++ b/modules/sdnq/kernels/triton_mm.py @@ -14,7 +14,7 @@ matmul_configs = [ for BK in [int(BK) for BK in os.environ.get("SDNQ_TRITON_MM_BLOCK_SIZE_K_LIST", "32,64,128").replace(" ","").split(",")] for GM in [int(GM) for GM in os.environ.get("SDNQ_TRITON_MM_GROUP_SIZE_M_LIST", "8").replace(" ","").split(",")] for w in [int(w) for w in os.environ.get("SDNQ_TRITON_MM_NUM_WARPS_LIST", "16" if torch.xpu.is_available() else "4").replace(" ","").split(",")] - for s in [int(s) for s in os.environ.get("SDNQ_TRITON_MM_NUM_STAGES_LIST", "1" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] + for s in [int(s) for s in os.environ.get("SDNQ_TRITON_MM_NUM_STAGES_LIST", "1,2" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] ] diff --git a/modules/sdnq/kernels/triton_scaled_mm.py b/modules/sdnq/kernels/triton_scaled_mm.py index c2c1d8689..c546b4ecf 100644 --- a/modules/sdnq/kernels/triton_scaled_mm.py +++ b/modules/sdnq/kernels/triton_scaled_mm.py @@ -14,7 +14,7 @@ matmul_configs = [ for BK in [int(BK) for BK in os.environ.get("SDNQ_TRITON_MM_BLOCK_SIZE_K_LIST", "32,64,128").replace(" ","").split(",")] for GM in [int(GM) for GM in os.environ.get("SDNQ_TRITON_MM_GROUP_SIZE_M_LIST", "8").replace(" ","").split(",")] for w in [int(w) for w in os.environ.get("SDNQ_TRITON_MM_NUM_WARPS_LIST", "16" if torch.xpu.is_available() else "4").replace(" ","").split(",")] - for s in [int(s) for s in os.environ.get("SDNQ_TRITON_MM_NUM_STAGES_LIST", "1" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] + for s in [int(s) for s in os.environ.get("SDNQ_TRITON_MM_NUM_STAGES_LIST", "1,2" if (torch.cuda.is_available() and torch.version.hip) else "2").replace(" ","").split(",")] ]