Skip to content
Merged
1 change: 1 addition & 0 deletions aiter/configs/a8w8_blockscale_group32_tuned_gemm.csv
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
gfx,cu_num,M,N,K,libtype,kernelId,splitK,us,kernelName,tflops,bw,errRatio
1 change: 1 addition & 0 deletions aiter/configs/a8w8_blockscale_group32_untuned_gemm.csv
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
M,N,K
1,225 changes: 1,225 additions & 0 deletions aiter/configs/model_configs/dsv41_a8w8_blockscale_group32_tuned_gemm.csv

Large diffs are not rendered by default.

15 changes: 15 additions & 0 deletions aiter/jit/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,13 @@ def mp_lock(
f"{AITER_ROOT_DIR}/aiter/configs/a8w8_blockscale_tuned_gemm.csv",
)

# Native E8M0 group32 scales have a different operand contract from the
# FP32 128x128 blockscale family, so shape-identical rows must stay separate.
AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32 = os.getenv(
"AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32",
f"{AITER_ROOT_DIR}/aiter/configs/a8w8_blockscale_group32_tuned_gemm.csv",
)

AITER_CONFIG_FMOE = os.getenv(
"AITER_CONFIG_FMOE",
f"{AITER_ROOT_DIR}/aiter/configs/tuned_fmoe.csv",
Expand Down Expand Up @@ -243,6 +250,14 @@ def AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_FILE(self):
"a8w8_blockscale_tuned_gemm",
)

@property
def AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32_FILE(self):
return self.get_config_file(
"AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32",
AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32,
"a8w8_blockscale_group32_tuned_gemm",
)

