Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
157 changes: 113 additions & 44 deletions aiter/aot/flydsl/gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,10 @@
way as runtime JIT config lookup.

Supported kernel families:
- ``flydsl_gemm2_*`` split-K HGEMM kernels
- ``flydsl_bpreshuflle_*`` a8w8 preshuffle GEMM kernels
- ``flydsl_gemm2_*`` split-K HGEMM kernels
- ``flydsl_bpreshuflle_*`` a8w8 preshuffle GEMM kernels
- ``flydsl_bpreshuffle_wmma_*`` gfx1250 a8w8 ptpc GEMM kernels
- ``flydsl_blockscale_bpreshuffle_wmma_*`` gfx1250 a8w8 blockscale GEMM kernels

Usage:
# Compile all unique FlyDSL GEMM kernels from default CSVs
Expand All @@ -34,7 +36,6 @@
import re
import sys
import time
from typing import Dict, Optional

import flydsl.expr as fx

Expand All @@ -47,7 +48,12 @@
run_jobs_parallel,
)
from aiter.jit.core import AITER_CONFIGS
from aiter.ops.flydsl.blockscale_bpreshuffle_gemm_gfx1250 import parse_wmma_kernel_name
from aiter.ops.flydsl.blockscale_bpreshuffle_gemm_gfx1250 import (
parse_wmma_kernel_name as parse_blockscale_wmma_kernel_name,
)
from aiter.ops.flydsl.bpreshuffle_gemm_gfx1250 import (
parse_wmma_kernel_name as parse_ptpc_wmma_kernel_name,
)
from aiter.ops.flydsl.gemm_kernels import (
SPLIT_K_SEMAPHORE_MAX_LEN,
get_flydsl_splitk_hgemm_kernel_params,
Expand Down Expand Up @@ -83,7 +89,7 @@
}


def _parse_bool(value: Optional[str]) -> bool:
def _parse_bool(value: str | None) -> bool:
if value is None:
return False
normalized = value.strip().lower()
Expand All @@ -96,7 +102,7 @@ def _parse_bool(value: Optional[str]) -> bool:
raise ValueError(f"Expected True/False, got {value!r}")


def _parse_preshuffle_kernel_name(name: str) -> Optional[Dict]:
def _parse_preshuffle_kernel_name(name: str) -> dict | None:
m = _PRESHUFFLE_RE.fullmatch(name)
if m is None:
return None
Expand Down Expand Up @@ -147,10 +153,15 @@ def parse_csv(csv_path: str):
if kernel_name.startswith("flydsl_bpreshuflle_"):
params = _parse_preshuffle_kernel_name(kernel_name)
elif kernel_name.startswith("flydsl_blockscale_bpreshuffle_wmma_"):
params = parse_wmma_kernel_name(kernel_name)
params = parse_blockscale_wmma_kernel_name(kernel_name)
if params is not None:
params = dict(params)
params["kind"] = "blockscale_wmma"
elif kernel_name.startswith("flydsl_bpreshuffle_wmma_"):
params = parse_ptpc_wmma_kernel_name(kernel_name)
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)
if params is not None:
Expand Down Expand Up @@ -379,12 +390,13 @@ def _compile_blockscale_wmma_to_cache(
cluster_n: int,
**kwargs,
):
del kwargs
del kwargs, split_k

import torch

from aiter.ops.flydsl.kernels.gemm_fp8fp4_gfx1250 import compile_blockscale_gemm
from aiter.ops.flydsl.kernels.tensor_shim import _run_compiled
from aiter.ops.flydsl.kernels.gemm_a8w8_blockscale_gfx1250 import (
launch_gemm_a8w8_bsc_col,
)

