Skip to content
Merged
184 changes: 109 additions & 75 deletions aiter/aot/flydsl/gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
way as runtime JIT config lookup.

Supported kernel families:
- ``flydsl_gemm2_*`` split-K HGEMM kernels
- ``flydsl_hgemm_*`` gfx950 A16W16 GEMM kernels
- ``flydsl_bpreshuflle_*`` a8w8 preshuffle GEMM kernels
- ``flydsl_bpreshuffle_8w_*`` gfx950 8-wave a8w8 ptpc GEMM kernels
- ``flydsl_bpreshuffle_wmma_*`` gfx1250 a8w8 ptpc GEMM kernels
Expand All @@ -38,7 +38,9 @@
import re
import sys
import time
from contextlib import nullcontext

import flydsl.compiler as flyc
import flydsl.expr as fx

from aiter.aot.flydsl.common import (
Expand All @@ -59,9 +61,17 @@
)
from aiter.ops.flydsl.gemm_kernels import (
SPLIT_K_SEMAPHORE_MAX_LEN,
get_flydsl_splitk_hgemm_kernel_params,
get_flydsl_hgemm_kernel_params,
)
from aiter.ops.flydsl.kernels.common import run_cached
from aiter.ops.flydsl.kernels.gemm_a16w16_gfx950 import (
GEMM_A16W16_DTYPE_BF16,
GEMM_A16W16_DTYPE_FP16,
GEMM_A16W16_DTYPE_FP32,
_dynamic_tensor_arg,
gemm_a16w16_gfx950,
make_gemm_a16w16_param_and_validate,
)
from aiter.ops.flydsl.kernels.hgemm_dispatch import compile_flydsl_hgemm_kernel
from aiter.ops.flydsl.kernels.preshuffle_gemm import compile_preshuffle_gemm
from aiter.ops.flydsl.mxfp8_128_bpreshuffle_gemm_gfx1250 import (
BLOCK_K as SCALE_BLOCK_SIZE,
Expand Down Expand Up @@ -195,8 +205,8 @@ def parse_csv(csv_path: str):
if params is not None:
params = dict(params)
params["kind"] = "ptpc_wmma"
elif kernel_name.startswith("flydsl_gemm"):
params = get_flydsl_splitk_hgemm_kernel_params(kernel_name)
elif kernel_name.startswith("flydsl_hgemm"):
params = get_flydsl_hgemm_kernel_params(kernel_name)
if params is not None:
params = dict(params)
params["kind"] = "hgemm"
Expand Down Expand Up @@ -236,6 +246,8 @@ def _torch_dtype_for_kernel(dtype_name: str):
"bf16": torch.bfloat16,
"f16": torch.float16,
"fp16": torch.float16,
"f32": torch.float32,
"fp32": torch.float32,
}
if dtype_name not in mapping:
raise ValueError(f"Unsupported torch dtype name for GEMM AOT: {dtype_name!r}")
Expand All @@ -260,87 +272,108 @@ def _compile_hgemm_to_cache(
k: int,
dtype: str,
out_dtype: str,
tile_m: int,
tile_n: int,
tile_k: int,
block_m: int,
block_n: int,
block_k: int,
stages: int,
split_k: int,
block_m_warps: int,
block_n_warps: int,
block_k_warps: int,
n_tile_repeat: int = 1,
persistent_n_tiles: int = 1,
waves_per_eu: int = 0,
b_to_lds_unroll: int = 0,
async_copy: bool,
b_to_lds: bool,
b_preshuffle: bool,
c_to_lds: bool,
m_waves: int,
n_waves: int,
k_waves: int,
group_m: int,
use_half_tile_interleaved: bool,
has_bias: bool,
target_gfx: str,
kernel_family: str = "hgemm",
has_bias: bool = False,
**kwargs,
):
del kwargs, out_dtype
del kwargs

import torch