@property
def AITER_CONFIG_FMOE_FILE(self):
return self.get_config_file(
Expand Down
162 changes: 114 additions & 48 deletions aiter/ops/gemm_op_a8w8.py
Original file line number Diff line number Diff line change
Expand Up @@ -803,8 +803,28 @@ def _blockscale_triton(
x_scale: Tensor,
w_scale: Tensor,
dtype: torch.dtype,
*,
group32: bool = False,
split_k: int | None = None,
) -> Tensor:
"""Run the triton kernel on CK's inputs: row-major x_scale, (N, K) weight, JIT per arch."""
"""Run Triton on unshuffled weights with the selected scale format."""
if group32:
from aiter.ops.triton.gemm.basic.gemm_a8w8_blockscale_group32 import (
gemm_a8w8_blockscale_group32,
)

assert WQ.ndim == w_scale.ndim == 2, "Expected matrix weights and scales"
group_n = 32 if w_scale.shape[0] == -(-WQ.shape[0] // 32) else 1
return gemm_a8w8_blockscale_group32(
XQ,
WQ,
x_scale,
w_scale,
dtype=dtype,
weight_group_rows=group_n,
split_k=split_k,
)

from aiter.ops.triton.gemm.basic.gemm_a8w8_blockscale import (
gemm_a8w8_blockscale as _gemm_a8w8_blockscale_triton,
)
Expand All @@ -821,77 +841,123 @@ def gemm_a8w8_blockscale_fake(
w_scale: Tensor,
dtype: torch.dtype = dtypes.bf16,
isBpreshuffled=False,
split_k: int | None = None,
) -> torch.Tensor:
m = XQ.shape[0]
n = WQ.shape[0]
Y = torch.empty(m, n, dtype=dtype, device=XQ.device)
return Y


@torch_compile_guard(gen_fake=gemm_a8w8_blockscale_fake)
@torch_compile_guard(mutates_args=[], gen_fake=gemm_a8w8_blockscale_fake)
def gemm_a8w8_blockscale(
XQ: Tensor,
WQ: Tensor,
x_scale: Tensor,
w_scale: Tensor,
dtype: torch.dtype = dtypes.bf16,
isBpreshuffled: bool = False,
split_k: int | None = None,
) -> torch.Tensor:
assert dtype in [
dtypes.bf16,
dtypes.fp16,
], f"Output {dtype=} is currently not supported in gemm_a8w8"
"""Blockscaled A8W8 GEMM with configuration-first backend dispatch.

Native E8M0 group32 scales (typed tensors or uint8 views) use
AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32;
FP32 128x128 scales use AITER_CONFIG_GEMM_A8W8_BLOCKSCALE. Both tables are
queried through get_CKGEMM_config before choosing a fallback. Group32
currently supports libtype="triton", also its default on a config miss.
Triton tile and split-K parameters come from its own tuning tables.
For native group32 operands, split_k optionally overrides the positive
partition count without changing configured backend selection.
"""
is_group32 = (
x_scale.dtype in (dtypes.fp8_e8m0, torch.uint8)
and w_scale.dtype in (dtypes.fp8_e8m0, torch.uint8)
and XQ.ndim == WQ.ndim == 2
and x_scale.shape == (XQ.shape[0], XQ.shape[1] // 32)
and w_scale.shape
in (
(WQ.shape[0], WQ.shape[1] // 32),
(-(-WQ.shape[0] // 32), WQ.shape[1] // 32),
)
)
# A malformed byte-scale layout must not fall through to FP32 CK dispatch.
assert (
is_group32 or x_scale.dtype == w_scale.dtype == dtypes.fp32
), "Expected E8M0 group32 scale shapes (typed or uint8), or FP32 128x128 scales"
assert (
split_k is None or is_group32
), "split_k override requires native group32 operands"
assert dtype in (dtypes.bf16, dtypes.fp16) or (
is_group32 and dtype == dtypes.fp32
), f"Output {dtype=} is currently not supported in gemm_a8w8"
m = XQ.shape[0]
n = WQ.shape[0]
k = XQ.shape[1]
Y = torch.empty(m, n, dtype=dtype, device=XQ.device)
if isBpreshuffled:
assert not is_group32, "Group32 FP8 GEMM requires native weights"
if get_gfx() in ["gfx950"] and m >= 16 and k >= 512 and dtype == dtypes.bf16:
Y = torch.empty(m, n, dtype=dtype, device=XQ.device)
return gfx950_a8w8_blockscale_ASM(XQ, WQ, x_scale, w_scale, Y)
else:
assert 0, "asm kernel only support B preshuffle and m >= 16"
else:
if not _hip_blockscale_supported():
return _blockscale_triton(XQ, WQ, x_scale, w_scale, dtype)
config = get_CKGEMM_config(
m, n, k, AITER_CONFIGS.AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_FILE

tuned_file = (
AITER_CONFIGS.AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32_FILE
if is_group32
else AITER_CONFIGS.AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_FILE
)
config = get_CKGEMM_config(m, n, k, tuned_file)
if config is not None:
libtype = config["libtype"]
if libtype == "triton":
return _blockscale_triton(
XQ, WQ, x_scale, w_scale, dtype, group32=is_group32, split_k=split_k
)
# CK/CKTile currently consume FP32 128x128 scales. A misconfigured
# group32 row must fail instead of reinterpreting its E8M0 bytes.
assert not is_group32, f"Unsupported libtype {libtype} for group32 GEMM"
splitK = int(config.get("splitK", 0))
kernelName = str(config.get("kernelName", ""))
Y = torch.empty(m, n, dtype=dtype, device=XQ.device)
if libtype == "ck":
return gemm_a8w8_blockscale_ck(
XQ,
WQ,
x_scale,
w_scale,
Y,
splitK=splitK,
kernelName=kernelName,
)
elif libtype == "cktile":
return gemm_a8w8_blockscale_cktile(
XQ,
WQ,
x_scale,
w_scale,
Y,
splitK=splitK,
kernelName=kernelName,
)
else:
assert 0, f"Unsupported libtype {libtype} for gemm_a8w8_blockscale"

if is_group32 or not _hip_blockscale_supported():
return _blockscale_triton(
XQ, WQ, x_scale, w_scale, dtype, group32=is_group32, split_k=split_k
)
if config is not None:
libtype = config["libtype"]
splitK = int(config.get("splitK", 0))
kernelName = str(config.get("kernelName", ""))
if libtype == "ck":
return gemm_a8w8_blockscale_ck(
XQ,
WQ,
x_scale,
w_scale,
Y,
splitK=splitK,
kernelName=kernelName,
)
elif libtype == "cktile":
return gemm_a8w8_blockscale_cktile(
XQ,
WQ,
x_scale,
w_scale,
Y,
splitK=splitK,
kernelName=kernelName,
)
else:
assert 0, f"Unsupported libtype {libtype} for gemm_a8w8_blockscale"
min_m = _BLOCKSCALE_TRITON_FALLBACK_MIN_M.get(get_gfx())
if min_m is not None and m >= min_m:
return _blockscale_triton(XQ, WQ, x_scale, w_scale, dtype)
try:
return gemm_a8w8_blockscale_ck(XQ, WQ, x_scale, w_scale, Y)
except RuntimeError as e:
raise RuntimeError(
f"gemm_a8w8_blockscale failed for shape M={m}, N={n}, K={k}, "
f"{dtype=}, config={config}: {e}"
) from e
min_m = _BLOCKSCALE_TRITON_FALLBACK_MIN_M.get(get_gfx())
if min_m is not None and m >= min_m:
return _blockscale_triton(XQ, WQ, x_scale, w_scale, dtype)
Y = torch.empty(m, n, dtype=dtype, device=XQ.device)
try:
return gemm_a8w8_blockscale_ck(XQ, WQ, x_scale, w_scale, Y)
except RuntimeError as e:
raise RuntimeError(
f"gemm_a8w8_blockscale failed for shape M={m}, N={n}, K={k}, "
f"{dtype=}, config={config}: {e}"
) from e


def flatmm_a8w8_blockscale_ASM(
Expand Down
Loading
Loading