diff --git a/aiter/configs/model_configs/dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv b/aiter/configs/model_configs/dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv new file mode 100644 index 0000000000..00b2cf3765 --- /dev/null +++ b/aiter/configs/model_configs/dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv @@ -0,0 +1,78 @@ +gfx,b,m,n,k,libtype,kernelId,splitK,us,kernelName,tflops,bw,errRatio +gfx950,2,1,1024,4096,opus,311,1,7.8,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,2.2,1076.95,0.0015 +gfx950,2,4,1024,4096,opus,311,1,7.84,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,8.6,1075.87,0.0015 +gfx950,2,8,1024,4096,opus,311,1,8.08,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,16.6,1050.11,0.0015 +gfx950,2,16,1024,4096,opus,311,1,8.37,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,32.1,1025.7,0.0015 +gfx950,2,32,1024,4096,opus,311,1,8.4,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,63.9,1045.26,0.0015 +gfx950,2,64,1024,4096,opus,311,1,8.53,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,125.9,1075.6,0.0015 +gfx950,2,128,1024,4096,opus,321,1,10.76,opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x32x256_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,199.5,925.53,0.0015 +gfx950,2,256,1024,4096,opus,324,1,11.9082,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,360.7,968.6,0.0002 +gfx950,2,512,1024,4096,opus,324,1,18.3635,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,467.8,799.42,0.0003 +gfx950,2,1024,1024,4096,opus,326,1,20.9909,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x64x256_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,818.4,999.08,0.0 +gfx950,2,1536,1024,4096,opus,325,1,36.81,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x128x128_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,700.1,740.68,0.0015 +gfx950,2,2048,1024,4096,opus,325,1,39.05,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x128x128_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,879.9,859.29,0.0015 +gfx950,2,2560,1024,4096,opus,158,1,56.53,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,759.7,704.83,0.0015 +gfx950,2,3072,1024,4096,opus,158,1,57.16,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,901.6,807.13,0.0015 +gfx950,2,3584,1024,4096,opus,158,1,58.34,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1030.7,898.72,0.0015 +gfx950,2,4096,1024,4096,opus,158,1,69.04,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,995.3,850.51,0.0015 +gfx950,2,8192,1024,4096,opus,158,1,85.01,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1616.8,1282.88,0.0015 +gfx950,2,16384,1024,4096,opus,158,1,145.45,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1889.8,1441.83,0.0015 +gfx950,2,32768,1024,4096,opus,158,1,261.61,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2101.4,1571.21,0.0015 +gfx950,4,1,1024,4096,opus,311,1,7.74,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,4.3,2169.4,0.0015 +gfx950,4,4,1024,4096,opus,311,1,7.89,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,17.0,2138.74,0.0015 +gfx950,4,8,1024,4096,opus,311,1,7.85,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,34.2,2161.65,0.0015 +gfx950,4,16,1024,4096,opus,311,1,8.39,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,64.0,2045.46,0.0015 +gfx950,4,32,1024,4096,opus,321,1,10.34,opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x32x256_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,103.8,1698.12,0.0015 +gfx950,4,64,1024,4096,opus,321,1,11.87,opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x32x256_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,180.9,1546.14,0.0015 +gfx950,4,128,1024,4096,opus,324,1,13.801,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,311.2,1443.59,0.0001 +gfx950,4,256,1024,4096,opus,324,1,17.6019,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,488.0,1310.58,0.0003 +gfx950,4,512,1024,4096,opus,326,1,25.2556,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x64x256_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,680.2,1162.52,0.0 +gfx950,4,1024,1024,4096,opus,325,1,39.32,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x128x128_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,873.9,1066.79,0.0015 +gfx950,4,1536,1024,4096,opus,158,1,57.68,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,893.5,945.32,0.0015 +gfx950,4,2048,1024,4096,opus,158,1,62.04,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1107.7,1081.75,0.0015 +gfx950,4,2560,1024,4096,opus,158,1,65.06,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1320.3,1224.92,0.0015 +gfx950,4,3072,1024,4096,opus,158,1,66.57,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1548.4,1386.08,0.0015 +gfx950,4,3584,1024,4096,opus,158,1,78.52,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1531.5,1335.4,0.0015 +gfx950,4,4096,1024,4096,opus,158,1,76.9,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1787.2,1527.13,0.0015 +gfx950,4,8192,1024,4096,opus,158,1,135.76,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2024.7,1606.55,0.0015 +gfx950,4,16384,1024,4096,opus,158,1,268.67,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2046.2,1561.11,0.0015 +gfx950,4,32768,1024,4096,opus,158,1,516.35,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2129.4,1592.12,0.0015 +gfx950,8,1,1024,4096,opus,311,1,9.62,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,7.0,3493.57,0.0015 +gfx950,8,4,1024,4096,opus,311,1,9.74,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,27.6,3465.75,0.0015 +gfx950,8,8,1024,4096,opus,311,1,9.83,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,54.6,3452.88,0.0015 +gfx950,8,16,1024,4096,opus,311,1,10.13,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x32x512_1x2_16x16x128_1x128x128_wgpcu2_scaleprefetch,106.0,3388.57,0.0015 +gfx950,8,32,1024,4096,opus,321,1,11.56,opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x32x256_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,185.8,3038.69,0.0015 +gfx950,8,64,1024,4096,opus,324,1,15.76,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,272.5,2328.22,0.0015 +gfx950,8,128,1024,4096,opus,653,1,20.76,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x64x128_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,413.7,1918.99,0.0015 +gfx950,8,256,1024,4096,opus,326,1,27.81,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x64x256_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,617.7,1658.86,0.0015 +gfx950,8,512,1024,4096,opus,325,1,38.33,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x128x128_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,896.4,1531.97,0.0015 +gfx950,8,1024,1024,4096,opus,158,1,60.87,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1129.0,1378.22,0.0015 +gfx950,8,1536,1024,4096,opus,158,1,78.73,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1309.3,1385.18,0.0015 +gfx950,8,2048,1024,4096,opus,158,1,80.06,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1716.7,1676.46,0.0015 +gfx950,8,2560,1024,4096,opus,158,1,139.21,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1234.1,1144.93,0.0015 +gfx950,8,3072,1024,4096,opus,158,1,125.44,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1643.5,1471.24,0.0015 +gfx950,8,3584,1024,4096,opus,158,1,147.82,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1627.1,1418.74,0.0015 +gfx950,8,4096,1024,4096,opus,158,1,140.65,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1954.4,1670.0,0.0015 +gfx950,8,8192,1024,4096,opus,158,1,260.59,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2109.6,1673.91,0.0015 +gfx950,8,16384,1024,4096,opus,158,1,519.59,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2116.1,1614.48,0.0015 +gfx950,8,32768,1024,4096,opus,158,1,1010.34,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2176.5,1627.34,0.0015 +gfx950,16,1,1024,4096,opus,313,1,15.7137,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x64x256_1x2_16x16x128_1x128x128_wgpcu2,8.5,4276.98,0.0002 +gfx950,16,4,1024,4096,opus,313,1,15.955,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x64x256_1x2_16x16x128_1x128x128_wgpcu2,33.6,4230.78,0.0001 +gfx950,16,8,1024,4096,opus,313,1,16.9,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x64x256_1x2_16x16x128_1x128x128_wgpcu2,63.5,4017.02,0.0014 +gfx950,16,16,1024,4096,opus,313,1,17.81,opus_bmm_a8w8_mxscale_flatmm_splitk_256x16x64x256_1x2_16x16x128_1x128x128_wgpcu2,120.6,3856.31,0.0014 +gfx950,16,32,1024,4096,opus,321,1,19.47,opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x32x256_2x1_16x16x128_1x128x128_wgpcu2_scaleprefetch,220.6,3608.38,0.0015 +gfx950,16,48,1024,4096,opus,324,1,20.0475,opus_bmm_a8w8_mxscale_flatmm_splitk_256x64x32x256_2x1_16x16x128_1x128x128_wgpcu2_sfpreload,321.4,3582.86,0.0002 +gfx950,16,64,1024,4096,opus,326,1,21.5703,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x64x256_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,398.2,3402.84,0.0 +gfx950,16,128,1024,4096,opus,326,1,30.66,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x64x256_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,560.4,2599.37,0.0015 +gfx950,16,256,1024,4096,opus,325,1,40.86,opus_bmm_a8w8_mxscale_flatmm_splitk_256x128x128x128_2x1_16x16x128_1x128x128_wgpcu1_sfpreload,840.9,2258.39,0.0015 +gfx950,16,512,1024,4096,opus,158,1,68.25,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1006.9,1720.69,0.0015 +gfx950,16,1024,1024,4096,opus,158,1,92.61,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1484.1,1811.68,0.0015 +gfx950,16,1536,1024,4096,opus,158,1,155.88,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1322.5,1399.14,0.0015 +gfx950,16,2048,1024,4096,opus,158,1,152.51,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1802.3,1760.06,0.0015 +gfx950,16,2560,1024,4096,opus,158,1,227.19,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1512.4,1403.1,0.0015 +gfx950,16,3072,1024,4096,opus,158,1,215.94,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1909.4,1709.28,0.0015 +gfx950,16,3584,1024,4096,opus,158,1,288.43,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1667.8,1454.19,0.0015 +gfx950,16,4096,1024,4096,opus,158,1,265.33,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2071.9,1770.46,0.0015 +gfx950,16,8192,1024,4096,opus,158,1,510.07,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2155.6,1710.37,0.0015 +gfx950,16,16384,1024,4096,opus,158,1,1107.72,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,1985.2,1514.58,0.0015 +gfx950,16,32768,1024,4096,opus,158,1,2088.03,opus_bmm_a8w8_mxscale_pipeline_512x256x256x128_2x1_16x16x128_1x128x128_preload_sf,2106.3,1574.85,0.0015 diff --git a/aiter/jit/core.py b/aiter/jit/core.py index 43231ab3cd..7bb5ea38c0 100644 --- a/aiter/jit/core.py +++ b/aiter/jit/core.py @@ -69,7 +69,7 @@ def mp_lock( return ret # Could not acquire: another process holds the lock. Wait for it. # wait() returns True if the holder released normally (work done), - # or False if it broke a stale lock left by a dead/abandoned holder โ€” + # or False if it broke a stale lock left by a dead/abandoned holder -- # in which case we loop and try to acquire + build ourselves. if baton.wait(): if WaitFunc is not None: @@ -133,6 +133,18 @@ def mp_lock( f"{AITER_ROOT_DIR}/aiter/configs/bf16_tuned_batched_gemm.csv", ) +# fp8 e8m0 mxscale (block-scale) batched-GEMM tuned config. Its own family +# (scale type baked into the filename, matching the a8w8_/bf16_ split) so a +# future fp32 rowwise-scale variant lands in a separate CSV and never collides +# on key. The scale type is identified by the filename alone. The +# per-model tuned data currently lives under model_configs/ (e.g. +# dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv), merged in at runtime by +# get_config_file; this canonical path may not exist on disk. +AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE = os.getenv( + "AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE", + f"{AITER_ROOT_DIR}/aiter/configs/batched_gemm_a8w8_blockscale_mxscale_tuned.csv", +) + AITER_CONFIG_GEMM_BF16 = os.getenv( "AITER_CONFIG_GEMM_BF16", f"{AITER_ROOT_DIR}/aiter/configs/bf16_tuned_gemm.csv", @@ -214,6 +226,14 @@ def AITER_CONFIG_GEMM_BF16_FILE(self): "AITER_CONFIG_GEMM_BF16", AITER_CONFIG_GEMM_BF16, "bf16_tuned_gemm" ) + @property + def AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE_FILE(self): + return self.get_config_file( + "AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE", + AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE, + "batched_gemm_a8w8_blockscale_mxscale_tuned", + ) + def update_config_files(self, file_path: str, merge_name: str): path_list = file_path.split(os.pathsep) if file_path else [] if len(path_list) <= 1: @@ -269,11 +289,10 @@ def update_config_files(self, file_path: str, merge_name: str): merge_df["_tag"] = merge_df["_tag"].fillna("") ## get keys from untuned file to drop_duplicates - untuned_name = ( - re.sub(r"(?:_)?tuned$", r"\1untuned", merge_name) - if re.search(r"(?:_)?tuned$", merge_name) - else merge_name.replace("tuned", "untuned") - ) + # Turn the tuned-file base name into its untuned sibling by rewriting the + # LAST "tuned" token (handles both mid-string names like + # "a8w8_tuned_gemm" and trailing ones like "..._mxscale_tuned"). + untuned_name = "untuned".join(merge_name.rsplit("tuned", 1)) untuned_path = f"{AITER_ROOT_DIR}/aiter/configs/{untuned_name}.csv" if os.path.exists(untuned_path): untunedf = pd.read_csv(untuned_path) @@ -283,6 +302,11 @@ def update_config_files(self, file_path: str, merge_name: str): if "gfx" in merge_df.columns and "gfx" not in keys: keys.append("gfx") dedup_keys = keys + ["_tag"] if has_tag else keys + # Only key on columns actually present in the merged frame. Most + # families carry cu_num, but some (e.g. the mxscale batched-GEMM + # table) key on gfx and never carry cu_num; keeping a missing column + # in the subset would raise inside pandas' duplicated(). + dedup_keys = [k for k in dedup_keys if k in merge_df.columns] duplicated_mask = merge_df.duplicated(subset=dedup_keys, keep=False) if duplicated_mask.any(): dup_count = int(duplicated_mask.sum()) diff --git a/aiter/jit/optCompilerConfig.json b/aiter/jit/optCompilerConfig.json index 49303f0686..d6c202c64a 100644 --- a/aiter/jit/optCompilerConfig.json +++ b/aiter/jit/optCompilerConfig.json @@ -320,7 +320,8 @@ "module_deepgemm_opus": { "srcs": [ "f'{AITER_CSRC_DIR}/pybind/opus_gemm_pybind.cu'", - "f'{AITER_CSRC_DIR}/opus_gemm/opus_gemm.cu'" + "f'{AITER_CSRC_DIR}/opus_gemm/opus_gemm.cu'", + "f'{AITER_CSRC_DIR}/opus_gemm/opus_bmm.cu'" ], "flags_extra_cc": [], "flags_extra_hip": [ diff --git a/aiter/ops/batched_gemm_op_a8w8.py b/aiter/ops/batched_gemm_op_a8w8.py index 65f6194140..2130490d6c 100644 --- a/aiter/ops/batched_gemm_op_a8w8.py +++ b/aiter/ops/batched_gemm_op_a8w8.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: MIT -# Copyright (C) 2024-2025, Advanced Micro Devices, Inc. All rights reserved. +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. import functools @@ -16,7 +16,9 @@ ) from ..jit.utils.chip_info import get_cu_num from ..jit.utils.chip_info import get_gfx_runtime as get_gfx +from ..jit.utils.torch_guard import torch_compile_guard from ..utility import dtypes +from .gemm_op_common import get_padded_m def gen_batched_gemm_a8w8_fake_tensors( @@ -86,7 +88,7 @@ def get_CKBatchedGEMM_config( get_CKBatchedGEMM_config.has_gfx = True else: logger.warning( - f"{AITER_CONFIGS.AITER_CONFIG_A8W8_BATCHED_GEMM_FILE} has no 'gfx' column โ€” " + f"{AITER_CONFIGS.AITER_CONFIG_A8W8_BATCHED_GEMM_FILE} has no 'gfx' column -- " "falling back to cu_num-only key. Re-run the tuner or migrate the CSV." ) get_CKBatchedGEMM_config.ck_batched_gemm_dict = ( @@ -147,6 +149,205 @@ def batched_gemm_a8w8_CK( return batched_gemm_a8w8(XQ, WQ, x_scale, w_scale, Y, bias, splitK) +# --------------------------------------------------------------------------- +# Shared tuned-CSV lookup for the mxscale batched GEMM. +# +# Shaped like tuned_gemm.py's multi-backend lookup: this layer locates the row +# and never interprets the kernel identifier, since that differs per backend +# (opus names kernels with an integer kernelId, flydsl with a kernelName). The +# row comes back whole, libtype included, so a caller can dispatch on it; +# libtype also filters up front for CSVs that carry one row per (shape, backend) +# rather than a single cross-backend winner per shape. + +# Tuner bookkeeping rather than selection inputs, so the lookup log drops them +# and stays readable. +_TUNED_PERF_COLUMNS = ("us", "tflops", "bw", "errRatio") + + +@functools.cache +def _load_mxscale_bmm_tuned(libtype: str | None = None) -> dict: + """{(gfx,b,m,n,k): row} from the mxscale BMM tuned CSV; {} if it is missing.""" + path = AITER_CONFIGS.AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE_FILE + try: + df = pd.read_csv(path).drop_duplicates() + except FileNotFoundError: + logger.warning("mxscale BMM tuned CSV not found at %s", path) + return {} + if libtype is not None and "libtype" in df.columns: + df = df[df["libtype"] == libtype] + return df.set_index(["gfx", "b", "m", "n", "k"]).to_dict("index") + + +@functools.lru_cache(maxsize=1024) +def lookup_mxscale_bmm_config( + b: int, m: int, n: int, k: int, *, libtype: str | None = None +): + """Exact tuned row for this shape, else one at a padded M. + + Same exact-then-two-granularities walk over the shared C++ getPaddedM that + the CK / asm / a16w16 lookups use. A bucket table built from the CSV's own M + values was the alternative and bought nothing: over every M up to the + largest tuned one, both cover the same shapes and reach the same kernel on + 131070 of 131072 M, so this keeps the one rounding rule the repo already has. + + Cached per shape like get_CKGEMM_config, and for the same reason: getPaddedM + is a ctypes hop into C++ at ~10us, and the padded levels run on every call + whose M is not itself a tuned row. DPA+MTP decode is exactly that case (M is + the ragged token count a rank happened to get), and paying it once per layer + per step cost ~1% end-to-end before this. The row is shared, so callers must + treat it as read-only. + + Returns the row, or None when no level hits. The log prints the row whole + instead of named fields, so a backend gets its own kernel identifier + reported without this layer knowing which column holds it. + """ + gfx = get_gfx() + path = AITER_CONFIGS.AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE_FILE + tuned = _load_mxscale_bmm_tuned(libtype) + + row, padded_m = None, m + for gl in (None, 0, 1): + padded_m = m if gl is None else get_padded_m(m, n, k, gl) + row = tuned.get((gfx, b, padded_m, n, k)) + if row is not None: + break + + if row is None: + logger.info( + f"shape is B:{b}, M:{m}, N:{n}, K:{k}, not found tuned/padded config " + f"in {path}, the caller will fall back!" + ) + return None + + if AITER_LOG_TUNED_CONFIG: + cfg = {c: v for c, v in row.items() if c not in _TUNED_PERF_COLUMNS} + if padded_m == m: + logger.info( + f"shape is B:{b}, M:{m}, N:{n}, K:{k}, is tuned on gfx = {gfx} " + f"in {path}, config is {cfg}!" + ) + else: + logger.info( + f"shape is B:{b}, M:{m}, N:{n}, K:{k}, exact miss on gfx = {gfx}; " + f"using padded_M: {padded_m} config {cfg} from {path}!" + ) + return row + + +# --------------------------------------------------------------------------- +# fp8 e8m0 mxscale (block-scale) batched GEMM -- public entry for the family. +# +# This file is the per-family (a8w8 batched) public surface, not a CK-only +# file: like aiter/ops/gemm_op_a8w8.py hosts gemm_a8w8 (CK rowwise) + +# gemm_a8w8_blockscale (ck/cktile/triton/asm) side by side and lazy-imports +# backend impls, we host the mxscale batched entry here too. The concrete +# kernels stay in their backend dirs (opus -> aiter.ops.opus.bmm_op). +# +# Dispatch follows tuned_gemm.mm: look the shape up once here, then let the +# winning row's libtype pick the backend, which is why the lookup runs +# unfiltered -- the tuner writes one winning row per shape and its libtype says +# who won. A second backend then only has to add rows and a branch below; it +# does not repeat the lookup. + +# Untuned shapes go to opus: it is the backend carrying a shape heuristic for +# rows the CSV does not have. +_MXSCALE_BMM_DEFAULT_LIBTYPE = "opus" + + +def _batched_gemm_a8w8_mxscale_impl( + x: Tensor, + wo_a: Tensor, + x_scale: Tensor, + w_scale: Tensor, + dtype: torch.dtype = dtypes.bf16, +) -> Tensor: + """Eager tuned-CSV lookup + libtype dispatch; returns token-major [M, G, N]. + + Kept unwrapped (plain Python) so tests can introspect the real dispatch + (which kernelId a shape resolves to) on meta tensors. The public + ``batched_gemm_a8w8_mxscale`` is the torch.compile-guarded custom op over + this; a caller that must write into its own (e.g. batch-major) output buffer + calls the opus backend (``aiter.ops.opus.bmm_op.bmm_a8w8_mxscale_opus``) + directly, which keeps the ``out=`` argument. + """ + from .opus.bmm_op import bmm_a8w8_mxscale_opus + + m, g, k = int(x.shape[0]), int(x.shape[1]), int(x.shape[2]) + n = int(wo_a.shape[1]) + + cfg = lookup_mxscale_bmm_config(g, m, n, k) + libtype = cfg["libtype"] if cfg is not None else _MXSCALE_BMM_DEFAULT_LIBTYPE + if libtype != "opus": + raise NotImplementedError( + f"tuned row for B:{g}, M:{m}, N:{n}, K:{k} wants libtype " + f"{libtype!r}, which has no batched mxscale backend here yet" + ) + + # Reading opus columns is this branch's job; whether that kernel can run + # this M, and what to do when it cannot, is the backend's. + return bmm_a8w8_mxscale_opus( + x, + wo_a, + x_scale, + w_scale, + None, + dtype=dtype, + kernelId=int(cfg["kernelId"]) if cfg is not None else None, + splitK=int(cfg["splitK"]) if cfg is not None else None, + ) + + +def _batched_gemm_a8w8_mxscale_fake( + x: Tensor, + wo_a: Tensor, + x_scale: Tensor, + w_scale: Tensor, + dtype: torch.dtype = dtypes.bf16, +) -> Tensor: + # token-major [M, G, N]; mirrors the eager allocation in bmm_a8w8_mxscale_opus. + return torch.empty( + (x.shape[0], x.shape[1], wo_a.shape[1]), + dtype=dtype, + device=x.device, + ) + + +@torch_compile_guard(mutates_args=[], gen_fake=_batched_gemm_a8w8_mxscale_fake) +def batched_gemm_a8w8_mxscale( + x: Tensor, + wo_a: Tensor, + x_scale: Tensor, + w_scale: Tensor, + dtype: torch.dtype = dtypes.bf16, +) -> Tensor: + """fp8 e8m0 mxscale (128x128 block-scale) batched GEMM. + + mmajor DSV4 wo_a layout (matches the opus kernels + op test): + + * ``x`` : [M, G, K] fp8 activation (per-token e8m0; transposed view + of batch-major [G, M, K]). + * ``wo_a`` : [G, N, K] fp8 weight (batch-major). + * ``x_scale`` : [M, G, K/128] uint8 e8m0 activation scale. + * ``w_scale`` : [G, N/128, K/128] uint8 e8m0 weight scale. + + Returns a fresh **token-major** [M, G, N] output. This entry is + torch.compile-guarded (registered as an ``aiter::`` custom op with a meta + kernel), so a framework can call it inside a compiled graph without the + tuned-CSV lookup / heuristic being traced. A caller that must write into its + own preallocated (e.g. batch-major) buffer uses + ``aiter.ops.opus.bmm_op.bmm_a8w8_mxscale_opus`` directly (it keeps ``out=``). + + Note this is *microscaling* (e8m0) block scale -- distinct from + ``gemm_a8w8_blockscale`` which uses fp32 block scale. Scale type is baked + into the name so a future fp32-block batched variant stays separate. + + The shape is looked up in the tuned CSV and the winning row's libtype picks + the backend. No kernel override lives on this entry: how a kernel is named is + backend-specific, so pin one at the backend (aiter.ops.opus.bmm_op). + """ + return _batched_gemm_a8w8_mxscale_impl(x, wo_a, x_scale, w_scale, dtype=dtype) + + def gen_batched_gemm_a8w8_tune_fake_tensors( XQ: Tensor, WQ: Tensor, diff --git a/aiter/ops/opus/__init__.py b/aiter/ops/opus/__init__.py index b543f3f5f8..44f42b466d 100644 --- a/aiter/ops/opus/__init__.py +++ b/aiter/ops/opus/__init__.py @@ -37,6 +37,7 @@ def _stub(*_args, **_kwargs): if _arch_ok: + from .bmm_op import bmm_a8w8_mxscale_opus from .gemm_op_a16w16 import ( gemm_a16w16_opus, opus_gemm_a16w16_tune, @@ -55,6 +56,7 @@ def opus_gemm_a8w8_blockscale_bpreshuffle_tune(*args, **kwargs): # it and silently disable the 30+ subsequent op imports. gemm_a16w16_opus = _make_unsupported_arch_stub("gemm_a16w16_opus") opus_gemm_a16w16_tune = _make_unsupported_arch_stub("opus_gemm_a16w16_tune") + bmm_a8w8_mxscale_opus = _make_unsupported_arch_stub("bmm_a8w8_mxscale_opus") opus_gemm_a8w8_blockscale_bpreshuffle_tune = _make_unsupported_arch_stub( "opus_gemm_a8w8_blockscale_bpreshuffle_tune" ) @@ -62,6 +64,7 @@ def opus_gemm_a8w8_blockscale_bpreshuffle_tune(*args, **kwargs): __all__ = [ + "bmm_a8w8_mxscale_opus", "gemm_a16w16_opus", "opus_gemm_a8w8_blockscale_bpreshuffle_tune", "opus_gemm_a16w16_tune", diff --git a/aiter/ops/opus/bmm_op.py b/aiter/ops/opus/bmm_op.py new file mode 100644 index 0000000000..7f79c5d765 --- /dev/null +++ b/aiter/ops/opus/bmm_op.py @@ -0,0 +1,175 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. +"""Opus batched-BMM Python bindings. + +This module is intentionally separate from `gemm_op_a16w16.py`: BMM callers use +batch-in-the-middle or grouped layouts (for example DSV4 `wo_a`) while the +underlying kernels still live in the shared opus GEMM backend. +""" + +import functools + +import torch + +from ...jit.core import compile_ops + + +def _gen_bmm_a8w8_scale_fake_tensors( + x: torch.Tensor, + wo_a: torch.Tensor, + Y: torch.Tensor, + x_scale: torch.Tensor, + w_scale: torch.Tensor, + splitK: int = 2, + kernelId: int = 0, +) -> None: + # In-place mutation of ``Y``; fake must mirror the void C++ op (full arg + # list + None return) so torch.compile registers a mutating op, not a + # tensor-producing one. + return None + + +# mmajor fp8 e8m0 mxscale BMM raw binding: x/Y are [M, batch, *], wo_a + w_scale +# batch-major (zero-copy DSV4 wo_a). kid-dispatched; driven by +# bmm_a8w8_mxscale_opus below. +@compile_ops( + "module_deepgemm_opus", + fc_name="opus_bmm_a8w8_mxscale", + gen_fake=_gen_bmm_a8w8_scale_fake_tensors, + develop=True, +) +def _opus_bmm_a8w8_mxscale_raw( + x: torch.Tensor, + wo_a: torch.Tensor, + Y: torch.Tensor, + x_scale: torch.Tensor, + w_scale: torch.Tensor, + splitK: int = 2, + kernelId: int = 0, +) -> None: + # In-place: result written into ``Y``, void return (``-> None`` keeps it + # torch.compile-safe as a mutating op). Callers read ``Y``. + ... + + +# ---- Shape-driven mxscale flatmm BMM (tuned row + heuristic fallback) ------ +# The raw binding has no tuning of its own (kernelId=0 -> slow k32 fused). This +# wrapper adds selection: explicit kernelId -> verbatim; else the tuned row the +# family entry looked up; else M-split for large unaligned M; else a coarse M/G +# heuristic. + + +@functools.cache +def _mxscale_kid_m_align() -> dict[int, int]: + """kid -> M multiple its launcher requires (1 == it masks a partial M tile). + + Comes from the codegen instance table, which is also what the tuner filters + candidates on. This used to be a hand-kept kid allowlist here and a second + hand-kept m_align column in the tuner, and the two disagreed: kid326 was + dispatched at unaligned M by this file while the tuner never tuned it there, + which cost ~9% at the wo_a decode shapes. + """ + from csrc.opus_gemm.opus_gemm_common import a8w8_mxscale_bmm_kernel_lists + + return { + int(kid): int(inst.m_align) + for fam in a8w8_mxscale_bmm_kernel_lists + for kid, inst in fam.items() + } + + +def _kid_runs_m(kid: int, m: int) -> bool: + """True iff kid's launcher accepts this M (unknown kid -> assume it does not). + + Only a tuned row found at a padded M can name a kernel that rejects the + real, smaller M, so this is what an incoming id is checked against below. + No tuned winner needs alignment today (all 11 mask their partial M tile), + but 10 of the 45 codegen instances require M % 128 or % 256, so a re-tune + can put one in the CSV. + """ + align = _mxscale_kid_m_align().get(int(kid)) + return align is not None and m % align == 0 + + +def _heuristic_mxscale_kid(g: int, m: int, n: int, k: int) -> int: + """Coarse M/G kid picker for shapes not in the tuned CSV. + + kid 158 (512x256 preload pipeline) for large-M/high-G, falling back to kid 150 + (256x256 plain) for K>8192 where 158 early-returns; kid 320/640 for small-M; + kid 653 the general strong mid/small-M pick; kid 0 (k32 fused) for shapes that + are not tile-aligned in N or K. + """ + + def div(a: int, b: int) -> bool: + return a % b == 0 + + if div(n, 256) and div(k, 128) and (m >= 2048 or (m >= 1024 and g >= 8)): + # Large M: the preload pipeline (kid158) is the tuned winner across this + # whole region (CSV picks 158 for every aligned m>=2048). No M alignment + # needed -- the pipeline family masks its partial trailing tile via buffer + # OOB. kid158 stages the SFA/SFB scales into LDS and early-returns for + # K>8192 (SFA_K_MAX), so gate the preload pick at K<=8192 and fall back to + # the plain 256x256 (kid150) for K>8192. Measured on g=2,n=1024,k=4096: + # kid150 was 34-51% slower than 158 at the untuned m=2560/3072/3584 + # buckets, and on unaligned M a single kid158 launch beats the sub-tile + # kid653 by 13-34% (g2/m2624, g8/m1000, g16/m600). + return 158 if 4096 <= k <= 8192 else 150 + # Sub-tile M: B_M=32/64 tiles mask partial M via buffer OOB, so run any M + # (no m-alignment needed -- verified 653/321/... run arbitrary unaligned M). + if m < 64: + return 640 if (div(n, 64) and div(k, 256)) else 653 + if m <= 256 and k <= 1024 and div(n, 32) and div(k, 256): + return 320 + if div(n, 64) and div(k, 128): + return 653 + return 0 # nothing tile-aligned: k32 fused runs arbitrary shapes + + +def bmm_a8w8_mxscale_opus( + x: torch.Tensor, + wo_a: torch.Tensor, + x_scale: torch.Tensor, + w_scale: torch.Tensor, + out: torch.Tensor | None = None, + dtype: torch.dtype = torch.bfloat16, + kernelId: int | None = None, + splitK: int | None = None, +) -> torch.Tensor: + """Opus fp8 e8m0 mxscale (block-scale) BMM by kernel id. + + mmajor DSV4 wo_a layout: ``x`` [M, G, K] fp8, ``wo_a`` [G, N, K] fp8, + ``x_scale`` [M, G, K/128], ``w_scale`` [G, N/128, K/128], ``out`` optional + [M, G, N]. Returns the [M, G, N] output. + + ``kernelId`` None falls back to the shape heuristic: the tuned CSV is read + one layer up, in batched_gemm_a8w8_mxscale, which hands the tuned id down. + An id this backend cannot run at this M gets the heuristic too, so the + caller never has to know the alignment rules; _opus_bmm_a8w8_mxscale_raw is + the entry that launches an id verbatim. ``splitK`` defaults to 1. + """ + m, g, k = int(x.shape[0]), int(x.shape[1]), int(x.shape[2]) + n = int(wo_a.shape[1]) + + if out is not None: + Y = out + else: + Y = torch.empty((m, g, n), dtype=dtype, device=x.device) + + # A tuned row found at a padded M can name a kernel whose launcher rejects + # the real, smaller M; drop its splitK along with it and let the heuristic + # pick instead of letting the launcher throw. + if kernelId is not None and not _kid_runs_m(int(kernelId), m): + kernelId = splitK = None + if kernelId is None: + kernelId = _heuristic_mxscale_kid(g, m, n, k) + if splitK is None: + splitK = 1 + + _opus_bmm_a8w8_mxscale_raw(x, wo_a, Y, x_scale, w_scale, int(splitK), int(kernelId)) + return Y + + +__all__ = [ + "_opus_bmm_a8w8_mxscale_raw", + "bmm_a8w8_mxscale_opus", +] diff --git a/csrc/include/rocm_ops.hpp b/csrc/include/rocm_ops.hpp index 96e99ffcdb..37a2f6e5e8 100644 --- a/csrc/include/rocm_ops.hpp +++ b/csrc/include/rocm_ops.hpp @@ -299,6 +299,18 @@ namespace py = pybind11; py::arg("kernelId") = 0, \ py::arg("splitK") = 0); +#define OPUS_BMM_A8W8_MXSCALE_PYBIND \ + m.def("opus_bmm_a8w8_mxscale", \ + &opus_bmm_a8w8_mxscale, \ + "mmajor fp8 e8m0 mxscale (block-scale) BMM with native " \ + "scaled MFMA; kid-dispatched flatmm split-K backend", \ + py::arg("O"), \ + py::arg("wo_a"), \ + py::arg("Y"), \ + py::arg("x_scale"), \ + py::arg("w_scale"), \ + py::arg("splitK") = 2, \ + py::arg("kernelId") = 0); #define OPUS_GEMM_A8W8_BLOCKSCALE_BPRESHUFFLE_TUNE_PYBIND \ m.def("opus_gemm_a8w8_blockscale_bpreshuffle_tune", \ &opus_gemm_a8w8_blockscale_bpreshuffle_tune, \ diff --git a/csrc/opus_gemm/README.md b/csrc/opus_gemm/README.md index 89085d4ab8..0fcc1e6619 100644 --- a/csrc/opus_gemm/README.md +++ b/csrc/opus_gemm/README.md @@ -15,6 +15,7 @@ This directory holds the C++ / JIT build inputs only. | `opus_gemm_common.py` | Kernel instance metadata โ€” all kids (a16w16 split-barrier, flatmm, flatmm_splitk) live here | | `gen_instances.py` | JIT codegen driver; `--tune_file` bakes the tuned CSV into `opus_gemm_lookup.h` | | `opus_gemm_tune.py` | Offline tuner CLI (see `aiter/ops/opus/README.md` ยง3 for usage) | +| `opus_bmm_mxscale_tune.py` | Offline tuner CLI for the a8w8 mxscale BMM (DSV4 wo_a), writing `dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv`. Its candidate pool is `_TUNE_POLICY` (kid -> split-K factors); tile shape, kernelName and M alignment come from `opus_gemm_common.py`, so a kid is never tuned on a shape its launcher rejects | | `include/opus_gemm.h`, `include/opus_gemm_arch.cuh` | Cross-arch declarations + `OpusGfxArch` enum + `opus_get_arch_info()` probe | | `include/opus_gemm_common.cuh`, `include/opus_gemm_utils.cuh` | Cross-arch traits umbrella + opus.hpp shim | | `include/gfx950/*.cuh` | gfx950-specific pipelines (a16w16 split-barrier / flatmm / flatmm_splitk, a8w8 noscale / scale), traits, splitk reduce, heuristic dispatch (`opus_a16w16_heuristic_dispatch_gfx950`), and the dispatch glue (`opus_gemm_arch_gfx950.cuh`). | diff --git a/csrc/opus_gemm/codegen/gen_instances_gfx950.py b/csrc/opus_gemm/codegen/gen_instances_gfx950.py index 52b0e8a32b..c161795be7 100644 --- a/csrc/opus_gemm/codegen/gen_instances_gfx950.py +++ b/csrc/opus_gemm/codegen/gen_instances_gfx950.py @@ -20,13 +20,22 @@ # ---------------- gfx950 arch-override maps ---------------- PIPELINE_HEADER_MAP = { - "a8w8_scale": "gfx950/opus_gemm_pipeline_a8w8_scale_gfx950.cuh", + "a8w8_scale": "gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh", + "a8w8_mxscale": "gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh", "a8w8": "gfx950/opus_gemm_pipeline_a8w8_noscale_gfx950.cuh", "a16w16": "gfx950/opus_gemm_pipeline_a16w16_gfx950.cuh", "a16w16_flatmm": "gfx950/opus_gemm_pipeline_a16w16_flatmm_gfx950.cuh", "a16w16_flatmm_splitk": "gfx950/opus_gemm_pipeline_a16w16_flatmm_splitk_gfx950.cuh", "a16w16_persistent": "gfx950/opus_gemm_pipeline_a16w16_persistent_gfx950.cuh", "a16w16_mono_tile": "gfx950/opus_gemm_pipeline_a16w16_mono_tile_gfx950.cuh", + "a8w8_mxscale_bmm_flatmm_splitk": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_minterleave": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_fused": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_pipeline": "gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh", + "a8w8_mxscale_bmm_mouter": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_mouter_tunable": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_wave8n2": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", + "a8w8_mxscale_bmm_wave4m2_selfload": "gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh", } # 4g_safe sibling pipelines: only defined for the a16w16-family tags that have @@ -40,22 +49,41 @@ TRAITS_HEADER_MAP = { "a8w8_scale": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", "a8w8": "gfx950/opus_gemm_traits_a8w8_noscale_gfx950.cuh", "a16w16": "gfx950/opus_gemm_traits_a16w16_gfx950.cuh", "a16w16_flatmm": "gfx950/opus_gemm_traits_a16w16_gfx950.cuh", "a16w16_flatmm_splitk": "gfx950/opus_gemm_traits_a16w16_gfx950.cuh", "a16w16_persistent": "gfx950/opus_gemm_traits_a16w16_gfx950.cuh", "a16w16_mono_tile": "gfx950/opus_gemm_traits_a16w16_gfx950.cuh", + "a8w8_mxscale_bmm_flatmm_splitk": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_minterleave": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_fused": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_pipeline": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_mouter": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_mouter_tunable": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_wave8n2": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", + "a8w8_mxscale_bmm_wave4m2_selfload": "gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh", } KERNEL_FUNC_MAP = { "a8w8_scale": "gemm_a8w8_scale_kernel", + "a8w8_mxscale": "gemm_a8w8_scale_kernel", "a8w8": "gemm_a8w8_noscale_kernel", "a16w16": "gemm_a16w16_kernel", "a16w16_flatmm": "gemm_a16w16_flatmm_kernel", "a16w16_flatmm_splitk": "gemm_a16w16_flatmm_splitk_kernel", "a16w16_persistent": "gemm_a16w16_persistent_kernel", "a16w16_mono_tile": "gemm_a16w16_mono_tile_kernel_gfx950", + "a8w8_mxscale_bmm_flatmm_splitk": "gemm_a8w8_mxscale_flatmm_splitk_kernel", + "a8w8_mxscale_bmm_minterleave": "gemm_a8w8_mxscale_flatmm_minterleave_kernel", + "a8w8_mxscale_bmm_fused": "gemm_a8w8_mxscale_flatmm_splitk_kernel", + # pipeline: default; the emit fn selects the real kernel per-kid from flags. + "a8w8_mxscale_bmm_pipeline": "gemm_a8w8_scale_kernel", + "a8w8_mxscale_bmm_mouter": "gemm_a8w8_mxscale_flatmm_splitk_mouter_kernel", + "a8w8_mxscale_bmm_mouter_tunable": "gemm_a8w8_mxscale_flatmm_splitk_mouter_kernel", + "a8w8_mxscale_bmm_wave8n2": "gemm_a8w8_mxscale_flatmm_splitk_wave8n2_kernel", + "a8w8_mxscale_bmm_wave4m2_selfload": "gemm_a8w8_mxscale_flatmm_splitk_wave4m2_selfload_kernel", } KERNEL_FUNC_MAP_4G_SAFE = { @@ -66,22 +94,64 @@ TRAITS_NAME_MAP = { "a8w8_scale": "opus_gemm_a8w8_scale_traits_gfx950", + "a8w8_mxscale": "opus_gemm_a8w8_scale_traits_gfx950", "a8w8": "opus_gemm_a8w8_noscale_traits_gfx950", "a16w16": "opus_gemm_a16w16_traits_gfx950", "a16w16_flatmm": "opus_gemm_a16w16_flatmm_traits_gfx950", "a16w16_flatmm_splitk": "opus_flatmm_splitk_traits_gfx950", "a16w16_persistent": "opus_gemm_a16w16_persistent_traits_gfx950", "a16w16_mono_tile": "opus_gemm_a16w16_mono_tile_traits_gfx950", + "a8w8_mxscale_bmm_flatmm_splitk": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_minterleave": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_fused": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_pipeline": "opus_gemm_a8w8_scale_traits_gfx950", + "a8w8_mxscale_bmm_mouter": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_mouter_tunable": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_wave8n2": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", + "a8w8_mxscale_bmm_wave4m2_selfload": "opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950", } KARGS_NAME_MAP = { "a8w8_scale": "opus_gemm_scale_kargs_gfx950", + "a8w8_mxscale": "opus_gemm_scale_kargs_gfx950", "a8w8": "opus_gemm_noscale_kargs_gfx950", "a16w16": "opus_gemm_noscale_kargs_gfx950", "a16w16_flatmm": "opus_gemm_flatmm_kargs_gfx950", "a16w16_flatmm_splitk": "opus_gemm_flatmm_splitk_kargs_gfx950", "a16w16_persistent": "opus_gemm_persistent_kargs_gfx950", "a16w16_mono_tile": "opus_gemm_mono_tile_kargs_gfx950", + "a8w8_mxscale_bmm_flatmm_splitk": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_minterleave": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_fused": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_pipeline": "opus_gemm_scale_kargs_gfx950", + "a8w8_mxscale_bmm_mouter": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_mouter_tunable": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_wave8n2": "opus_gemm_scale_splitk_kargs_gfx950", + "a8w8_mxscale_bmm_wave4m2_selfload": "opus_gemm_scale_splitk_kargs_gfx950", +} + + +def splitk_reduce_extra_device_instantiations(): + # gfx950 carries a second reduce kernel: the mmajor BMM reduce used by the + # a8w8_mxscale BMM split-K launchers (VEC=8/BLOCK=128, explicit C strides, + # no bias fold). Those launchers <<<>>> it from their fused host TU and only + # see a forward decl, so exactly one TU must own the device kernel plus its + # host stub. It lives in the same splitk_reduce_gfx950.cuh as the baseline + # reduce, so it rides along in this TU; that keeps opus_bmm.cu out of the + # device pass entirely, matching opus_gemm.cu. + return ( + "// mmajor BMM reduce (a8w8_mxscale split-K launchers)\n" + "template __global__ void opus_bmm_splitk_reduce_kernel<__bf16, 8, 128>(\n" + " const opus_splitk_ws_handle*, __bf16*,\n" + " int, int, int, int, int, int, int, int);\n" + "template __global__ void opus_bmm_splitk_reduce_kernel(\n" + " const opus_splitk_ws_handle*, float*,\n" + " int, int, int, int, int, int, int, int);\n" + ) + + +SPLITK_REDUCE_EXTRA_MAP = { + "device_instantiations": splitk_reduce_extra_device_instantiations, } register_arch_map("gfx950", "pipeline_header", PIPELINE_HEADER_MAP) @@ -89,6 +159,7 @@ register_arch_map("gfx950", "kernel_func", KERNEL_FUNC_MAP) register_arch_map("gfx950", "traits_name", TRAITS_NAME_MAP) register_arch_map("gfx950", "kargs_name", KARGS_NAME_MAP) +register_arch_map("gfx950", "splitk_reduce_extra", SPLITK_REDUCE_EXTRA_MAP) # ---------------- gfx950 validators ---------------- @@ -710,7 +781,7 @@ def gen_scale_instance( template using {k.name}_Traits = {traits_name}<{k.BLOCK_SIZE}, opus::seq<{k.B_M}, {k.B_N}, {k.B_K}>, - opus::tuple<{da}, {db}, D_C, fp32_t, fp32_t>, + opus::tuple<{da}, {db}, D_C, fp32_t, {"unsigned char" if k.kernel_tag == "a8w8_mxscale" else "fp32_t"}>, opus::seq<{k.VEC_A}, {k.VEC_B}, {k.VEC_C}>, opus::seq<{k.GROUP_M}, {k.GROUP_N}, {k.GROUP_K}>>; """ @@ -786,6 +857,83 @@ def gen_scale_instance( Path(os.path.join(cg.impl_path, f"{k.name}.cuh")).write_text(INSTANCE_IMPL) record_one_instantiation(cg, k, kernel_func, kargs_name, A8W8_SCALE_HOST_EXTRA) + # "_mmajor" sibling: A(XQ)/Y are [M, batch, *] (dim0=M, dim1=batch) and + # x_scale is [M, batch, K/GROUP_K] (per-token M) so the DSV4 wo_a activation + # o=[num_tokens, n_groups, K] feeds in with NO caller-side transpose. Weight + # (WQ) and its scale (w_scale) stay batch-major [batch, N, K] / + # [batch, N/GROUP_N, K/GROUP_K]. Same kernel/traits; the launcher just reads + # A/Y/sfa strides from the tensors instead of hardcoding batch-major. + INSTANCE_IMPL_MMAJOR = f""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +template +void +{k.name}_mmajor( + aiter_tensor_t &XQ, + aiter_tensor_t &WQ, + aiter_tensor_t &Y, + std::optional x_scale, + std::optional w_scale) +{{{{ + int M = XQ.size(0); + int batch = XQ.size(1); + int N = WQ.size(1); + int K = XQ.size(2); + + int GROUP_N = {k.GROUP_N}; + int GROUP_K = {k.GROUP_K}; + int num_groups_n = N / GROUP_N; + int num_groups_k = K / GROUP_K; + + {kargs_name} kargs{{}}; + kargs.ptr_a = XQ.data_ptr(); + kargs.ptr_b = WQ.data_ptr(); + kargs.ptr_c = Y.data_ptr(); + kargs.m = M; + kargs.n = N; + kargs.k = K; + kargs.batch = batch; + // mmajor A/Y (dim0=M, dim1=batch); weight WQ stays batch-major. + kargs.stride_a = (int)XQ.stride(0); + kargs.stride_b = (int)WQ.stride(1); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_a_batch = (int)XQ.stride(1); + kargs.stride_b_batch = (int)WQ.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + + kargs.ptr_sfa = x_scale.value().data_ptr(); + kargs.ptr_sfb = w_scale.value().data_ptr(); + // x_scale mmajor [M, batch, num_groups_k]; w_scale batch-major. + kargs.stride_sfa = (int)x_scale.value().stride(0); + kargs.stride_sfa_batch = (int)x_scale.value().stride(1); + kargs.stride_sfb = num_groups_k; + kargs.stride_sfb_batch = num_groups_n * num_groups_k; + + int num_tiles_m = (M + {k.B_M} - 1) / {k.B_M}; + int num_tiles_n = (N + {k.B_N} - 1) / {k.B_N}; + dim3 grid(num_tiles_m * num_tiles_n, 1, batch); + dim3 block({k.BLOCK_SIZE}); + + auto stream = aiter::getCurrentHIPStream(); + {kernel_func}<{k.name}_Traits><<>>(kargs); + +}}}} +#endif // launcher only on regular host pass +""" + with open(os.path.join(cg.impl_path, f"{k.name}.cuh"), "a") as _f: + _f.write(INSTANCE_IMPL_MMAJOR) + + for CDtype in k.output_dtypes: + host_decl_mmajor = ( + f"template void\n" + f"{k.name}_mmajor<{CDtype}>(\n" + f" aiter_tensor_t &XQ,\n" + f" aiter_tensor_t &WQ,\n" + f" aiter_tensor_t &Y{A8W8_SCALE_HOST_EXTRA});\n" + ) + cg._host_instantiations.append( + {"kid_name": k.name, "dtype": CDtype, "host_decl": host_decl_mmajor} + ) + def gen_noscale_instance_gfx950( cg, @@ -1427,11 +1575,1293 @@ def gen_flatmm_splitk_instance( record_one_instantiation(cg, k, kernel_func, kargs_name, A16W16_TUNE_HOST_EXTRA) +def _assert_m_align(k, tile_mult): + """Tie the declared m_align to the M guard the launcher body actually emits. + + `tile_mult` is the B_M multiple the body below hardcodes in its AITER_CHECK, + or 0 when it emits no M check because the kernel masks the partial tile. + OpusGemmInstance.m_align is what the tuner's candidate filter and the + runtime's padded-M lookup read, so a guard edit that forgets to update + _BMM_M_ALIGN_TILES must fail the build rather than silently teach the two + consumers a wrong alignment. + """ + expect = k.B_M * tile_mult if tile_mult else 1 + assert k.m_align == expect, ( + f"{k.name}: launcher guards M % {expect} == 0 but m_align says " + f"{k.m_align}; fix _BMM_M_ALIGN_TILES in opus_gemm_common.py" + ) + + +# Body of the a8w8_mxscale BMM flatmm split-K launcher (mmajor layout), a +# faithful port of opus_bmm_a8w8_mxscale_flatmm_splitk_impl() in +# opus_bmm.cu. Written with @@TOKEN@@ placeholders + .replace() (NOT an +# f-string) so the C++ body keeps plain single braces and stays trivially +# reviewable against the hand-written original. +# +# Templated on D_C only to satisfy the codegen host-decl machinery +# (one instantiation); the body ignores D_C and branches on Y.dtype() +# at runtime with native __bf16/float, exactly like the original. The fused +# reduce path (splitK==2 counter variant) is intentionally NOT ported -- the +# fused-reduce kid stays monolithic in opus_bmm.cu. +_BMM_MXSCALE_SPLITK_LAUNCHER_BODY = r""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +// mmajor: O/Y are [M, batch, *] (dim0=M, dim1=batch); wo_a stays batch-major +// [batch, N, K]. Caller (opus_bmm.cu switch) does dtype/arch/common checks. +template +void +@@NAME@@( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int splitK) +{ + using Traits = @@NAME@@_Traits; + constexpr bool DIRECT_ONLY = @@DIRECT@@; + constexpr bool PREFETCH_SCALE = @@PREFETCH@@; + constexpr bool PRELOAD_SF_LDS = @@PRELOAD@@; + + AITER_CHECK(splitK >= 1, "splitK must be >= 1"); + if constexpr (DIRECT_ONLY) { + AITER_CHECK(splitK == 1, "@@NAME@@ consumer-self-load kernel requires splitK == 1"); + } + + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + // No M alignment at any tile size: A and SFA are bounded to the tile's valid row + // count, the split_k==1 store bounds C the same way, split_k>1 partials go to a + // workspace sized for padded_M, and both reducers touch only rows < M. + AITER_CHECK(N % Traits::B_N == 0, + "@@NAME@@ requires N % ", Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); + + const int split_k = splitK; + const bool no_split_k = (split_k == 1); + const int total_iters = K / Traits::B_K; + const int iters_full = (total_iters + split_k - 1) / split_k; + const int last_loops = total_iters - (split_k - 1) * iters_full; + AITER_CHECK(last_loops >= Traits::prefetch_k_iter, + "@@NAME@@ requires every split to have at least ", + Traits::prefetch_k_iter, " K-tiles; K=", K, + " gives total_iters=", total_iters, ", splitK=", split_k, + ", last split loops=", last_loops); + + const int num_tiles_m = (M + Traits::B_M - 1) / Traits::B_M; + const int num_tiles_n = (N + Traits::B_N - 1) / Traits::B_N; + const int padded_M = num_tiles_m * Traits::B_M; + const int padded_N = num_tiles_n * Traits::B_N; + const size_t partial_bytes = (size_t)split_k * (size_t)batch + * (size_t)padded_M * (size_t)padded_N * sizeof(float); + + auto stream = aiter::getCurrentHIPStream(); + + opus_gemm_scale_splitk_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ws_handle = nullptr; + kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch; + kargs.split_k = split_k; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_ws = padded_N; + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.stride_ws_batch = padded_M * padded_N; + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + + dim3 grid_main(num_tiles_m * num_tiles_n * split_k, 1, batch); + dim3 block_main(Traits::BLOCK_SIZE); + if (no_split_k) { + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } + return; + } + + if constexpr (!DIRECT_ONLY) { + extern opus_splitk_ws_handle* opus_splitk_ws_get(hipStream_t, bool); + hipStreamCaptureStatus capture_status = hipStreamCaptureStatusNone; + HIP_CALL(hipStreamIsCapturing(stream, &capture_status)); + const bool capturing = (capture_status != hipStreamCaptureStatusNone); + auto* ws_handle = opus_splitk_ws_get(stream, /*allow_create=*/!capturing); + + const size_t ws_bytes = partial_bytes; + if (ws_handle->ptr == nullptr || ws_bytes > ws_handle->bytes) { + AITER_CHECK(!capturing, + "splitk workspace grow inside HIP graph capture is not supported"); + void* new_ptr = nullptr; + const size_t kGrowAlign = (size_t)4 * 1024 * 1024; + size_t grow_bytes = ((ws_bytes + kGrowAlign - 1) / kGrowAlign) * kGrowAlign; + HIP_CALL(hipMalloc(&new_ptr, grow_bytes)); + if (ws_handle->ptr != nullptr) { + HIP_CALL(hipDeviceSynchronize()); + HIP_CALL(hipFree(ws_handle->ptr)); + } + ws_handle->ptr = new_ptr; + ws_handle->bytes = grow_bytes; + } + kargs.ws_handle = ws_handle; + + // Pass all 4 template args explicitly (D_OUT=void: the split-K main kernel + // writes an fp32 workspace, so its output dtype is irrelevant; the reduce + // kernel casts to the runtime Y dtype). The fused host TU only sees a + // no-default forward decl of @@KERNEL@@, so relying on the template's + // default args here would fail overload resolution ("no matching function"). + @@KERNEL@@ + <<>>(kargs); + + constexpr int REDUCE_VEC = 8; + constexpr int REDUCE_BS = 128; + dim3 grid_reduce((N + REDUCE_VEC * REDUCE_BS - 1) / (REDUCE_VEC * REDUCE_BS), + batch * M, 1); + dim3 block_reduce(REDUCE_BS); + const int y_stride_c = (int)Y.stride(0); + const int y_stride_c_batch = (int)Y.stride(1); + if (Y.dtype() == AITER_DTYPE_bf16) { + opus_bmm_splitk_reduce_kernel<__bf16, REDUCE_VEC, REDUCE_BS> + <<>>( + ws_handle, reinterpret_cast<__bf16*>(Y.data_ptr()), + split_k, M, N, batch, padded_M, padded_N, + y_stride_c, y_stride_c_batch); + } else { + opus_bmm_splitk_reduce_kernel + <<>>( + ws_handle, reinterpret_cast(Y.data_ptr()), + split_k, M, N, batch, padded_M, padded_N, + y_stride_c, y_stride_c_batch); + } + } +} +#endif // launcher only on regular host pass +""" + + +def gen_bmm_mxscale_flatmm_splitk_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + """gfx950 a8w8_mxscale BMM flatmm split-K launcher emit. + + Differs from the GEMM emitters: + * traits alias is NOT templated on D_C (fp32 workspace is fixed); the + launcher is templated on D_C only for the host-decl machinery. + * launcher signature is (O, wo_a, Y, x_scale, w_scale, int splitK) with + the mmajor layout, matching opus_bmm.cu's _impl. + * custom device-instantiation matrix over the kernel's (D_OUT, DIRECT_ONLY, + PREFETCH_SCALE) template params (the standard record_one_instantiation + assumes a single-template-arg > kernel). + * the split-K reduce kernel (opus_bmm_splitk_reduce_kernel) is declared in + the a8w8_scale traits header and instantiated once in opus_bmm.cu, so it + is NOT re-instantiated here. + """ + _, fwd_decl_kargs_tpl, fwd_decl_kargs_fnarg = kargs_template_vars( + k.kernel_tag, kargs_name + ) + + # Non-templated traits alias: fp32 split-K workspace is fixed; the workspace + # tuple slot 4 (scale) is `unsigned char` for the e8m0 mxscale path. + traits_aliases = f""" +using {k.name}_Traits = {traits_name}<{k.BLOCK_SIZE}, + opus::seq<{k.B_M}, {k.B_N}, {k.B_K}>, + opus::tuple<{da}, {db}, fp32_t, fp32_t, unsigned char>, + opus::seq<{k.VEC_A}, {k.VEC_B}, {k.VEC_C}>, + opus::seq<{k.GROUP_M}, {k.GROUP_N}, {k.GROUP_K}>, + {k.WG_PER_CU}>; +""" + + preamble = instance_impl_preamble() + host_tu_split = instance_impl_host_tu_split( + traits_header, + pipeline_header, + fwd_decl_kargs_tpl, + kernel_func, + fwd_decl_kargs_fnarg, + ) + + # Forward-declare the split-K reduce kernel. On the fused host TU pass + # host_tu_split only pulls in the (light) traits header -- not the pipeline + # header that defines this kernel -- so the launcher body's <<<...>>> call + # needs a visible declaration. On the non-fused device pass the pipeline + # header (via splitk_reduce_gfx950.cuh) provides a compatible definition, so + # this is just a harmless redeclaration there. opus_splitk_ws_handle is a + # complete type in both passes via the included traits/pipeline header. + reduce_fwd_decl = """ +template +__global__ void opus_bmm_splitk_reduce_kernel( + const opus_splitk_ws_handle* __restrict__ ws_handle, + D_OUT* __restrict__ out, + int split_k, int M, int N, int batch, + int padded_M, int padded_N, + int stride_c, int stride_c_batch); +""" + + launcher = ( + _BMM_MXSCALE_SPLITK_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", kernel_func) + .replace("@@DIRECT@@", "true" if k.direct_only else "false") + .replace("@@PREFETCH@@", "true" if k.prefetch_scale else "false") + .replace("@@PRELOAD@@", "true" if k.preload_sf else "false") + ) + + INSTANCE_IMPL = ( + f"{preamble}\n{host_tu_split}\n{reduce_fwd_decl}\n{traits_aliases}\n{launcher}" + ) + Path(os.path.join(cg.impl_path, f"{k.name}.cuh")).write_text(INSTANCE_IMPL) + + # Host instantiation(s): launcher templated on D_C; a single stub. + # (XQ/WQ/Y positional names in _make_host_decl map to O/wo_a/Y by type.) + host_extra = ( + ",\n aiter_tensor_t &x_scale," + "\n aiter_tensor_t &w_scale," + "\n int splitK" + ) + for dtype in k.output_dtypes: + host_decl = ( + f"template void\n" + f"{k.name}<{dtype}>(\n" + f" aiter_tensor_t &O,\n" + f" aiter_tensor_t &wo_a,\n" + f" aiter_tensor_t &Y{host_extra});\n" + ) + cg._host_instantiations.append( + {"kid_name": k.name, "dtype": dtype, "host_decl": host_decl} + ) + + # Device instantiation matrix: split-1 direct-store variants for both Y + # dtypes, plus (non-direct kids only) the fp32-workspace variant used by the + # split-K > 1 path (kernel default D_OUT=void). + direct = "true" if k.direct_only else "false" + prefetch = "true" if k.prefetch_scale else "false" + preload = "true" if k.preload_sf else "false" + + def _dev(dtype_tag, d_out, dir_flag, pfk_flag): + decl = ( + f"template __global__ void {kernel_func}<\n" + f" {k.name}_Traits, {d_out}, {dir_flag}, {pfk_flag}, {preload}>({kargs_name});\n" + ) + cg._device_instantiations.append( + {"kid_name": k.name, "dtype": dtype_tag, "device_decl": decl} + ) + + _dev("bf16", "__bf16", direct, prefetch) + _dev("fp32", "float", direct, prefetch) + if not k.direct_only: + # Split-K > 1 workspace path: host launches . DIRECT_ONLY is false here (direct kids + # never take the workspace path), but PREFETCH_SCALE / PRELOAD_SF_LDS must + # match the kid, else the instantiation is missing for + # prefetch/preload kids -> undefined symbol at load. + _dev("void", "void", "false", prefetch) + + +_BMM_MXSCALE_MINTERLEAVE_LAUNCHER_BODY = r""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +// M-tile interleaved launcher: MI=2 consecutive M tiles per WG share the B +// stream (requires M % (MI*B_M) == 0). splitK arg is unused (must be 1). mmajor: +// O/Y are [M, batch, *] (dim0=M, dim1=batch); wo_a stays batch-major [batch,N,K]. +// Caller (opus_bmm.cu dispatch) does dtype/arch/common checks. +template +void +@@NAME@@( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int /*splitK*/) +{ + using Traits = @@NAME@@_Traits; + constexpr bool SKIP_SCALE_WAIT = @@SKIP@@; + constexpr int MI = 2; + + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + AITER_CHECK(M % (MI * Traits::B_M) == 0, + "@@NAME@@ requires M % ", (MI * Traits::B_M), " == 0, got ", M); + AITER_CHECK(N % Traits::B_N == 0, + "@@NAME@@ requires N % ", Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); + const int total_iters = K / Traits::B_K; + AITER_CHECK(total_iters >= Traits::prefetch_k_iter, + "@@NAME@@ requires at least ", Traits::prefetch_k_iter, + " K-tiles, got ", total_iters); + + auto stream = aiter::getCurrentHIPStream(); + + opus_gemm_scale_splitk_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ws_handle = nullptr; + kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch; + const int num_tiles_m = M / Traits::B_M; + const int num_tiles_n = N / Traits::B_N; + kargs.split_k = MI; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + + const int split_m = num_tiles_m / MI; // M-tile groups (WGs along M) + constexpr int NUM_XCD = 8; + const int m_grp_per_xcd = (split_m + NUM_XCD - 1) / NUM_XCD; + kargs.stride_ws = split_m; + kargs.stride_ws_batch = m_grp_per_xcd; + dim3 grid_main(NUM_XCD * m_grp_per_xcd * num_tiles_n, 1, batch); + dim3 block_main(Traits::BLOCK_SIZE); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" + + +def gen_bmm_mxscale_minterleave_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + """gfx950 a8w8_mxscale BMM M-tile-interleaved launcher emit (kids 162/163). + + Sibling of gen_bmm_mxscale_flatmm_splitk_instance: + * kernel template is (no DIRECT_ONLY/ + PREFETCH_SCALE/PRELOAD_SF_LDS axes, no split-K workspace/reduce path). + * MI=2 is baked in the launcher; splitK is ignored (must be 1). + * device instantiation matrix is just (D_OUT in {bf16, float}) x the kid's + fixed SKIP_SCALE_WAIT flag. + """ + _, fwd_decl_kargs_tpl, fwd_decl_kargs_fnarg = kargs_template_vars( + k.kernel_tag, kargs_name + ) + + # Non-templated traits alias: identical geometry/tuple to the flatmm split-K + # family (fp32 workspace slot, unsigned char e8m0 scale slot). + traits_aliases = f""" +using {k.name}_Traits = {traits_name}<{k.BLOCK_SIZE}, + opus::seq<{k.B_M}, {k.B_N}, {k.B_K}>, + opus::tuple<{da}, {db}, fp32_t, fp32_t, unsigned char>, + opus::seq<{k.VEC_A}, {k.VEC_B}, {k.VEC_C}>, + opus::seq<{k.GROUP_M}, {k.GROUP_N}, {k.GROUP_K}>, + {k.WG_PER_CU}>; +""" + + preamble = instance_impl_preamble() + host_tu_split = instance_impl_host_tu_split( + traits_header, + pipeline_header, + fwd_decl_kargs_tpl, + kernel_func, + fwd_decl_kargs_fnarg, + ) + + launcher = ( + _BMM_MXSCALE_MINTERLEAVE_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", kernel_func) + .replace("@@SKIP@@", "true" if k.skip_scale_wait else "false") + ) + + INSTANCE_IMPL = f"{preamble}\n{host_tu_split}\n{traits_aliases}\n{launcher}" + Path(os.path.join(cg.impl_path, f"{k.name}.cuh")).write_text(INSTANCE_IMPL) + + # Host instantiation: launcher templated on D_C; single stub. + host_extra = ( + ",\n aiter_tensor_t &x_scale," + "\n aiter_tensor_t &w_scale," + "\n int splitK" + ) + for dtype in k.output_dtypes: + host_decl = ( + f"template void\n" + f"{k.name}<{dtype}>(\n" + f" aiter_tensor_t &O,\n" + f" aiter_tensor_t &wo_a,\n" + f" aiter_tensor_t &Y{host_extra});\n" + ) + cg._host_instantiations.append( + {"kid_name": k.name, "dtype": dtype, "host_decl": host_decl} + ) + + # Device instantiation matrix: for both Y + # dtypes (the kid's SKIP_SCALE_WAIT is fixed). + skip = "true" if k.skip_scale_wait else "false" + + def _dev(dtype_tag, d_out): + decl = ( + f"template __global__ void {kernel_func}<\n" + f" {k.name}_Traits, {d_out}, {skip}>({kargs_name});\n" + ) + cg._device_instantiations.append( + {"kid_name": k.name, "dtype": dtype_tag, "device_decl": decl} + ) + + _dev("bf16", "__bf16") + _dev("fp32", "float") + + +def _bmm_specialized_traits_alias(k, traits_name, da, db): + """Non-templated traits alias shared by all a8w8_mxscale BMM specialized + pipelines (fp32 workspace slot, unsigned char e8m0 scale slot).""" + return f""" +using {k.name}_Traits = {traits_name}<{k.BLOCK_SIZE}, + opus::seq<{k.B_M}, {k.B_N}, {k.B_K}>, + opus::tuple<{da}, {db}, fp32_t, fp32_t, unsigned char>, + opus::seq<{k.VEC_A}, {k.VEC_B}, {k.VEC_C}>, + opus::seq<{k.GROUP_M}, {k.GROUP_N}, {k.GROUP_K}>, + {k.WG_PER_CU}>; +""" + + +def _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + preamble, + host_tu_split, + launcher, + dev_flag_suffix, + emit_device=True, +): + """Shared tail for BMM specialized-pipeline emits: write impl/{name}.cuh + (preamble + host-TU split + traits alias + inlined launcher), then register + one host stub and the (bf16, fp32) device instantiation pair. + + dev_flag_suffix is the comma-prefixed template-arg tail after D_OUT in the + kernel instantiation (e.g. ", true, false, ..." for the wave families, + "" for wave8n2). + + emit_device=False emits only the host launcher (used by mouter_tunable, + which reuses the identical gemm_..._mouter_kernel + specializations already emitted by the mouter family -- emitting them again + under a different alias name would be a duplicate-symbol ODR violation). + """ + traits_aliases = _bmm_specialized_traits_alias(k, traits_name, da, db) + INSTANCE_IMPL = f"{preamble}\n{host_tu_split}\n{traits_aliases}\n{launcher}" + Path(os.path.join(cg.impl_path, f"{k.name}.cuh")).write_text(INSTANCE_IMPL) + + host_extra = ( + ",\n aiter_tensor_t &x_scale," + "\n aiter_tensor_t &w_scale," + "\n int splitK" + ) + for dtype in k.output_dtypes: + host_decl = ( + f"template void\n" + f"{k.name}<{dtype}>(\n" + f" aiter_tensor_t &O,\n" + f" aiter_tensor_t &wo_a,\n" + f" aiter_tensor_t &Y{host_extra});\n" + ) + cg._host_instantiations.append( + {"kid_name": k.name, "dtype": dtype, "host_decl": host_decl} + ) + + if not emit_device: + return + for dtype_tag, d_out in (("bf16", "__bf16"), ("fp32", "float")): + decl = ( + f"template __global__ void {kernel_func}<\n" + f" {k.name}_Traits, {d_out}{dev_flag_suffix}>({kargs_name});\n" + ) + cg._device_instantiations.append( + {"kid_name": k.name, "dtype": dtype_tag, "device_decl": decl} + ) + + +# Common launcher signature + shared checks/kargs preamble. mmajor: O/Y are +# [M, batch, *] (dim0=M, dim1=batch); wo_a stays batch-major. Caller does the +# dtype/arch/common checks (see opus_bmm.cu dispatch). +_BMM_SPEC_SIG = r""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +template +void +@@NAME@@( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int @@SPLITK_ARG@@) +{ + using Traits = @@NAME@@_Traits; +""" + +_BMM_SPEC_KARGS = r""" + auto stream = aiter::getCurrentHIPStream(); + + opus_gemm_scale_splitk_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ws_handle = nullptr; + kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_ws = N; + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.stride_ws_batch = M * N; + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); +""" + +# ---- wave8n2 (kid 132) ---- +_BMM_WAVE8N2_LAUNCHER_BODY = ( + _BMM_SPEC_SIG.replace("@@SPLITK_ARG@@", "/*splitK*/") + + r""" const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + constexpr int LOGICAL_B_N = Traits::B_N * 2; + AITER_CHECK(M % Traits::B_M == 0, + "@@NAME@@ requires M % ", Traits::B_M, " == 0, got ", M); + AITER_CHECK(N % LOGICAL_B_N == 0, + "@@NAME@@ requires N % ", LOGICAL_B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); +""" + + _BMM_SPEC_KARGS.replace( + "kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch;", + "kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch;\n kargs.split_k = 1;", + ) + + r""" + const int num_tiles_m = M / Traits::B_M; + const int num_tiles_n = N / LOGICAL_B_N; + dim3 grid_main(num_tiles_m * num_tiles_n, 1, batch); + dim3 block_main(512); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" +) + + +def gen_bmm_mxscale_wave8n2_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + launcher = _BMM_WAVE8N2_LAUNCHER_BODY.replace("@@NAME@@", k.name).replace( + "@@KERNEL@@", kernel_func + ) + _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + instance_impl_preamble(), + instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, kernel_func, fn + ), + launcher, + "", + ) + + +def _cppbool(v): + return "true" if v else "false" + + +# ---- wave4m2_selfload (kids 134/142/148) ---- +_BMM_WAVE4M2_LAUNCHER_BODY = ( + _BMM_SPEC_SIG.replace("@@SPLITK_ARG@@", "/*splitK*/") + + r""" constexpr bool SKIP_SCALE_WAIT = @@SSW@@; + constexpr bool PACK_SCALE_ON_DEMAND = @@PSOD@@; + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + constexpr int LOGICAL_B_M = Traits::B_M * 2; + AITER_CHECK(M % LOGICAL_B_M == 0, + "@@NAME@@ requires M % ", LOGICAL_B_M, " == 0, got ", M); + AITER_CHECK(N % Traits::B_N == 0, + "@@NAME@@ requires N % ", Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); +""" + + _BMM_SPEC_KARGS.replace( + "kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch;", + "kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch;\n kargs.split_k = 1;", + ) + + r""" + const int num_tiles_m = M / LOGICAL_B_M; + const int num_tiles_n = N / Traits::B_N; + dim3 grid_main(num_tiles_m * num_tiles_n, 1, batch); + dim3 block_main(256); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@< + Traits, __bf16, SKIP_SCALE_WAIT, PACK_SCALE_ON_DEMAND> + <<>>(kargs); + } else { + @@KERNEL@@< + Traits, float, SKIP_SCALE_WAIT, PACK_SCALE_ON_DEMAND> + <<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" +) + + +def gen_bmm_mxscale_wave4m2_selfload_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + launcher = ( + _BMM_WAVE4M2_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", kernel_func) + .replace("@@SSW@@", _cppbool(k.skip_scale_wait)) + .replace("@@PSOD@@", _cppbool(k.pack_scale_on_demand)) + ) + suffix = f", {_cppbool(k.skip_scale_wait)}, {_cppbool(k.pack_scale_on_demand)}" + _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + instance_impl_preamble(), + instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, kernel_func, fn + ), + launcher, + suffix, + ) + + +# ---- mouter (kids 131/144) + mouter_tunable (kids 160/161) ---- +# Shared persistent-mouter kernel; the two families differ only in how m_per_wg +# is derived (heuristic vs API-splitK sweep). XCD-aware grid remap is identical. +_BMM_MOUTER_CHECKS = r""" constexpr bool SKIP_SCALE_WAIT = @@SSW@@; + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + AITER_CHECK(M % Traits::B_M == 0, + "@@NAME@@ requires M % ", Traits::B_M, " == 0, got ", M); + AITER_CHECK(N % Traits::B_N == 0, + "@@NAME@@ requires N % ", Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); + const int total_iters = K / Traits::B_K; + AITER_CHECK(total_iters >= Traits::prefetch_k_iter, + "@@NAME@@ requires at least ", Traits::prefetch_k_iter, + " K-tiles, got ", total_iters); +""" + +_BMM_MOUTER_KARGS = r""" + auto stream = aiter::getCurrentHIPStream(); + + opus_gemm_scale_splitk_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ws_handle = nullptr; + kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch; + const int num_tiles_m = M / Traits::B_M; + const int num_tiles_n = N / Traits::B_N; +""" + +_BMM_MOUTER_TAIL = r""" kargs.split_k = m_per_wg; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_ws = N; + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.stride_ws_batch = M * N; + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + + const int split_m = (num_tiles_m + m_per_wg - 1) / m_per_wg; + constexpr int NUM_XCD = 8; + const int m_grp_per_xcd = (split_m + NUM_XCD - 1) / NUM_XCD; + kargs.stride_ws = split_m; + kargs.stride_ws_batch = m_grp_per_xcd; + dim3 grid_main(NUM_XCD * m_grp_per_xcd * num_tiles_n, 1, batch); + dim3 block_main(Traits::BLOCK_SIZE); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" + +_BMM_MOUTER_LAUNCHER_BODY = ( + _BMM_SPEC_SIG.replace("@@SPLITK_ARG@@", "/*splitK*/") + + _BMM_MOUTER_CHECKS + + _BMM_MOUTER_KARGS + + " const int m_per_wg = (num_tiles_m >= 16) ? 2 : 1;\n" + + _BMM_MOUTER_TAIL +) + +# Tunable variant: API splitK is repurposed as m_per_wg (clamped to [1, +# num_tiles_m]); reuses the same mouter kernel. +_BMM_MOUTER_TUNABLE_LAUNCHER_BODY = ( + _BMM_SPEC_SIG.replace("@@SPLITK_ARG@@", "splitK") + + _BMM_MOUTER_CHECKS + + _BMM_MOUTER_KARGS + + " int m_per_wg = splitK;\n" + " if (m_per_wg > num_tiles_m) m_per_wg = num_tiles_m;\n" + " if (m_per_wg < 1) m_per_wg = 1;\n" + _BMM_MOUTER_TAIL +) + + +def gen_bmm_mxscale_mouter_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + launcher = ( + _BMM_MOUTER_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", kernel_func) + .replace("@@SSW@@", _cppbool(k.skip_scale_wait)) + ) + _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + instance_impl_preamble(), + instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, kernel_func, fn + ), + launcher, + f", {_cppbool(k.skip_scale_wait)}", + ) + + +def gen_bmm_mxscale_mouter_tunable_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + launcher = ( + _BMM_MOUTER_TUNABLE_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", kernel_func) + .replace("@@SSW@@", _cppbool(k.skip_scale_wait)) + ) + # host-only: device instantiations are shared with the mouter family. + _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + instance_impl_preamble(), + instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, kernel_func, fn + ), + launcher, + f", {_cppbool(k.skip_scale_wait)}", + emit_device=False, + ) + + +# ---- pipeline (kids 150/158/151/152) ---- +# Dual bf16/fp32 traits (output dtype baked into the traits tuple slot 3), +# non-splitk scale kargs, BLOCK_SIZE 512. Flags pick one of four scale kernels. +_BMM_PIPELINE_LAUNCHER_BODY = r""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +// mmajor: O/Y are [M, batch, *] (dim0=M, dim1=batch); wo_a stays batch-major. +// splitK must be 1 (checked by caller). Caller does dtype/arch/common checks. +template +void +@@NAME@@( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int /*splitK*/) +{ + using Bf16Traits = @@NAME@@_Bf16Traits; + using Fp32Traits = @@NAME@@_Fp32Traits; + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + // No M alignment requirement: the kernel bounds its A / sfa / C buffers to the + // tile's valid row window, so a partial trailing M tile is masked by buffer OOB. + AITER_CHECK(N % Bf16Traits::B_N == 0, + "@@NAME@@ requires N % ", Bf16Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Bf16Traits::B_K == 0, + "@@NAME@@ requires K % ", Bf16Traits::B_K, " == 0, got ", K); +@@K1024_CHECK@@ + opus_gemm_scale_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ptr_c = Y.data_ptr(); + kargs.m = M; + kargs.n = N; + kargs.k = K; + kargs.batch = batch; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + + const int num_tiles_m = (M + Bf16Traits::B_M - 1) / Bf16Traits::B_M; + const int num_tiles_n = N / Bf16Traits::B_N; + dim3 grid_main(num_tiles_m * num_tiles_n, 1, batch); + dim3 block_main(Bf16Traits::BLOCK_SIZE); + auto stream = aiter::getCurrentHIPStream(); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@<<>>(kargs); + } else { + @@KERNEL@@<<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" + + +def _bmm_pipeline_dual_traits_alias(k, traits_name): + def one(suffix, out_dtype): + return ( + f"using {k.name}_{suffix} = {traits_name}<{k.BLOCK_SIZE},\n" + f" opus::seq<{k.B_M}, {k.B_N}, {k.B_K}>,\n" + f" opus::tuple,\n" + f" opus::seq<{k.VEC_A}, {k.VEC_B}, {k.VEC_C}>,\n" + f" opus::seq<{k.GROUP_M}, {k.GROUP_N}, {k.GROUP_K}>>;\n" + ) + + return "\n" + one("Bf16Traits", "bf16_t") + one("Fp32Traits", "fp32_t") + + +def gen_bmm_mxscale_pipeline_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + if k.preload_sf_lds: + real_kernel = "gemm_a8w8_scale_preload_sf_kernel" + elif k.k1024_lb1: + real_kernel = "gemm_a8w8_scale_k1024_lb1_kernel" + elif k.k1024_only: + real_kernel = "gemm_a8w8_scale_k1024_kernel" + else: + real_kernel = "gemm_a8w8_scale_kernel" + + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + k1024_check = "" + if k.k1024_only or k.k1024_lb1: + k1024_check = ( + f' AITER_CHECK(K == 1024, "{k.name} requires K == 1024, got ", K);\n' + ) + launcher = ( + _BMM_PIPELINE_LAUNCHER_BODY.replace("@@NAME@@", k.name) + .replace("@@KERNEL@@", real_kernel) + .replace("@@K1024_CHECK@@", k1024_check) + ) + + traits_aliases = _bmm_pipeline_dual_traits_alias(k, traits_name) + host_tu = instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, real_kernel, fn + ) + INSTANCE_IMPL = ( + f"{instance_impl_preamble()}\n{host_tu}\n{traits_aliases}\n{launcher}" + ) + Path(os.path.join(cg.impl_path, f"{k.name}.cuh")).write_text(INSTANCE_IMPL) + + host_extra = ( + ",\n aiter_tensor_t &x_scale," + "\n aiter_tensor_t &w_scale," + "\n int splitK" + ) + for dtype in k.output_dtypes: + host_decl = ( + f"template void\n{k.name}<{dtype}>(\n" + f" aiter_tensor_t &O,\n aiter_tensor_t &wo_a,\n" + f" aiter_tensor_t &Y{host_extra});\n" + ) + cg._host_instantiations.append( + {"kid_name": k.name, "dtype": dtype, "host_decl": host_decl} + ) + + for dtype_tag, traits_suffix in (("bf16", "Bf16Traits"), ("fp32", "Fp32Traits")): + decl = ( + f"template __global__ void {real_kernel}<\n" + f" {k.name}_{traits_suffix}>({kargs_name});\n" + ) + cg._device_instantiations.append( + {"kid_name": k.name, "dtype": dtype_tag, "device_decl": decl} + ) + + +# ---- fused (kid 100) ---- +# Fused-reduce split-K path (counter/atomic variant). Same 256x32x128x128 wg2 +# traits + gemm_a8w8_mxscale_flatmm_splitk_kernel device symbols as standard kid 0/32 -> host-only emit. +_BMM_FUSED_LAUNCHER_BODY = r""" +#if !defined(__HIP_DEVICE_COMPILE__) && !defined(__HIPCC_RTC__) +// mmajor fused-reduce launcher: the main kernel accumulates partials into the +// Y buffer directly via an atomic tile counter (no separate reduce kernel). +// Caller (opus_bmm.cu dispatch) does dtype/arch/common checks. +template +void +@@NAME@@( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int splitK) +{ + using Traits = @@NAME@@_Traits; + AITER_CHECK(splitK >= 1, "splitK must be >= 1"); + + const int M = O.size(0); + const int batch = O.size(1); + const int N = wo_a.size(1); + const int K = O.size(2); + // No M alignment: the partial tile is masked in-kernel (see the launcher above). + AITER_CHECK(N % Traits::B_N == 0, + "@@NAME@@ requires N % ", Traits::B_N, " == 0, got ", N); + AITER_CHECK(K % Traits::B_K == 0, + "@@NAME@@ requires K % ", Traits::B_K, " == 0, got ", K); + + const int split_k = splitK; + const bool no_split_k = (split_k == 1); + const int total_iters = K / Traits::B_K; + const int iters_full = (total_iters + split_k - 1) / split_k; + const int last_loops = total_iters - (split_k - 1) * iters_full; + AITER_CHECK(last_loops >= Traits::prefetch_k_iter, + "@@NAME@@ requires every split to have at least ", + Traits::prefetch_k_iter, " K-tiles; K=", K, + " gives total_iters=", total_iters, ", splitK=", split_k, + ", last split loops=", last_loops); + + const int num_tiles_m = (M + Traits::B_M - 1) / Traits::B_M; + const int num_tiles_n = (N + Traits::B_N - 1) / Traits::B_N; + const int padded_M = num_tiles_m * Traits::B_M; + const int padded_N = num_tiles_n * Traits::B_N; + const size_t partial_bytes = (size_t)split_k * (size_t)batch + * (size_t)padded_M * (size_t)padded_N * sizeof(float); + const size_t counter_offset = (partial_bytes + 255) & ~((size_t)255); + const size_t counter_bytes = (size_t)batch * (size_t)num_tiles_m + * (size_t)num_tiles_n * sizeof(int); + + auto stream = aiter::getCurrentHIPStream(); + + opus_gemm_scale_splitk_kargs_gfx950 kargs{}; + kargs.ptr_a = O.data_ptr(); + kargs.ptr_b = wo_a.data_ptr(); + kargs.ws_handle = nullptr; + kargs.m = M; kargs.n = N; kargs.k = K; kargs.batch = batch; + kargs.split_k = split_k; + kargs.stride_a = (int)O.stride(0); + kargs.stride_b = (int)wo_a.stride(1); + kargs.stride_ws = padded_N; + kargs.stride_a_batch = (int)O.stride(1); + kargs.stride_b_batch = (int)wo_a.stride(0); + kargs.stride_ws_batch = padded_M * padded_N; + kargs.ptr_sfa = x_scale.data_ptr(); + kargs.ptr_sfb = w_scale.data_ptr(); + kargs.stride_sfa = (int)x_scale.stride(0); + kargs.stride_sfa_batch = (int)x_scale.stride(1); + kargs.stride_sfb = (int)w_scale.stride(1); + kargs.stride_sfb_batch = (int)w_scale.stride(0); + + dim3 grid_main(num_tiles_m * num_tiles_n * split_k, 1, batch); + dim3 block_main(Traits::BLOCK_SIZE); + if (no_split_k) { + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } + return; + } + + extern opus_splitk_ws_handle* opus_splitk_ws_get(hipStream_t, bool); + hipStreamCaptureStatus capture_status = hipStreamCaptureStatusNone; + HIP_CALL(hipStreamIsCapturing(stream, &capture_status)); + const bool capturing = (capture_status != hipStreamCaptureStatusNone); + auto* ws_handle = opus_splitk_ws_get(stream, /*allow_create=*/!capturing); + + const size_t ws_bytes = counter_offset + counter_bytes; + if (ws_handle->ptr == nullptr || ws_bytes > ws_handle->bytes) { + AITER_CHECK(!capturing, + "splitk workspace grow inside HIP graph capture is not supported"); + void* new_ptr = nullptr; + const size_t kGrowAlign = (size_t)4 * 1024 * 1024; + size_t grow_bytes = ((ws_bytes + kGrowAlign - 1) / kGrowAlign) * kGrowAlign; + HIP_CALL(hipMalloc(&new_ptr, grow_bytes)); + if (ws_handle->ptr != nullptr) { + HIP_CALL(hipDeviceSynchronize()); + HIP_CALL(hipFree(ws_handle->ptr)); + } + ws_handle->ptr = new_ptr; + ws_handle->bytes = grow_bytes; + } + kargs.ws_handle = ws_handle; + + kargs.ptr_c = Y.data_ptr(); + kargs.stride_c = (int)Y.stride(0); + kargs.stride_c_batch = (int)Y.stride(1); + kargs.counter_offset_bytes = counter_offset; + HIP_CALL(hipMemsetAsync(static_cast(ws_handle->ptr) + counter_offset, + 0, counter_bytes, stream)); + if (Y.dtype() == AITER_DTYPE_bf16) { + @@KERNEL@@ + <<>>(kargs); + } else { + @@KERNEL@@ + <<>>(kargs); + } +} +#endif // launcher only on regular host pass +""" + + +def gen_bmm_mxscale_fused_instance( + cg, + k, + pipeline_header, + traits_header, + kernel_func, + da, + db, + traits_name, + kargs_name, + kargs_template_vars, + instance_impl_preamble, + instance_impl_host_tu_split, + record_one_instantiation, + **_unused, +): + _, tpl, fn = kargs_template_vars(k.kernel_tag, kargs_name) + launcher = _BMM_FUSED_LAUNCHER_BODY.replace("@@NAME@@", k.name).replace( + "@@KERNEL@@", kernel_func + ) + # host-only: device symbols are shared + # with the standard flatmm split-K kid 0/32 (same traits) and emitted there. + _emit_bmm_specialized( + cg, + k, + kernel_func, + traits_name, + kargs_name, + da, + db, + instance_impl_preamble(), + instance_impl_host_tu_split( + traits_header, pipeline_header, tpl, kernel_func, fn + ), + launcher, + "", + emit_device=False, + ) + + # ---------- Self-register at import time ---------- register_emit("gfx950", "a16w16_persistent", gen_persistent_instance) register_emit("gfx950", "a8w8_scale", gen_scale_instance) +register_emit("gfx950", "a8w8_mxscale", gen_scale_instance) register_emit("gfx950", "a16w16", gen_noscale_instance_gfx950) register_emit("gfx950", "a8w8", gen_noscale_instance_gfx950) register_emit("gfx950", "a16w16_mono_tile", gen_mono_tile_instance) register_emit("gfx950", "a16w16_flatmm", gen_flatmm_instance) register_emit("gfx950", "a16w16_flatmm_splitk", gen_flatmm_splitk_instance) + + +def _register_bmm_emit(kernel_tag, fn, launcher_tile_mult): + """register_emit + the m_align cross-check for the BMM families. + + `launcher_tile_mult` is the B_M multiple this family's launcher body + hardcodes in its M guard, or 0 for the bodies that mask a partial M tile and + emit no M check. Checking it at emit time is what stops the guard and + OpusGemmInstance.m_align -- which the tuner and the runtime dispatch both + read -- from drifting apart. + """ + + def emit(cg, k, **kwargs): + _assert_m_align(k, launcher_tile_mult) + return fn(cg, k, **kwargs) + + emit.__name__ = fn.__name__ + register_emit("gfx950", kernel_tag, emit) + + +# _BMM_MXSCALE_SPLITK_LAUNCHER_BODY / _BMM_PIPELINE_LAUNCHER_BODY / +# _BMM_FUSED_LAUNCHER_BODY emit no M check ("No M alignment ..."); minterleave +# guards MI(=2)*B_M, wave4m2 guards LOGICAL_B_M(=2*B_M), the rest guard B_M. +_register_bmm_emit( + "a8w8_mxscale_bmm_flatmm_splitk", gen_bmm_mxscale_flatmm_splitk_instance, 0 +) +_register_bmm_emit( + "a8w8_mxscale_bmm_minterleave", gen_bmm_mxscale_minterleave_instance, 2 +) +_register_bmm_emit("a8w8_mxscale_bmm_fused", gen_bmm_mxscale_fused_instance, 0) +_register_bmm_emit("a8w8_mxscale_bmm_pipeline", gen_bmm_mxscale_pipeline_instance, 0) +_register_bmm_emit("a8w8_mxscale_bmm_mouter", gen_bmm_mxscale_mouter_instance, 1) +_register_bmm_emit( + "a8w8_mxscale_bmm_mouter_tunable", gen_bmm_mxscale_mouter_tunable_instance, 1 +) +_register_bmm_emit("a8w8_mxscale_bmm_wave8n2", gen_bmm_mxscale_wave8n2_instance, 1) +_register_bmm_emit( + "a8w8_mxscale_bmm_wave4m2_selfload", gen_bmm_mxscale_wave4m2_selfload_instance, 2 +) diff --git a/csrc/opus_gemm/gen_instances.py b/csrc/opus_gemm/gen_instances.py index 29c7b91d3b..3f726e8dd3 100644 --- a/csrc/opus_gemm/gen_instances.py +++ b/csrc/opus_gemm/gen_instances.py @@ -29,6 +29,7 @@ HEURISTIC_DEFAULT_KIDS, OpusGemmInstance, a8w8_kernels_list, + a8w8_mxscale_bmm_kernel_lists, a8w8_scale_kernels_list, a16w16_flatmm_kernels_list, a16w16_flatmm_splitk_kernels_list, @@ -138,6 +139,15 @@ def _kernel_func_for(k): INPUT_DTYPE_MAP = { "a8w8_scale": ("fp8_t", "fp8_t"), + "a8w8_mxscale": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_flatmm_splitk": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_fused": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_minterleave": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_mouter": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_mouter_tunable": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_pipeline": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_wave8n2": ("fp8_t", "fp8_t"), + "a8w8_mxscale_bmm_wave4m2_selfload": ("fp8_t", "fp8_t"), "a8w8": ("fp8_t", "fp8_t"), "a8w8_blockscale_bpreshuffle_singlebuf": ("fp8_t", "fp8_t"), **{tag: ("bf16_t", "bf16_t") for tag in _A16W16_TAGS}, @@ -172,6 +182,45 @@ def _kernel_func_for(k): def _kargs_template_vars(kernel_tag, kargs_name): + # a8w8_mxscale BMM flatmm splitK kernel has two extra compile-time booleans + # (DIRECT_ONLY, PREFETCH_SCALE) plus a non-void D_OUT after Traits. The fused + # host TU must forward-declare all four template params so the launcher body + # (which launches gemm_a8w8_mxscale_flatmm_splitk_kernel) compiles without pulling in the device pipeline header. + if kernel_tag in ( + "a8w8_mxscale_bmm_flatmm_splitk", + "a8w8_mxscale_bmm_fused", + ): + return ( + "", + ", typename D_OUT, bool DIRECT_ONLY, bool PREFETCH_SCALE, bool PRELOAD_SF_LDS", + kargs_name, + ) + # BMM M-tile-interleaved kernel: . The + # fused host TU must forward-declare all three template params so the launcher + # body's gemm_a8w8_mxscale_flatmm_minterleave_kernel + # <<<...>>> call compiles without the device pipeline header. + if kernel_tag == "a8w8_mxscale_bmm_minterleave": + return "", ", typename D_OUT, bool SKIP_SCALE_WAIT", kargs_name + # BMM specialized pipelines: forward-declare the exact kernel template params + # so the fused host TU's <<<...>>> call compiles against only the traits header. + if kernel_tag == "a8w8_mxscale_bmm_pipeline": + # scale-pipeline kernels are templated on a single Traits (output dtype is + # baked into the traits tuple) -> no extra template params. + return "", "", kargs_name + if kernel_tag in ( + "a8w8_mxscale_bmm_mouter", + "a8w8_mxscale_bmm_mouter_tunable", + ): + return "", ", typename D_OUT, bool SKIP_SCALE_WAIT", kargs_name + if kernel_tag == "a8w8_mxscale_bmm_wave8n2": + return "", ", typename D_OUT", kargs_name + if kernel_tag == "a8w8_mxscale_bmm_wave4m2_selfload": + return ( + "", + ", typename D_OUT, bool SKIP_SCALE_WAIT, bool PACK_SCALE_ON_DEMAND", + kargs_name, + ) # Paired W3 kernels: fn arg 'Kargs' so deduction keeps host/device mangling. if kernel_tag in _NOSPLIT or kernel_tag in _SPLITK: return f", {kargs_name}", ", typename Kargs", "Kargs" @@ -638,6 +687,61 @@ def _emit_map(f, macro_name, ctype): f.write(header) _emit_map(f, "GENERATE_A8W8_TUNE_LOOKUP_BF16", "bf16_t") + def gen_bmm_mxscale_tune_lookup(self, kernels_dict): + """Emit opus_bmm_mxscale_tune_lookup.h: int-kid -> launcher map for the + a8w8_mxscale BMM flatmm split-K family (gfx950-only). + + Mirrors gen_a8w8_tune_lookup, but the kid->name mapping lives in + a8w8_mxscale_bmm_flatmm_splitk_kernels_list (kid-keyed), NOT in + kernels_dict (which is name-keyed for the BMM family so gen_manifest_head + can dedup identical geometries). We iterate the kid-keyed source directly + so every switch kid keeps its historical number, even when two kids share + one launcher symbol (e.g. 0 and 32 -> same geometry -> same &launcher). + + The launcher templates static_assert D_C == float (Y=bf16 is produced by + the reduce kernel from an fp32 workspace), so only the fp32_t + specialization is instantiated -> emit a single fp32 macro. The dispatch + wrapper in opus_bmm.cu combines this with the hand-written specialized + pipelines (mouter / wave*n* / minterleave / pipeline / fused). + """ + header = """#pragma once +// SPDX-License-Identifier: MIT +// Copyright (C) 2025-2026, Advanced Micro Devices, Inc. All rights reserved. +// +// Auto-generated. Do not edit. See gen_instances.py:gen_bmm_mxscale_tune_lookup. +// +// fp32-workspace flat map for a8w8_mxscale BMM flatmm split-K tuning (gfx950). +// See opus_bmm.cu opus_bmm_a8w8_mxscale_tune_dispatch(). +""" + entry = """\ + {{ {kid}, &{kernel_name} }}, \\ +""" + + # The specialized-pipeline families (minterleave, ...) share the same + # int-kid -> launcher map and uniform launcher signature; concatenate + # their kid-keyed source lists so every kid keeps its historical number. + def _emit_map(f, macro_name, ctype): + f.write(f"#define {macro_name}(CTYPE) \\\n") + rows = [ + (kid, k.name) + for kernels in a8w8_mxscale_bmm_kernel_lists + for kid, k in kernels.items() + if ctype in k.output_dtypes + ] + rows.sort(key=lambda row: row[0]) + for index, (kid, name) in enumerate(rows): + line = entry.format(kid=kid, kernel_name=name) + if index == len(rows) - 1: + line = line.rstrip().rstrip("\\").rstrip() + "\n" + f.write(line) + f.write("\n") + + with open( + os.path.join(self.working_path, "opus_bmm_mxscale_tune_lookup.h"), "w" + ) as f: + f.write(header) + _emit_map(f, "GENERATE_BMM_MXSCALE_FLATMM_SPLITK_LOOKUP_FP32", "fp32_t") + def gen_manifest_head(self, kernels_dict): # Forward declarations for every launcher symbol the dispatcher references. MANIFEST_HEAD = """#pragma once @@ -677,11 +781,36 @@ def gen_manifest_head(self, kernels_dict): aiter_tensor_t &Y, std::optional bias, int splitK); +""" + # a8w8_mxscale BMM flatmm split-K launcher: mmajor layout with two fp8 + # scale tensors + an int splitK, dispatched by the hand-written + # opus_bmm.cu switch (not the (M,N,K) lookup table). + MANIFEST_BMM_MXSCALE_SPLITK = """ +template +void +{kernel_name}( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int splitK); """ with open(os.path.join(self.working_path, "opus_gemm_manifest.h"), "w") as f: f.write(MANIFEST_HEAD) for k in kernels_dict.values(): - if k.kernel_tag in A16W16_TUNE_TAGS: + if k.kernel_tag in ( + "a8w8_mxscale_bmm_flatmm_splitk", + "a8w8_mxscale_bmm_fused", + "a8w8_mxscale_bmm_minterleave", + "a8w8_mxscale_bmm_mouter", + "a8w8_mxscale_bmm_mouter_tunable", + "a8w8_mxscale_bmm_pipeline", + "a8w8_mxscale_bmm_wave8n2", + "a8w8_mxscale_bmm_wave4m2_selfload", + ): + f.write(MANIFEST_BMM_MXSCALE_SPLITK.format(kernel_name=k.name)) + elif k.kernel_tag in A16W16_TUNE_TAGS: f.write(MANIFEST_NOSCALE_4ARG.format(kernel_name=k.name)) elif k.kernel_tag in NOSCALE_TAGS: f.write(MANIFEST_NOSCALE_3ARG.format(kernel_name=k.name)) @@ -902,6 +1031,7 @@ def gen_instances(self, kernels_dict): self.gen_manifest_head(kernels_dict) self.gen_a16w16_tune_lookup(kernels_dict) self.gen_a8w8_tune_lookup(kernels_dict) + self.gen_bmm_mxscale_tune_lookup(kernels_dict) def get_tune_dict(tune_dict_csv): @@ -1181,6 +1311,20 @@ def _expand_tune_paths(spec): # Build the per-kid dict that drives codegen. kdict = {kid: kernels_list[kid] for kid in sorted(S)} + # a8w8_mxscale BMM flatmm split-K family (gfx950-only). These live in the + # opus_bmm.cu switch's PRIVATE kid namespace (ints 0/32/64/128/...), which + # collides with the global integer kids in kernels_list/S, so we never put + # them in S. Instead merge them into kdict keyed by kernel NAME: gen_instance + # only reads the value (k), gen_manifest_head emits by k.name, and every + # lookup/tune emitter gates on isinstance(key,int|tuple)+tag so the string + # keys are skipped. Name-keying also auto-dedups kids with identical geometry + # (e.g. switch kids 0 and 32 -> one launcher symbol). Always emitted (like + # a8w8_mxscale) so the opus_bmm dispatch never hits a missing symbol. + if target_arches is None or "gfx950" in target_arches: + for _bmm_list in a8w8_mxscale_bmm_kernel_lists: + for _bmm_k in _bmm_list.values(): + kdict[_bmm_k.name] = _bmm_k + print( f"[opus gen_instances] subset compile: |S|={len(S)} kids " f"(CSV={len(csv_kids)}, sidecar={len(sidecar_kids)}, heuristic={len(HEURISTIC_DEFAULT_KIDS)})" diff --git a/csrc/opus_gemm/include/gfx950/opus_bmm_launchers_a8w8_mxscale_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_bmm_launchers_a8w8_mxscale_gfx950.cuh new file mode 100644 index 0000000000..afc4a2d979 --- /dev/null +++ b/csrc/opus_gemm/include/gfx950/opus_bmm_launchers_a8w8_mxscale_gfx950.cuh @@ -0,0 +1,48 @@ +// SPDX-License-Identifier: MIT +// Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. +#pragma once + +// Shared host-side helpers for the a8w8 mxscale BMM kernel families. The +// per-kid host launchers are now fully codegen'd (impl/*.cuh, compiled in +// all_instances_host_gfx950.cu) as inlined template functions, mirroring the +// opus_gemm module. The old hand-written `*_impl` launcher templates that used +// to live here have been removed; only the common shape/dtype check remains. +// The device-kernel definitions are provided by the flatmm split-K pipeline +// header (included ahead of this one in opus_bmm.cu), so no forward +// declarations are needed here. Host pass only. +#ifndef __HIP_DEVICE_COMPILE__ + +#include "opus_bmm.h" +#include "opus_gemm_arch.cuh" +#include "opus_build_archs.h" +#include "opus_gemm_utils.cuh" // bf16_t / fp32_t +#include "aiter_stream.h" + +#include + +static void opus_bmm_a8w8_common_checks(aiter_tensor_t &O, aiter_tensor_t &wo_a, + aiter_tensor_t &Y, const char *who) +{ + aiter_detail::g_aiter_can_throw = true; + AITER_CHECK(O.dim() == 3 && wo_a.dim() == 3 && Y.dim() == 3, + who, ": O/wo_a/Y must be 3D " + "([M,batch,K] / [batch,N,K] / [M,batch,N])"); + AITER_CHECK(O.dtype() == AITER_DTYPE_fp8 && wo_a.dtype() == AITER_DTYPE_fp8, + who, ": O and wo_a must be fp8"); + AITER_CHECK(Y.dtype() == AITER_DTYPE_fp32 || Y.dtype() == AITER_DTYPE_bf16, + who, ": Y must be fp32 or bf16"); + // The kernels index A/B along K with unit stride (kargs carries only M/N/batch + // strides, never a K stride), so K must be the innermost contiguous dim. The + // batch axis position is free -- it is fully described by stride_*_batch -- so + // any batch layout (m-major [M,batch,K], batch-major view, ...) is accepted as + // long as K stays contiguous. Reject anything else here (host-side, once per + // launch) rather than silently producing wrong results. + AITER_CHECK(O.stride(2) == 1, who, + ": O (x) must be K-contiguous (stride(2)==1); got stride ", + (long)O.stride(2)); + AITER_CHECK(wo_a.stride(2) == 1, who, + ": wo_a must be K-contiguous (stride(2)==1); got stride ", + (long)wo_a.stride(2)); +} + +#endif // __HIP_DEVICE_COMPILE__ diff --git a/csrc/opus_gemm/include/gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh new file mode 100644 index 0000000000..89521a1f1d --- /dev/null +++ b/csrc/opus_gemm/include/gfx950/opus_bmm_pipeline_a8w8_mxscale_gfx950.cuh @@ -0,0 +1,1034 @@ +// SPDX-License-Identifier: MIT +// Copyright (C) 2025-2026, Advanced Micro Devices, Inc. All rights reserved. +#pragma once + +// BMM a8w8 mxscale (e8m0) GEMM pipeline: scale-accumulation helpers plus all +// batched GEMM kernels (main / K1024 / preload-SFA / split-K). +// Reuses the shared a8w8_scale layout infrastructure from the base header. +#include "opus_gemm_pipeline_a8w8_scale_gfx950.cuh" +// pack_e8m0x4 (broadcast e8m0 -> x4 word) is shared via opus_gemm_utils.cuh +// (pulled in transitively), so both this header and the flatmm split-K pipeline +// reference one definition instead of a per-header copy. + +#ifdef __HIP_DEVICE_COMPILE__ + +template +OPUS_D void mma_scale_accum(Mma& mma, const VA& v_a, const VB& v_b, + const VSFA& v_sfa, const VSFB& v_sfb, VC& v_c) { + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + if constexpr (std::is_same_v) { + // DSV4 scale is 128-block. The gfx950 scaled MFMA consumes 32-block + // E8M0 scale bytes; replicate one checkpoint byte across all four + // subblocks in the packed scale word to preserve 128-block semantics. + static_assert(T::B_K == T::GROUP_K, "e8m0 path assumes one K scale block per B_K"); + static_assert(T::HALF_B_N == T::GROUP_N, "e8m0 path assumes one B scale per half-tile"); + if constexpr (T::E_M == 1) { + const int scale_a = pack_e8m0x4(v_sfa[0]); + const int scale_b = pack_e8m0x4(v_sfb[0]); + v_c = mma(v_a, v_b, v_c, scale_a, scale_b, 0_I, 0_I); + } else { + using MMA = typename Mma::MMA; + constexpr int a_len = Mma::mma_a_len; + constexpr int b_len = Mma::mma_b_len; + constexpr int c_len = Mma::mma_c_len; + constexpr int rep_n_per_scale = T::GROUP_N / (T::W_N * T::T_N); + static_assert(T::GROUP_N % (T::W_N * T::T_N) == 0); + opus::static_for([&](auto im_c) { + constexpr int im = decltype(im_c)::value; + opus::static_for([&](auto in_c) { + constexpr int in = decltype(in_c)::value; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + const int scale_a = pack_e8m0x4(v_sfa[im * T::E_K + ik]); + const int scale_b = + pack_e8m0x4(v_sfb[(in / rep_n_per_scale) * T::E_K + ik]); + constexpr int i_tile_a = im * T::E_K + ik; + constexpr int i_tile_b = in * T::E_K + ik; + constexpr int i_tile_c = im * T::E_N + in; + auto s_a = opus::slice(v_a, + opus::number{}, + opus::number{}); + auto s_b = opus::slice(v_b, + opus::number{}, + opus::number{}); + auto s_c = opus::slice(v_c, + opus::number{}, + opus::number{}); + s_c = MMA{}(s_a, s_b, s_c, scale_a, scale_b, 0_I, 0_I); + opus::set_slice(v_c, s_c, + opus::number{}, + opus::number{}); + }); + }); + }); + } + } else { + typename Mma::vtype_c v_mma = mma(v_a, v_b, 0, 0); + scale_c_tile(v_mma, v_sfa, v_sfb, v_c); + } +} + +#endif // __HIP_DEVICE_COMPILE__ (scale-accum helpers) + +// ============================================================================ +// Hand-tuned GEMM kernel with block-scale (a8w8 + scale 1x128x128) +// Kernel definition visible on both passes (host pass needs it for stub generation). +// ============================================================================ + +template +__device__ __forceinline__ void gemm_a8w8_scale_kernel_impl(opus_gemm_scale_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_C = typename T::D_C; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + + const int grid_dim_x = opus::grid_size_x() / opus::block_size_x(); + int wgid = (opus::block_id_y() * grid_dim_x) + opus::block_id_x(); + // L2-locality rasterization (Triton-style GROUP_M grouping): process a panel + // of GROUP_M m-tiles across all n-tiles before advancing, iterating m-tiles + // fastest within the panel. This keeps each B[n_tile] (~1 MiB weights) hot in + // L2 across the panel's GROUP_M reuses, recovering high-G / large-M throughput. + constexpr int GROUP_M = 16; + const int num_tiles_m = ceil_div(kargs.m, T::B_M); + const int num_tiles_n = ceil_div(kargs.n, T::B_N); + const int tiles_per_group = GROUP_M * num_tiles_n; + + // A batch swizzle here (advance batch once per panel, spreading the C drain + // over more memory channels) was 15% faster in isolation but 1.5% slower in + // DPA serving: it pays for those channels with the GROUP_M reuse above, and a + // real step arrives with L2 contended. See opus_bmm.md. + const int group_id = wgid / tiles_per_group; + const int first_m = group_id * GROUP_M; + const int local = wgid - group_id * tiles_per_group; + const int m_remaining = num_tiles_m - first_m; + const int group_rows = m_remaining < GROUP_M ? m_remaining : GROUP_M; + int row = (first_m + (local % group_rows)) * T::B_M; + int col = (local / group_rows) * T::B_N; + + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + + // Base offsets in 64-bit: with a batch-in-the-middle layout stride_*_batch is + // M*K (A) / M*N (C), which overflows int32 well before the 4 GiB buffer limit. + // + // OOB masking for partial M tiles: bound A / sfa / C to this tile's valid row + // window so lanes past M read 0 and their stores are dropped by num_records. + // Any M then runs on a B_M tile (N and K stay divisible, so B and sfb need no + // bound); the garbage accumulated for masked rows is never stored. + // + // Clamp to B_M rather than the full (M - row) span: stride_a = batch*K here, + // so rows_avail*stride_a would overflow the 32-bit num_records field and wrap + // on a large-M / high-batch shape. Each WG owns one B_M tile and the base is + // already at `row`, so the clamp still masks the tail. + const int rows_left = kargs.m - row; + const int rows_avail = rows_left < T::B_M ? rows_left : T::B_M; + const unsigned int a_bytes = + (unsigned int)rows_avail * (unsigned int)kargs.stride_a * sizeof(D_A); + const unsigned int c_bytes = + (unsigned int)rows_avail * (unsigned int)kargs.stride_c * sizeof(D_C); + const unsigned int sfa_bytes = + (unsigned int)ceil_div(rows_avail, T::GROUP_M) * (unsigned int)kargs.stride_sfa * sizeof(D_SF); + + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + (size_t)batch_id*kargs.stride_a_batch + (size_t)row*kargs.stride_a, a_bytes); + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + (size_t)batch_id*kargs.stride_b_batch + (size_t)col*kargs.stride_b); + auto g_c = make_gmem(reinterpret_cast(kargs.ptr_c) + (size_t)batch_id*kargs.stride_c_batch + (size_t)row*kargs.stride_c + col, c_bytes); + + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + (size_t)batch_id*kargs.stride_sfa_batch + (size_t)(row/T::GROUP_M)*kargs.stride_sfa, sfa_bytes); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + (size_t)batch_id*kargs.stride_sfb_batch + (size_t)(col/T::GROUP_N)*kargs.stride_sfb); + + int wave_id_m = wave_id % T::T_M; + int wave_id_n = wave_id / T::T_M; + + auto u_ga = make_layout_ga(lane_id, wave_id_m, wave_id_n, kargs.stride_a); + auto u_sa = make_layout_sa(lane_id, wave_id_m, wave_id_n); + auto u_ra = make_layout_ra(lane_id, wave_id_m); + auto u_gb = make_layout_gb(lane_id, wave_id_m, wave_id_n, kargs.stride_b); + auto u_sb = make_layout_sb(lane_id, wave_id_m, wave_id_n); + auto u_rb = make_layout_rb(lane_id, wave_id_n); + + auto u_sfa = make_layout_sfa(lane_id, wave_id_m, kargs.stride_sfa); + + constexpr int smem_a_byte = T::smem_m_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_A); + __shared__ char smem_a[smem_a_byte * 4]; + smem s_a[2][2] = { + {make_smem(reinterpret_cast(smem_a)), + make_smem(reinterpret_cast(smem_a + smem_a_byte))}, + {make_smem(reinterpret_cast(smem_a + 2 * smem_a_byte)), + make_smem(reinterpret_cast(smem_a + 3 * smem_a_byte))} + }; + constexpr int smem_b_byte = T::smem_n_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_B); + __shared__ char smem_b[smem_b_byte * 4]; + smem s_b[2][2] = { + {make_smem(reinterpret_cast(smem_b)), + make_smem(reinterpret_cast(smem_b + smem_b_byte))}, + {make_smem(reinterpret_cast(smem_b + 2 * smem_b_byte)), + make_smem(reinterpret_cast(smem_b + 3 * smem_b_byte))} + }; + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + constexpr int ELEM_C = decltype(mma)::elem_c; + + typename decltype(mma)::vtype_a v_a[2]; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c[2][2]; + clear(v_c[0][0]); + clear(v_c[0][1]); + clear(v_c[1][0]); + clear(v_c[1][1]); + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + vtype_sfa v_sfa[2][2]; + vtype_sfb v_sfb[2][2]; + + auto a_offset = [&](int half_tile_m, int tile_k) { + return half_tile_m * T::HALF_B_M * kargs.stride_a + tile_k * T::B_K; + }; + auto b_offset = [&](int half_tile_n, int tile_k) { + return half_tile_n * T::HALF_B_N * kargs.stride_b + tile_k * T::B_K; + }; + auto sfa_offset = [&](int half_tile_m, int tile_k) { + return half_tile_m * (T::HALF_B_M / T::GROUP_M) * kargs.stride_sfa + tile_k * (T::B_K / T::GROUP_K); + }; + auto sfb_offset = [&](int half_tile_n, int tile_k) { + return half_tile_n * (T::HALF_B_N / T::GROUP_N) * kargs.stride_sfb + tile_k * (T::B_K / T::GROUP_K); + }; + + // kid157: preload the whole A-scale panel into LDS once, then read per-tile + // A-scale from LDS (ds_read/lgkmcnt) in the main loop instead of a per-tile + // global buffer_load_b8 (vmcnt) every K iteration. The panel is a compact + // [B_M/GROUP_M rows][K/B_K K-blocks] row-major byte tile (GROUP_M==1 and + // B_K==GROUP_K for this traits). The LDS buffer is sized for a compile-time + // K upper bound (SFA_K_MAX); the actual packed K-tile count is a runtime + // value so any K<=SFA_K_MAX (and K%B_K==0) works. SFA_K_MAX=8192 keeps the + // panel <=16 KiB, so total LDS stays 1 WG/CU. + constexpr int SFA_K_MAX = 8192; + constexpr int SFA_K_TILES_MAX = PRELOAD_SFA_LDS ? (SFA_K_MAX / T::B_K) : 1; + constexpr int SFA_ROWS = T::B_M / T::GROUP_M; + constexpr int SFA_LDS_BYTES = + PRELOAD_SFA_LDS ? (SFA_ROWS * SFA_K_TILES_MAX * (int)sizeof(D_SF)) : 1; + // 16B-aligned so the panel fill below can land ds_write_b128; a bare char + // array is only byte-aligned as far as the language is concerned. + __shared__ alignas(16) char smem_sfa[SFA_LDS_BYTES]; + D_SF* s_sfa_ptr = reinterpret_cast(smem_sfa); + // Runtime packed K-tile count (== loops); used as the compact LDS M-row + // stride so the read layout reuses make_layout_sfa with stride_sfa replaced. + const int sfa_k_tiles = PRELOAD_SFA_LDS ? (kargs.k / T::B_K) : 1; + auto u_sfa_lds = make_layout_sfa(lane_id, wave_id_m, sfa_k_tiles); + auto sfa_lds_offset = [&](int half_tile_m, int tile_k) { + return half_tile_m * (T::HALF_B_M / T::GROUP_M) * sfa_k_tiles + + tile_k * (T::B_K / T::GROUP_K); + }; + auto load_sfa = [&](int half_tile_m, int tile_k) { + if constexpr (PRELOAD_SFA_LDS) { + auto s = make_smem(s_sfa_ptr + sfa_lds_offset(half_tile_m, tile_k)); + return load(s, u_sfa_lds); + } else { + return load(g_sfa, u_sfa, sfa_offset(half_tile_m, tile_k)); + } + }; + + // kid158: same idea as PRELOAD_SFA_LDS but for the B (block) scale. SFB is + // tiny (B_N/GROUP_N N-groups * K/B_K K-tiles, block-shared across M) so the + // panel is a few dozen bytes; the win is purely removing the per-K-tile SFB + // global buffer_load from the steady-state vmcnt gate. Read layout mirrors + // sfb_offset with stride_sfb replaced by the compact per-N-group K length. + constexpr int SFB_K_MAX = 8192; + constexpr int SFB_K_TILES_MAX = PRELOAD_SFB_LDS ? (SFB_K_MAX / T::B_K) : 1; + constexpr int SFB_SPK = T::B_K / T::GROUP_K; // scales per K-tile + constexpr int SFB_NG_PER_HALF = T::HALF_B_N / T::GROUP_N; // N-groups per half-n + constexpr int SFB_ROWS = 2 * SFB_NG_PER_HALF; // N-groups in B_N tile + constexpr int SFB_LDS_BYTES = + PRELOAD_SFB_LDS ? (SFB_ROWS * SFB_K_TILES_MAX * SFB_SPK * (int)sizeof(D_SF)) : 1; + __shared__ char smem_sfb[SFB_LDS_BYTES]; + D_SF* s_sfb_ptr = reinterpret_cast(smem_sfb); + const int sfb_k_scales = PRELOAD_SFB_LDS ? ((kargs.k / T::B_K) * SFB_SPK) : 1; + auto sfb_lds_offset = [&](int half_tile_n, int tile_k) { + return half_tile_n * SFB_NG_PER_HALF * sfb_k_scales + tile_k * SFB_SPK; + }; + auto load_sfb = [&](int half_tile_n, int tile_k) { + if constexpr (PRELOAD_SFB_LDS) { + auto s = make_smem(s_sfb_ptr + sfb_lds_offset(half_tile_n, tile_k)); + return load(s, 0); + } else { + return load(g_sfb, sfb_offset(half_tile_n, tile_k)); + } + }; + // A preloaded panel is read from LDS and issues no vm ops, so it must drop out + // of every vmcnt threshold below: over-counting retires the wait early and lets + // the barrier release while the A/B async_loads are still landing. + constexpr int SFA_VM = PRELOAD_SFA_LDS ? 0 : T::sfa_buffer_load_insts; + constexpr int SFB_VM = PRELOAD_SFB_LDS ? 0 : T::sfb_buffer_load_insts; + + if constexpr (K1024_ONLY) { + static_assert(T::B_K == 128, "K1024_ONLY expects eight 128-wide K tiles"); + if (kargs.k != 1024) return; + } + if constexpr (PRELOAD_SFA_LDS) { + if (kargs.k > SFA_K_MAX || (kargs.k % T::B_K) != 0) return; + } + if constexpr (PRELOAD_SFB_LDS) { + if (kargs.k > SFB_K_MAX || (kargs.k % T::B_K) != 0) return; + } + const int loops = K1024_ONLY ? 8 : ceil_div(kargs.k, T::B_K); + int tic = 0, toc = 1; + + // kid158: issue the B-scale fetch before the A panel fill so the two global round + // trips overlap -- the A fill's own vmcnt(0) retires this load too. The panel is + // under one byte per thread, so one predicated load covers it and the value can + // sit in a register across the A fill. + using sfb_reg_t = decltype(load<1>(g_sfb, 0)); + sfb_reg_t sfb_val{}; + bool sfb_take = false; + if constexpr (PRELOAD_SFB_LDS) { + static_assert(SFB_ROWS * SFB_K_TILES_MAX * SFB_SPK <= T::BLOCK_SIZE, + "B-scale panel must fit one byte per thread"); + const int tid = opus::thread_id_x(); + sfb_take = tid < SFB_ROWS * sfb_k_scales; + if (sfb_take) { + const int ng = tid / sfb_k_scales; + const int ks = tid - ng * sfb_k_scales; + sfb_val = load<1>(g_sfb, ng * kargs.stride_sfb + ks); + } + } + + // kid157: one-shot cooperative fill of the A-scale panel into LDS, published by + // the barrier below. Byte-at-a-time was 16 iterations per thread at K=4096, each + // stalling on its own vmcnt(0). A chunk must not span two M rows nor land + // unaligned, so the width has to divide both sfa_k_tiles and stride_sfa. + if constexpr (PRELOAD_SFA_LDS) { + auto s_sfa = make_smem(s_sfa_ptr); + const int tid = opus::thread_id_x(); + const int sfa_total = SFA_ROWS * sfa_k_tiles; + auto fill = [&](auto vec_c) { + constexpr int VEC = decltype(vec_c)::value; + for (int idx = tid * VEC; idx < sfa_total; idx += T::BLOCK_SIZE * VEC) { + const int m = idx / sfa_k_tiles; + const int kt = idx - m * sfa_k_tiles; + s_sfa.template store( + load(g_sfa, m * kargs.stride_sfa + kt), idx); + } + }; + const int widths = sfa_k_tiles | kargs.stride_sfa; + if ((widths & 15) == 0) fill(number<16>{}); + else if ((widths & 3) == 0) fill(number<4>{}); + else fill(number<1>{}); + } + + // Land the B scale fetched above; its latency is already spent by now. + if constexpr (PRELOAD_SFB_LDS) { + if (sfb_take) { + make_smem(s_sfb_ptr).template store<1>(sfb_val, opus::thread_id_x()); + } + } + + // One barrier for both panels: draining after each fill in turn cost two full + // global round trips, since B could not issue until A's barrier released. The + // panels live in disjoint LDS, so the fills need no ordering between them. + // s_barrier does not retire LDS traffic, hence the explicit lgkmcnt wait. + if constexpr (PRELOAD_SFA_LDS || PRELOAD_SFB_LDS) { + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } + + // Prologue + v_sfa[tic][0] = load_sfa(0, 0); + v_sfb[tic][0] = load_sfb(0, 0); + async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, 0)); + async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, 0)); + v_sfa[tic][1] = load_sfa(1, 0); + v_sfb[tic][1] = load_sfb(1, 0); + async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, 0)); + async_load(g_b, s_b[tic][1].ptr, u_gb, u_sb, b_offset(1, 0)); + + if (wave_id_n == 1) __builtin_amdgcn_s_barrier(); + + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + v_sfa[toc][0] = load_sfa(0, 1); + v_sfb[toc][0] = load_sfb(0, 1); + async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, 1)); + async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, 1)); + async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, 1)); + + s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + SFA_VM + SFB_VM>{}); + __builtin_amdgcn_s_barrier(); + + v_a[0] = load(s_a[tic][0], u_ra); + __builtin_amdgcn_s_barrier(); + + // Main loop + for(int tile = 0; tile < loops - 2; tile += 2) { + // First tile + v_sfb[toc][1] = load_sfb(1, tile + 1); + v_b = load(s_b[tic][0], u_rb); + async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); + s_waitcnt_lgkmcnt(number{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][1] = load_sfa(1, tile + 1); + v_a[1] = load(s_a[tic][1], u_ra); + async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, tile + 2)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfb[tic][0] = load_sfb(0, tile + 2); + v_b = load(s_b[tic][1], u_rb); + async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, tile + 2)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[tic][0] = load_sfa(0, tile + 2); + v_a[0] = load(s_a[toc][0], u_ra); + async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, tile + 2)); + s_waitcnt_vmcnt(number< + 2 * T::a_buffer_load_insts + + T::b_buffer_load_insts + + 2 * SFA_VM + + SFB_VM>{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + // Second tile + v_sfb[tic][1] = load_sfb(1, tile + 2); + v_b = load(s_b[toc][0], u_rb); + async_load( + g_b, s_b[tic][1].ptr, u_gb, u_sb, + b_offset(1, tile + 2)); + s_waitcnt_lgkmcnt(number{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[toc][0], v_sfb[toc][0], v_c[0][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[tic][1] = load_sfa(1, tile + 2); + v_a[1] = load(s_a[toc][1], u_ra); + async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, tile + 3)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[toc][1], v_sfb[toc][0], v_c[1][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfb[toc][0] = load_sfb(0, tile + 3); + v_b = load(s_b[toc][1], u_rb); + async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, tile + 3)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[toc][0], v_sfb[toc][1], v_c[0][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][0] = load_sfa(0, tile + 3); + v_a[0] = load(s_a[tic][0], u_ra); + async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, tile + 3)); + s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + 2 * SFA_VM + SFB_VM>{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[toc][1], v_sfb[toc][1], v_c[1][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + } + + // Epilogue + { + int tile = loops - 2; + + v_sfb[toc][1] = load_sfb(1, tile + 1); + v_b = load(s_b[tic][0], u_rb); + async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][1] = load_sfa(1, tile + 1); + v_a[1] = load(s_a[tic][1], u_ra); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_b = load(s_b[tic][1], u_rb); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + tic ^= 1; + toc ^= 1; + } + + { + v_a[0] = load(s_a[tic][0], u_ra); + v_b = load(s_b[tic][0], u_rb); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_a[1] = load(s_a[tic][1], u_ra); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_b = load(s_b[tic][1], u_rb); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + } + + if (wave_id_n == 0) __builtin_amdgcn_s_barrier(); + + // Store results to global memory + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, wave_id_n, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + + auto c_offset = [&](int half_tile_m, int half_tile_n) { + return half_tile_m * T::HALF_B_M * kargs.stride_c + half_tile_n * T::HALF_B_N; + }; + + store(g_c, v_c[0][0], u_gc, c_offset(0, 0)); + store(g_c, v_c[0][1], u_gc, c_offset(0, 1)); + store(g_c, v_c[1][0], u_gc, c_offset(1, 0)); + store(g_c, v_c[1][1], u_gc, c_offset(1, 1)); +#else + // Non-gfx950 device pass: empty stub. a8w8 is gfx950-only; the host + // launcher symbol must still exist for the unconditional dispatcher + // reference, but the body uses gfx950-only intrinsics. +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, 2) void gemm_a8w8_scale_kernel(opus_gemm_scale_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + gemm_a8w8_scale_kernel_impl(kargs); +#else + // Non-gfx950 device pass: empty stub. +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, 2) void gemm_a8w8_scale_k1024_kernel(opus_gemm_scale_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + gemm_a8w8_scale_kernel_impl(kargs); +#else + // Non-gfx950 device pass: empty stub. +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, 1) void gemm_a8w8_scale_k1024_lb1_kernel(opus_gemm_scale_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + gemm_a8w8_scale_kernel_impl(kargs); +#else + // Non-gfx950 device pass: empty stub. +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// EXPERIMENTAL (kid158): kid150 + both the A (per-token) and B (block) scale panels preloaded into +// LDS, so the steady-state loop reads both SFA and SFB from LDS (ds_read) and the +// per-K-tile SFA/SFB global buffer_loads are removed from the vmcnt gate entirely. +// Supports any K<=8192 (K%B_K==0); LDS panels sized for the compile-time upper +// bound, packed K-tile count resolved at runtime. +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, 2) +void gemm_a8w8_scale_preload_sf_kernel(opus_gemm_scale_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + gemm_a8w8_scale_kernel_impl(kargs); +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// Split-K main kernel: computes one K partition into fp32 workspace. +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, 2) void gemm_a8w8_scale_splitk_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) { +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_C = typename T::D_C; + static_assert(std::is_same_v, "splitK main writes fp32 workspace"); + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + + int wgid_full = opus::block_id_x(); + int split_id = wgid_full % kargs.split_k; + int wgid = wgid_full / kargs.split_k; + const int num_tiles_n = ceil_div(kargs.n, T::B_N); + int row = (wgid / num_tiles_n) * T::B_M; + int col = (wgid % num_tiles_n) * T::B_N; + + const int total_iters = ceil_div(kargs.k, T::B_K); + const int iters_full = ceil_div(total_iters, kargs.split_k); + int loops = (split_id < kargs.split_k - 1) + ? iters_full + : (total_iters - (kargs.split_k - 1) * iters_full); + if (loops <= 0) return; + int k_start = split_id * iters_full * T::B_K; + int sf_start = split_id * iters_full * (T::B_K / T::GROUP_K); + + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + + // 64-bit base offsets (see the non-splitK path above): batch_id*stride_*_batch + // overflows int32 for large-M batch-in-the-middle layouts. + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + (size_t)batch_id*kargs.stride_a_batch + (size_t)row*kargs.stride_a + k_start); + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + (size_t)batch_id*kargs.stride_b_batch + (size_t)col*kargs.stride_b + k_start); + auto g_c = make_gmem(reinterpret_cast(kargs.ws_handle->ptr) + (size_t)split_id * kargs.batch * kargs.stride_ws_batch + (size_t)batch_id * kargs.stride_ws_batch + (size_t)row * kargs.stride_ws + col); + + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + (size_t)batch_id*kargs.stride_sfa_batch + (size_t)(row/T::GROUP_M)*kargs.stride_sfa + sf_start); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + (size_t)batch_id*kargs.stride_sfb_batch + (size_t)(col/T::GROUP_N)*kargs.stride_sfb + sf_start); + + int wave_id_m = wave_id % T::T_M; + int wave_id_n = wave_id / T::T_M; + + auto u_ga = make_layout_ga(lane_id, wave_id_m, wave_id_n, kargs.stride_a); + auto u_sa = make_layout_sa(lane_id, wave_id_m, wave_id_n); + auto u_ra = make_layout_ra(lane_id, wave_id_m); + auto u_gb = make_layout_gb(lane_id, wave_id_m, wave_id_n, kargs.stride_b); + auto u_sb = make_layout_sb(lane_id, wave_id_m, wave_id_n); + auto u_rb = make_layout_rb(lane_id, wave_id_n); + + auto u_sfa = make_layout_sfa(lane_id, wave_id_m, kargs.stride_sfa); + + constexpr int smem_a_byte = T::smem_m_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_A); + __shared__ char smem_a[smem_a_byte * 4]; + smem s_a[2][2] = { + {make_smem(reinterpret_cast(smem_a)), + make_smem(reinterpret_cast(smem_a + smem_a_byte))}, + {make_smem(reinterpret_cast(smem_a + 2 * smem_a_byte)), + make_smem(reinterpret_cast(smem_a + 3 * smem_a_byte))} + }; + constexpr int smem_b_byte = T::smem_n_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_B); + __shared__ char smem_b[smem_b_byte * 4]; + smem s_b[2][2] = { + {make_smem(reinterpret_cast(smem_b)), + make_smem(reinterpret_cast(smem_b + smem_b_byte))}, + {make_smem(reinterpret_cast(smem_b + 2 * smem_b_byte)), + make_smem(reinterpret_cast(smem_b + 3 * smem_b_byte))} + }; + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + constexpr int ELEM_C = decltype(mma)::elem_c; + + typename decltype(mma)::vtype_a v_a[2]; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c[2][2]; + clear(v_c[0][0]); + clear(v_c[0][1]); + clear(v_c[1][0]); + clear(v_c[1][1]); + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + vtype_sfa v_sfa[2][2]; + vtype_sfb v_sfb[2][2]; + + auto a_offset = [&](int half_tile_m, int tile_k) { + return half_tile_m * T::HALF_B_M * kargs.stride_a + tile_k * T::B_K; + }; + auto b_offset = [&](int half_tile_n, int tile_k) { + return half_tile_n * T::HALF_B_N * kargs.stride_b + tile_k * T::B_K; + }; + auto sfa_offset = [&](int half_tile_m, int tile_k) { + return half_tile_m * (T::HALF_B_M / T::GROUP_M) * kargs.stride_sfa + tile_k * (T::B_K / T::GROUP_K); + }; + auto sfb_offset = [&](int half_tile_n, int tile_k) { + return half_tile_n * (T::HALF_B_N / T::GROUP_N) * kargs.stride_sfb + tile_k * (T::B_K / T::GROUP_K); + }; + + int tic = 0, toc = 1; + + // Prologue + v_sfa[tic][0] = load(g_sfa, u_sfa, sfa_offset(0, 0)); + v_sfb[tic][0] = load(g_sfb, sfb_offset(0, 0)); + async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, 0)); + async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, 0)); + v_sfa[tic][1] = load(g_sfa, u_sfa, sfa_offset(1, 0)); + v_sfb[tic][1] = load(g_sfb, sfb_offset(1, 0)); + async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, 0)); + async_load(g_b, s_b[tic][1].ptr, u_gb, u_sb, b_offset(1, 0)); + + if (wave_id_n == 1) __builtin_amdgcn_s_barrier(); + + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + v_sfa[toc][0] = load(g_sfa, u_sfa, sfa_offset(0, 1)); + v_sfb[toc][0] = load(g_sfb, sfb_offset(0, 1)); + async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, 1)); + async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, 1)); + async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, 1)); + + s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); + __builtin_amdgcn_s_barrier(); + + v_a[0] = load(s_a[tic][0], u_ra); + __builtin_amdgcn_s_barrier(); + + // Main loop + for(int tile = 0; tile < loops - 2; tile += 2) { + // First tile + v_sfb[toc][1] = load(g_sfb, sfb_offset(1, tile + 1)); + v_b = load(s_b[tic][0], u_rb); + async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); + s_waitcnt_lgkmcnt(number{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 1)); + v_a[1] = load(s_a[tic][1], u_ra); + async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, tile + 2)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfb[tic][0] = load(g_sfb, sfb_offset(0, tile + 2)); + v_b = load(s_b[tic][1], u_rb); + async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, tile + 2)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[tic][0] = load(g_sfa, u_sfa, sfa_offset(0, tile + 2)); + v_a[0] = load(s_a[toc][0], u_ra); + async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, tile + 2)); + s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + 2 * T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + // Second tile + v_sfb[tic][1] = load(g_sfb, sfb_offset(1, tile + 2)); + v_b = load(s_b[toc][0], u_rb); + async_load(g_b, s_b[tic][1].ptr, u_gb, u_sb, b_offset(1, tile + 2)); + s_waitcnt_lgkmcnt(number{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[toc][0], v_sfb[toc][0], v_c[0][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[tic][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 2)); + v_a[1] = load(s_a[toc][1], u_ra); + async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, tile + 3)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[toc][1], v_sfb[toc][0], v_c[1][0]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfb[toc][0] = load(g_sfb, sfb_offset(0, tile + 3)); + v_b = load(s_b[toc][1], u_rb); + async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, tile + 3)); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[toc][0], v_sfb[toc][1], v_c[0][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][0] = load(g_sfa, u_sfa, sfa_offset(0, tile + 3)); + v_a[0] = load(s_a[tic][0], u_ra); + async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, tile + 3)); + s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + 2 * T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[toc][1], v_sfb[toc][1], v_c[1][1]); + sched_barrier_pairs<2, 0, 0>(); + sched_barrier_pairs<1, 2, 0>(); + sched_barrier_pairs<5, 4, 0>(); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + } + + // Epilogue + { + int tile = loops - 2; + + v_sfb[toc][1] = load(g_sfb, sfb_offset(1, tile + 1)); + v_b = load(s_b[tic][0], u_rb); + async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_sfa[toc][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 1)); + v_a[1] = load(s_a[tic][1], u_ra); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_b = load(s_b[tic][1], u_rb); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + tic ^= 1; + toc ^= 1; + } + + { + v_a[0] = load(s_a[tic][0], u_ra); + v_b = load(s_b[tic][0], u_rb); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_a[1] = load(s_a[tic][1], u_ra); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + + v_b = load(s_b[tic][1], u_rb); + __builtin_amdgcn_s_barrier(); + + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_scale_accum(mma, v_a[0], v_b, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); + mma_scale_accum(mma, v_a[1], v_b, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_sched_barrier(0); + } + + if (wave_id_n == 0) __builtin_amdgcn_s_barrier(); + + // Store results to global memory + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, wave_id_n, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, opus::make_tuple(kargs.stride_ws, 1_I), p_coord_c); + + auto c_offset = [&](int half_tile_m, int half_tile_n) { + return half_tile_m * T::HALF_B_M * kargs.stride_ws + half_tile_n * T::HALF_B_N; + }; + + store(g_c, v_c[0][0], u_gc, c_offset(0, 0)); + store(g_c, v_c[0][1], u_gc, c_offset(0, 1)); + store(g_c, v_c[1][0], u_gc, c_offset(1, 0)); + store(g_c, v_c[1][1], u_gc, c_offset(1, 1)); +#else + // Non-gfx950 device pass: empty stub. a8w8 is gfx950-only; the host + // launcher symbol must still exist for the unconditional dispatcher + // reference, but the body uses gfx950-only intrinsics. +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} diff --git a/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh new file mode 100644 index 0000000000..685490e96e --- /dev/null +++ b/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_mxscale_flatmm_splitk_gfx950.cuh @@ -0,0 +1,1996 @@ +// SPDX-License-Identifier: MIT +// Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. +// +// gfx950 fp8/e8m0 flatmm split-K pipeline for decode-oriented BMM. +// +// This is the first B_K=128 version: one K iteration maps to one DSv4 +// checkpoint scale block, and one consumer-wave M tile maps to one per-row A +// scale. The host launcher keeps this v1 on divisible decode shapes. +#pragma once + +#include "opus_gemm_traits_a8w8_scale_gfx950.cuh" +// opus_bmm_splitk_reduce_kernel: the split-K > 1 path's fp32->Y reduce. Lives +// in the shared reduce header so the codegen'd BMM launcher (which #includes +// this pipeline header on the non-fused device pass) can launch it. +#include "splitk_reduce_gfx950.cuh" + +#ifdef __HIP_DEVICE_COMPILE__ + +// ============================================================================ +// Layout helpers. Suffixed with _mxsk to avoid ODR collisions with the bf16 +// flatmm splitK helpers when both headers are included in a build. +// ============================================================================ + +template +inline __device__ auto make_layout_gmem_group_load_mxsk(int lane_id, int wave_id, int stride) { + constexpr int threads_k = T::LOAD_GROUP_K / T::VEC_A; + constexpr int threads_m_per_wave = opus::get_warp_size() / threads_k; + constexpr int interlanegroup_m = threads_m_per_wave / T::LOAD_GROUP_M_LANE; + constexpr int repeat_m = T::slots / WAVES; + + constexpr auto g_block_shape = opus::make_tuple( + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}); + + constexpr auto g_block_dim = opus::make_tuple( + opus::make_tuple(opus::p_dim{}, opus::y_dim{}, opus::p_dim{}, opus::p_dim{}), + opus::make_tuple(opus::p_dim{}, opus::y_dim{})); + + return opus::make_layout<0>( + g_block_shape, + opus::unfold_x_stride(g_block_dim, g_block_shape, opus::tuple{stride, 1_I}), + opus::unfold_p_coord(g_block_dim, + opus::tuple{lane_id / threads_k / T::LOAD_GROUP_M_LANE, + wave_id % WAVES, + (lane_id / threads_k) % T::LOAD_GROUP_M_LANE, + lane_id % threads_k})); +} + +template +inline __device__ auto make_layout_smem_group_load_mxsk(int lane_id, int wave_id) { + constexpr int repeat_m = T::slots / WAVES; + + constexpr auto s_block_shape = opus::make_tuple( + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}); + + constexpr auto s_block_dim = opus::make_tuple( + opus::make_tuple(opus::y_dim{}, opus::p_dim{}), + opus::make_tuple(opus::p_dim{}, opus::y_dim{})); + + return opus::make_layout<0>( + s_block_shape, + opus::unfold_x_stride(s_block_dim, s_block_shape, + opus::tuple{T::smem_linear_wave_per_async_load + T::smem_padding, 1_I}), + opus::unfold_p_coord(s_block_dim, opus::tuple{wave_id % WAVES, lane_id})); +} + +template +inline __device__ auto make_layout_ra_mxsk(int lane_id, int wave_id_m) { + constexpr int threads_k = opus::get_warp_size() / T::W_M; + constexpr int threads_m_per_wave = opus::get_warp_size() / threads_k; + constexpr int interlanegroup_m = threads_m_per_wave / T::LOAD_GROUP_M_LANE; + constexpr int per_block_load = T::slots * (T::smem_linear_wave_per_async_load + T::smem_padding); + constexpr int m_block_stride = T::NUM_LOAD_GROUPS_PER_BK * per_block_load; + + constexpr auto ra_block_shape = opus::make_tuple( + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number<2>{}, + opus::number{}, + opus::number{}); + + constexpr auto ra_block_dim = opus::make_tuple( + opus::make_tuple(opus::y_dim{}), + opus::make_tuple(opus::p_dim{}), + opus::make_tuple(opus::y_dim{}), + opus::make_tuple(opus::p_dim{}, opus::p_dim{}, opus::p_dim{}, + opus::y_dim{}, opus::p_dim{}, opus::y_dim{})); + + auto lane_id_m = lane_id % T::W_M; + + return opus::make_layout<0>( + ra_block_shape, + opus::unfold_x_stride(ra_block_dim, ra_block_shape, + opus::tuple{opus::number{}, + opus::number{}, + opus::number{}, + 1_I}), + opus::unfold_p_coord(ra_block_dim, + opus::tuple{lane_id_m % T::slots, + wave_id_m, + lane_id_m / T::slots, + lane_id_m % T::LOAD_GROUP_M_LANE, + lane_id / T::W_M})); +} + +template +inline __device__ auto make_layout_rb_mxsk(int lane_id) { + constexpr int grpk_b = opus::get_warp_size() / T::W_N; + constexpr int interlanegroup_n = T::W_N / T::LOAD_GROUP_N_LANE; + constexpr int loops_b = interlanegroup_n / T::slots; + constexpr int tiles_per_block_n = T::LOAD_GROUP_N / T::W_N; + constexpr int num_blocks_n = T::COM_REP_N / tiles_per_block_n; + constexpr int per_block_load = T::slots * (T::smem_linear_wave_per_async_load + T::smem_padding); + constexpr int n_block_stride = T::NUM_LOAD_GROUPS_PER_BK * per_block_load; + constexpr int n_intra_stride = T::LOAD_GROUP_N_LANE * 2 * grpk_b * T::VEC_B; + + constexpr auto rb_block_shape = opus::make_tuple( + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + opus::number<2>{}, + opus::number{}, + opus::number{}); + + constexpr auto rb_block_dim = opus::make_tuple( + opus::make_tuple(opus::y_dim{}), + opus::make_tuple(opus::p_dim{}), + opus::make_tuple(opus::y_dim{}, opus::p_dim{}), + opus::make_tuple(opus::y_dim{}), + opus::make_tuple(opus::p_dim{}, opus::y_dim{}, opus::p_dim{}, opus::y_dim{})); + + auto lane_id_n = lane_id % T::W_N; + + return opus::make_layout<0>( + rb_block_shape, + opus::unfold_x_stride(rb_block_dim, rb_block_shape, + opus::tuple{opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}, + 1_I}), + opus::unfold_p_coord(rb_block_dim, + opus::tuple{lane_id_n % T::slots, + lane_id_n / T::slots, + lane_id_n % T::LOAD_GROUP_N_LANE, + lane_id / T::W_N})); +} + +template +inline __device__ auto make_layout_sfa_mxsk(int lane_id, int wave_id_m, int stride_sfa) { + constexpr auto sfa_block_shape = opus::make_tuple( + opus::number{}, + opus::number{}, + opus::number{}, + opus::number{}); + + constexpr auto sfa_block_dim = opus::make_tuple( + opus::make_tuple(opus::y_dim{}, opus::p_dim{}, opus::p_dim{}), + opus::make_tuple(opus::y_dim{})); + + return opus::make_layout( + sfa_block_shape, + opus::unfold_x_stride(sfa_block_dim, sfa_block_shape, + opus::tuple{stride_sfa, 1_I}), + opus::unfold_p_coord(sfa_block_dim, + opus::tuple{wave_id_m, lane_id % T::W_M})); +} + +// pack_e8m0x4 (broadcast e8m0 -> x4 word) is shared via opus_gemm_utils.cuh. + +// Per-subtile scaled-MFMA loop -- the shared "else" body used whenever the +// register tile spans more than one MX scale group. The MMA issue pattern is +// identical for every scale layout; only where each subtile's packed scale comes +// from differs, so that is injected via scale_a_of(im, ik) / scale_b_of(in, ik). +// Providers receive compile-time (opus::number<>) subtile indices and return the +// packed-int32 e8m0x4 scale. This lets plain block-scale and shuffled/preloaded +// block-scale reuse one implementation. +// OPSEL == false (default): scale_a_of/scale_b_of return a broadcast-packed x4 +// e8m0 word (all 4 bytes equal) and every MFMA selects byte 0 -- legacy +// scalar behavior. +// OPSEL == true: scale_a_of/scale_b_of return a K-packed word holding the +// COM_REP_K distinct K-group e8m0 bytes (byte ik == the ik-th K group) and are +// K-independent; each MFMA selects its own byte through the compile-time +// scale_op_sel == ik. This drops the per-subtile broadcast pack and shrinks the +// K-direction scale register footprint to one word per M / N-scale group. +template +OPUS_D void mma_mxscale_subtile_loop(const VA& v_a, const VB& v_b, VC& v_c, + ScaleAOf&& scale_a_of, ScaleBOf&& scale_b_of) { + using MMA = typename Mma::MMA; + constexpr int a_len = Mma::mma_a_len; + constexpr int b_len = Mma::mma_b_len; + constexpr int c_len = Mma::mma_c_len; + opus::static_for([&](auto im_c) { + constexpr int im = decltype(im_c)::value; + opus::static_for([&](auto in_c) { + constexpr int in = decltype(in_c)::value; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + const int scale_a = scale_a_of(im_c, ik_c); + const int scale_b = scale_b_of(in_c, ik_c); + constexpr int i_tile_a = (im * T::COM_REP_K + ik); + constexpr int i_tile_b = (in * T::COM_REP_K + ik); + constexpr int i_tile_c = im * T::COM_REP_N + in; + auto s_a = opus::slice(v_a, + opus::number{}, + opus::number{}); + auto s_b = opus::slice(v_b, + opus::number{}, + opus::number{}); + auto s_c = opus::slice(v_c, + opus::number{}, + opus::number{}); + if constexpr (OPSEL) + s_c = MMA{}(s_a, s_b, s_c, scale_a, scale_b, ik_c, ik_c); + else + s_c = MMA{}(s_a, s_b, s_c, scale_a, scale_b, 0_I, 0_I); + opus::set_slice(v_c, s_c, + opus::number{}, + opus::number{}); + }); + }); + }); +} + +// How the per-subtile scales are materialized in the multi-scale-group path: +// preload -- pack every distinct scale once up front, then index the packed +// registers in the loop (fewer pack ops when a scale is reused +// across COM_REP_N / COM_REP_M). +// on_demand -- pack each subtile's scale inline (lower register pressure). +// opsel -- pack the COM_REP_K K-group scales into one word per M / N-scale +// group (native e8m0x4, no broadcast) and select the K byte per +// MFMA via the hardware scale_op_sel immediate. Fewest scale ALU +// ops + smallest K-direction scale register footprint. +enum class mxscale_pack { preload, on_demand, opsel }; + +template +OPUS_D void mma_mxscale_tiled(Mma& mma, const VA& v_a, const VB& v_b, + const VSFA& v_sfa, const VSFB& v_sfb, VC& v_c) { + static_assert(std::is_same_v); + static_assert((T::COM_REP_M == 1 || T::COM_REP_M == 2 || T::COM_REP_M == 4) + && (T::COM_REP_K == 1 || T::COM_REP_K == 2 || T::COM_REP_K == 4)); + static_assert(T::B_K % T::GROUP_K == 0); + constexpr int rep_n_per_scale = T::GROUP_N / (T::W_N * T::T_N); + static_assert(rep_n_per_scale > 0 && T::GROUP_N % (T::W_N * T::T_N) == 0); + // Whole register tile in a single scale group -> one (scale_a, scale_b) pair + // -> a single tiled-mma call covers the tile. + if constexpr (T::COM_REP_M == 1 && T::COM_REP_N <= rep_n_per_scale && T::COM_REP_K == 1) { + const int scale_a = pack_e8m0x4(v_sfa[0]); + const int scale_b = pack_e8m0x4(v_sfb[0]); + v_c = mma(v_a, v_b, v_c, scale_a, scale_b, 0_I, 0_I); + } else if constexpr (MODE == mxscale_pack::opsel) { + // One word per M-subtile / N-scale-group holding the COM_REP_K K-group + // e8m0 bytes; the subtile loop picks byte ik via scale_op_sel == ik. + // NOTE: reference path only. With the vec-wide (dword) scale load below, + // the shift/or here folds away, but op_sel packing still measures on par + // with or slightly slower than preload's broadcast pack across the tuned + // shapes, so preload stays the default. Kept for experimentation. + opus::vector_t packed_sfa; + opus::vector_t packed_sfb; + opus::static_for([&](auto im_c) { + constexpr int im = decltype(im_c)::value; + int w = 0; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + w |= (static_cast(v_sfa[im * T::SCALES_PER_BK + ik]) & 0xFF) << (8 * ik); + }); + packed_sfa[im] = w; + }); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + int w = 0; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + w |= (static_cast(v_sfb[ng * T::SCALES_PER_BK + ik]) & 0xFF) << (8 * ik); + }); + packed_sfb[ng] = w; + }); + mma_mxscale_subtile_loop(v_a, v_b, v_c, + [&](auto im_c, auto) { + return packed_sfa[decltype(im_c)::value]; + }, + [&](auto in_c, auto) { + return packed_sfb[decltype(in_c)::value / rep_n_per_scale]; + }); + } else if constexpr (MODE == mxscale_pack::preload) { + opus::vector_t packed_sfa; + opus::vector_t packed_sfb; + opus::static_for([&](auto im_c) { + constexpr int im = decltype(im_c)::value; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + packed_sfa[im * T::COM_REP_K + ik] = + pack_e8m0x4(v_sfa[im * T::SCALES_PER_BK + ik]); + }); + }); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + opus::static_for([&](auto ik_c) { + constexpr int ik = decltype(ik_c)::value; + packed_sfb[ng * T::COM_REP_K + ik] = + pack_e8m0x4(v_sfb[ng * T::SCALES_PER_BK + ik]); + }); + }); + mma_mxscale_subtile_loop(v_a, v_b, v_c, + [&](auto im_c, auto ik_c) { + return packed_sfa[decltype(im_c)::value * T::COM_REP_K + decltype(ik_c)::value]; + }, + [&](auto in_c, auto ik_c) { + return packed_sfb[(decltype(in_c)::value / rep_n_per_scale) * T::COM_REP_K + + decltype(ik_c)::value]; + }); + } else { + mma_mxscale_subtile_loop(v_a, v_b, v_c, + [&](auto im_c, auto ik_c) { + return pack_e8m0x4( + v_sfa[decltype(im_c)::value * T::SCALES_PER_BK + decltype(ik_c)::value]); + }, + [&](auto in_c, auto ik_c) { + return pack_e8m0x4( + v_sfb[(decltype(in_c)::value / rep_n_per_scale) * T::SCALES_PER_BK + + decltype(ik_c)::value]); + }); + } +} + +// Scale-packing strategy for the default multi-scale-group accum path. Flip +// between preload and opsel here to A/B the hardware scale_op_sel byte-select. +inline constexpr mxscale_pack MXSCALE_ACCUM_MODE = mxscale_pack::preload; + +// Thin wrappers preserving the original entry points / call sites. +template +OPUS_D void mma_mxscale_flatmm_accum(Mma& mma, const VA& v_a, const VB& v_b, + const VSFA& v_sfa, const VSFB& v_sfb, VC& v_c) { + mma_mxscale_tiled(mma, v_a, v_b, v_sfa, v_sfb, v_c); +} + +template +OPUS_D void mma_mxscale_flatmm_accum_on_demand(Mma& mma, const VA& v_a, const VB& v_b, + const VSFA& v_sfa, const VSFB& v_sfb, VC& v_c) { + mma_mxscale_tiled(mma, v_a, v_b, v_sfa, v_sfb, v_c); +} + +#endif // __HIP_DEVICE_COMPILE__ + +// ============================================================================ +// Main kernel: 4-wave flatmm splitK, fp32 workspace output. +// ============================================================================ + +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, Traits::WG_PER_CU) +void gemm_a8w8_mxscale_flatmm_splitk_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_C = typename T::D_C; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + static_assert(std::is_same_v, "flatmm splitK main writes fp32 workspace"); + static_assert(!DIRECT_ONLY || !std::is_void_v, + "DIRECT_ONLY requires an output dtype for direct Y stores"); + // tileN (T_M=1, T_N=2 consumer N-split) is only wired through the non-DIRECT + // producer/consumer path (barrier-synced; split_k==1 still stores Y directly + // via the branch below). The persistent DIRECT_ONLY schedule keeps its + // original tileM-only consumer mapping, so do not instantiate it for tileN. + static_assert(!(DIRECT_ONLY && T::IS_TILE_N), + "tileN is not supported by the DIRECT_ONLY persistent kernel"); + // PRELOAD_SF_LDS stages the scale panels for the barrier-synced + // producer/consumer schedule; it is not wired into the DIRECT_ONLY + // consumer-self-load path. + static_assert(!(PRELOAD_SF_LDS && DIRECT_ONLY), + "PRELOAD_SF_LDS is only supported by the non-DIRECT_ONLY schedule"); + + int wgid_full = opus::block_id_x(); + int split_id = 0; + int wgid = wgid_full; + if constexpr (!DIRECT_ONLY) { + split_id = wgid_full % kargs.split_k; + wgid = wgid_full / kargs.split_k; + } + const int num_tiles_m = ceil_div(kargs.m, T::B_M); + int row = (wgid % num_tiles_m) * T::B_M; + int col = (wgid / num_tiles_m) * T::B_N; + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + + const int total_iters = ceil_div(kargs.k, T::B_K); + int my_loops = total_iters; + int k_start = 0; + int sf_start = 0; + if constexpr (!DIRECT_ONLY) { + const int iters_full = ceil_div(total_iters, kargs.split_k); + my_loops = (split_id < kargs.split_k - 1) + ? iters_full + : (total_iters - (kargs.split_k - 1) * iters_full); + k_start = split_id * iters_full * T::B_K; + sf_start = split_id * iters_full * (T::B_K / T::GROUP_K); + } + if (my_loops < T::prefetch_k_iter) return; + + // OOB masking for partial M tiles: bound the A / sfa / C buffers to the + // valid row window so lanes mapping to rows >= M read 0 and their stores are + // dropped by the buffer's num_records bound. This lets any M run on a B_M + // tile without requiring M % B_M == 0 (N/K stay divisible so B and the K + // axis need no bound). + // + // Each WG owns exactly one B_M row tile and the buffer base is already at + // `row`, so the bound only needs to cover this tile's own rows -- clamp to + // B_M. Using the full (M - row) span would set num_records to + // rows_avail*stride_a, and with batch-in-the-middle stride_a = batch*K, so a + // large-M / high-batch shape would overflow the 32-bit buffer-descriptor + // num_records (4 GiB) field and silently wrap, corrupting the OOB bound. + // min(rows_avail, B_M) still masks the partial-M tail correctly. + // rows_avail >= 1 always (row < M by construction). + const int rows_left = kargs.m - row; + const int rows_avail = rows_left < T::B_M ? rows_left : T::B_M; + const unsigned int a_bytes = + (unsigned int)rows_avail * (unsigned int)kargs.stride_a * sizeof(D_A); + // 64-bit base offsets: batch_id*stride_*_batch (= M*K for a batch-in-the- + // middle A layout) overflows int32 for large M well before the 4 GiB buffer + // limit, so cast the batch/row products to size_t to keep the base exact. + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + + (size_t)batch_id * kargs.stride_a_batch + (size_t)row * kargs.stride_a + k_start, + a_bytes); + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + + (size_t)batch_id * kargs.stride_b_batch + (size_t)col * kargs.stride_b + k_start); + const bool direct_store = DIRECT_ONLY || (!std::is_void_v && kargs.split_k == 1); + const int stride_c_main = direct_store ? kargs.stride_c : kargs.stride_ws; + const unsigned int sfa_bytes = + (unsigned int)rows_avail * (unsigned int)kargs.stride_sfa * sizeof(D_SF); + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + + (size_t)batch_id * kargs.stride_sfa_batch + + (size_t)row * kargs.stride_sfa + sf_start, + sfa_bytes); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + + (size_t)batch_id * kargs.stride_sfb_batch + + (size_t)(col / T::GROUP_N) * kargs.stride_sfb + sf_start); + + int role = ((wave_id & 1) ^ ((wgid >> 8) & 1)); + + constexpr int smem_slot_factor = DIRECT_ONLY ? 2 : 1; + __shared__ char smem_a[smem_slot_factor * T::prefetch_k_iter * T::NUM_LOAD_GROUPS_PER_BM + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + __shared__ char smem_b[smem_slot_factor * T::prefetch_k_iter * T::NUM_LOAD_GROUPS_PER_BN + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + + // PRELOAD_SF_LDS (kid324): stage the A per-token scale (SFA) and B block + // scale (SFB) panels for this split's whole K range into LDS once, then read + // them from LDS (ds_read/lgkmcnt) in the consumer's per-K-tile scale fetch + // instead of a per-tile global buffer_load (vmcnt) that gates every MMA. The + // panels are compact byte tiles: SFA is [B_M/GROUP_M rows][loops*SCALES_PER_BK] + // and SFB is [N_SCALE_GROUPS rows][loops*SCALES_PER_BK], both row-major with a + // runtime per-row stride == the packed K-scale count. The LDS buffer is sized + // for a compile-time K upper bound (SFA_K_MAX); the actual packed count is a + // runtime value so any K<=SFA_K_MAX (K%B_K==0) works. SFA_K_MAX=8192 keeps the + // combined panel <=~4.2 KiB, well inside the WG_PER_CU=2 LDS headroom. + constexpr int SFA_K_MAX = 8192; + constexpr int SFA_K_TILES_MAX = PRELOAD_SF_LDS ? (SFA_K_MAX / T::B_K) : 1; + constexpr int SF_SCALES_MAX = SFA_K_TILES_MAX * T::SCALES_PER_BK; + constexpr int SFA_ROWS = T::B_M / T::GROUP_M; + constexpr int SF_LDS_ELEMS = + PRELOAD_SF_LDS ? ((SFA_ROWS + T::N_SCALE_GROUPS) * SF_SCALES_MAX) : 1; + // 16B-aligned so the panel fill below can land ds_write_b128; a byte array is + // only byte-aligned as far as the language is concerned. + __shared__ alignas(16) D_SF smem_sf[SF_LDS_ELEMS]; + + auto smem_a_at = [&](int slot_k, int m_block, int k_group) -> D_A* { + return reinterpret_cast(smem_a + + ((slot_k * T::NUM_LOAD_GROUPS_PER_BM + m_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + auto smem_b_at = [&](int slot_k, int n_block, int k_group) -> D_B* { + return reinterpret_cast(smem_b + + ((slot_k * T::NUM_LOAD_GROUPS_PER_BN + n_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + + auto a_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_M * kargs.stride_a + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + auto b_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_N * kargs.stride_b + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + const int loops = my_loops; + // Runtime packed K-scale count per SFA/SFB row (== loops * SCALES_PER_BK) and + // the two LDS panel base pointers. SFB is packed immediately after SFA using + // the runtime SFA size so both stay compact regardless of K. + const int sf_k_scales = loops * T::SCALES_PER_BK; + D_SF* s_sfa_ptr = smem_sf; + D_SF* s_sfb_ptr = smem_sf + SFA_ROWS * sf_k_scales; + constexpr int mb_a = T::a_buffer_load_insts; + constexpr int mb_b = T::b_buffer_load_insts; + constexpr int mb = mb_a + mb_b; + + if constexpr (DIRECT_ONLY) { + __shared__ int b_ready[T::prefetch_k_iter]; + if (opus::thread_id_x() < T::prefetch_k_iter) { + b_ready[opus::thread_id_x()] = -1; + } + s_waitcnt_lgkmcnt(0_I); // retire the init writes before other waves read them + __builtin_amdgcn_s_barrier(); + if ((wave_id & 1) == 0) return; + + int wave_id_m = wave_id / 2; + int wave_id_n_cons = 0; + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, 0, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, 0); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, 0, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, 0); + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c; + clear(v_c); + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + + auto issue_a_tile = [&](int loop_k) { + const int slot = wave_id_m * T::prefetch_k_iter + (loop_k % T::prefetch_k_iter); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(loop_k, m, kg)); + }); + }); + }; + + auto issue_b_tile = [&](int loop_k) { + const int slot = loop_k % T::prefetch_k_iter; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(loop_k, n, kg)); + }); + }); + }; + + auto load_scales = [&](int loop_k, vtype_sfa& v_sfa, vtype_sfb& v_sfb) { + const int scale_base = loop_k * T::SCALES_PER_BK; + v_sfa = load(g_sfa, u_sfa, scale_base); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + s_waitcnt_vmcnt(0_I); + }; + + auto do_mma = [&](const auto& va, const auto& vb, + const vtype_sfa& v_sfa, const vtype_sfb& v_sfb) { + __builtin_amdgcn_s_setprio(1); + mma_mxscale_flatmm_accum(mma, va, vb, v_sfa, v_sfb, v_c); + __builtin_amdgcn_s_setprio(0); + }; + + issue_a_tile(0); + if (wave_id_m == 0) { + issue_b_tile(0); + } + for (int k = 0; k < loops; ++k) { + const int a_slot = wave_id_m * T::prefetch_k_iter + (k % T::prefetch_k_iter); + const int b_slot = k % T::prefetch_k_iter; + s_waitcnt_vmcnt(0_I); + if (wave_id_m == 0) { + reinterpret_cast(b_ready)[b_slot] = k; + } else { + volatile int* ready = reinterpret_cast(b_ready); + while (ready[b_slot] != k) { + } + } + + auto sa = make_smem(smem_a_at(a_slot, 0, 0)); + auto sb = make_smem(smem_b_at(b_slot, 0, 0)); + v_a = load(sa, u_ra); + v_b = load(sb, u_rb); + s_waitcnt_lgkmcnt(0_I); + + vtype_sfa v_sfa; + vtype_sfb v_sfb; + load_scales(k, v_sfa, v_sfb); + if (k + 1 < loops) { + issue_a_tile(k + 1); + if (wave_id_m == 0) { + issue_b_tile(k + 1); + } + } + do_mma(v_a, v_b, v_sfa, v_sfb); + } + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + wave_id_n_cons, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr, + (unsigned int)rows_avail * (unsigned int)kargs.stride_c * sizeof(D_OUT)); + store(g_out, v_c, u_gc, 0); + return; + } + + // PRELOAD_SF_LDS: cooperative one-shot fill of the SFA + SFB panels into LDS. + // Executed by all BLOCK_SIZE threads (both producer and consumer waves) before + // the producer/consumer role split, with a barrier publishing the panels for + // the consumer reads below. Grid-stride over the compact scalar byte counts. + // OOB rows (partial-M tail) read 0 via g_sfa's num_records bound and are + // never consumed. Bail if this split's K exceeds the compile-time LDS bound + // (the dispatch never selects kid324 for such K, so this only guards misuse). + if constexpr (PRELOAD_SF_LDS) { + if (loops > SFA_K_TILES_MAX) return; + const int tid = opus::thread_id_x(); + auto sm_sfa = make_smem(s_sfa_ptr); + auto sm_sfb = make_smem(s_sfb_ptr); + const int sfa_total = SFA_ROWS * sf_k_scales; + const int sfb_total = T::N_SCALE_GROUPS * sf_k_scales; + // Copy the widest chunk the panel geometry allows. Byte-at-a-time is 16 + // grid-stride iterations per thread for the 128-row SFA panel at K=4096 + // (kid325/326), and the fill sits in front of a barrier, so its latency is + // exposed rather than overlapped. A chunk must not span two panel rows and + // its source offset must stay naturally aligned, so the width has to divide + // both sf_k_scales and the row stride; hence the short-K fallbacks. + auto fill = [&](auto vec_c, auto sm, auto g, int stride, int total) { + constexpr int VEC = decltype(vec_c)::value; + for (int idx = tid * VEC; idx < total; idx += T::BLOCK_SIZE * VEC) { + const int r = idx / sf_k_scales; + const int kt = idx - r * sf_k_scales; + sm.template store(load(g, r * stride + kt), idx); + } + }; + auto fill_panel = [&](auto sm, auto g, int stride, int total) { + const int widths = sf_k_scales | stride; + if ((widths & 15) == 0) fill(number<16>{}, sm, g, stride, total); + else if ((widths & 3) == 0) fill(number<4>{}, sm, g, stride, total); + else fill(number<1>{}, sm, g, stride, total); + }; + fill_panel(sm_sfa, g_sfa, kargs.stride_sfa, sfa_total); + fill_panel(sm_sfb, g_sfb, kargs.stride_sfb, sfb_total); + // vmcnt retires the global reads feeding the panel; lgkmcnt retires the + // ds_writes that actually publish it. s_barrier does neither on its own. + s_waitcnt_vmcnt(0_I); + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } + + if (role == 0) { + int wave_id_prod = wave_id / 2; + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + + opus::static_for([&](auto p_c) { + constexpr int p = decltype(p_c)::value; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(p, m, kg), u_ga, u_sa, a_offset(p, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(p, n, kg), u_gb, u_sb, b_offset(p, n, kg)); + }); + }); + }); + + opus::static_for([&](auto i_c) { + constexpr int p = T::prefetch_k_iter - 1 - decltype(i_c)::value; + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + }); + + if constexpr (T::prefetch_k_iter == 3) { + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + for (int i = T::prefetch_k_iter - 1; i < loops - 1; i++) { + int issue_k = i + 1; + int slot = issue_k % T::prefetch_k_iter; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(issue_k, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(issue_k, n, kg)); + }); + }); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + } + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } else { + for (int i = T::prefetch_k_iter - 2; i < loops - 2; i++) { + int issue_k = i + 2; + int slot = issue_k % T::prefetch_k_iter; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(issue_k, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(issue_k, n, kg)); + }); + }); + s_waitcnt_vmcnt(number<2 * mb>{}); + __builtin_amdgcn_s_barrier(); + } + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } + } else { + // Consumer waves. tileM: two waves split M (wave_id_m in {0,1}, single + // N column-block). tileN: two waves split N (single M-wave, wave_id_n + // in {0,1}); each consumer reads its own 16-col B group from smem, and + // the C partition / rb layout follow via T_N=2. wave_id_n_cons is 0 for + // tileM, so the shared `n_block = wave_id_n_cons` smem-B base below is + // bit-identical for the existing tileM kids. + int wave_id_m = T::IS_TILE_N ? 0 : (wave_id / 2); + int wave_id_n_cons = T::IS_TILE_N ? (wave_id / 2) : 0; + // Consumer B smem N-group base. Each consumer N-wave owns COM_REP_N + // contiguous 16-col load groups (rb reads num_blocks_n=COM_REP_N from + // this base). tileM: wave_id_n_cons=0 -> nbc=0 (bit-identical). + const int nbc = wave_id_n_cons * T::COM_REP_N; + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + // LDS read layout for the preloaded SFA panel: same lane/wave mapping as + // u_sfa but with the compact per-row K-scale count as the row stride. + auto u_sfa_lds = make_layout_sfa_mxsk(lane_id, wave_id_m, sf_k_scales); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a0, v_a1; + typename decltype(mma)::vtype_b v_b0, v_b1; + typename decltype(mma)::vtype_c v_c; + clear(v_c); + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + constexpr int ds_read_insts = T::a_ds_read_insts + T::b_ds_read_insts; + + auto load_scale_regs = [&](int loop_k, vtype_sfa& v_sfa, vtype_sfb& v_sfb) { + const int scale_base = loop_k * T::SCALES_PER_BK; + if constexpr (PRELOAD_SF_LDS) { + // Read this K-tile's scales from the preloaded LDS panels + // (ds_read / lgkmcnt) instead of a per-tile global buffer_load. + // Vec = SCALES_PER_BK so the contiguous per-M-row K bytes come in + // as one dword (ds_read_b32) instead of SCALES_PER_BK byte reads. + auto sm_a = make_smem(s_sfa_ptr + scale_base); + v_sfa = load(sm_a, u_sfa_lds); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sm_b = make_smem(s_sfb_ptr + ng * sf_k_scales + scale_base); + auto sfb = load(sm_b, 0); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + } else { + // Vec = SCALES_PER_BK: the contiguous per-M-row K-scale bytes are + // read as one dword (buffer_load_b32) rather than SCALES_PER_BK + // separate buffer_load_ubyte. SFB already loads b32 the same way. + v_sfa = load(g_sfa, u_sfa, scale_base); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + } + }; + + auto do_scaled_mma = [&](const auto& va, const auto& vb, + const vtype_sfa& v_sfa, const vtype_sfb& v_sfb) { + if constexpr (PRELOAD_SF_LDS) { + // Scales come from LDS now, so wait on lgkmcnt rather than vmcnt. + // This also drains the just-issued next-buffer A/B ds_reads, still + // a win against the global scale-load stall it replaces. + s_waitcnt_lgkmcnt(0_I); + } else { + s_waitcnt_vmcnt(0_I); + } + __builtin_amdgcn_s_setprio(1); + mma_mxscale_flatmm_accum(mma, va, vb, v_sfa, v_sfb, v_c); + __builtin_amdgcn_s_setprio(0); + }; + + auto scaled_mma = [&](const auto& va, const auto& vb, int loop_k) { + vtype_sfa v_sfa; + vtype_sfb v_sfb; + load_scale_regs(loop_k, v_sfa, v_sfb); + do_scaled_mma(va, vb, v_sfa, v_sfb); + }; + + auto wait_lgkm_then_scaled_mma = + [&](const auto& va, const auto& vb, int loop_k, auto lgkm_cnt) { + if constexpr (PREFETCH_SCALE) { + vtype_sfa v_sfa; + vtype_sfb v_sfb; + load_scale_regs(loop_k, v_sfa, v_sfb); + s_waitcnt_lgkmcnt(lgkm_cnt); + do_scaled_mma(va, vb, v_sfa, v_sfb); + } else { + s_waitcnt_lgkmcnt(lgkm_cnt); + scaled_mma(va, vb, loop_k); + } + }; + + __builtin_amdgcn_s_barrier(); + { + auto sa0 = make_smem(smem_a_at(0, 0, 0)); + auto sb0 = make_smem(smem_b_at(0, nbc, 0)); + v_a0 = load(sa0, u_ra); + v_b0 = load(sb0, u_rb); + } + + opus::static_for([&](auto i_c) { + constexpr int p = decltype(i_c)::value + 1; + constexpr int cur = (p - 1) & 1; + constexpr int nxt = p & 1; + __builtin_amdgcn_s_barrier(); + auto sa_p = make_smem(smem_a_at(p, 0, 0)); + auto sb_p = make_smem(smem_b_at(p, nbc, 0)); + if constexpr (nxt == 0) { + v_a0 = load(sa_p, u_ra); + v_b0 = load(sb_p, u_rb); + } else { + v_a1 = load(sa_p, u_ra); + v_b1 = load(sb_p, u_rb); + } + if constexpr (cur == 0) { + wait_lgkm_then_scaled_mma(v_a0, v_b0, p - 1, number{}); + } else { + wait_lgkm_then_scaled_mma(v_a1, v_b1, p - 1, number{}); + } + }); + + constexpr int L = (T::prefetch_k_iter - 2) & 1; + int k = T::prefetch_k_iter - 1; + for (; k + 1 < loops - 1; k += 2) { + __builtin_amdgcn_s_barrier(); + { + int slot = k % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, nbc, 0)); + if constexpr (L == 0) { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } else { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } + } + if constexpr (L == 0) { + wait_lgkm_then_scaled_mma(v_a0, v_b0, k - 1, number{}); + } else { + wait_lgkm_then_scaled_mma(v_a1, v_b1, k - 1, number{}); + } + + __builtin_amdgcn_s_barrier(); + { + int slot = (k + 1) % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, nbc, 0)); + if constexpr (L == 0) { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } else { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } + } + if constexpr (L == 0) { + wait_lgkm_then_scaled_mma(v_a1, v_b1, k, number{}); + } else { + wait_lgkm_then_scaled_mma(v_a0, v_b0, k, number{}); + } + } + + bool last_in_buf1 = (L != 0); + if (k < loops - 1) { + __builtin_amdgcn_s_barrier(); + { + int slot = k % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, nbc, 0)); + if constexpr (L == 0) { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } else { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } + } + if constexpr (L == 0) { + wait_lgkm_then_scaled_mma(v_a0, v_b0, k - 1, number{}); + } else { + wait_lgkm_then_scaled_mma(v_a1, v_b1, k - 1, number{}); + } + last_in_buf1 = (L == 0); + k++; + } + + __builtin_amdgcn_s_barrier(); + int last_slot = (loops - 1) % T::prefetch_k_iter; + auto sa_last = make_smem(smem_a_at(last_slot, 0, 0)); + auto sb_last = make_smem(smem_b_at(last_slot, nbc, 0)); + if (last_in_buf1) { + v_a0 = load(sa_last, u_ra); + v_b0 = load(sb_last, u_rb); + wait_lgkm_then_scaled_mma(v_a1, v_b1, loops - 2, number{}); + wait_lgkm_then_scaled_mma(v_a0, v_b0, loops - 1, 0_I); + } else { + v_a1 = load(sa_last, u_ra); + v_b1 = load(sb_last, u_rb); + wait_lgkm_then_scaled_mma(v_a0, v_b0, loops - 2, number{}); + wait_lgkm_then_scaled_mma(v_a1, v_b1, loops - 1, 0_I); + } + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + wave_id_n_cons, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(stride_c_main, 1_I), p_coord_c); + // tileN with COM_REP_N>1: two consumer waves split B_N and each replicates + // COM_REP_N N-repeats. The generic swap_ab C partition nests the register + // N-repeat (expd_n) OUTSIDE the consumer-wave tile (tile_n), which transposes + // the (wave, n-rep) -> column map (each wave writes a strided, interleaved + // column set) whenever BOTH T_N>1 and COM_REP_N>1 -- the accumulators are + // correct but land in swapped output columns. Consumer wave w computes n-rep + // j from B column-group (w*COM_REP_N + j) (see nbc = wave_id_n_cons*COM_REP_N + // + the num_blocks_n rb read), so store each n-rep slice to that contiguous + // column group with a single-N-tile (E_N=1,T_N=1) layout and a scalar column + // offset. tileM (T_N=1) and tileN COM_REP_N==1 keep the original single store + // (SPLIT_N_STORE=false), so they stay bit-identical. + constexpr bool SPLIT_N_STORE = T::IS_TILE_N && (T::COM_REP_N > 1); + constexpr int C_LEN = decltype(mma)::mma_c_len; + auto mma_c1 = make_tiled_mma( + seq{}, seq{}, + seq{}, mfma_adaptor_swap_ab{}); + auto p_coord_c1 = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + 0, lane_id / mma.grpn_c); + auto u_gc1 = partition_layout_c(mma_c1, + opus::make_tuple(stride_c_main, 1_I), p_coord_c1); + auto store_c = [&](auto& g) { + if constexpr (SPLIT_N_STORE) { + opus::static_for([&](auto j_c) { + constexpr int j = decltype(j_c)::value; + auto vj = opus::slice(v_c, opus::number{}, + opus::number{}); + store(g, vj, u_gc1, + (wave_id_n_cons * T::COM_REP_N + j) * T::W_N); + }); + } else { + store(g, v_c, u_gc, 0); + } + }; + if constexpr (!std::is_void_v) { + if (kargs.split_k == 1) { + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr, + (unsigned int)rows_avail * (unsigned int)kargs.stride_c * sizeof(D_OUT)); + store_c(g_out); + } else { + D_C* ws_c_ptr = reinterpret_cast(kargs.ws_handle->ptr) + + (size_t)split_id * kargs.batch * kargs.stride_ws_batch + + (size_t)batch_id * kargs.stride_ws_batch + + (size_t)row * kargs.stride_ws + + (size_t)col; + auto g_c = make_gmem(ws_c_ptr); + store_c(g_c); + } + } else { + D_C* ws_c_ptr = reinterpret_cast(kargs.ws_handle->ptr) + + (size_t)split_id * kargs.batch * kargs.stride_ws_batch + + (size_t)batch_id * kargs.stride_ws_batch + + (size_t)row * kargs.stride_ws + + (size_t)col; + auto g_c = make_gmem(ws_c_ptr); + store_c(g_c); + } + } + + if constexpr (!std::is_void_v) { + if (kargs.split_k == 1) return; + + __shared__ int fused_do_reduce; + if (opus::thread_id_x() == 0) { + fused_do_reduce = 0; + } + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_barrier(); + __builtin_amdgcn_fence(__ATOMIC_RELEASE, "agent"); + __builtin_amdgcn_s_barrier(); + + int* counters = reinterpret_cast( + reinterpret_cast(kargs.ws_handle->ptr) + kargs.counter_offset_bytes); + const int num_tiles = num_tiles_m * ceil_div(kargs.n, T::B_N); + const int tile_id = batch_id * num_tiles + wgid; + if (opus::thread_id_x() == 0) { + const int old = __atomic_fetch_add(counters + tile_id, 1, __ATOMIC_ACQ_REL); + fused_do_reduce = (old == kargs.split_k - 1); + } + // Every thread branches on this below, so lane 0's write has to be retired + // (not merely issued) before the barrier lets the rest read it. + s_waitcnt_lgkmcnt(0_I); + __builtin_amdgcn_s_barrier(); + + if (fused_do_reduce) { + const D_C* ws_base = reinterpret_cast(kargs.ws_handle->ptr); + D_OUT* out = reinterpret_cast(kargs.ptr_c); + const size_t split_stride = (size_t)kargs.batch * (size_t)kargs.stride_ws_batch; + for (int i = int(opus::thread_id_x()); i < T::B_M * T::B_N; i += T::BLOCK_SIZE) { + const int mi = i / T::B_N; + const int ni = i - mi * T::B_N; + if (row + mi >= kargs.m) continue; // skip OOB rows of a partial M tile + float acc = 0.0f; + const size_t base = (size_t)batch_id * (size_t)kargs.stride_ws_batch + + (size_t)(row + mi) * (size_t)kargs.stride_ws + + (size_t)(col + ni); + for (int s = 0; s < kargs.split_k; ++s) { + acc += static_cast(ws_base[(size_t)s * split_stride + base]); + } + const size_t out_idx = (size_t)batch_id * (size_t)kargs.stride_c_batch + + (size_t)(row + mi) * (size_t)kargs.stride_c + + (size_t)(col + ni); + out[out_idx] = static_cast(acc); + } + __builtin_amdgcn_s_barrier(); + if (opus::thread_id_x() == 0) { + counters[tile_id] = 0; + } + } + } +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// Direct-store persistent M-outer kernel. Each WG owns one N tile and a small +// run of M tiles, reusing the same B tile stream across the outer loop without +// increasing the per-tile accumulator footprint. +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, Traits::WG_PER_CU) +void gemm_a8w8_mxscale_flatmm_splitk_mouter_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + + const int num_tiles_m = ceil_div(kargs.m, T::B_M); + const int num_tiles_n = ceil_div(kargs.n, T::B_N); + const int m_per_wg = kargs.split_k; + int bid = opus::block_id_x(); + constexpr int NUM_XCD = 8; + int xcd_id = __builtin_amdgcn_readfirstlane(bid % NUM_XCD); + int pos_xcd = __builtin_amdgcn_readfirstlane(bid / NUM_XCD); + int tile_n_id = __builtin_amdgcn_readfirstlane(pos_xcd % num_tiles_n); + int m_grp_local = __builtin_amdgcn_readfirstlane(pos_xcd / num_tiles_n); + int m_grp = __builtin_amdgcn_readfirstlane(xcd_id * kargs.stride_ws_batch + m_grp_local); + if (m_grp >= kargs.stride_ws) return; + int tile_m_lo = m_grp * m_per_wg; + int tile_m_hi = tile_m_lo + m_per_wg; + int col = tile_n_id * T::B_N; + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + int role = ((wave_id & 1) ^ ((bid >> 8) & 1)); + + const int loops = kargs.k / T::B_K; + if (loops < T::prefetch_k_iter) return; + + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + + (size_t)batch_id * kargs.stride_b_batch + (size_t)col * kargs.stride_b); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + + (size_t)batch_id * kargs.stride_sfb_batch + + (size_t)(col / T::GROUP_N) * kargs.stride_sfb); + + __shared__ char smem_a[T::prefetch_k_iter * T::NUM_LOAD_GROUPS_PER_BM + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + __shared__ char smem_b[T::prefetch_k_iter * T::NUM_LOAD_GROUPS_PER_BN + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + + auto smem_a_at = [&](int slot_k, int m_block, int k_group) -> D_A* { + return reinterpret_cast(smem_a + + ((slot_k * T::NUM_LOAD_GROUPS_PER_BM + m_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + auto smem_b_at = [&](int slot_k, int n_block, int k_group) -> D_B* { + return reinterpret_cast(smem_b + + ((slot_k * T::NUM_LOAD_GROUPS_PER_BN + n_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + + auto b_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_N * kargs.stride_b + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + constexpr int mb_a = T::a_buffer_load_insts; + constexpr int mb_b = T::b_buffer_load_insts; + constexpr int mb = mb_a + mb_b; + + if (role == 0) { + int wave_id_prod = wave_id / 2; + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + + for (int tile_m = tile_m_lo; tile_m < tile_m_hi && tile_m < num_tiles_m; ++tile_m) { + int row = tile_m * T::B_M; + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + + (size_t)batch_id * kargs.stride_a_batch + + (size_t)row * kargs.stride_a); + auto a_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_M * kargs.stride_a + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + opus::static_for([&](auto p_c) { + constexpr int p = decltype(p_c)::value; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(p, m, kg), u_ga, u_sa, a_offset(p, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(p, n, kg), u_gb, u_sb, b_offset(p, n, kg)); + }); + }); + }); + + opus::static_for([&](auto i_c) { + constexpr int p = T::prefetch_k_iter - 1 - decltype(i_c)::value; + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + }); + + if constexpr (T::prefetch_k_iter == 3) { + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + for (int i = T::prefetch_k_iter - 1; i < loops - 1; i++) { + int issue_k = i + 1; + int slot = issue_k % T::prefetch_k_iter; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(issue_k, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(issue_k, n, kg)); + }); + }); + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + } + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } else { + for (int i = T::prefetch_k_iter - 2; i < loops - 2; i++) { + int issue_k = i + 2; + int slot = issue_k % T::prefetch_k_iter; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(issue_k, m, kg)); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(issue_k, n, kg)); + }); + }); + s_waitcnt_vmcnt(number<2 * mb>{}); + __builtin_amdgcn_s_barrier(); + } + s_waitcnt_vmcnt(number{}); + __builtin_amdgcn_s_barrier(); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } + __builtin_amdgcn_s_barrier(); + } + } else { + int wave_id_m = wave_id / 2; + int wave_id_n_cons = 0; + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a0, v_a1; + typename decltype(mma)::vtype_b v_b0, v_b1; + typename decltype(mma)::vtype_c v_c; + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + constexpr int ds_read_insts = T::a_ds_read_insts + T::b_ds_read_insts; + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + wave_id_n_cons, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + + for (int tile_m = tile_m_lo; tile_m < tile_m_hi && tile_m < num_tiles_m; ++tile_m) { + int row = tile_m * T::B_M; + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + + (size_t)batch_id * kargs.stride_sfa_batch + + (size_t)row * kargs.stride_sfa); + clear(v_c); + + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + auto scaled_mma = [&](const auto& va, const auto& vb, int loop_k) { + const int scale_base = loop_k * T::SCALES_PER_BK; + vtype_sfa v_sfa = load(g_sfa, u_sfa, scale_base); + vtype_sfb v_sfb; + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + if constexpr (!SKIP_SCALE_WAIT) { + s_waitcnt_vmcnt(0_I); + } + __builtin_amdgcn_s_setprio(1); + mma_mxscale_flatmm_accum(mma, va, vb, v_sfa, v_sfb, v_c); + __builtin_amdgcn_s_setprio(0); + }; + + __builtin_amdgcn_s_barrier(); + { + auto sa0 = make_smem(smem_a_at(0, 0, 0)); + auto sb0 = make_smem(smem_b_at(0, 0, 0)); + v_a0 = load(sa0, u_ra); + v_b0 = load(sb0, u_rb); + } + + opus::static_for([&](auto i_c) { + constexpr int p = decltype(i_c)::value + 1; + constexpr int cur = (p - 1) & 1; + constexpr int nxt = p & 1; + __builtin_amdgcn_s_barrier(); + auto sa_p = make_smem(smem_a_at(p, 0, 0)); + auto sb_p = make_smem(smem_b_at(p, 0, 0)); + if constexpr (nxt == 0) { + v_a0 = load(sa_p, u_ra); + v_b0 = load(sb_p, u_rb); + } else { + v_a1 = load(sa_p, u_ra); + v_b1 = load(sb_p, u_rb); + } + s_waitcnt_lgkmcnt(number{}); + if constexpr (cur == 0) scaled_mma(v_a0, v_b0, p - 1); + else scaled_mma(v_a1, v_b1, p - 1); + }); + + constexpr int L = (T::prefetch_k_iter - 2) & 1; + int k = T::prefetch_k_iter - 1; + for (; k + 1 < loops - 1; k += 2) { + __builtin_amdgcn_s_barrier(); + { + int slot = k % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, 0, 0)); + if constexpr (L == 0) { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } else { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } + } + s_waitcnt_lgkmcnt(number{}); + if constexpr (L == 0) scaled_mma(v_a0, v_b0, k - 1); + else scaled_mma(v_a1, v_b1, k - 1); + + __builtin_amdgcn_s_barrier(); + { + int slot = (k + 1) % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, 0, 0)); + if constexpr (L == 0) { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } else { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } + } + s_waitcnt_lgkmcnt(number{}); + if constexpr (L == 0) scaled_mma(v_a1, v_b1, k); + else scaled_mma(v_a0, v_b0, k); + } + + bool last_in_buf1 = (L != 0); + if (k < loops - 1) { + __builtin_amdgcn_s_barrier(); + { + int slot = k % T::prefetch_k_iter; + auto sa_k = make_smem(smem_a_at(slot, 0, 0)); + auto sb_k = make_smem(smem_b_at(slot, 0, 0)); + if constexpr (L == 0) { + v_a1 = load(sa_k, u_ra); + v_b1 = load(sb_k, u_rb); + } else { + v_a0 = load(sa_k, u_ra); + v_b0 = load(sb_k, u_rb); + } + } + s_waitcnt_lgkmcnt(number{}); + if constexpr (L == 0) scaled_mma(v_a0, v_b0, k - 1); + else scaled_mma(v_a1, v_b1, k - 1); + last_in_buf1 = (L == 0); + k++; + } + + __builtin_amdgcn_s_barrier(); + int last_slot = (loops - 1) % T::prefetch_k_iter; + auto sa_last = make_smem(smem_a_at(last_slot, 0, 0)); + auto sb_last = make_smem(smem_b_at(last_slot, 0, 0)); + if (last_in_buf1) { + v_a0 = load(sa_last, u_ra); + v_b0 = load(sb_last, u_rb); + s_waitcnt_lgkmcnt(number{}); + scaled_mma(v_a1, v_b1, loops - 2); + s_waitcnt_lgkmcnt(0_I); + scaled_mma(v_a0, v_b0, loops - 1); + } else { + v_a1 = load(sa_last, u_ra); + v_b1 = load(sb_last, u_rb); + s_waitcnt_lgkmcnt(number{}); + scaled_mma(v_a0, v_b0, loops - 2); + s_waitcnt_lgkmcnt(0_I); + scaled_mma(v_a1, v_b1, loops - 1); + } + + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr); + store(g_out, v_c, u_gc, 0); + __builtin_amdgcn_s_barrier(); + } + } +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// -- M-tile interleaved direct-store kernel (correctness-first v1) ------------ +// +// Each WG owns one N tile and MI=2 consecutive M tiles that share the SAME B +// tile stream. Both M tiles' A operands are resident in LDS simultaneously +// (smem_a0 / smem_a1) while B is loaded once (smem_b). Per K iteration the +// consumer issues MFMA for tile0 then tile1 back-to-back into two independent +// accumulators, so the MFMA instruction stream is ~MIx longer across a single +// prologue/epilogue, and B global traffic is halved. +// +// v1 is single-buffered (load-then-compute, barrier-bracketed) for provable +// deadlock-freedom and correctness; the software-pipelined perf version is a +// follow-up. Producer waves (role 0) load; consumer waves (role 1) MFMA. +template +__global__ __launch_bounds__(Traits::BLOCK_SIZE, Traits::WG_PER_CU) +void gemm_a8w8_mxscale_flatmm_minterleave_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + + constexpr int MI = 2; // M tiles interleaved per WG + constexpr int SB = 2; // double LDS buffer per operand (prefetch k+1) + + const int num_tiles_m = ceil_div(kargs.m, T::B_M); + const int num_tiles_n = ceil_div(kargs.n, T::B_N); + int bid = opus::block_id_x(); + constexpr int NUM_XCD = 8; + int xcd_id = __builtin_amdgcn_readfirstlane(bid % NUM_XCD); + int pos_xcd = __builtin_amdgcn_readfirstlane(bid / NUM_XCD); + int tile_n_id = __builtin_amdgcn_readfirstlane(pos_xcd % num_tiles_n); + int m_grp_local = __builtin_amdgcn_readfirstlane(pos_xcd / num_tiles_n); + int m_grp = __builtin_amdgcn_readfirstlane(xcd_id * kargs.stride_ws_batch + m_grp_local); + if (m_grp >= kargs.stride_ws) return; + int tile_m0 = m_grp * MI; + if (tile_m0 + MI > num_tiles_m) return; // host guarantees M % (MI*B_M) == 0 + int col = tile_n_id * T::B_N; + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + int role = (wave_id & 1); + + const int loops = kargs.k / T::B_K; + if (loops < T::prefetch_k_iter) return; + + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + + (size_t)batch_id * kargs.stride_b_batch + (size_t)col * kargs.stride_b); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + + (size_t)batch_id * kargs.stride_sfb_batch + + (size_t)(col / T::GROUP_N) * kargs.stride_sfb); + + __shared__ char smem_a[MI * SB * T::NUM_LOAD_GROUPS_PER_BM + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + __shared__ char smem_b[SB * T::NUM_LOAD_GROUPS_PER_BN + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + + auto smem_a_at = [&](int mi, int slot_k, int m_block, int k_group) -> D_A* { + return reinterpret_cast(smem_a + + (((mi * SB + slot_k) * T::NUM_LOAD_GROUPS_PER_BM + m_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + auto smem_b_at = [&](int slot_k, int n_block, int k_group) -> D_B* { + return reinterpret_cast(smem_b + + ((slot_k * T::NUM_LOAD_GROUPS_PER_BN + n_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + + auto b_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_N * kargs.stride_b + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + if (role == 0) { + int wave_id_prod = wave_id / 2; + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, wave_id_prod, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, wave_id_prod); + + auto a_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_M * kargs.stride_a + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + const D_A* base_a = reinterpret_cast(kargs.ptr_a) + + (size_t)batch_id * kargs.stride_a_batch; + + auto issue_loads = [&](int k, int slot) { + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto mi_c) { + constexpr int mi = decltype(mi_c)::value; + auto g_a = make_gmem(base_a + (size_t)(tile_m0 + mi) * T::B_M * kargs.stride_a); + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(mi, slot, m, kg), u_ga, u_sa, a_offset(k, m, kg)); + }); + }); + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(k, n, kg)); + }); + }); + }; + + // Prologue: preload slot 0 (K-tile 0), then stream while consumers MFMA. + issue_loads(0, 0); + s_waitcnt_vmcnt(0_I); + for (int k = 0; k < loops; ++k) { + __builtin_amdgcn_s_barrier(); // R(k): slot k ready + if (k + 1 < loops) { + issue_loads(k + 1, (k + 1) % SB); + s_waitcnt_vmcnt(0_I); // slot k+1 fully loaded before releasing slot k + } + __builtin_amdgcn_s_barrier(); // F(k): consumer done reading slot k + } + } else { + int wave_id_m = wave_id / 2; + int wave_id_n_cons = 0; + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a[MI]; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c[MI]; + + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + wave_id_n_cons, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + + // Per-tile A-scale gmem + layout. + auto g_sfa = [&](int mi) { + return make_gmem(reinterpret_cast(kargs.ptr_sfa) + + (size_t)batch_id * kargs.stride_sfa_batch + + (size_t)(tile_m0 + mi) * T::B_M * kargs.stride_sfa); + }; + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + + opus::static_for([&](auto mi_c) { clear(v_c[decltype(mi_c)::value]); }); + + for (int k = 0; k < loops; ++k) { + __builtin_amdgcn_s_barrier(); // R(k): wait producer data ready + const int slot = k % SB; + auto sb0 = make_smem(smem_b_at(slot, 0, 0)); + v_b = load(sb0, u_rb); + opus::static_for([&](auto mi_c) { + constexpr int mi = decltype(mi_c)::value; + auto sa = make_smem(smem_a_at(mi, slot, 0, 0)); + v_a[mi] = load(sa, u_ra); + }); + s_waitcnt_lgkmcnt(0_I); + + const int scale_base = k * T::SCALES_PER_BK; + vtype_sfb v_sfb; + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + opus::static_for([&](auto mi_c) { + constexpr int mi = decltype(mi_c)::value; + auto gsfa = g_sfa(mi); + vtype_sfa v_sfa = load(gsfa, u_sfa, scale_base); + if constexpr (!SKIP_SCALE_WAIT) s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_mxscale_flatmm_accum(mma, v_a[mi], v_b, v_sfa, v_sfb, v_c[mi]); + __builtin_amdgcn_s_setprio(0); + }); + __builtin_amdgcn_s_barrier(); // (C2) done reading slot + } + + opus::static_for([&](auto mi_c) { + constexpr int mi = decltype(mi_c)::value; + int row = (tile_m0 + mi) * T::B_M; + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr); + store(g_out, v_c[mi], u_gc, 0); + }); + } +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// 8-wave split-accumulator direct-store kernel. +// +// Logical tile: 128x256x128. Internally this is two independent 128x128 +// accumulator groups computed by consumer waves {4,5} and {6,7}. Producer +// waves {0,1} load shared A and B phase 0, producer waves {2,3} load B phase 1. +// This keeps each consumer's v_c identical to the proven 128x128 WG1 kernel +// while halving the logical N workgroup count. +template +__global__ __launch_bounds__(512, 1) +void gemm_a8w8_mxscale_flatmm_splitk_wave8n2_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + static_assert(T::B_M == 128 && T::B_N == 128 && T::B_K == 128, + "wave8n2 builds a logical 128x256 tile from 128x128 traits"); + + constexpr int N_PHASES = 2; + int wgid = opus::block_id_x(); + const int num_tiles_m = ceil_div(kargs.m, T::B_M); + int row = (wgid % num_tiles_m) * T::B_M; + int col_base = (wgid / num_tiles_m) * (T::B_N * N_PHASES); + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + const int loops = kargs.k / T::B_K; + if (loops < 1) return; + + constexpr int WAVE8N2_PREFETCH_SLOTS = 2; + __shared__ char smem_a[WAVE8N2_PREFETCH_SLOTS * T::NUM_LOAD_GROUPS_PER_BM + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + __shared__ char smem_b[WAVE8N2_PREFETCH_SLOTS * N_PHASES * T::NUM_LOAD_GROUPS_PER_BN + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + + auto smem_a_at = [&](int slot, int m_block, int k_group) -> D_A* { + return reinterpret_cast(smem_a + + ((slot * T::NUM_LOAD_GROUPS_PER_BM + m_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + auto smem_b_at = [&](int slot, int phase, int n_block, int k_group) -> D_B* { + return reinterpret_cast(smem_b + + (((slot * N_PHASES + phase) * T::NUM_LOAD_GROUPS_PER_BN + n_block) + * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + + if (wave_id < 4) { + const int phase = wave_id / 2; + const int prod_wave = wave_id & 1; + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + + (size_t)batch_id * kargs.stride_a_batch + + (size_t)row * kargs.stride_a); + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + + (size_t)batch_id * kargs.stride_b_batch + + (size_t)(col_base + phase * T::B_N) * kargs.stride_b); + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, prod_wave, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, prod_wave); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, prod_wave, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, prod_wave); + + auto a_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_M * kargs.stride_a + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + auto b_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_N * kargs.stride_b + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + auto issue_tile = [&](int loop_k) { + const int slot = loop_k & 1; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + if (phase == 0) { + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m, kg), u_ga, u_sa, a_offset(loop_k, m, kg)); + }); + } + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, phase, n, kg), u_gb, u_sb, b_offset(loop_k, n, kg)); + }); + }); + }; + + issue_tile(0); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + + for (int k = 0; k < loops; ++k) { + if (k + 1 < loops) { + issue_tile(k + 1); + } + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + } + return; + } + + const int consumer = wave_id - 4; + const int phase = consumer / 2; + const int wave_id_m = consumer & 1; + const int col = col_base + phase * T::B_N; + + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + + (size_t)batch_id * kargs.stride_sfa_batch + + (size_t)row * kargs.stride_sfa); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + + (size_t)batch_id * kargs.stride_sfb_batch + + (size_t)(col / T::GROUP_N) * kargs.stride_sfb); + + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c; + clear(v_c); + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + + __builtin_amdgcn_s_barrier(); + { + auto sa = make_smem(smem_a_at(0, 0, 0)); + auto sb = make_smem(smem_b_at(0, phase, 0, 0)); + v_a = load(sa, u_ra); + v_b = load(sb, u_rb); + s_waitcnt_lgkmcnt(0_I); + } + + for (int k = 0; k < loops; ++k) { + const int scale_base = k * T::SCALES_PER_BK; + vtype_sfa v_sfa = load(g_sfa, u_sfa, scale_base); + vtype_sfb v_sfb; + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_setprio(1); + mma_mxscale_flatmm_accum(mma, v_a, v_b, v_sfa, v_sfb, v_c); + __builtin_amdgcn_s_setprio(0); + __builtin_amdgcn_s_barrier(); + if (k + 1 < loops) { + const int slot = (k + 1) & 1; + auto sa = make_smem(smem_a_at(slot, 0, 0)); + auto sb = make_smem(smem_b_at(slot, phase, 0, 0)); + v_a = load(sa, u_ra); + v_b = load(sb, u_rb); + s_waitcnt_lgkmcnt(0_I); + } + } + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + 0, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr); + store(g_out, v_c, u_gc, 0); +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} + +// 4-wave self-load split-accumulator direct-store kernel with M reuse. +// +// Logical tile: 256x128x128. Two independent 128x128 accumulator groups cover +// adjacent M tiles and share one B tile. This targets large-M shapes where B +// reuse matters more than reducing N workgroups. +template +__global__ __launch_bounds__(256, 1) +void gemm_a8w8_mxscale_flatmm_splitk_wave4m2_selfload_kernel(opus_gemm_scale_splitk_kargs_gfx950 kargs) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + + using T = opus::remove_cvref_t; + using D_A = typename T::D_A; + using D_B = typename T::D_B; + using D_ACC = typename T::D_ACC; + using D_SF = typename T::D_SF; + static_assert(T::B_M == 128 && T::B_N == 128 && T::B_K == 128, + "wave4m2 selfload builds a logical 256x128 tile from 128x128 traits"); + + constexpr int M_PHASES = 2; + constexpr int PREFETCH_SLOTS = 2; + int wgid = opus::block_id_x(); + const int logical_b_m = T::B_M * M_PHASES; + const int num_tiles_m = ceil_div(kargs.m, logical_b_m); + int row_base = (wgid % num_tiles_m) * logical_b_m; + int col = (wgid / num_tiles_m) * T::B_N; + int batch_id = opus::block_id_z(); + int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); + int lane_id = opus::thread_id_x() % get_warp_size(); + const int loops = kargs.k / T::B_K; + if (loops < 1) return; + + const int m_phase = wave_id / 2; + const int wave_id_m = wave_id & 1; + const int row = row_base + m_phase * T::B_M; + + auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + + (size_t)batch_id * kargs.stride_a_batch + + (size_t)row * kargs.stride_a); + auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + + (size_t)batch_id * kargs.stride_b_batch + + (size_t)col * kargs.stride_b); + auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + + (size_t)batch_id * kargs.stride_sfa_batch + + (size_t)row * kargs.stride_sfa); + auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + + (size_t)batch_id * kargs.stride_sfb_batch + + (size_t)(col / T::GROUP_N) * kargs.stride_sfb); + + __shared__ char smem_a[PREFETCH_SLOTS * M_PHASES * T::NUM_LOAD_GROUPS_PER_BM + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + __shared__ char smem_b[PREFETCH_SLOTS * T::NUM_LOAD_GROUPS_PER_BN + * T::NUM_LOAD_GROUPS_PER_BK * T::smem_per_group_load_size]; + + auto smem_a_at = [&](int slot, int phase, int m_block, int k_group) -> D_A* { + return reinterpret_cast(smem_a + + (((slot * M_PHASES + phase) * T::NUM_LOAD_GROUPS_PER_BM + m_block) + * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + auto smem_b_at = [&](int slot, int n_block, int k_group) -> D_B* { + return reinterpret_cast(smem_b + + ((slot * T::NUM_LOAD_GROUPS_PER_BN + n_block) * T::NUM_LOAD_GROUPS_PER_BK + k_group) + * T::smem_per_group_load_size); + }; + + auto a_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_M * kargs.stride_a + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + auto b_offset = [&](int loop_k_idx, int group_load_idx, int k_group) { + return group_load_idx * T::LOAD_GROUP_N * kargs.stride_b + + (loop_k_idx * T::NUM_LOAD_GROUPS_PER_BK + k_group) * T::LOAD_GROUP_K; + }; + + auto u_ga = make_layout_gmem_group_load_mxsk(lane_id, wave_id_m, kargs.stride_a); + auto u_sa = make_layout_smem_group_load_mxsk(lane_id, wave_id_m); + auto u_gb = make_layout_gmem_group_load_mxsk(lane_id, wave_id_m, kargs.stride_b); + auto u_sb = make_layout_smem_group_load_mxsk(lane_id, wave_id_m); + auto u_ra = make_layout_ra_mxsk(lane_id, wave_id_m); + auto u_rb = make_layout_rb_mxsk(lane_id); + auto u_sfa = make_layout_sfa_mxsk(lane_id, wave_id_m, kargs.stride_sfa); + + auto mma = make_tiled_mma( + seq{}, + seq{}, + seq{}, + mfma_adaptor_swap_ab{}); + + typename decltype(mma)::vtype_a v_a; + typename decltype(mma)::vtype_b v_b; + typename decltype(mma)::vtype_c v_c; + clear(v_c); + using vtype_sfa = vector_t; + using vtype_sfb = vector_t; + + auto issue_tile = [&](int loop_k) { + const int slot = loop_k & 1; + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + opus::static_for([&](auto m_c) { + constexpr int m = decltype(m_c)::value; + async_load(g_a, smem_a_at(slot, m_phase, m, kg), u_ga, u_sa, a_offset(loop_k, m, kg)); + }); + if (m_phase == 0) { + opus::static_for([&](auto n_c) { + constexpr int n = decltype(n_c)::value; + async_load(g_b, smem_b_at(slot, n, kg), u_gb, u_sb, b_offset(loop_k, n, kg)); + }); + } + }); + }; + + auto load_scales = [&](int loop_k, vtype_sfa& v_sfa, vtype_sfb& v_sfb) { + const int scale_base = loop_k * T::SCALES_PER_BK; + v_sfa = load(g_sfa, u_sfa, scale_base); + opus::static_for([&](auto ng_c) { + constexpr int ng = decltype(ng_c)::value; + auto sfb = load(g_sfb, ng * kargs.stride_sfb + scale_base); + opus::static_for([&](auto kg_c) { + constexpr int kg = decltype(kg_c)::value; + v_sfb[ng * T::SCALES_PER_BK + kg] = sfb[kg]; + }); + }); + if constexpr (!SKIP_SCALE_WAIT) { + s_waitcnt_vmcnt(0_I); + } + }; + + issue_tile(0); + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + { + auto sa = make_smem(smem_a_at(0, m_phase, 0, 0)); + auto sb = make_smem(smem_b_at(0, 0, 0)); + v_a = load(sa, u_ra); + v_b = load(sb, u_rb); + s_waitcnt_lgkmcnt(0_I); + } + + for (int k = 0; k < loops; ++k) { + vtype_sfa v_sfa; + vtype_sfb v_sfb; + load_scales(k, v_sfa, v_sfb); + if (k + 1 < loops) { + issue_tile(k + 1); + } + __builtin_amdgcn_s_setprio(1); + if constexpr (PACK_SCALE_ON_DEMAND) { + mma_mxscale_flatmm_accum_on_demand(mma, v_a, v_b, v_sfa, v_sfb, v_c); + } else { + mma_mxscale_flatmm_accum(mma, v_a, v_b, v_sfa, v_sfb, v_c); + } + __builtin_amdgcn_s_setprio(0); + if (k + 1 < loops) { + s_waitcnt_vmcnt(0_I); + __builtin_amdgcn_s_barrier(); + const int slot = (k + 1) & 1; + auto sa = make_smem(smem_a_at(slot, m_phase, 0, 0)); + auto sb = make_smem(smem_b_at(slot, 0, 0)); + v_a = load(sa, u_ra); + v_b = load(sb, u_rb); + s_waitcnt_lgkmcnt(0_I); + } + } + + auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, + 0, lane_id / mma.grpn_c); + auto u_gc = partition_layout_c(mma, + opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); + D_OUT* out_ptr = reinterpret_cast(kargs.ptr_c) + + (size_t)batch_id * kargs.stride_c_batch + + (size_t)row * kargs.stride_c + + (size_t)col; + auto g_out = make_gmem(out_ptr); + store(g_out, v_c, u_gc, 0); +#endif // __gfx950__ +#endif // __HIP_DEVICE_COMPILE__ +} diff --git a/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_scale_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_scale_gfx950.cuh index a166093749..c157489226 100644 --- a/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_scale_gfx950.cuh +++ b/csrc/opus_gemm/include/gfx950/opus_gemm_pipeline_a8w8_scale_gfx950.cuh @@ -167,384 +167,4 @@ inline __device__ auto make_layout_sfa(int lane_id, int wave_id_m, int stride_sf opus::unfold_x_stride(sfa_block_dim, sfa_block_shape, opus::tuple{stride_sfa, 1_I}), opus::unfold_p_coord(sfa_block_dim, opus::tuple{wave_id_m, lane_id % T::W_M})); } - -#endif // __HIP_DEVICE_COMPILE__ (layout functions) - -// ============================================================================ -// Hand-tuned GEMM kernel with block-scale (a8w8 + scale 1x128x128) -// Kernel definition visible on both passes (host pass needs it for stub generation). -// ============================================================================ - -template -__global__ __launch_bounds__(Traits::BLOCK_SIZE, 2) void gemm_a8w8_scale_kernel(opus_gemm_scale_kargs_gfx950 kargs) { -#ifdef __HIP_DEVICE_COMPILE__ -#if defined(__gfx950__) - using namespace opus; - - using T = opus::remove_cvref_t; - using D_A = typename T::D_A; - using D_B = typename T::D_B; - using D_C = typename T::D_C; - using D_ACC = typename T::D_ACC; - using D_SF = typename T::D_SF; - - const int grid_dim_x = opus::grid_size_x() / opus::block_size_x(); - int wgid = (opus::block_id_y() * grid_dim_x) + opus::block_id_x(); - const int num_tiles_n = ceil_div(kargs.n, T::B_N); - int row = (wgid / num_tiles_n) * T::B_M; - int col = (wgid % num_tiles_n) * T::B_N; - - int batch_id = opus::block_id_z(); - int wave_id = __builtin_amdgcn_readfirstlane(opus::thread_id_x() / get_warp_size()); - int lane_id = opus::thread_id_x() % get_warp_size(); - - auto g_a = make_gmem(reinterpret_cast(kargs.ptr_a) + batch_id*kargs.stride_a_batch + row*kargs.stride_a); - auto g_b = make_gmem(reinterpret_cast(kargs.ptr_b) + batch_id*kargs.stride_b_batch + col*kargs.stride_b); - auto g_c = make_gmem(reinterpret_cast(kargs.ptr_c) + batch_id*kargs.stride_c_batch + row*kargs.stride_c + col); - - auto g_sfa = make_gmem(reinterpret_cast(kargs.ptr_sfa) + batch_id*kargs.stride_sfa_batch + static_cast(row/T::GROUP_M)*kargs.stride_sfa); - auto g_sfb = make_gmem(reinterpret_cast(kargs.ptr_sfb) + batch_id*kargs.stride_sfb_batch + static_cast(col/T::GROUP_N)*kargs.stride_sfb); - - int wave_id_m = wave_id % T::T_M; - int wave_id_n = wave_id / T::T_M; - - auto u_ga = make_layout_ga(lane_id, wave_id_m, wave_id_n, kargs.stride_a); - auto u_sa = make_layout_sa(lane_id, wave_id_m, wave_id_n); - auto u_ra = make_layout_ra(lane_id, wave_id_m); - auto u_gb = make_layout_gb(lane_id, wave_id_m, wave_id_n, kargs.stride_b); - auto u_sb = make_layout_sb(lane_id, wave_id_m, wave_id_n); - auto u_rb = make_layout_rb(lane_id, wave_id_n); - - auto u_sfa = make_layout_sfa(lane_id, wave_id_m, kargs.stride_sfa); - - constexpr int smem_a_byte = T::smem_m_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_A); - __shared__ char smem_a[smem_a_byte * 4]; - smem s_a[2][2] = { - {make_smem(reinterpret_cast(smem_a)), - make_smem(reinterpret_cast(smem_a + smem_a_byte))}, - {make_smem(reinterpret_cast(smem_a + 2 * smem_a_byte)), - make_smem(reinterpret_cast(smem_a + 3 * smem_a_byte))} - }; - constexpr int smem_b_byte = T::smem_n_rep * (T::smem_linear_wave + T::smem_padding) * sizeof(D_B); - __shared__ char smem_b[smem_b_byte * 4]; - smem s_b[2][2] = { - {make_smem(reinterpret_cast(smem_b)), - make_smem(reinterpret_cast(smem_b + smem_b_byte))}, - {make_smem(reinterpret_cast(smem_b + 2 * smem_b_byte)), - make_smem(reinterpret_cast(smem_b + 3 * smem_b_byte))} - }; - - auto mma = make_tiled_mma( - seq{}, - seq{}, - seq{}, - mfma_adaptor_swap_ab{}); - constexpr int ELEM_C = decltype(mma)::elem_c; - - typename decltype(mma)::vtype_a v_a[2]; - typename decltype(mma)::vtype_b v_b; - typename decltype(mma)::vtype_c v_c[2][2], v_mma; - clear(v_c[0][0]); - clear(v_c[0][1]); - clear(v_c[1][0]); - clear(v_c[1][1]); - - using vtype_sfa = vector_t; - using vtype_sfb = vector_t; - vtype_sfa v_sfa[2][2]; - vtype_sfb v_sfb[2][2]; - - auto a_offset = [&](int half_tile_m, int tile_k) { - return half_tile_m * T::HALF_B_M * kargs.stride_a + tile_k * T::B_K; - }; - auto b_offset = [&](int half_tile_n, int tile_k) { - return half_tile_n * T::HALF_B_N * kargs.stride_b + tile_k * T::B_K; - }; - auto sfa_offset = [&](int half_tile_m, int tile_k) { - return half_tile_m * (T::HALF_B_M / T::GROUP_M) * kargs.stride_sfa + tile_k * (T::B_K / T::GROUP_K); - }; - auto sfb_offset = [&](int half_tile_n, int tile_k) { - return half_tile_n * (T::HALF_B_N / T::GROUP_N) * kargs.stride_sfb + tile_k * (T::B_K / T::GROUP_K); - }; - - const int loops = ceil_div(kargs.k, T::B_K); - int tic = 0, toc = 1; - - // Prologue - v_sfa[tic][0] = load(g_sfa, u_sfa, sfa_offset(0, 0)); - v_sfb[tic][0] = load(g_sfb, sfb_offset(0, 0)); - async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, 0)); - async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, 0)); - v_sfa[tic][1] = load(g_sfa, u_sfa, sfa_offset(1, 0)); - v_sfb[tic][1] = load(g_sfb, sfb_offset(1, 0)); - async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, 0)); - async_load(g_b, s_b[tic][1].ptr, u_gb, u_sb, b_offset(1, 0)); - - if (wave_id_n == 1) __builtin_amdgcn_s_barrier(); - - s_waitcnt_vmcnt(number{}); - __builtin_amdgcn_s_barrier(); - - v_sfa[toc][0] = load(g_sfa, u_sfa, sfa_offset(0, 1)); - v_sfb[toc][0] = load(g_sfb, sfb_offset(0, 1)); - async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, 1)); - async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, 1)); - async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, 1)); - - s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); - __builtin_amdgcn_s_barrier(); - - v_a[0] = load(s_a[tic][0], u_ra); - __builtin_amdgcn_s_barrier(); - - // Main loop - for(int tile = 0; tile < loops - 2; tile += 2) { - // First tile - v_sfb[toc][1] = load(g_sfb, sfb_offset(1, tile + 1)); - v_b = load(s_b[tic][0], u_rb); - async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); - s_waitcnt_lgkmcnt(number{}); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfa[toc][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 1)); - v_a[1] = load(s_a[tic][1], u_ra); - async_load(g_a, s_a[tic][0].ptr, u_ga, u_sa, a_offset(0, tile + 2)); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfb[tic][0] = load(g_sfb, sfb_offset(0, tile + 2)); - v_b = load(s_b[tic][1], u_rb); - async_load(g_b, s_b[tic][0].ptr, u_gb, u_sb, b_offset(0, tile + 2)); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfa[tic][0] = load(g_sfa, u_sfa, sfa_offset(0, tile + 2)); - v_a[0] = load(s_a[toc][0], u_ra); - async_load(g_a, s_a[tic][1].ptr, u_ga, u_sa, a_offset(1, tile + 2)); - s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + 2 * T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - // Second tile - v_sfb[tic][1] = load(g_sfb, sfb_offset(1, tile + 2)); - v_b = load(s_b[toc][0], u_rb); - async_load(g_b, s_b[tic][1].ptr, u_gb, u_sb, b_offset(1, tile + 2)); - s_waitcnt_lgkmcnt(number{}); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[toc][0], v_sfb[toc][0], v_c[0][0]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfa[tic][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 2)); - v_a[1] = load(s_a[toc][1], u_ra); - async_load(g_a, s_a[toc][0].ptr, u_ga, u_sa, a_offset(0, tile + 3)); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[toc][1], v_sfb[toc][0], v_c[1][0]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfb[toc][0] = load(g_sfb, sfb_offset(0, tile + 3)); - v_b = load(s_b[toc][1], u_rb); - async_load(g_b, s_b[toc][0].ptr, u_gb, u_sb, b_offset(0, tile + 3)); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[toc][0], v_sfb[toc][1], v_c[0][1]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfa[toc][0] = load(g_sfa, u_sfa, sfa_offset(0, tile + 3)); - v_a[0] = load(s_a[tic][0], u_ra); - async_load(g_a, s_a[toc][1].ptr, u_ga, u_sa, a_offset(1, tile + 3)); - s_waitcnt_vmcnt(number<2 * T::a_buffer_load_insts + T::b_buffer_load_insts + 2 * T::sfa_buffer_load_insts + T::sfb_buffer_load_insts>{}); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[toc][1], v_sfb[toc][1], v_c[1][1]); - sched_barrier_pairs<2, 0, 0>(); - sched_barrier_pairs<1, 2, 0>(); - sched_barrier_pairs<5, 4, 0>(); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - } - - // Epilogue - { - int tile = loops - 2; - - v_sfb[toc][1] = load(g_sfb, sfb_offset(1, tile + 1)); - v_b = load(s_b[tic][0], u_rb); - async_load(g_b, s_b[toc][1].ptr, u_gb, u_sb, b_offset(1, tile + 1)); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_sfa[toc][1] = load(g_sfa, u_sfa, sfa_offset(1, tile + 1)); - v_a[1] = load(s_a[tic][1], u_ra); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_b = load(s_b[tic][1], u_rb); - s_waitcnt_vmcnt(number{}); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - tic ^= 1; - toc ^= 1; - } - - { - v_a[0] = load(s_a[tic][0], u_ra); - v_b = load(s_b[tic][0], u_rb); - s_waitcnt_vmcnt(number{}); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][0], v_c[0][0]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_a[1] = load(s_a[tic][1], u_ra); - s_waitcnt_vmcnt(0_I); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][0], v_c[1][0]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - - v_b = load(s_b[tic][1], u_rb); - __builtin_amdgcn_s_barrier(); - - s_waitcnt_lgkmcnt(0_I); - __builtin_amdgcn_s_setprio(1); - v_mma = mma(v_a[0], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][0], v_sfb[tic][1], v_c[0][1]); - v_mma = mma(v_a[1], v_b, 0, 0); - scale_c_tile(v_mma, v_sfa[tic][1], v_sfb[tic][1], v_c[1][1]); - __builtin_amdgcn_s_setprio(0); - __builtin_amdgcn_s_barrier(); - __builtin_amdgcn_sched_barrier(0); - } - - if (wave_id_n == 0) __builtin_amdgcn_s_barrier(); - - // Store results to global memory - auto p_coord_c = opus::make_tuple(wave_id_m, lane_id % mma.grpn_c, wave_id_n, lane_id / mma.grpn_c); - auto u_gc = partition_layout_c(mma, opus::make_tuple(kargs.stride_c, 1_I), p_coord_c); - - auto c_offset = [&](int half_tile_m, int half_tile_n) { - return half_tile_m * T::HALF_B_M * kargs.stride_c + half_tile_n * T::HALF_B_N; - }; - - store(g_c, v_c[0][0], u_gc, c_offset(0, 0)); - store(g_c, v_c[0][1], u_gc, c_offset(0, 1)); - store(g_c, v_c[1][0], u_gc, c_offset(1, 0)); - store(g_c, v_c[1][1], u_gc, c_offset(1, 1)); -#else - // Non-gfx950 device pass: empty stub. a8w8 is gfx950-only; the host - // launcher symbol must still exist for the unconditional dispatcher - // reference, but the body uses gfx950-only intrinsics. -#endif // __gfx950__ #endif // __HIP_DEVICE_COMPILE__ -} diff --git a/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a16w16_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a16w16_gfx950.cuh index 27fd6580f5..6842902c4c 100644 --- a/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a16w16_gfx950.cuh +++ b/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a16w16_gfx950.cuh @@ -203,7 +203,7 @@ struct opus_gemm_a16w16_flatmm_traits_gfx950 { static constexpr int T_N = 1; // compute-wave count along N static constexpr int T_K = 1; // compute-wave count along K - // โ”€โ”€ Warp-spec 4-wave pipeline constraints โ”€โ”€ + // -- Warp-spec 4-wave pipeline constraints -- static_assert(T_K == 1, "flatmm requires T_K=1"); static_assert(T_M == 2, "flatmm requires T_M=2 (ra layout depends on it)"); static_assert(T_N == 1, "flatmm requires T_N=1 (consumer waves share N slab)"); @@ -447,7 +447,7 @@ struct opus_flatmm_splitk_traits_gfx950 { // Pipeline: opus_gemm_pipeline_a16w16_persistent_gfx950.cuh (ported from the // standalone reference kernel gemm_a16w16_8wave_mouter.cc). // -// Layout: each WG handles m_per_wg tile_m ร— 1 tile_n (M outer loop). Within +// Layout: each WG handles m_per_wg tile_m x 1 tile_n (M outer loop). Within // one XCD, consecutive launch-wave WGs share the same m_grp and span all 8 // tile_n stripes; this lets the A tile stay resident in L2 across 8 N tiles // for the duration of one m_grp. @@ -690,7 +690,7 @@ struct opus_gemm_a16w16_mono_tile_traits_gfx950 { static constexpr int VEC_B = opus::get<1>(VEC{}); static constexpr int VEC_C = opus::get<2>(VEC{}); - // โ”€โ”€ Locked tile/wave geometry (kernel-internal constants) โ”€โ”€ + // -- Locked tile/wave geometry (kernel-internal constants) -- static constexpr int T_M = 2; static constexpr int T_N = 4; static constexpr int T_K = 1; @@ -702,13 +702,13 @@ struct opus_gemm_a16w16_mono_tile_traits_gfx950 { static_assert(BLOCK_SIZE == 512, "mono_tile requires BLOCK_SIZE = 512 (8 waves * 64 lanes)"); - // โ”€โ”€ Locked vector widths โ”€โ”€ + // -- Locked vector widths -- static_assert(VEC_A == 8 && VEC_B == 8 && VEC_C == 8, "mono_tile requires VEC_A = VEC_B = VEC_C = 8"); static_assert(VEC_A == 16 / sizeof(D_A), "mono_tile VEC_A must equal 16 / sizeof(D_A) (= 8 for bf16)"); - // โ”€โ”€ Block tile divisibility โ”€โ”€ + // -- Block tile divisibility -- static_assert(B_M % (W_M * T_M) == 0, "mono_tile requires B_M divisible by W_M * T_M = 32"); static_assert(B_N % (W_N * T_N) == 0, @@ -716,7 +716,7 @@ struct opus_gemm_a16w16_mono_tile_traits_gfx950 { static_assert(B_K % (W_K * T_K) == 0, "mono_tile requires B_K divisible by W_K * T_K = 32"); - // โ”€โ”€ Derived MMA repeat counts โ”€โ”€ + // -- Derived MMA repeat counts -- static constexpr int E_M = B_M / (W_M * T_M); static constexpr int E_N = B_N / (W_N * T_N); static constexpr int E_K = B_K / (W_K * T_K); @@ -726,7 +726,7 @@ struct opus_gemm_a16w16_mono_tile_traits_gfx950 { "mono_tile requires E_N divisible by (T_N / T_M) = 2 " "-> B_N % 128 == 0 with the locked T_M=2,T_N=4 geometry"); - // โ”€โ”€ LDS layout โ”€โ”€ + // -- LDS layout -- static constexpr int smem_linear_wave = 64 * 16 / sizeof(D_A); // 512 for bf16 static_assert(smem_linear_wave % B_K == 0, "mono_tile requires B_K to divide smem_linear_wave (=512 for bf16)"); @@ -752,8 +752,8 @@ struct opus_gemm_a16w16_mono_tile_traits_gfx950 { static_assert(smem_sub_e_m > 0 && (E_M % smem_sub_e_m) == 0, "mono_tile: E_M must be divisible by smem_sub / (W_M/T_N)"); - // โ”€โ”€ Buffer / ds_read instruction counts (mirror kernel_traits in the - // upstream template; recomputed here so codegen can sanity-print). โ”€โ”€ + // -- Buffer / ds_read instruction counts (mirror kernel_traits in the + // upstream template; recomputed here so codegen can sanity-print). -- static constexpr int a_buffer_load_insts = B_M * B_K / (BLOCK_SIZE * VEC_A); static constexpr int b_buffer_load_insts = B_N * B_K / (BLOCK_SIZE * VEC_B); static constexpr int a_ds_read_insts = (E_M * E_K * W_M * W_K) / (64 * VEC_A); diff --git a/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh b/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh index f59a8d020c..c341e04a80 100644 --- a/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh +++ b/csrc/opus_gemm/include/gfx950/opus_gemm_traits_a8w8_scale_gfx950.cuh @@ -6,6 +6,7 @@ #pragma once #include "../opus_gemm_utils.cuh" +#include "opus_gemm_traits_a16w16_gfx950.cuh" // opus_splitk_ws_handle template +struct opus_gemm_a8w8_mxscale_flatmm_splitk_traits_gfx950 { + using BLOCK = opus::remove_cvref_t; + using DTYPE = opus::remove_cvref_t; + using VEC = opus::remove_cvref_t; + using GROUP = opus::remove_cvref_t; + + static constexpr int BLOCK_SIZE = BLOCK_SIZE_; + + static constexpr int B_M = opus::get<0>(BLOCK{}); + static constexpr int B_N = opus::get<1>(BLOCK{}); + static constexpr int B_K = opus::get<2>(BLOCK{}); + + using D_A = opus::tuple_element_t<0, DTYPE>; + using D_B = opus::tuple_element_t<1, DTYPE>; + using D_C = opus::tuple_element_t<2, DTYPE>; + using D_ACC = opus::tuple_element_t<3, DTYPE>; + using D_SF = opus::tuple_element_t<4, DTYPE>; + static_assert(std::is_same::value); + static_assert(std::is_same_v, "mxscale flatmm splitK expects fp8 A/B"); + static_assert(std::is_same_v, "mxscale flatmm splitK main writes fp32 workspace"); + static_assert(std::is_same_v, "mxscale flatmm splitK accumulates in fp32"); + static_assert(std::is_same_v, "mxscale flatmm splitK consumes e8m0 uint8 scales"); + + // 4 waves per WG: 2 producer waves + 2 consumer waves. + // + // Two consumer-wave layouts, selected at compile time from B_M: + // tileM (B_M >= 32): consumers split M (T_M=2, T_N=1). Default; used by + // all pre-existing kids (64/128 rows). Bit-identical to the original. + // tileN (B_M == 16): consumers split N (T_M=1, T_N=2). A 16-row A tile + // maps to a single MFMA M-wave, so small-M / decode BMM shapes stop + // over-computing a fat B_M tile (the ~10 us floor on kid320=64x32 for + // M<=32 came from computing 64 rows for 8 valid ones). The two consumer + // waves instead each own half of B_N. + static constexpr bool IS_TILE_N = (B_M == 16); + static constexpr int T_M = IS_TILE_N ? 1 : 2; + static constexpr int T_N = IS_TILE_N ? 2 : 1; + static constexpr int T_K = 1; + static_assert(T_K == 1); + static_assert(BLOCK_SIZE == 256, "flatmm splitK requires 4 wave64 waves"); +#if !defined(__HIP_DEVICE_COMPILE__) || defined(__gfx950__) + static_assert(BLOCK_SIZE == 4 * opus::get_warp_size(), + "flatmm splitK requires exactly four waves"); +#endif + + static constexpr int W_M = 16; + static constexpr int W_N = 16; + static constexpr int W_K = 128; + + static constexpr int VEC_A = opus::get<0>(VEC{}); + static constexpr int VEC_B = opus::get<1>(VEC{}); + static constexpr int VEC_C = opus::get<2>(VEC{}); + + static constexpr int GROUP_M = opus::get<0>(GROUP{}); + static constexpr int GROUP_N = opus::get<1>(GROUP{}); + static constexpr int GROUP_K = opus::get<2>(GROUP{}); + static_assert(GROUP_M == 1 && GROUP_N == 128 && GROUP_K == 128); + static_assert(B_K % GROUP_K == 0, + "flatmm K tile must contain whole DSv4 scale blocks"); + + // async group load geometry; fp8-specific B_K=128 path uses one MFMA per + // LOAD_GROUP_K, unlike a16w16 flatmm where LOAD_GROUP_K=W_K*2. + // tileN uses 16-wide A/B async-load groups so that (a) a B_M=16 A tile is a + // single load group and (b) LOAD_GROUP_M == LOAD_GROUP_N keeps a single + // `slots` value valid for both A and B (no A/B slot decoupling needed), and + // B_N=32 splits into two 16-col groups -- one per consumer N-wave. + static constexpr int LOAD_GROUP_M = IS_TILE_N ? 16 : 32; + static constexpr int LOAD_GROUP_N = IS_TILE_N ? 16 : 32; + static constexpr int LOAD_GROUP_K = W_K; + static constexpr int LOAD_GROUP_M_LANE = 1; + static constexpr int LOAD_GROUP_N_LANE = 1; + static constexpr int NUM_LOAD_GROUPS_PER_BM = B_M / LOAD_GROUP_M; + static constexpr int NUM_LOAD_GROUPS_PER_BN = B_N / LOAD_GROUP_N; + static constexpr int NUM_LOAD_GROUPS_PER_BK = B_K / LOAD_GROUP_K; + static_assert(NUM_LOAD_GROUPS_PER_BM * LOAD_GROUP_M == B_M); + static_assert(NUM_LOAD_GROUPS_PER_BN * LOAD_GROUP_N == B_N); + static_assert(NUM_LOAD_GROUPS_PER_BK == B_K / GROUP_K); + + static constexpr int COM_REP_M = B_M / (W_M * T_M); + static constexpr int COM_REP_N = B_N / (W_N * T_N); + static constexpr int COM_REP_K = B_K / (W_K * T_K); + static_assert(COM_REP_M == 1 || COM_REP_M == 2 || COM_REP_M == 4, + "mxscale flatmm splitK supports 16 (tileN) / 32 / 64 / 128 rows per tile"); + static_assert(COM_REP_N >= 1, "B_N must be a multiple of W_N*T_N"); + // tileN splits B_N across two consumer waves, so B_N must contain 2*W_N cols + // and every N scale group must be wave-splittable without straddling. + static_assert(!IS_TILE_N || (B_N % (W_N * T_N) == 0), + "tileN requires B_N divisible by W_N*T_N (=32)"); + static_assert(COM_REP_K == NUM_LOAD_GROUPS_PER_BK); + static_assert(B_N <= 2 * GROUP_N, + "mxscale flatmm splitK supports up to two 128-column B scale blocks"); + static_assert(GROUP_N % B_N == 0 || B_N % GROUP_N == 0, + "B tile must align with 128-column B scale blocks"); + static constexpr int SCALES_PER_BK = B_K / GROUP_K; + static constexpr int N_SCALE_GROUPS = (B_N + GROUP_N - 1) / GROUP_N; + + static_assert(VEC_A == 16 / sizeof(D_A)); + static_assert(VEC_B == 16 / sizeof(D_B)); + static constexpr int smem_linear_wave_per_async_load = opus::get_warp_size() * 16 / sizeof(D_A); + static constexpr int smem_sub = smem_linear_wave_per_async_load / LOAD_GROUP_K; + static constexpr int slots = LOAD_GROUP_M / smem_sub; + static constexpr int smem_padding = 2 * 16 / sizeof(D_A); + static constexpr int smem_per_group_load_size = + slots * (smem_linear_wave_per_async_load + smem_padding) * sizeof(D_A); + + static constexpr int WG_PER_CU = WG_PER_CU_; + static constexpr int LDS_SIZE_TOTAL = 163840; + static constexpr int max_lds_size_per_wg = LDS_SIZE_TOTAL / WG_PER_CU_; + static constexpr int per_block_iter_lds_size = + (NUM_LOAD_GROUPS_PER_BM + NUM_LOAD_GROUPS_PER_BN) + * NUM_LOAD_GROUPS_PER_BK * smem_per_group_load_size; + static constexpr int prefetch_k_iter = max_lds_size_per_wg / per_block_iter_lds_size; + static_assert(prefetch_k_iter >= 3, + "flatmm splitK pipeline requires at least 3 LDS prefetch slots"); + + static constexpr int a_buffer_load_insts = NUM_LOAD_GROUPS_PER_BM * NUM_LOAD_GROUPS_PER_BK * slots / 2; + static constexpr int b_buffer_load_insts = NUM_LOAD_GROUPS_PER_BN * NUM_LOAD_GROUPS_PER_BK * slots / 2; + static constexpr int a_ds_read_insts = (COM_REP_M * COM_REP_K * W_M * W_K) / (opus::get_warp_size() * VEC_A); + static constexpr int b_ds_read_insts = (COM_REP_N * COM_REP_K * W_N * W_K) / (opus::get_warp_size() * VEC_B); + static constexpr int mma_insts = COM_REP_M * COM_REP_N * COM_REP_K; +}; diff --git a/csrc/opus_gemm/include/gfx950/splitk_reduce_gfx950.cuh b/csrc/opus_gemm/include/gfx950/splitk_reduce_gfx950.cuh index 327c6bbf3b..0a4f4a9365 100644 --- a/csrc/opus_gemm/include/gfx950/splitk_reduce_gfx950.cuh +++ b/csrc/opus_gemm/include/gfx950/splitk_reduce_gfx950.cuh @@ -110,7 +110,7 @@ __global__ void splitk_reduce_kernel( const int b = bm_id / M; const int m = bm_id - b * M; - // โ”€โ”€ Bias prefetch (per-N vector load) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ + // -- Bias prefetch (per-N vector load) ---------------------------------- // Bias is per-output-feature [N] (F.linear convention). Each thread // loads VEC bias values at its own n_base. Fired before the split-K // accumulation so the vmem loads overlap. @@ -200,7 +200,7 @@ __global__ void splitk_reduce_kernel( c_idx + g * STEP); }); } else if (n_base < N) { - // Tail path: decompose valid โˆˆ [1, VEC-1] into descending + // Tail path: decompose valid ? [1, VEC-1] into descending // power-of-2 chunks so we emit dwordx4/dwordx2/dword/short // instead of VEC scalar stores. // Ref: demon_gcn/opus_gemm/mxfp8_e8m0/gemm_mxfp_a8w8_1d1d.hpp @@ -257,3 +257,71 @@ __global__ void splitk_reduce_kernel( #endif // __gfx950__ #endif // __HIP_DEVICE_COMPILE__ } + +// -- opus_bmm_splitk_reduce_kernel ------------------------------------------ +// mmajor batched-matmul split-K reduce. Sums the fp32 split-K workspace along +// the split axis and casts to D_OUT, writing an mmajor C tensor whose row +// (M) and batch strides are passed explicitly (stride_c / stride_c_batch). +// Distinct from splitk_reduce_kernel above (no bias fold, mmajor C strides, +// VEC/BLOCK defaults tuned for the BMM launcher's 8x128 grid). Lives here so +// the codegen'd a8w8_mxscale BMM launcher (.device.cu / fused host TU) can +// reference it; the explicit instantiations are emitted alongside the baseline +// reduce into splitk_reduce_gfx950.device.cu (gen_instances_gfx950.py's +// splitk_reduce_extra hook). +template +__global__ void opus_bmm_splitk_reduce_kernel( + const opus_splitk_ws_handle* __restrict__ ws_handle, + D_OUT* __restrict__ out, + int split_k, int M, int N, int batch, + int padded_M, int padded_N, + int stride_c, int stride_c_batch) +{ +#ifdef __HIP_DEVICE_COMPILE__ +#if defined(__gfx950__) + using namespace opus; + constexpr int STEP = 16 / sizeof(D_OUT); + static_assert(VEC % STEP == 0); + + const int bm_id = int(block_id_y()); + const int n_base = (int(block_id_x()) * BLOCK + int(thread_id_x())) * VEC; + if (bm_id >= batch * M || n_base >= N) return; + const int b = bm_id / M; + const int m = bm_id - b * M; + + const float* __restrict__ workspace = + reinterpret_cast(ws_handle->ptr); + const long split_stride = (long)batch * padded_M * padded_N; + const int base = b * padded_M * padded_N + m * padded_N + n_base; + auto g_ws = make_gmem(workspace, (unsigned int)(split_stride * split_k * sizeof(float))); + vector_t acc; + static_for([&](auto i) { acc[i.value] = 0.0f; }); + for (int s = 0; s < split_k; ++s) { + #pragma unroll + for (int g = 0; g < VEC / 4; ++g) { + auto v = g_ws.template load<4>((long)s * split_stride + base + g * 4); + static_for<4>([&](auto j) { acc[g * 4 + j.value] += v[j.value]; }); + } + } + + vector_t out_v; + static_for([&](auto i) { out_v[i.value] = static_cast(acc[i.value]); }); + const int c_idx = b * stride_c_batch + m * stride_c + n_base; + const size_t c_records = + (size_t)(M - 1) * (size_t)stride_c + + (size_t)(batch - 1) * (size_t)stride_c_batch + + (size_t)N; + auto g_c = make_gmem(out, (unsigned int)(c_records * sizeof(D_OUT))); + if (n_base + VEC <= N) { + static_for([&](auto group) { + constexpr int off = group.value * STEP; + g_c.template store(slice(out_v, number{}, number{}), c_idx + off); + }); + } else { + #pragma unroll + for (int i = 0; i < VEC; ++i) { + if (n_base + i < N) g_c.template store<1>(out_v[i], c_idx + i); + } + } +#endif +#endif +} diff --git a/csrc/opus_gemm/include/opus_bmm.h b/csrc/opus_gemm/include/opus_bmm.h new file mode 100644 index 0000000000..ba9c7a5514 --- /dev/null +++ b/csrc/opus_gemm/include/opus_bmm.h @@ -0,0 +1,20 @@ +// SPDX-License-Identifier: MIT +// Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. +#pragma once + +#include "aiter_tensor.h" + +// Opus BMM public C++ API. These frontends use BMM/grouped layouts (for example +// DSV4 wo_a) while reusing the shared opus GEMM backend kernels. + +// fp8 e8m0 mxscale (block-scale) BMM (zero-copy DSV4 wo_a): O/Y are [M, batch, +// *], wo_a/w_scale batch-major. Y dtype in {fp32, bf16}. dim0=M, dim1=batch (K +// contiguous); the batch axis memory position is otherwise free (see host +// stride checks). kid-dispatched; driven by bmm_a8w8_mxscale_opus (Python). +void opus_bmm_a8w8_mxscale(aiter_tensor_t& O, + aiter_tensor_t& wo_a, + aiter_tensor_t& Y, + aiter_tensor_t& x_scale, + aiter_tensor_t& w_scale, + int splitK, + int kernelId); diff --git a/csrc/opus_gemm/include/opus_gemm_utils.cuh b/csrc/opus_gemm/include/opus_gemm_utils.cuh index 2f6d24f0ba..2895039626 100644 --- a/csrc/opus_gemm/include/opus_gemm_utils.cuh +++ b/csrc/opus_gemm/include/opus_gemm_utils.cuh @@ -94,4 +94,14 @@ inline __device__ void scale_c_tile( }); } +// Broadcast one e8m0 exponent byte across all four 32-block slots of a packed +// scale word (128-block DSV4 scale -> 32-block gfx950 scaled-MFMA semantics). +// Shared by the a8w8_mxscale BMM pipeline and the flatmm split-K pipeline (both +// #included in opus_bmm.cu), so it lives here once instead of a per-header copy. +template +OPUS_D int pack_e8m0x4(S scale) { + const int e = static_cast(scale); + return e * 0x01010101; +} + #endif // __HIP_DEVICE_COMPILE__ diff --git a/csrc/opus_gemm/opus_bmm.cu b/csrc/opus_gemm/opus_bmm.cu new file mode 100644 index 0000000000..5b94ac0bf0 --- /dev/null +++ b/csrc/opus_gemm/opus_bmm.cu @@ -0,0 +1,85 @@ +// SPDX-License-Identifier: MIT +// Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. + +// Host-side BMM frontends. These expose BMM/grouped-layout APIs while reusing +// the generated opus GEMM backend launcher symbols. +// +// Host pass only, like opus_gemm.cu. Every device symbol this module needs is +// codegen'd: the per-tile compute kernels into _C{void,bf16,fp32}.device.cu +// and opus_bmm_splitk_reduce_kernel<{__bf16,float}, 8, 128> into +// splitk_reduce_gfx950.device.cu (gen_instances_gfx950.py's splitk_reduce_extra +// hook), which owns both the device kernel and the host __device_stub__ the +// codegen'd split-K launchers reference. +#ifndef __HIP_DEVICE_COMPILE__ + +#include "opus_bmm.h" +#include "opus_gemm_arch.cuh" +#include "opus_build_archs.h" +#include "opus_gemm_manifest.h" +#include "opus_bmm_mxscale_tune_lookup.h" // GENERATE_BMM_MXSCALE_FLATMM_SPLITK_LOOKUP_FP32 +#include "opus_gemm_utils.cuh" // bf16_t / fp32_t +#include "aiter_stream.h" +#include "gfx950/opus_bmm_launchers_a8w8_mxscale_gfx950.cuh" + +#include + +#ifdef OPUS_BUILD_HAS_GFX950 +namespace opus_bmm_detail { + +// Uniform kid->launcher fn-pointer type. Every kid is codegen'd (no hand-written +// adapters), so this namespace only holds the shared type. +using OpusBmmMxscaleFlatmmSplitkKernel = void (*)( + aiter_tensor_t &, aiter_tensor_t &, aiter_tensor_t &, + aiter_tensor_t &, aiter_tensor_t &, int /*splitK*/); +} // namespace opus_bmm_detail + +// Table-driven kid -> launcher dispatch. Launchers come from the generated +// GENERATE_BMM_MXSCALE_FLATMM_SPLITK_LOOKUP_FP32 macro; unknown / untuned kids +// fall back to the 32x128x128 wg2 baseline. +static opus_bmm_detail::OpusBmmMxscaleFlatmmSplitkKernel +opus_bmm_a8w8_mxscale_tune_dispatch(int id) +{ + using namespace opus_bmm_detail; + static const std::unordered_map kTune = { + GENERATE_BMM_MXSCALE_FLATMM_SPLITK_LOOKUP_FP32(fp32_t) + }; + auto it = kTune.find(id); + if (it != kTune.end()) + return it->second; + return &opus_bmm_a8w8_mxscale_flatmm_splitk_256x32x128x128_2x1_16x16x128_1x128x128_wgpcu2; +} +#endif // OPUS_BUILD_HAS_GFX950 + +void opus_bmm_a8w8_mxscale( + aiter_tensor_t &O, + aiter_tensor_t &wo_a, + aiter_tensor_t &Y, + aiter_tensor_t &x_scale, + aiter_tensor_t &w_scale, + int splitK, + int kernelId) +{ + // Common dtype/shape validation + arch gate, done once here so the codegen'd + // launchers (which omit these to stay lean) and the fused kid 100 wrapper share + // one check. The _impl still re-checks internally (idempotent). + opus_bmm_a8w8_common_checks(O, wo_a, Y, + "opus_bmm_a8w8_mxscale"); +#ifndef OPUS_BUILD_HAS_GFX950 + AITER_CHECK(false, + "opus_bmm_a8w8_mxscale requires " + "OPUS_BUILD_HAS_GFX950"); +#else + { + const auto &arch_info = opus_get_arch_info(); + AITER_CHECK(arch_info.arch == OpusGfxArch::Gfx950, + "opus_bmm_a8w8_mxscale is gfx950-only; " + "current device ", arch_info.dev, " has gcnArchName='", + arch_info.name, "'"); + } + // Single table lookup instead of a ~40-case switch (see opus_gemm.cu). + opus_bmm_a8w8_mxscale_tune_dispatch(kernelId)( + O, wo_a, Y, x_scale, w_scale, splitK); +#endif // OPUS_BUILD_HAS_GFX950 +} + +#endif // !__HIP_DEVICE_COMPILE__ diff --git a/csrc/opus_gemm/opus_bmm_mxscale_tune.py b/csrc/opus_gemm/opus_bmm_mxscale_tune.py new file mode 100644 index 0000000000..0f6388df29 --- /dev/null +++ b/csrc/opus_gemm/opus_bmm_mxscale_tune.py @@ -0,0 +1,483 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2026, Advanced Micro Devices, Inc. All rights reserved. +"""Framework tuner for the opus fp8 e8m0 mxscale flatmm split-K BMM (DSV4 wo_a). + +Wired into the canonical :class:`GemmCommonTuner`, so it runs like the other +aiter GEMM tuners: multi-GPU via ``mp_tuner``, standard ``-i/--untune_file`` / +``-o/--tune_file`` CLI, batching, and the shared post-process / CSV writer. + +The candidate pool lives here (``_TUNE_POLICY``). Per kid it holds only the +split-K factors to sweep; tile geometry, kernelName and the M alignment come from +the codegen instance table, so a kid cannot be tuned on a shape its launcher +rejects. That alignment used to be a second hand-maintained column and was wrong +in both directions -- it hid kid326, which is really arbitrary-M, from every +unaligned shape while the runtime dispatched it there anyway. + +Runtime schema (what the tuner emits, and what the runtime reads back): + gfx,b,m,n,k,libtype,kernelId,splitK,us,kernelName,tflops,bw,errRatio +``aiter/ops/batched_gemm_op_a8w8.py:lookup_mxscale_bmm_config`` indexes on +``["gfx","b","m","n","k"]``, dispatches to a backend on the winning row's +``libtype``, and ``bmm_op.py`` reads ``kernelId`` / ``splitK`` off that row, so +those columns must match exactly. + +Verification (the part that catches column-transpose / scale defects): + * inputs are *signed* and have *per-128-K-block varied magnitude* + (``randn * 2**randint(-4,4)`` per block) so the e8m0 128-block scales span + many exponents. Uniform non-negative ``rand()/10`` data hides a pure output + column permutation (kid312/313 measured ~0.007 there but ~0.7-1.0 on real + signed data) -- see the opus_bmm.md root-cause note. + * reference is a dequantized fp32 einsum. + * gate: ``mp_tuner`` runs ``checkAllclose(rtol=1e-2, atol=1e-2)`` and + ``post_process`` keeps the fastest candidate whose mismatch fraction is + ``<= --errRatio`` (default 0.02). A still-broken tileN COM_REP_N>1 kernel + measures ~0.5 here and is rejected; the fp8 e8m0 quant floor is ~1e-4. + +Usage (gfx950 only; the repo root must be on PYTHONPATH so the edited/rebuilt +tree wins over any installed aiter): + cd && PYTHONPATH=$PWD \\ + python3 csrc/opus_gemm/opus_bmm_mxscale_tune.py -g 16 -m 1,16,64 -n 1024 -k 4096 + + # re-tune every shape already in the shipped CSV, write a diffable copy: + ... opus_bmm_mxscale_tune.py + + # overwrite the shipped tuned CSV in place: + ... opus_bmm_mxscale_tune.py --apply + + # from an untuned CSV (columns: b,m,n,k -- or g,m,n,k), 8-way parallel: + ... opus_bmm_mxscale_tune.py -i my_untuned.csv -o /tmp/out.csv --mp 8 +""" + +import os +import sys +from typing import Any, ClassVar + +import pandas as pd +import torch + +from aiter import dtypes, logger +from aiter.ops.opus.bmm_op import _opus_bmm_a8w8_mxscale_raw +from aiter.utility.base_tuner import GemmCommonTuner, TunerCommon +from aiter.utility.mp_tuner import mp_tuner + +# Neither op_tests nor this directory is a package, so put both on sys.path. This +# also has to hold in the spawned mp_tuner subprocesses, which re-import this +# module top-to-bottom. +_HERE = os.path.dirname(os.path.abspath(__file__)) +_REPO = os.path.abspath(os.path.join(_HERE, "..", "..")) +_OPTESTS = os.path.join(_REPO, "op_tests") +for _p in (_HERE, _OPTESTS): + if _p not in sys.path: + sys.path.insert(0, _p) + +# opus_gemm_common is pure python (stdlib only), so importing the codegen kid +# table here does not pull in the build. +from opus_gemm_common import a8w8_mxscale_bmm_kernel_lists +from test_opus_a8w8_bmm import ( + GROUP, + _quant_block_e8m0, + _quant_per_token_e8m0, + run_torch, +) + +# kid -> OpusGemmInstance. Kids are disjoint across the BMM families today; +# assert so a future collision (which the codegen dedups by launcher name +# downstream) is caught here instead of silently tuning one of the two. +_CODEGEN_BMM = {} +for _fam in a8w8_mxscale_bmm_kernel_lists: + for _kid, _inst in _fam.items(): + assert ( + _kid not in _CODEGEN_BMM + ), f"bmm kid {_kid} collides across codegen families; disambiguate by name" + _CODEGEN_BMM[_kid] = _inst + +# Split-K sweep for the flatmm_splitk family. Small-M / few-tile shapes (the G16 +# wo_a decode: 16 batch * n1024 * k4096) underfill the CUs at splitK=1, so split-K +# (fp32-workspace partials + fused reduce tail) can win by exposing parallelism +# along K. The correctness gate drops any combo a kernel mishandles, so an +# over-broad sweep is safe, just slower. +_SK = [1, 2, 4, 8] + +# Tuning policy: kid -> splitK list. The ONLY hand-maintained per-kid metadata -- +# it decides which kids to sweep and with which split-K factors, not their +# geometry and not their M alignment. Tile shape, kernelName and m_align all come +# from the codegen instance, so this cannot drift from what compiles. kid 0 (the +# heuristic default) is intentionally not tuned. +_TUNE_POLICY = { + # flatmm_splitk family: the M=16/32 last-mile tiles, the mid-M SFA/SFB-preload + # tiles and the 64x* tiles. All are split-K capable via the fused reduce tail, + # except kid646 whose persistent DIRECT_ONLY schedule requires splitK == 1. + 32: _SK, + 64: _SK, + 138: _SK, + 139: _SK, + 256: _SK, + 311: _SK, + 312: _SK, + 313: _SK, + 314: _SK, + 316: _SK, + 317: _SK, + 318: _SK, + 319: _SK, + 320: _SK, + 321: _SK, + 322: _SK, + 323: _SK, + 324: _SK, + 326: _SK, + 327: _SK, + 640: _SK, + 642: _SK, + 646: [1], + 650: _SK, + 653: _SK, + # fused single-tile launcher. + 100: [1], + # pipeline family; kid158 preloads both the per-token SFA and the block SFB + # panel into LDS. + 149: [1], + 150: [1], + 151: [1], + 152: [1], + 158: [1], + # monolithic mouter / wave pipelines. + 131: [1], + 132: [1], + 134: [1], + 142: [1], + 144: [1], + 148: [1], + 160: [1], + 161: [1], + # minterleave only exists in split-K form. + 162: [2, 4, 8], + 163: [2, 4, 8], + # 128x128x128 tiles, splitK=1 only and deliberately so. They are the largest + # BMM tile (COM_REP_M=4 x COM_REP_N=8 -> 32 C fragments, 128 fp32 C values per + # lane) at 512 VGPRs / occupancy 1. At splitK=1 they run the Cbf16 + # direct-output kernel and are strong -- kid325 wins 5 shipped wo_a rows. At + # splitK>1 they switch to the Cvoid fp32-workspace kernel, which spills and has + # never won (g2/m256: best kid325 split-K is 23.6us against the 14.5us winner), + # and which is also where the clang-22 gfx950 greedy-VGPR miscompile lives (one + # C-fragment dword left unmaterialized under --amdgpu-mfma-vgpr-form). + 128: [1], + 137: [1], + 325: [1], +} + +# Only the flatmm_splitk (non-direct) and minterleave launchers honor splitK>1. +# Any other family sweeping it is a policy bug, so fail loudly at import. +for _kid, _sks in _TUNE_POLICY.items(): + if any(s > 1 for s in _sks): + _tag = _CODEGEN_BMM[_kid].kernel_tag + assert ( + _tag == "a8w8_mxscale_bmm_flatmm_splitk" + and not getattr(_CODEGEN_BMM[_kid], "direct_only", False) + ) or _tag == "a8w8_mxscale_bmm_minterleave", ( + f"kid {_kid} ({_tag}) is not split-K capable but sweeps {_sks}" + ) + + +def _applicable(kid, g, m, n, k): + """Split-K factors worth trying for this kid on this shape ([] == skip it).""" + k_inst = _CODEGEN_BMM[kid] + if n % k_inst.B_N or k % k_inst.B_K or m % k_inst.m_align: + return [] + return _TUNE_POLICY[kid] + + +SHIPPED_CSV = os.path.join( + _REPO, + "aiter", + "configs", + "model_configs", + "dsv4_batched_gemm_a8w8_blockscale_mxscale_tuned.csv", +) +DEFAULT_OUT = os.path.join(_REPO, "dsv4_bmm_mxscale_retuned.csv") + + +# --------------------------------------------------------------------------- +# mp_tuner hooks (module-level so the spawn workers can import them by name). +# --------------------------------------------------------------------------- +def _gen_varied(shape, k, device): + """Signed, per-128-K-block varied-magnitude bf16 (mirrors _block_varied).""" + x = torch.randn(shape, dtype=dtypes.fp32, device=device) + amp = torch.exp2(torch.randint(-4, 4, (k // GROUP,), device=device).float()) + return (x * amp.repeat_interleave(GROUP)).to(dtypes.bf16) + + +def gen_bmm_mxscale_data(batch, m, n, k, seed, out_dtype, device="cuda"): + """Return the 6-tuple mp_tuner indexes into: + + 0 O_in [m,g,k] fp8 (mmajor transposed view, K contiguous) + 1 W_mx [g,n,k] fp8 (batch-major) + 2 Y [m,g,n] out_dtype output buffer + 3 xs_in [m,g,k/128] uint8 e8m0 per-token scale (mmajor view) + 4 ws_mx [g,n/128,k/128] uint8 e8m0 128x128-block scale + 5 ref [m,g,n] out_dtype dequant fp32 einsum reference + """ + torch.manual_seed(seed) + O_bf16 = _gen_varied((batch, m, k), k, device) + W_bf16 = _gen_varied((batch, n, k), k, device) + O_mx, xs_mx, xs_fp32 = _quant_per_token_e8m0(O_bf16) + W_mx, ws_mx, ws_fp32 = _quant_block_e8m0(W_bf16) + O_in = O_mx.transpose(0, 1) # [m,g,k] + xs_in = xs_mx.transpose(0, 1) # [m,g,k/128] + Y = torch.empty((m, batch, n), dtype=out_dtype, device=device) + ref = run_torch(O_mx, W_mx, xs_fp32, ws_fp32).transpose(0, 1).to(out_dtype) + return (O_in, W_mx, Y, xs_in, ws_mx, ref) + + +def run_bmm_mxscale_bench(O_in, W_mx, Y, xs_in, ws_mx, kernelId, splitK): + """Tuner bench func: run the kid in-place, return Y for checkAllclose.""" + _opus_bmm_a8w8_mxscale_raw(O_in, W_mx, Y, xs_in, ws_mx, splitK, kernelId) + return Y + + +def _bmm_ref_passthrough(ref): + """ref_func: the fp32 reference is precomputed in gen_data (slot 5).""" + return ref + + +# --------------------------------------------------------------------------- +# Tuner +# --------------------------------------------------------------------------- +class OpusBmmMxscaleTuner(GemmCommonTuner): + ARG_DEFAULTS: ClassVar[dict[str, Any]] = { + **GemmCommonTuner.ARG_DEFAULTS, + "tune_file": DEFAULT_OUT, + "untune_file": "", + # Fraction-of-mismatch (rtol=atol=1e-2) accept threshold. Correct kids + # sit at the ~1e-4 fp8 e8m0 quant floor; a column-transposed kid is ~0.5. + "errRatio": 0.02, + "batch": 100, + } + + KEYS: ClassVar[list[str]] = ["gfx", "b", "m", "n", "k"] + RESULTS: ClassVar[list[str]] = [ + "libtype", + "kernelId", + "splitK", + "us", + "kernelName", + "tflops", + "bw", + "errRatio", + ] + + def __init__(self): + # Bypass GemmCommonTuner.__init__ (it force-swaps "M"/"N" in the key, + # which assumes the uppercase gptoss schema). Go straight to the + # grandparent with our lowercase batched schema. + TunerCommon.__init__( + self, + "OpusBmmMxscaleTuner", + self.KEYS, + self.RESULTS, + description="Tune opus fp8 e8m0 mxscale flatmm split-K BMM (DSV4 wo_a)", + ) + # sort N before M like the GEMM tuners (cosmetic ordering of the CSV). + self.sort_keys = ["gfx", "b", "n", "m", "k"] + + # --- schema helpers ----------------------------------------------------- + def getKernelName(self, kernelId): + k_inst = _CODEGEN_BMM.get(int(kernelId)) + return k_inst.name if k_inst else None + + def calculate(self, results, bpes=None): + info, time, _err = results + if time == self.INVALID_TIME: + return 0, 0 + _gfx, b, m, n, k = info[0] + us_s = time * 1e-6 + tflops = round(2 * b * m * n * k / us_s / 1e12, 1) + # fp8 A + fp8 W + bf16 out. + bw = round((b * m * k + b * n * k + 2 * b * m * n) / us_s / 1e9, 2) + return tflops, bw + + def result_to_df(self, results): + rows = [] + for info, time, err in results: + keys, kernelId, splitK, kernelName = info + resolved = kernelName or self.getKernelName(kernelId) + tflops, bw = self.calculate((info, time, err)) + row = dict(zip(self.keys, keys)) + row.update( + { + "libtype": "opus", + "kernelId": int(kernelId), + "splitK": int(splitK), + "us": time, + "kernelName": "None" if resolved is None else str(resolved), + "tflops": tflops, + "bw": bw, + "errRatio": err, + } + ) + rows.append(row) + return pd.DataFrame(rows, columns=self.columns) + + # --- CLI ---------------------------------------------------------------- + def _setup_specific_arguments(self): + # Free the base "-k/--splitK" store_true so we can reuse -k for the K dim. + for action in list(self.parser._actions): + if "-k" in action.option_strings or "--splitK" in action.option_strings: + self.parser._actions.remove(action) + for s in action.option_strings: + self.parser._option_string_actions.pop(s, None) + for grp in self.parser._action_groups: + if action in grp._group_actions: + grp._group_actions.remove(action) + break + + def _intlist(s): + return [int(x) for x in str(s).split(",") if x != ""] + + self.parser.add_argument( + "-g", + "--batch_g", + type=_intlist, + default=None, + help="comma list of batch g (e.g. 2,8,16)", + ) + self.parser.add_argument( + "-m", + "--M", + type=_intlist, + default=None, + help="comma list of M (e.g. 1,16,64)", + ) + self.parser.add_argument( + "-n", + "--N", + type=_intlist, + default=[1024], + help="comma list of N (default 1024)", + ) + self.parser.add_argument( + "-k", + "--K", + type=_intlist, + default=[4096], + help="comma list of K (default 4096)", + ) + self.parser.add_argument( + "--apply", + action="store_true", + default=False, + help="overwrite the shipped tuned CSV in place", + ) + + # --- shape sourcing ----------------------------------------------------- + def _shapes_from_shipped(self): + try: + df = pd.read_csv(SHIPPED_CSV) + except FileNotFoundError: + return [] + return sorted( + {(int(r.b), int(r.m), int(r.n), int(r.k)) for _, r in df.iterrows()} + ) + + def pre_process(self, args): + if args.apply: + args.tune_file = SHIPPED_CSV + + gfx = self.get_gfx() + if args.batch_g and args.M: + shapes = [ + (g, m, n, k) + for g in args.batch_g + for m in args.M + for n in args.N + for k in args.K + ] + elif args.untune_file and os.path.exists(args.untune_file): + df = pd.read_csv(args.untune_file) + df.columns = [c.strip().lower() for c in df.columns] + bcol = "b" if "b" in df.columns else "g" + shapes = [ + (int(r[bcol]), int(r["m"]), int(r["n"]), int(r["k"])) + for _, r in df.iterrows() + ] + else: + logger.info( + "no -g/-m and no untune_file; re-tuning shapes from %s", SHIPPED_CSV + ) + shapes = self._shapes_from_shipped() + + self.untunedf = pd.DataFrame( + [{"gfx": gfx, "b": g, "m": m, "n": n, "k": k} for (g, m, n, k) in shapes], + columns=self.keys, + ) + self.tunedf = self.get_tuned_gemm_list(args.tune_file) + + # Skip shapes already present in the tuned CSV (unless --all forces retune). + if not args.all and len(self.tunedf) and len(self.untunedf): + td = self.tunedf + if "gfx" not in td.columns: + td = td.assign(gfx=gfx) + have = set(td[self.keys].apply(lambda r: tuple(r), axis=1).tolist()) + mask = self.untunedf.apply(lambda r: tuple(r) in have, axis=1) + if args.verbose and mask.any(): + logger.info("skipping %d already-tuned shapes", int(mask.sum())) + self.untunedf = self.untunedf[~mask].reset_index(drop=True) + + # --- tuning ------------------------------------------------------------- + def tune(self, untunedf, tunedf, args): + gfx = self.get_gfx() + out_dtype = dtypes.bf16 + perf_kwargs = {"num_warmup": args.warmup, "num_iters": args.iters} + + task = [] + tasks_data = [] + for seed, i in enumerate(range(len(untunedf)), start=1): + b = int(untunedf.loc[i, "b"]) + m = int(untunedf.loc[i, "m"]) + n = int(untunedf.loc[i, "n"]) + k = int(untunedf.loc[i, "k"]) + info_keys = (gfx, b, m, n, k) + + n_cand = 0 + for kid in _TUNE_POLICY: + for sk in _applicable(kid, b, m, n, k): + info = (info_keys, kid, sk, "") + task.append( + ( + info, + gen_bmm_mxscale_data, + (b, m, n, k, seed, out_dtype), + run_bmm_mxscale_bench, + ([0, 1, 2, 3, 4], kid, sk), + perf_kwargs, + _bmm_ref_passthrough, + ([5],), + {}, + None, + 1e-2, # rtol + 1e-2, # atol + None, # compare_fn + None, # max_abs_delta + [2], # output_keys: NaN-init Y to catch partial writes + ) + ) + n_cand += 1 + tasks_data.append((n_cand, ())) + + if not task: + return [] + return mp_tuner( + task, + tasks_data, + args.mp, + False, + args.shape_grouped, + args.errRatio, + timeout=args.timeout, + verbose=args.verbose, + ) + + +if __name__ == "__main__": + tuner = OpusBmmMxscaleTuner() + _args = tuner.parse_args() + tuner.run(_args, False) diff --git a/csrc/opus_gemm/opus_gemm_common.py b/csrc/opus_gemm/opus_gemm_common.py index d22c52766c..7ee5ffa1b4 100644 --- a/csrc/opus_gemm/opus_gemm_common.py +++ b/csrc/opus_gemm/opus_gemm_common.py @@ -83,10 +83,49 @@ class OpusGemmInstance: cluster_wg_m: int = 4 cluster_wg_n: int = 4 + # --- a8w8_mxscale BMM flatmm-splitK axes (kernel_tag == + # "a8w8_mxscale_bmm_flatmm_splitk"). The BMM main kernel template is + # gemm_a8w8_mxscale_flatmm_splitk_kernel + # so unlike a16w16 each kid carries two compile-time booleans in addition to + # the tile. direct_only == consumer-self-load direct-store (splitK==1 only); + # prefetch_scale == scale-prefetch variant; fused_reduce == splitK==2 fused + # tail-reduce launch path. These drive both the launcher body and the set of + # device instantiations gen_instances emits for the kid. + direct_only: bool = False + prefetch_scale: bool = False + fused_reduce: bool = False + # a8w8_mxscale BMM flatmm-splitK only: preload this split's SFA (per-token) + + # SFB (block) scale panels into LDS once, then read scales from LDS in the + # consumer instead of a per-K-tile global buffer_load. Maps to the kernel's + # 5th template bool PRELOAD_SF_LDS. + preload_sf: bool = False + # a8w8_mxscale BMM specialized-pipeline axis (minterleave / mouter / + # mouter_tunable / wave4m2_selfload families). Maps to the kernel's trailing + # `bool SKIP_SCALE_WAIT` template param: skip the s_waitcnt on the per-K-tile + # scale load (the scale is issued a tile ahead), trading a correctness margin + # for pipeline overlap. Drives both the launcher body and the device + # instantiation set for the kid. + skip_scale_wait: bool = False + # a8w8_mxscale BMM wave4m2_selfload family extra bool axis (kernel template + # order: ). + pack_scale_on_demand: bool = False + # a8w8_mxscale BMM pipeline family (kids 150/151/152): dual bf16/fp32 + # traits + one of the gemm_a8w8_scale_* kernels selected by these flags + # (all-false = plain scale kernel). + k1024_only: bool = False + k1024_lb1: bool = False + # a8w8_mxscale BMM pipeline family (kid158): preload BOTH SFA (per-token) and + # SFB (block) scale panels into LDS. Maps to the pipeline kernel + # gemm_a8w8_scale_preload_sf_kernel. + preload_sf_lds: bool = False + # Symbol root ("opus_gemm" for GEMM, "opus_bmm" for the batched frontends). + name_root: str = "opus_gemm" + @property def name(self) -> str: parts = [ - "opus_gemm", + self.name_root, "x".join(map(str, [self.BLOCK_SIZE, self.B_M, self.B_N, self.B_K])), "x".join(map(str, [self.T_M, self.T_N])), "x".join(map(str, [self.W_M, self.W_N, self.W_K])), @@ -96,7 +135,54 @@ def name(self) -> str: parts.insert(1, self.arch_prefix) # tag inserts shift right by one slot when arch_prefix is set tag_at = 1 + (1 if self.arch_prefix else 0) - if self.kernel_tag == "a16w16_flatmm": + if self.kernel_tag == "a8w8_mxscale_bmm_flatmm_splitk": + # opus_bmm_a8w8_mxscale_flatmm_splitk__wgpcu{N}[_selfload][_scaleprefetch] + parts.insert(tag_at, "a8w8_mxscale_flatmm_splitk") + parts.append(f"wgpcu{self.WG_PER_CU}") + if self.direct_only: + parts.append("selfload") + if self.prefetch_scale: + parts.append("scaleprefetch") + if self.preload_sf: + parts.append("sfpreload") + elif self.kernel_tag == "a8w8_mxscale_bmm_minterleave": + # opus_bmm_a8w8_mxscale_flatmm_minterleave__wgpcu{N}[_skip_scale_wait] + parts.insert(tag_at, "a8w8_mxscale_flatmm_minterleave") + parts.append(f"wgpcu{self.WG_PER_CU}") + if self.skip_scale_wait: + parts.append("skip_scale_wait") + elif self.kernel_tag == "a8w8_mxscale_bmm_fused": + parts.insert(tag_at, "a8w8_mxscale_flatmm_fused") + parts.append(f"wgpcu{self.WG_PER_CU}") + elif self.kernel_tag == "a8w8_mxscale_bmm_pipeline": + parts.insert(tag_at, "a8w8_mxscale_pipeline") + if self.k1024_only: + parts.append("k1024") + elif self.k1024_lb1: + parts.append("k1024lb1") + elif self.preload_sf_lds: + parts.append("preload_sf") + elif self.kernel_tag == "a8w8_mxscale_bmm_mouter": + parts.insert(tag_at, "a8w8_mxscale_flatmm_mouter") + parts.append(f"wgpcu{self.WG_PER_CU}") + if self.skip_scale_wait: + parts.append("ssw") + elif self.kernel_tag == "a8w8_mxscale_bmm_mouter_tunable": + parts.insert(tag_at, "a8w8_mxscale_flatmm_mouter_tunable") + parts.append(f"wgpcu{self.WG_PER_CU}") + if self.skip_scale_wait: + parts.append("ssw") + elif self.kernel_tag == "a8w8_mxscale_bmm_wave8n2": + parts.insert(tag_at, "a8w8_mxscale_flatmm_wave8n2") + parts.append(f"wgpcu{self.WG_PER_CU}") + elif self.kernel_tag == "a8w8_mxscale_bmm_wave4m2_selfload": + parts.insert(tag_at, "a8w8_mxscale_flatmm_wave4m2_selfload") + parts.append(f"wgpcu{self.WG_PER_CU}") + if self.skip_scale_wait: + parts.append("ssw") + if self.pack_scale_on_demand: + parts.append("psod") + elif self.kernel_tag == "a16w16_flatmm": parts.insert(tag_at, "flatmm") parts.append(f"wgpcu{self.WG_PER_CU}") elif self.kernel_tag == "a16w16_flatmm_splitk": @@ -142,6 +228,42 @@ def name(self) -> str: parts.append(f"cA{self.cachectl_a}cB{self.cachectl_b}") return "_".join(parts) + @property + def m_align(self) -> int: + """M multiple this kid's generated host guard enforces (1 == any M). + + The launcher family decides it, not the kid: see _BMM_M_ALIGN_TILES and + the AITER_CHECK blocks the matching launcher body in + codegen/gen_instances_gfx950.py emits. Consumers that pick a kid for a + shape (the tuner's candidate filter, the runtime's padded-M lookup) must + read it from here rather than keep their own list -- two hand-maintained + copies is exactly how kid326 ended up excluded from tuning while the + runtime dispatched it anyway. + """ + mult = _BMM_M_ALIGN_TILES.get(self.kernel_tag) + if mult is not None: + return self.B_M * mult if mult else 1 + # Non-BMM families: has_oob is the codegen flag that says whether the + # tail is masked, and opus_gemm_tune.py already gates on it this way. + return 1 if self.has_oob else self.B_M + + +# a8w8_mxscale BMM launcher family -> the B_M multiple its host guard requires, +# or 0 when the launcher masks a partial M tile and emits no M check at all. +# Mirrors the AITER_CHECK blocks in the launcher bodies of +# codegen/gen_instances_gfx950.py (_BMM_*_LAUNCHER_BODY); gen_instances asserts +# the two agree, so a guard edit that forgets this table fails the build. +_BMM_M_ALIGN_TILES = { + "a8w8_mxscale_bmm_flatmm_splitk": 0, + "a8w8_mxscale_bmm_pipeline": 0, + "a8w8_mxscale_bmm_fused": 0, + "a8w8_mxscale_bmm_minterleave": 2, # MI=2 M tiles per WG, baked in + "a8w8_mxscale_bmm_wave4m2_selfload": 2, # LOGICAL_B_M = B_M * 2 + "a8w8_mxscale_bmm_wave8n2": 1, + "a8w8_mxscale_bmm_mouter": 1, + "a8w8_mxscale_bmm_mouter_tunable": 1, +} + def _a16w16(bs, bm, bn, bk, tn, wm, wn, wk, has_oob=True, cachectl_a=0, cachectl_b=17): """Factory for a16w16 split-barrier kid instances. @@ -236,9 +358,286 @@ def _a16w16_flatmm(bm, bn, bk, wg_per_cu): # fmt: off # --- per-pipeline kernel instance lists --- a8w8_scale_kernels_list = { + # kid 1 (256x256) is the launcher hardcoded by opus_gemm.cu's + # opus_dispatch_scale (the only a8w8_scale GEMM path). The 128x256 sibling + # kid 720 was removed below. 1: OpusGemmInstance(512, 256, 256, 128, 4, 2, 16, 16, 128, 16, 16, 4, 1, 128, 128, "a8w8_scale", ["fp32_t"]), } +# Dead 128x256 scale GEMM tiles removed (no CSV/dispatch caller): +# - kid 720 (a8w8_scale, fp32 block-scale): only consumer was the removed +# opus_bmm_a8w8_scale mmajor path. +# - kid 710 (a8w8_mxscale, e8m0 block-scale): only consumer was the opus_bmm +# kid 149 hand-written adapter (via the _mmajor sibling), now replaced by +# the BMM-native a8w8_mxscale_bmm_pipeline 128x256 instance. +# Both were the same gemm_a8w8_scale_kernel specialization, differing only in +# scale dtype; opus_dispatch_scale still uses the 256x256 kid 1 above. + + +def _a8w8_mxscale_bmm_flatmm_splitk( + bm, bn, bk, wg_per_cu, direct_only=False, prefetch_scale=False, preload_sf=False +): + """fp8 e8m0 mxscale BATCHED matmul flatmm split-K tile. + + Backs opus_bmm_a8w8_mxscale(); the main kernel + (gemm_a8w8_mxscale_flatmm_splitk_kernel) writes an fp32 workspace and a + shared reduce kernel casts to the Y dtype (bf16/fp32), so output_dtypes is + fp32 workspace here. Locked geometry (matches the hand-written traits in + opus_bmm.cu): BLOCK_SIZE=256 (4 waves), T_M=2/T_N=1, MFMA 16x16x128 (fp8), + VEC=(16,16,4), GROUP=(1,128,128) (per-token M, 128x128 block scale). + direct_only / prefetch_scale are the two kernel compile-time booleans. + """ + # tileN (bm==16): consumers split N (T_M=1, T_N=2). tileM (bm>=32): split M + # (T_M=2, T_N=1). The real T_M/T_N is derived in the C++ traits from B_M; + # these values only drive the generated symbol name, so keep them honest. + t_m, t_n = (1, 2) if bm == 16 else (2, 1) + inst = OpusGemmInstance( + 256, # BLOCK_SIZE + bm, bn, bk, # BLOCK tile + t_m, t_n, # T_M, T_N (4-wave warp-spec; tileN=1,2 / tileM=2,1) + 16, 16, 128, # W_M, W_N, W_K (MFMA 16x16x128 fp8) -- name only + 16, 16, 4, # VEC_A, VEC_B, VEC_C + 1, 128, 128, # GROUP_M=1 (per-token), GROUP_N=GROUP_K=128 + "a8w8_mxscale_bmm_flatmm_splitk", + # Single host instantiation: the launcher is templated on D_C + # only to satisfy the codegen host-decl machinery; its body ignores D_C + # and branches on Y.dtype() at runtime (native __bf16/float), exactly + # like the hand-written _impl. The fp32 split-K workspace dtype is fixed + # inside the traits, and the reduce kernel casts to the runtime Y dtype. + ["fp32_t"], + wg_per_cu, + ) + inst.name_root = "opus_bmm" + inst.direct_only = direct_only + inst.prefetch_scale = prefetch_scale + inst.preload_sf = preload_sf + return inst + + +# fp8 e8m0 mxscale BMM flatmm split-K tiles. kid numbers preserved from the old +# opus_bmm.cu switch so existing tuned CSVs / heuristics keep working. Each kid = +# (B_M, B_N, B_K, WG_PER_CU, direct_only, prefetch_scale). Big-tile pipelines +# (mouter / minterleave / wave*n* / pipeline, kids 131/132/134/140-163/149-152) +# stay monolithic in opus_bmm.cu and are NOT migrated here. +_BMM_MXSCALE_SPLITK_TILES = { + # tileN (B_M=16): single 16-row MFMA M-wave so small-M/decode shapes (M<=32) + # don't over-compute a fat B_M tile. Targets the G=2 K=4096 M<=32 gap vs bf16. + 316: (16, 32, 256, 2, False, False), + 317: (16, 32, 256, 2, False, True), # scale prefetch + 318: (16, 32, 128, 2, False, False), + # prefetch-depth sweep: higher WG_PER_CU shrinks per-WG LDS -> shallower + # prefetch_k_iter + more occupancy (small-M/few-tile shapes want this). + 319: (16, 32, 256, 4, False, False), + 314: (16, 32, 512, 2, False, False), # fewer K-iters (8) per WG + # wider-N tileN: larger B_N raises COM_REP_N (more MFMA/iter) to hide + # ds_read+scale latency; WG_PER_CU keeps prefetch_k_iter >= 3. + 313: (16, 64, 256, 2, False, False), # COM_REP_N=2 + 312: (16, 128, 256, 1, False, False), # COM_REP_N=4 + # M=16/32 last-mile (G=2 N=1024 K=4096): 311 = wide-K tileN + scale prefetch; + # 321/323 = 32x32 tileM (exact M=32 fit, no OOB waste, COM_REP_N=2). + 311: (16, 32, 512, 2, False, True), + 321: (32, 32, 256, 2, False, True), + 323: (32, 32, 128, 2, False, True), + # fine tiles (small / mid M) + 320: (64, 32, 256, 2, False, False), + 322: (64, 32, 256, 1, False, False), + # kid324 = kid320 tile + SFA+SFB scale panels preloaded into LDS + # (PRELOAD_SF_LDS; wired via the preload-tiles dict below, not the 6-tuple). + # ATT on kid320 showed ~20% of consumer cycles stalled on vmcnt for the + # per-K-tile global scale load; staging both panels into LDS once (ds_read / + # lgkmcnt) breaks the mid-M valley: G4 K4096 M256 0.93->1.00x, M512 + # 0.94->1.01x, M192 0.91->0.98x vs bf16 (+8-26% TFLOPS over kid320, M128-1024). + # Other attempts (scaleprefetch, B_K=128/512, wg4, 64x64 splitK) all <= kid320. + 640: (32, 64, 256, 2, False, False), + 642: (32, 64, 256, 1, False, False), + 646: (32, 64, 256, 2, True, False), # consumer self-load (splitK==1) + 650: (64, 64, 128, 2, False, False), + 653: (64, 64, 128, 2, False, True), # scale prefetch + # No 64x64x256 kid: mirroring bf16's MT64x64x256 forces wg_per_cu=1 (LDS + # ~198KB), so at M=256 it runs half the WGs and lands 0.77x vs bf16. bf16 only + # wins it via stream-K (refills low tile count), which the flatmm pipeline lacks. + 128: (128, 128, 128, 1, False, False), + 137: (128, 128, 128, 1, False, True), # scale prefetch + 138: (64, 128, 256, 1, False, False), + 139: (128, 64, 256, 1, False, False), + # baseline tiles (guaranteed-runnable fallbacks; kid 0 is the heuristic default) + 256: (32, 256, 128, 1, False, False), + 64: (64, 128, 128, 2, False, False), + 0: (32, 128, 128, 2, False, False), + 32: (32, 128, 128, 2, False, False), +} +a8w8_mxscale_bmm_flatmm_splitk_kernels_list = { + kid: _a8w8_mxscale_bmm_flatmm_splitk(bm, bn, bk, wg, direct, prefetch) + for kid, (bm, bn, bk, wg, direct, prefetch) in _BMM_MXSCALE_SPLITK_TILES.items() +} + +# SFA/SFB-into-LDS preload variants (PRELOAD_SF_LDS). Kept in a separate dict so +# the base 6-tuple stays untouched; each entry is (B_M, B_N, B_K, WG_PER_CU) and +# always sets preload_sf=True (non-direct, non-prefetch). +_BMM_MXSCALE_SPLITK_PRELOAD_TILES = { + 324: (64, 32, 256, 2), # = kid320 + SFA/SFB scale panels preloaded to LDS + # mid-M wg1 tiles + SFA/SFB preload (same mechanism as kid324/kid158): staging + # both scale panels into LDS removes the per-K-tile global scale vmcnt load that + # gated the plain/scaleprefetch tiles. On K=4096 M256-2048 this wins +13-17% + # over the old kid137/653/139 picks (kid325 ships G2/M2048, G4/M1024, G8/M512, + # G16/M256; kid326 ships G8/M256). K=1024 gains are ~noise (few K-tiles). kid327 + # kept as a candidate but wins nothing robustly (clock-fragile at cold sclk). + 325: (128, 128, 128, 1), # = kid128/137 tile + preload + 326: (128, 64, 256, 1), # = kid139 tile + preload + 327: (64, 128, 256, 1), # = kid138 tile + preload +} +a8w8_mxscale_bmm_flatmm_splitk_kernels_list.update({ + kid: _a8w8_mxscale_bmm_flatmm_splitk(bm, bn, bk, wg, preload_sf=True) + for kid, (bm, bn, bk, wg) in _BMM_MXSCALE_SPLITK_PRELOAD_TILES.items() +}) + + +def _a8w8_mxscale_bmm_minterleave(bm, bn, bk, wg_per_cu, skip_scale_wait=False): + """fp8 e8m0 mxscale BATCHED matmul M-tile-interleaved tile. + + Backs opus_bmm_a8w8_mxscale() kids 162/163. The main kernel + (gemm_a8w8_mxscale_flatmm_minterleave_kernel) + processes MI=2 consecutive M tiles per WG (baked in the launcher, requires + M % (MI*B_M) == 0); splitK is unused (must be 1). Same locked geometry / + traits as the flatmm split-K family (BLOCK_SIZE=256, T_M=2/T_N=1, MFMA + 16x16x128, VEC=(16,16,4), GROUP=(1,128,128), fp32 workspace tuple slot). + """ + t_m, t_n = (1, 2) if bm == 16 else (2, 1) + inst = OpusGemmInstance( + 256, # BLOCK_SIZE + bm, bn, bk, # BLOCK tile + t_m, t_n, # T_M, T_N (name only) + 16, 16, 128, # W_M, W_N, W_K (name only) + 16, 16, 4, # VEC_A, VEC_B, VEC_C + 1, 128, 128, # GROUP_M=1 (per-token), GROUP_N=GROUP_K=128 + "a8w8_mxscale_bmm_minterleave", + ["fp32_t"], # single fp32 host stub; body branches on Y.dtype() + wg_per_cu, + ) + inst.name_root = "opus_bmm" + inst.skip_scale_wait = skip_scale_wait + return inst + + +# fp8 e8m0 mxscale BMM M-tile-interleaved tiles (kids 162/163). Fixed geometry +# m128n128k128 wg1; the only axis is SKIP_SCALE_WAIT. +_BMM_MXSCALE_MINTERLEAVE_TILES = { + # (B_M, B_N, B_K, WG_PER_CU, skip_scale_wait) + 162: (128, 128, 128, 1, False), + 163: (128, 128, 128, 1, True), # skip per-K-tile scale s_waitcnt +} +a8w8_mxscale_bmm_minterleave_kernels_list = { + kid: _a8w8_mxscale_bmm_minterleave(bm, bn, bk, wg, skip) + for kid, (bm, bn, bk, wg, skip) in _BMM_MXSCALE_MINTERLEAVE_TILES.items() +} + + +def _a8w8_mxscale_bmm_spec(tag, bm, bn, bk, wg_per_cu, **flags): + """Generic fp8 e8m0 mxscale BMM specialized-pipeline tile builder. + + Same locked geometry/traits family as the flatmm split-K kids (BLOCK_SIZE + 256, MFMA 16x16x128, VEC=(16,16,4), GROUP=(1,128,128), fp32 workspace tuple + slot). `tag` selects the kernel family (wave8n2 / wave4m2_selfload); + `flags` sets the family's compile-time axes. + """ + t_m, t_n = (1, 2) if bm == 16 else (2, 1) + inst = OpusGemmInstance( + 256, bm, bn, bk, t_m, t_n, 16, 16, 128, 16, 16, 4, 1, 128, 128, + tag, ["fp32_t"], wg_per_cu, + ) + inst.name_root = "opus_bmm" + for key, val in flags.items(): + setattr(inst, key, val) + return inst + + +# fused (kid 100): the only fused-reduce path (splitK counter variant). Same +# 256x32x128x128 wg2 traits as standard kid 0/32, so its device symbols resolve +# to the standard family's TUs -> host-only launcher emit. +a8w8_mxscale_bmm_fused_kernels_list = { + 100: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_fused", 32, 128, 128, 2), +} + +# pipeline (kids 149/150/151/152/158): BLOCK_SIZE 512, m{128,256}n256k128, dual +# bf16/fp32 traits (output dtype baked into the traits tuple), non-splitk scale +# kargs. One of the gemm_a8w8_scale_* kernels selected by flags. The wave +# layout (T_M/T_N/W_*) is derived inside opus_gemm_a8w8_scale_traits_gfx950 from +# BLOCK + , so only B_M/B_N/B_K matter here (the T_M/T_N passed to +# OpusGemmInstance are cosmetic for this tag). +def _a8w8_mxscale_bmm_pipeline(**flags): + inst = OpusGemmInstance( + 512, 256, 256, 128, 2, 1, 16, 16, 128, 16, 16, 4, 1, 128, 128, + "a8w8_mxscale_bmm_pipeline", ["fp32_t"], 1, + ) + inst.name_root = "opus_bmm" + for key, val in flags.items(): + setattr(inst, key, val) + return inst + + +a8w8_mxscale_bmm_pipeline_kernels_list = { + # kid 149: B_M=128 plain scale pipeline (m128n256k128). Same gemm_a8w8_scale_ + # kernel as kid 150, just half the M tile -> 2x output tiles -> fills more CUs + # on batched wo_a shapes. Was a hand-written cross-module adapter delegating + # to opus_gemm's a8w8_mxscale GEMM launcher; now BMM-native codegen. + 149: _a8w8_mxscale_bmm_pipeline(B_M=128), + 150: _a8w8_mxscale_bmm_pipeline(), + 151: _a8w8_mxscale_bmm_pipeline(k1024_only=True), + 152: _a8w8_mxscale_bmm_pipeline(k1024_lb1=True), + # kid158: preload BOTH SFA (per-token) and SFB (block) scale panels into LDS. + 158: _a8w8_mxscale_bmm_pipeline(preload_sf_lds=True), +} + +# mouter (kids 131/144) + mouter_tunable (kids 160/161): wg1 m128n128k128, +# 1 bool axis . Both share gemm_..._mouter_kernel, so the +# tunable variant reuses the mouter device instantiations (host-only emit). +a8w8_mxscale_bmm_mouter_kernels_list = { + 131: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_mouter", 128, 128, 128, 1), + 144: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_mouter", 128, 128, 128, 1, skip_scale_wait=True), +} +a8w8_mxscale_bmm_mouter_tunable_kernels_list = { + 160: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_mouter_tunable", 128, 128, 128, 1), + 161: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_mouter_tunable", 128, 128, 128, 1, skip_scale_wait=True), +} + +# wave8n2 (kid 132): wg1 m128n128k128, no compile-time flags (logical B_N = 256). +a8w8_mxscale_bmm_wave8n2_kernels_list = { + 132: _a8w8_mxscale_bmm_spec("a8w8_mxscale_bmm_wave8n2", 128, 128, 128, 1), +} + +# wave4m2_selfload (kids 134/142/148): wg1 m128n128k128, 2 bool axes +# (logical B_M = 128*2 = 256). +_BMM_WAVE4M2_TILES = { + # (ssw, psod) + 134: (False, False), + 142: (True, False), + 148: (True, True), +} +a8w8_mxscale_bmm_wave4m2_selfload_kernels_list = { + kid: _a8w8_mxscale_bmm_spec( + "a8w8_mxscale_bmm_wave4m2_selfload", 128, 128, 128, 1, + skip_scale_wait=ssw, pack_scale_on_demand=psod, + ) + for kid, (ssw, psod) in _BMM_WAVE4M2_TILES.items() +} + +# All name-keyed a8w8_mxscale BMM kernel families (gfx950-only). Kept as a tuple +# of the per-family kid-keyed dicts -- NOT merged into one dict, because int kids +# repeat across families and are deduped downstream by launcher NAME (see +# gen_instances.py). Single source of truth for both consumers there: the codegen +# kdict merge and the BMM int-kid tune-lookup emitter. +a8w8_mxscale_bmm_kernel_lists = ( + a8w8_mxscale_bmm_flatmm_splitk_kernels_list, + a8w8_mxscale_bmm_fused_kernels_list, + a8w8_mxscale_bmm_minterleave_kernels_list, + a8w8_mxscale_bmm_mouter_kernels_list, + a8w8_mxscale_bmm_mouter_tunable_kernels_list, + a8w8_mxscale_bmm_pipeline_kernels_list, + a8w8_mxscale_bmm_wave8n2_kernels_list, + a8w8_mxscale_bmm_wave4m2_selfload_kernels_list, +) + + a8w8_kernels_list = { 2: OpusGemmInstance(512, 256, 256, 128, 2, 4, 16, 16, 128, 16, 16, 4, 0, 0, 0, "a8w8", ["fp32_t"]), } diff --git a/csrc/pybind/opus_gemm_pybind.cu b/csrc/pybind/opus_gemm_pybind.cu index 4540a2cc60..e377b41ce0 100644 --- a/csrc/pybind/opus_gemm_pybind.cu +++ b/csrc/pybind/opus_gemm_pybind.cu @@ -9,12 +9,14 @@ #include "rocm_ops.hpp" #include "aiter_stream.h" #include "opus_gemm.h" +#include "opus_bmm.h" PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { AITER_SET_STREAM_PYBIND OPUS_GEMM_PYBIND; OPUS_GEMM_A16W16_TUNE_PYBIND; + OPUS_BMM_A8W8_MXSCALE_PYBIND; OPUS_GEMM_A8W8_BLOCKSCALE_BPRESHUFFLE_TUNE_PYBIND; OPUS_GEMM_WORKSPACE_INIT_PYBIND; } diff --git a/op_tests/test_opus_a8w8_bmm.py b/op_tests/test_opus_a8w8_bmm.py new file mode 100644 index 0000000000..1b95b5a8f7 --- /dev/null +++ b/op_tests/test_opus_a8w8_bmm.py @@ -0,0 +1,503 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2025-2026, Advanced Micro Devices, Inc. All rights reserved. +"""Regression + perf sweep for the opus fp8 e8m0 mxscale flatmm split-K BMM. + +Covers the mmajor DeepSeek-V4 wo_a path: O/Y are [M, G, *] (transposed views of +batch-major [G, M, *]); wo_a + w_scale stay batch-major. Activation scale is +per-token e8m0 (GROUP_M=1), weight scale is 128x128-block e8m0. Candidates are +kid 0 (always-runnable baseline) and the public dispatch path; the reference is +a dequantized fp32 einsum. Per-kid perf comparison / winner selection lives in +``csrc/opus_gemm/opus_bmm_mxscale_tune.py``. + +``--check-m-align`` runs a different check instead of the sweep: an every-kid +guard that OpusGemmInstance.m_align still matches launcher behaviour (see +``check_m_align``). It is kept out of the sweep because it deliberately provokes +launch failures and needs no timing. + +Usage: + python3 op_tests/test_opus_a8w8_bmm.py + python3 op_tests/test_opus_a8w8_bmm.py -s 512,1024,4096 -g 2 -d bf16 + python3 op_tests/test_opus_a8w8_bmm.py --check-m-align +""" + +import argparse +import itertools +import sys + +import pandas as pd +import torch + +import aiter +from aiter import dtypes +from aiter.jit.utils.chip_info import get_gfx +from aiter.ops.batched_gemm_op_a8w8 import lookup_mxscale_bmm_config +from aiter.ops.opus.bmm_op import _opus_bmm_a8w8_mxscale_raw, bmm_a8w8_mxscale_opus +from aiter.test_common import benchmark, checkAllclose, run_perftest + +torch.set_default_device("cuda") + +SUPPORTED_GFX = ["gfx950"] # fp8 e8m0 mxscale flatmm is gfx950-only +GROUP = 128 # GROUP_N == GROUP_K == 128; GROUP_M == 1 (per-token) +_DT = {"fp32": dtypes.fp32, "bf16": dtypes.bf16} + + +def _to_e8m0_scale(scale): + # Round scale up to a power of two so quantized fp8 values stay in range. + e = torch.ceil(torch.log2(scale.to(dtypes.fp32))).to(torch.int32) + 127 + e = torch.clamp(e, 0, 255).to(torch.uint8) + scale_pow2 = torch.exp2(e.to(dtypes.fp32) - 127.0) + return e, scale_pow2 + + +def _quant_per_token_e8m0(x_bf16): + """[G,M,K] bf16 -> fp8 + e8m0 x_scale [G,M,K/128] + fp32 scale.""" + G, M, K = x_bf16.shape + xb = x_bf16.to(dtypes.fp32).view(G, M, K // GROUP, GROUP) + raw = xb.abs().amax(dim=-1, keepdim=True).clamp(min=1e-8) / 448.0 + e8m0, scale = _to_e8m0_scale(raw) + q = (xb / scale).clamp(-448.0, 448.0).to(dtypes.fp8) + return q.view(G, M, K), e8m0.squeeze(-1), scale.squeeze(-1) + + +def _quant_block_e8m0(w_bf16): + """[G,N,K] bf16 -> fp8 + e8m0 w_scale [G,N/128,K/128] + fp32 scale.""" + G, N, K = w_bf16.shape + wb = w_bf16.to(dtypes.fp32).view(G, N // GROUP, GROUP, K // GROUP, GROUP) + raw = wb.abs().amax(dim=(2, 4), keepdim=True).clamp(min=1e-8) / 448.0 + e8m0, scale = _to_e8m0_scale(raw) + q = (wb / scale).clamp(-448.0, 448.0).to(dtypes.fp8) + return ( + q.view(G, N, K), + e8m0.view(G, N // GROUP, K // GROUP), + scale.view(G, N // GROUP, K // GROUP), + ) + + +def run_torch(O_fp8, W_fp8, x_scale, w_scale): + """Reference: dequant fp8 -> fp32 einsum -> [G,M,N]. Not timed.""" + G, M, K = O_fp8.shape + N = W_fp8.shape[1] + act = O_fp8.to(dtypes.fp32).view(G, M, K // GROUP, GROUP) + act = (act * x_scale.unsqueeze(-1)).view(G, M, K) + W = W_fp8.to(dtypes.fp32).view(G, N // GROUP, GROUP, K // GROUP, GROUP) + W = (W * w_scale.view(G, N // GROUP, 1, K // GROUP, 1)).view(G, N, K) + return torch.einsum("gmk,gnk->gmn", act, W).to(dtypes.fp32) + + +def _block_varied(shape, k): + """Signed random tensor whose per-128-K-block magnitude spans several powers + of two, so the e8m0 128-block scales cover many exponents. + + ``rand()/10`` (non-negative, near-uniform) is what let the shipped kid312/313 + tileN COM_REP_N>1 kernels pass this test at ~0.007 rel while silently + transposing output column groups: a pure column permutation over symmetric + positive columns barely moves any element, and the collapsed single block + scale hides scale-application bugs. Signed data makes swapped columns + uncorrelated (~100% element mismatch), and the varied amplitude exercises + real per-block scales -- together they turn this test into a real guard.""" + x = torch.randn(shape, dtype=dtypes.fp32) + amp = torch.exp2(torch.randint(-4, 4, (k // GROUP,), device=x.device).float()) + x = x * amp.repeat_interleave(GROUP) + return x.to(dtypes.bf16) + + +@benchmark() +def test_mxscale_bmm(g, m, n, k, dtype): + ydt = _DT[dtype] + # Canonical batch-major tensors, then feed the kernel transposed (mmajor) + # views exactly like the DSV4 wo_a call does (zero-copy, no contiguous copy). + O_bf16 = _block_varied((g, m, k), k) + W_bf16 = _block_varied((g, n, k), k) + O_mx, xs_mx, xs_fp32 = _quant_per_token_e8m0(O_bf16) + W_mx, ws_mx, ws_fp32 = _quant_block_e8m0(W_bf16) + + O_in = O_mx.transpose(0, 1) # [m,g,k] view + xs_in = xs_mx.transpose(0, 1) # [m,g,k/128] view + ref = run_torch(O_mx, W_mx, xs_fp32, ws_fp32).transpose(0, 1) # [m,g,n] + y_shape = (m, g, n) + + def _call(kid): + Y = torch.empty(y_shape, dtype=ydt) + _opus_bmm_a8w8_mxscale_raw(O_in, W_mx, Y, xs_in, ws_mx, 1, kid) + return Y + + # Correctness-focused: kid 0 (k32 fused) is a fixed baseline with no + # tile-alignment requirement (always runnable), plus the public dispatch path + # end to end. Per-kid perf comparison / winner selection lives in + # csrc/opus_gemm/opus_bmm_mxscale_tune.py, not here. + candidates = {"kid0_k32_fused": (lambda: _call(0), ref)} + + # Public backend-neutral entry: no kernelId -> per-(g,m,n,k) tuned-CSV + # lookup + heuristic fallback + libtype backend routing. Exercises the + # whole aiter.batched_gemm_a8w8_mxscale -> bmm_a8w8_mxscale_opus path end + # to end (not the raw binding). + candidates["auto (batched_gemm_a8w8_mxscale)"] = ( + lambda: aiter.batched_gemm_a8w8_mxscale(O_in, W_mx, xs_in, ws_mx, dtype=ydt), + ref, + ) + + flops = 2.0 * g * m * n * k + # fp8 A + fp8 W + e8m0 scales (uint8) + output. + nbytes = ( + g * m * k + + g * n * k + + g * m * (k // GROUP) + + g * (n // GROUP) * (k // GROUP) + + m * g * n * torch.empty((), dtype=ydt).element_size() + ) + + ret = {"gfx": get_gfx()} + for name, (fn, fn_ref) in candidates.items(): + out, us = run_perftest(fn) + err = checkAllclose( + fn_ref.to(dtypes.fp32), + out.to(dtypes.fp32), + rtol=1e-2, + atol=1e-2, + msg=f"mxscale_bmm {name} g={g} m={m} n={n} k={k}", + ) + ret[f"{name} us"] = us + ret[f"{name} TFLOPS"] = flops / us / 1e6 + ret[f"{name} TB/s"] = nbytes / us / 1e6 + ret[f"{name} err"] = err + return ret + + +@benchmark() +def test_mxscale_bmm_batch_first(g, m, n, k, dtype): + """Batch-leading (batch-major) round trip. + + The caller's natural DSV4 buffers are batch-major: batch is the *first* + (outermost-in-memory) dimension -- activation/output are [G, M, *], weight + is [G, N, K]. They are handed to the kernel as zero-copy [M, G, *] + transposed views (dim0=M, dim1=batch), and the result is written straight + back into a batch-major [G, M, N] buffer through its [M, G, N] view. + + This is the stride path the dropped ``_mmajor`` suffix used to over-claim: + the batch axis sits at an arbitrary (here outermost) memory position while + only K (inputs) and N (output) stay contiguous. Same tuned CSV / heuristic + entries must serve it. Correctness is checked in the caller's native + [G, M, N] order. + """ + ydt = _DT[dtype] + O_bf16 = _block_varied((g, m, k), k) + W_bf16 = _block_varied((g, n, k), k) + O_mx, xs_mx, xs_fp32 = _quant_per_token_e8m0(O_bf16) + W_mx, ws_mx, ws_fp32 = _quant_block_e8m0(W_bf16) + + O_in = O_mx.transpose(0, 1) # [m, g, k] view (K contiguous) + xs_in = xs_mx.transpose(0, 1) # [m, g, k/128] view + ref = run_torch(O_mx, W_mx, xs_fp32, ws_fp32) # [g, m, n] batch-major + + def _call_raw(kid): + # Batch-major output buffer; hand the kernel its [m, g, n] view so the + # store lands at Y.stride(1) (batch) = m*n (outermost), N contiguous. + Yb = torch.empty((g, m, n), dtype=ydt) + _opus_bmm_a8w8_mxscale_raw(O_in, W_mx, Yb.transpose(0, 1), xs_in, ws_mx, 1, kid) + return Yb # [g, m, n] + + def _call_auto(): + # Same tuned-CSV lookup the public entry does, but writing into a + # caller-owned batch-major buffer -- which the guarded public entry no + # longer exposes (it returns fresh token-major), so drive the opus + # backend directly with the looked-up kid + the batch-major out= view. + Yb = torch.empty((g, m, n), dtype=ydt) + cfg = lookup_mxscale_bmm_config(g, m, n, k) + bmm_a8w8_mxscale_opus( + O_in, + W_mx, + xs_in, + ws_mx, + out=Yb.transpose(0, 1), + dtype=ydt, + kernelId=int(cfg["kernelId"]) if cfg is not None else None, + splitK=int(cfg["splitK"]) if cfg is not None else None, + ) + return Yb + + # Correctness-focused: kid 0 (always runnable) as the batch-major baseline, + # plus the backend dispatch path writing into the batch-major buffer via + # out=. Per-kid perf sweep lives in csrc/opus_gemm/opus_bmm_mxscale_tune.py. + candidates = {"kid0_k32_fused": (lambda: _call_raw(0), ref)} + candidates["auto (bmm_a8w8_mxscale_opus)"] = (_call_auto, ref) + + flops = 2.0 * g * m * n * k + # fp8 A + fp8 W + e8m0 scales (uint8) + output. + nbytes = ( + g * m * k + + g * n * k + + g * m * (k // GROUP) + + g * (n // GROUP) * (k // GROUP) + + m * g * n * torch.empty((), dtype=ydt).element_size() + ) + + ret = {"gfx": get_gfx()} + for name, (fn, fn_ref) in candidates.items(): + out, us = run_perftest(fn) + err = checkAllclose( + fn_ref.to(dtypes.fp32), + out.to(dtypes.fp32), + rtol=1e-2, + atol=1e-2, + msg=f"mxscale_bmm_batch_first {name} g={g} m={m} n={n} k={k}", + ) + ret[f"{name} us"] = us + ret[f"{name} TFLOPS"] = flops / us / 1e6 + ret[f"{name} TB/s"] = nbytes / us / 1e6 + ret[f"{name} err"] = err + return ret + + +# --- tileN column-map regression guard ------------------------------------ +# These COM_REP_N>1 kernels previously transposed output column groups. Keep +# them out of the narrow perf table, but always exercise both output layouts +# with signed, varied-block-scale data so the bug cannot silently return. +_TILEN_REGRESSION_KIDS = (312, 313) +_TILEN_REGRESSION_SHAPE = (2, 16, 128, 1024) # G, M, N, K; accepts both kids +_TILEN_REGRESSION_ERR_TOL = 0.003 + + +def check_tilen_column_map(): + """Check kid312/313 column mapping for token- and batch-major output.""" + g, m, n, k = _TILEN_REGRESSION_SHAPE + O_mx, xs_mx, xs_fp32 = _quant_per_token_e8m0(_block_varied((g, m, k), k)) + W_mx, ws_mx, ws_fp32 = _quant_block_e8m0(_block_varied((g, n, k), k)) + O_in = O_mx.transpose(0, 1) + xs_in = xs_mx.transpose(0, 1) + ref = run_torch(O_mx, W_mx, xs_fp32, ws_fp32).transpose(0, 1) + failures = [] + + for kid in _TILEN_REGRESSION_KIDS: + for layout in ("token-major", "batch-major"): + if layout == "token-major": + out = torch.full((m, g, n), float("nan"), dtype=dtypes.bf16) + else: + out = torch.full((g, m, n), float("nan"), dtype=dtypes.bf16).transpose( + 0, 1 + ) + _opus_bmm_a8w8_mxscale_raw(O_in, W_mx, out, xs_in, ws_mx, 1, kid) + torch.cuda.synchronize() + delta = (out.to(dtypes.fp32) - ref).abs() + rows = delta.flatten(1).mean(1) / (ref.abs().flatten(1).mean(1) + 1e-9) + err = rows.max().item() + if not (err <= _TILEN_REGRESSION_ERR_TOL): + failures.append( + f"kid {kid} {layout}: worst row rel err {err:.4f} " + f"> {_TILEN_REGRESSION_ERR_TOL}" + ) + + assert not failures, "tileN column-map regression:\n " + "\n ".join(failures) + return len(_TILEN_REGRESSION_KIDS) * 2 + + +# --- m_align guard --------------------------------------------------------- +# Straddles every tile boundary in the family (B_M is 16/32/64/128/256) and every +# declared m_align (1 / B_M / 2*B_M), with aligned and unaligned M on both sides. +_ALIGN_MS = [1, 17, 48, 64, 96, 127, 128, 129, 200, 255, 256, 512] +_ALIGN_G = 2 +# (N, K) candidates: the second entry serves the k1024-only pipeline kids. +_ALIGN_SHAPES = [(1024, 4096), (1024, 1024)] +_ALIGN_ERR_TOL = 0.003 # e8m0 quant floor is ~0.0014; same gate the tuner uses + + +def _align_kids(): + # Imported here so the sweep and the many scripts reusing the helpers above + # never need the codegen package on sys.path. + from csrc.opus_gemm.opus_gemm_common import a8w8_mxscale_bmm_kernel_lists + + return { + int(kid): inst + for fam in a8w8_mxscale_bmm_kernel_lists + for kid, inst in fam.items() + } + + +_ALIGN_INPUTS = {} + + +def _align_inputs(m, n, k): + """Quantized inputs + fp32 reference for one shape, shared across kids.""" + key = (m, n, k) + if key not in _ALIGN_INPUTS: + g = _ALIGN_G + O_mx, xs_mx, xs_fp32 = _quant_per_token_e8m0(_block_varied((g, m, k), k)) + W_mx, ws_mx, ws_fp32 = _quant_block_e8m0(_block_varied((g, n, k), k)) + ref = run_torch(O_mx, W_mx, xs_fp32, ws_fp32).transpose(0, 1) + _ALIGN_INPUTS[key] = ( + O_mx.transpose(0, 1), + W_mx, + xs_mx.transpose(0, 1), + ws_mx, + ref, + ) + return _ALIGN_INPUTS[key] + + +def _align_run(kid, m, n, k): + """Return (ok, rel_err). ok False means the launcher refused the shape.""" + O_in, W_mx, xs_in, ws_mx, ref = _align_inputs(m, n, k) + # NaN-filled so a row the kernel never writes shows up as nan, not as a + # plausible value that a mean error would dilute. + Y = torch.full((m, _ALIGN_G, n), float("nan"), dtype=dtypes.bf16) + try: + _opus_bmm_a8w8_mxscale_raw(O_in, W_mx, Y, xs_in, ws_mx, 1, kid) + torch.cuda.synchronize() + except RuntimeError: + # The launcher's AITER_CHECK on M surfaces here. Deliberately not a + # blanket except: a harness bug must fail loudly, not read as a refusal. + return False, 0.0 + d = (Y.to(dtypes.fp32) - ref).abs() + # Per-row, not global: one wrong row out of a long M barely moves the mean. + rows = d.flatten(1).mean(1) / (ref.abs().flatten(1).mean(1) + 1e-9) + return True, rows.max().item() + + +def _align_pick_shape(kid, inst): + """First (N, K) this kid accepts at an aligned M, or None if it accepts none.""" + for n, k in _ALIGN_SHAPES: + if n % inst.B_N or k % inst.B_K: + continue + if _align_run(kid, max(inst.m_align, inst.B_M), n, k)[0]: + return n, k + return None + + +def check_m_align(): + """Assert OpusGemmInstance.m_align matches what each mxscale BMM kid does. + + m_align says which M values a kid's launcher accepts (1 == it masks a partial + M tile). Both the runtime's padded-M lookup (aiter/ops/opus/bmm_op.py) and a + tuner's candidate filter act on it, so a wrong value is not merely cosmetic: + too strict hides the fastest kernel from tuning (kid326 lost ~9% at the DSV4 + wo_a decode shapes that way, while the runtime dispatched it at those very + M), too loose makes both propose a kid whose launcher throws. + + For every kid this checks the declaration against observed behaviour: at an M + the declaration accepts, the launch must succeed and match the dequantized + fp32 reference; at an M it rejects, the launch must raise. Kids are never + silently skipped -- an unrunnable kid is reported. + """ + kids = _align_kids() + failures, unrunnable = [], [] + + for kid, inst in sorted(kids.items()): + shape = _align_pick_shape(kid, inst) + if shape is None: + unrunnable.append(kid) + continue + n, k = shape + align = inst.m_align + for m in _ALIGN_MS: + ok, err = _align_run(kid, m, n, k) + if m % align == 0: + if not ok: + failures.append(f"kid {kid}: m_align={align} but M={m} rejected") + elif not (err <= _ALIGN_ERR_TOL): + failures.append( + f"kid {kid}: M={m} accepted but worst row rel err " + f"{err:.4f} > {_ALIGN_ERR_TOL}" + ) + elif ok: + failures.append( + f"kid {kid}: m_align={align} claims M={m} unusable, " + f"but it ran (worst row rel err {err:.4f}) -- m_align too strict" + ) + + assert not unrunnable, ( + f"kids that ran on no test shape: {unrunnable}; extend _ALIGN_SHAPES so " + f"the guard keeps covering them" + ) + assert not failures, "m_align disagrees with the launcher:\n " + "\n ".join( + failures + ) + return len(kids) + + +def main(): + if get_gfx() not in SUPPORTED_GFX: + aiter.logger.warning( + "opus mxscale flatmm BMM unsupported on %s; skipping", get_gfx() + ) + return + + parser = argparse.ArgumentParser( + formatter_class=argparse.RawTextHelpFormatter, + description="opus fp8 e8m0 mxscale flatmm split-K BMM test", + ) + parser.add_argument( + "-d", + "--dtype", + type=str, + nargs="*", + default=["bf16"], + choices=["bf16", "fp32"], + help="output dtype(s) to sweep (default: bf16)", + ) + parser.add_argument( + "-g", + "--groups", + type=int, + nargs="*", + default=[2, 8], + help="batch group counts to sweep (DSV4 wo_a G; default: 2)", + ) + parser.add_argument( + "-s", + "--mnk", + type=dtypes.str2tuple, + nargs="*", + default=[ + (1, 1024, 4096), + (16, 1024, 4096), + (128, 1024, 4096), + (256, 1024, 4096), + (512, 1024, 4096), + (8192, 1024, 4096), + (16384, 1024, 4096), + ], + help="(m,n,k) shapes to sweep", + ) + parser.add_argument( + "--check-m-align", + action="store_true", + help="run the every-kid m_align guard instead of the perf sweep", + ) + args = parser.parse_args() + + n_tilen_checks = check_tilen_column_map() + aiter.logger.info( + "tileN column mapping passed for %d kid/layout combinations", n_tilen_checks + ) + + if args.check_m_align: + try: + n_kids = check_m_align() + except AssertionError as exc: + aiter.logger.error("m_align guard FAILED: %s", exc) + sys.exit(1) + aiter.logger.info( + "m_align matches launcher behaviour for all %d mxscale BMM kids", n_kids + ) + return + + for dtype in args.dtype: + df = [] + df_bf = [] + for g, (m, n, k) in itertools.product(args.groups, args.mnk): + df.append(test_mxscale_bmm(g, m, n, k, dtype)) + df_bf.append(test_mxscale_bmm_batch_first(g, m, n, k, dtype)) + aiter.logger.info( + "opus mxscale flatmm BMM summary (dtype=%s):\n%s", + dtype, + pd.DataFrame(df).to_markdown(index=False), + ) + aiter.logger.info( + "opus mxscale flatmm BMM batch-first (batch-major) summary " + "(dtype=%s):\n%s", + dtype, + pd.DataFrame(df_bf).to_markdown(index=False), + ) + + +if __name__ == "__main__": + main() diff --git a/op_tests/tuning_tests/test_config_shape_collision.py b/op_tests/tuning_tests/test_config_shape_collision.py index 9cf3343ba3..cd742169f5 100644 --- a/op_tests/tuning_tests/test_config_shape_collision.py +++ b/op_tests/tuning_tests/test_config_shape_collision.py @@ -62,6 +62,10 @@ ), ("AITER_CONFIG_A8W8_BATCHED_GEMM", "a8w8_tuned_batched_gemm"), ("AITER_CONFIG_BF16_BATCHED_GEMM", "bf16_tuned_batched_gemm"), + ( + "AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE", + "batched_gemm_a8w8_blockscale_mxscale_tuned", + ), ("AITER_CONFIG_GEMM_BF16", "bf16_tuned_gemm"), ("AITER_CONFIG_FMOE", "tuned_fmoe"), ("AITER_CONFIG_GROUPED_FMOE", "tuned_grouped_fmoe"), @@ -202,6 +206,12 @@ def test_a8w8_batched(self): def test_bf16_batched(self): self._check_family("AITER_CONFIG_BF16_BATCHED_GEMM", "bf16_tuned_batched_gemm") + def test_batched_gemm_a8w8_blockscale_mxscale(self): + self._check_family( + "AITER_CONFIG_BATCHED_GEMM_A8W8_BLOCKSCALE_MXSCALE", + "batched_gemm_a8w8_blockscale_mxscale_tuned", + ) + def test_bf16(self): self._check_family("AITER_CONFIG_GEMM_BF16", "bf16_tuned_gemm")