dev = torch.device("cpu")
torch_dtype = _torch_dtype_for_kernel(dtype)
if target_gfx != "gfx950":
raise ValueError(
f"The FlyDSL A16W16 kernel only supports gfx950, got {target_gfx}"
)

out = torch.empty((m, n), device=dev, dtype=torch_dtype)
a = torch.empty((m, k), device=dev, dtype=torch_dtype)
b = torch.empty((n, k), device=dev, dtype=torch_dtype)
bias = torch.empty((n,), device=dev, dtype=torch_dtype)
semaphore = torch.zeros(
(SPLIT_K_SEMAPHORE_MAX_LEN,),
device=dev,
dtype=torch.int32,
torch_dtype = _torch_dtype_for_kernel(dtype)
torch_out_dtype = _torch_dtype_for_kernel(out_dtype)
if torch_out_dtype not in (torch_dtype, torch.float32):
raise ValueError(
f"Unsupported output dtype {out_dtype!r} for input dtype {dtype!r}"
)
in_dtype_id = (
GEMM_A16W16_DTYPE_FP16
if torch_dtype == torch.float16
else GEMM_A16W16_DTYPE_BF16
)
signal = torch.zeros(
(SPLIT_K_SEMAPHORE_MAX_LEN,),
device=dev,
dtype=torch.int32,
out_dtype_id = (
GEMM_A16W16_DTYPE_FP32 if torch_out_dtype == torch.float32 else in_dtype_id
)
stream = fx.Stream(0)
config = {
"in_dtype_id": in_dtype_id,
"out_dtype_id": out_dtype_id,
"block_m": block_m,
"block_n": block_n,
"block_k": block_k,
"stages": stages,
"split_k": split_k,
"m_waves": m_waves,
"n_waves": n_waves,
"k_waves": k_waves,
"group_m": group_m,
"use_half_tile_interleaved": use_half_tile_interleaved,
"a_is_transposed": False,
"b_is_transposed": True,
"has_bias": has_bias,
}
param = make_gemm_a16w16_param_and_validate(m, n, k, config)
if param is None:
raise ValueError(
f"Invalid FlyDSL A16W16 config for M={m}, N={n}, K={k}: {config}"
)

exe = compile_flydsl_hgemm_kernel(
dtype,
n,
k,
kernel_family=kernel_family,
tile_m=tile_m,
tile_n=tile_n,
tile_k=tile_k,
stages=stages,
split_k=split_k,
block_m_warps=block_m_warps,
block_n_warps=block_n_warps,
block_k_warps=block_k_warps,
n_tile_repeat=n_tile_repeat,
persistent_n_tiles=persistent_n_tiles,
waves_per_eu=waves_per_eu,
b_to_lds_unroll=b_to_lds_unroll,
async_copy=async_copy,
b_to_lds=b_to_lds,
b_preshuffle=b_preshuffle,
c_to_lds=c_to_lds,
has_bias=has_bias,
)
# FlyDSL JIT does not accept None for tensor slots; pass real buffers for
# optional bias and split-K sync tensors.
launch_bias = bias if has_bias else b
_compile_executable_to_cache(
exe,
_ptr_view_safe(out),
_ptr_view_safe(a),
_ptr_view_safe(b),
_ptr_view_safe(launch_bias),
m,
_ptr_view_safe(semaphore),
_ptr_view_safe(signal),
stream,
)
# Layout-dynamic arguments make this compile independent of M/N/K. Small
# real CPU tensors avoid materializing model-sized buffers during AOT.
with compile_only_env():
dev = torch.device("cpu")
representative_extent = 8
a = torch.empty((1, representative_extent), device=dev, dtype=torch_dtype)
b = torch.empty(
(representative_extent, representative_extent),
device=dev,
dtype=torch_dtype,
).t()
out = torch.empty((1, representative_extent), device=dev, dtype=torch_out_dtype)
bias = torch.empty((representative_extent,), device=dev, dtype=torch_dtype)
semaphore = torch.zeros(
(SPLIT_K_SEMAPHORE_MAX_LEN,), device=dev, dtype=torch.int32
)
signal = torch.zeros(
(SPLIT_K_SEMAPHORE_MAX_LEN,), device=dev, dtype=torch.int32
)
stream = fx.Stream(0)
a_arg = _dynamic_tensor_arg(a, 1)
b_arg = _dynamic_tensor_arg(b, 0)
out_arg = _dynamic_tensor_arg(out, 1)
bias_arg = a_arg if not has_bias else _dynamic_tensor_arg(bias, 0)
dispatch_args = (
out_arg,
a_arg,
b_arg,
bias_arg,
semaphore,
signal,
split_k,
param,
stream,
)
run_cached(
gemm_a16w16_gfx950,
*dispatch_args,
constexpr_param=param,
compiler=flyc.compile,
dispatch_args=dispatch_args,
)


def _compile_preshuffle_to_cache(
Expand Down Expand Up @@ -646,9 +679,10 @@ def compile_one_config(

t0 = time.time()
try:
tensor_context = nullcontext() if kind == "hgemm" else FakeTensorMode()
with (
override_env("FLYDSL_GPU_ARCH", aot_arch),
FakeTensorMode(),
tensor_context,
):
if kind == "hgemm":
hgemm_kwargs = dict(kwargs)
Expand Down
10 changes: 5 additions & 5 deletions aiter/configs/bf16_tuned_gemm.csv
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ gfx950,256,1,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,13,5.
gfx950,256,1,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,6,12.8405,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,1.96,1961.15
gfx950,256,1,3072,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,5,13.6048,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,2.77,2776.02
gfx950,256,1,50016,6144,False,torch.bfloat16,torch.float32,False,False,opus,5008,0,122.4478,opus_gemm_512x128x256x64_2x4_16x16x32_0x0x0_4g_safe,0.0,5.02,5020.17
gfx950,256,1,6144,2048,False,torch.bfloat16,torch.float32,False,False,flydsl,3161,2,10.8611,flydsl_gemm5_abf16_wbf16_bf16_t16x64x128_split_k2_block_m_warp1_block_n_warp2_block_k_warp1_async_copyTrue_b_to_ldsTrue_b_preshuffleFalse_c_to_ldsFalse_gfx950,0.0438,2.32,2318.57
gfx950,256,1,6144,2048,False,torch.bfloat16,torch.float32,False,False,flydsl,87,1,6.9521,flydsl_hgemm_abf16_wbf16_fp32_t16x16x128x5_ks1_w1x1x1_bias0_ktail0_gm4_pft_gfx950,0.0,3.62,3622.24
gfx950,256,1,6144,3072,False,torch.bfloat16,torch.float32,False,False,asm,1,2,13.5868,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,2.78,2779.7
gfx950,256,2,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,16,5.2543,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,0.6,304.12
gfx950,256,2,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,6,12.8459,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,3.92,1961.61
Expand All @@ -25,20 +25,20 @@ gfx950,256,4,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,3,13,5.
gfx950,256,4,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,5,12.9009,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,7.8,1955.78
gfx950,256,4,3072,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,5,13.5136,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,11.17,2798.84
gfx950,256,4,50016,6144,False,torch.bfloat16,torch.float32,False,False,opus,2008,0,123.6511,opus_gemm_512x128x256x64_2x4_16x16x32_0x0x0_cA1cB17,0.0,19.88,4974.04
gfx950,256,4,6144,2048,False,torch.bfloat16,torch.float32,False,False,flydsl,4265,2,10.708,flydsl_gemm8_abf16_wbf16_bf16_t16x64x64_split_k2_block_m_warp1_block_n_warp2_block_k_warp1_async_copyTrue_b_to_ldsTrue_b_preshuffleFalse_c_to_ldsFalse_gfx950,0.0444,9.4,2356.31
gfx950,256,4,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,7384,2,13.0193,flydsl_gemm8_abf16_wbf16_bf16_t16x64x64_split_k2_block_m_warp1_block_n_warp2_block_k_warp1_async_copyTrue_b_to_ldsTrue_b_preshuffleFalse_c_to_ldsFalse_gfx950,0.0486,11.6,2905.11
gfx950,256,4,6144,2048,False,torch.bfloat16,torch.float32,False,False,flydsl,91,1,7.2059,flydsl_hgemm_abf16_wbf16_fp32_t16x16x128x6_ks1_w1x1x1_bias0_ktail0_gm4_pft_gfx950,0.0,13.97,3501.49
gfx950,256,4,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,251,1,8.6647,flydsl_hgemm_abf16_wbf16_fp32_t16x32x128x6_ks1_w1x2x1_bias0_ktail0_gm0_pft_gfx950,0.0,17.43,4365.12
gfx950,256,8,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,16,5.25,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,2.4,318.71
gfx950,256,8,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,6,12.8808,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,15.63,1963.92
gfx950,256,8,3072,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,5,13.7111,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,22.03,2763.91
gfx950,256,8,50016,6144,False,torch.bfloat16,torch.float32,False,False,opus,2008,0,124.408,opus_gemm_512x128x256x64_2x4_16x16x32_0x0x0_cA1cB17,0.0,39.52,4947.39
gfx950,256,8,6144,2048,False,torch.bfloat16,torch.float32,False,False,asm,1,2,10.8198,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,18.61,2338.02
gfx950,256,8,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,6292,2,13.2449,flydsl_gemm6_abf16_wbf16_bf16_t16x64x128_split_k2_block_m_warp1_block_n_warp2_block_k_warp1_async_copyTrue_b_to_ldsTrue_b_preshuffleFalse_c_to_ldsFalse_gfx950,0.0485,22.8,2861.19
gfx950,256,8,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,241,1,8.6517,flydsl_hgemm_abf16_wbf16_fp32_t16x32x128x5_ks1_w1x2x1_bias0_ktail0_gm0_pft_gfx950,0.0,34.91,4380.2
gfx950,256,16,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,13,5.3721,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,4.68,330.14
gfx950,256,16,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,3,6,13.0667,_ZN5aiter39bf16gemm_fp32bf16_tn_48x64_splitk_cleanE,0.0,30.82,1946.01
gfx950,256,16,3072,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,5,13.7282,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,44.0,2771.2
gfx950,256,16,50016,6144,False,torch.bfloat16,torch.float32,False,False,opus,204,2,131.0394,opus_gemm_flatmm_splitk_256x32x128x64_2x1_16x16x32_0x0x0_wgpcu2,0.0,75.04,4703.88
gfx950,256,16,6144,2048,False,torch.bfloat16,torch.float32,False,False,asm,1,2,10.987,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,36.65,2314.37
gfx950,256,16,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,6292,2,13.9957,flydsl_gemm6_abf16_wbf16_bf16_t16x64x128_split_k2_block_m_warp1_block_n_warp2_block_k_warp1_async_copyTrue_b_to_ldsTrue_b_preshuffleFalse_c_to_ldsFalse_gfx950,0.049,43.15,2718.24
gfx950,256,16,6144,3072,False,torch.bfloat16,torch.float32,False,False,flydsl,270,1,9.2768,flydsl_hgemm_abf16_wbf16_fp32_t16x32x128x8_ks1_w1x2x1_bias0_ktail0_gm0_pft_gfx950,0.0,65.11,4100.95
gfx950,256,32,128,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,16,5.6787,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,8.86,347.66
gfx950,256,32,2048,6144,False,torch.bfloat16,torch.float32,False,False,asm,1,6,13.2436,_ZN5aiter39bf16gemm_fp32bf16_tn_32x64_splitk_cleanE,0.0,60.81,1939.81
gfx950,256,32,3072,6144,False,torch.bfloat16,torch.float32,False,False,asm,3,5,14.302,_ZN5aiter39bf16gemm_fp32bf16_tn_48x64_splitk_cleanE,0.0,84.46,2680.64
Expand Down
Loading
Loading