dev = torch.device("cpu")
k_blocks = (k + 127) // 128
Expand All @@ -395,54 +407,101 @@ def _compile_blockscale_wmma_to_cache(
out = torch.empty((m, n), device=dev, dtype=torch.bfloat16)
stream = fx.Stream(0)

exe = compile_blockscale_gemm(
N=n,
K=k,
data_format="fp8",
tile_m=tile_m,
tile_n=tile_n,
tile_k=tile_k,
m_warp=m_warp,
n_warp=n_warp,
num_buffers=num_buffers,
waves_per_eu=None,
cluster_m=cluster_m,
cluster_n=cluster_n,
out_dtype="bf16",
split_k=split_k,
scale_block_k=128,
scale_block_n=128,
ascale_layout="col_major",
)
with compile_only_env():
_run_compiled(
exe,
launch_gemm_a8w8_bsc_col(
out,
xq.view(torch.uint8),
wq.view(torch.uint8),
_ptr_view_safe(xq),
_ptr_view_safe(wq),
a_scale.view(torch.uint8),
b_scale.view(torch.uint8),
m,
stream,
n,
k,
a_scale.numel() // a_scale.stride(0),
xq.stride(0),
out.stride(0),
a_scale.stride(1),
a_scale.numel() // a_scale.stride(0),
tile_m,
tile_n,
tile_k,
m_warp,
n_warp,
0,
num_buffers,
cluster_m,
cluster_n,
)


def _compile_ptpc_wmma_to_cache(
*,
m: int,
n: int,
k: int,
tile_m: int,
tile_n: int,
tile_k: int,
m_warp: int,
n_warp: int,
num_buffers: int,
split_k: int,
cluster_m: int,
cluster_n: int,
**kwargs,
):
del kwargs, split_k

import torch

from aiter.ops.flydsl.kernels.gemm_a8w8_ptpc_gfx1250 import launch_gemm_a8w8_ptpc

dev = torch.device("cpu")
xq = torch.empty((m, k), device=dev, dtype=torch.uint8)
wq = torch.empty((n, k), device=dev, dtype=torch.uint8)
scale_a = torch.empty((max(m, 1),), device=dev, dtype=torch.float32)
scale_b = torch.empty((max(n, 1),), device=dev, dtype=torch.float32)
out = torch.empty((m, n), device=dev, dtype=torch.bfloat16)
stream = fx.Stream(0)

with compile_only_env():
launch_gemm_a8w8_ptpc(
out,
_ptr_view_safe(xq),
_ptr_view_safe(wq),
scale_a,
scale_b,
m,
stream,
n,
k,
xq.stride(0),
out.stride(0),
tile_m,
tile_n,
tile_k,
m_warp,
n_warp,
0,
num_buffers,
cluster_m,
cluster_n,
)


def job_arch(kind: str, cu_num: int = 0) -> str:
"""Target arch a job would compile for -- shared by dispatch and ARCH filtering."""
if kind in ("blockscale_wmma", "ptpc_wmma"):
return "gfx1250"
return cu_num_to_arch(cu_num, default=GEMM_AOT_ARCH_DEFAULT)


def compile_one_config(
kernel_name: str, kind: str, m: int, n: int, k: int, cu_num: int = 0, **kwargs
) -> dict:
"""Compile one GEMM kernel configuration and save it to cache."""
from torch._subclasses.fake_tensor import FakeTensorMode

aot_arch = (
"gfx1250"
if kind == "blockscale_wmma"
else cu_num_to_arch(cu_num, default=GEMM_AOT_ARCH_DEFAULT)
)
aot_arch = job_arch(kind, cu_num)
shape_str = f"{kernel_name} M={m} N={n} K={k}"
result = {
"kernel_name": kernel_name,
Expand All @@ -466,6 +525,8 @@ def compile_one_config(
_compile_preshuffle_to_cache(m=m, n=n, k=k, **kwargs)
elif kind == "blockscale_wmma":
_compile_blockscale_wmma_to_cache(m=m, n=n, k=k, **kwargs)
elif kind == "ptpc_wmma":
_compile_ptpc_wmma_to_cache(m=m, n=n, k=k, **kwargs)
else:
raise ValueError(f"Unknown GEMM AOT kind: {kind}")

Expand Down Expand Up @@ -501,13 +562,20 @@ def main():
cache_dir = os.path.expanduser(
os.environ.get("FLYDSL_RUNTIME_CACHE_DIR", "~/.flydsl/cache")
)
arch = os.environ.get("ARCH") or os.environ.get("GPU_ARCHS") or "(auto-detect)"
arch = os.environ.get("ARCH") or os.environ.get("GPU_ARCHS")

all_jobs = collect_aot_jobs(csv_paths, parse_csv)
if arch:
# GPU_ARCHS may be a ';'- or ','-separated list (e.g. "gfx942;gfx950").
arch_set = {a.strip() for a in re.split(r"[;,]", arch) if a.strip()}
n_before = len(all_jobs)
all_jobs = [j for j in all_jobs if job_arch(j["kind"], j["cu_num"]) in arch_set]
print(f"[aiter] ARCH={arch}: {len(all_jobs)}/{n_before} jobs match")
Comment thread
aoli26 marked this conversation as resolved.

hgemm_jobs = [j for j in all_jobs if j["kind"] == "hgemm"]
preshuffle_jobs = [j for j in all_jobs if j["kind"] == "preshuffle"]
blockscale_wmma_jobs = [j for j in all_jobs if j["kind"] == "blockscale_wmma"]
ptpc_wmma_jobs = [j for j in all_jobs if j["kind"] == "ptpc_wmma"]

print("=" * 72)
print("FlyDSL GEMM AOT Pre-compilation")
Expand All @@ -517,10 +585,10 @@ def main():
print(f" HGEMM jobs: {len(hgemm_jobs)}")
print(f" Preshuffle jobs: {len(preshuffle_jobs)}")
print(f" Blockscale wmma jobs: {len(blockscale_wmma_jobs)}")
print(f" PTPC wmma jobs: {len(ptpc_wmma_jobs)}")
print(f" Total jobs: {len(all_jobs)}")
print(" Compile arch: (from cu_num)")
print(f" Cache dir: {cache_dir}")
print(f" Target arch: {arch}")
print(f" Target arch: {arch or '(all archs found in CSVs)'}")
print("=" * 72)

total_t0 = time.time()
Expand All @@ -529,7 +597,8 @@ def main():
# separate serial passes per kind.
print(f"\n--- Compiling {len(all_jobs)} kernels ---")
results = run_jobs_parallel(
compile_one_config, hgemm_jobs + preshuffle_jobs + blockscale_wmma_jobs
compile_one_config,
hgemm_jobs + preshuffle_jobs + blockscale_wmma_jobs + ptpc_wmma_jobs,
)

total_elapsed = time.time() - total_t0
Expand Down
Loading
Loading