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
17 changes: 17 additions & 0 deletions python/sglang/kernels/ops/gemm/gfx95_batched_gemm_bf16_fp8_grid.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,10 @@
_CACHE_MODIFIER = ".cg"
# 8 splits of 256-wide K steps fill the machine up to this many rows
_SPLIT_K, _SPLIT_K_BLOCK_K, _SPLIT_K_MAX_M = 8, 256, 64
# T > _SPLIT_K_MAX_M: "hipblaslt" (strided-batched bf16 + fake-quant, default) or "triton" (the grid kernel)
import os as _os

_LARGE_M_ROUTE = _os.environ.get("SGLANG_OPT_WO_A_LARGE_M_ROUTE", "hipblaslt")


@triton.jit
Expand Down Expand Up @@ -317,6 +321,19 @@ def batched_gemm_bf16_fp8_grid(
assert _split_k_applies(T, D, R), (T, D, R)
_batched_gemm_split_k(x, w, out, fp8_grid, eps)
return out
if _LARGE_M_ROUTE == "hipblaslt":
# Above the split-K ceiling the 16x32 Triton tile re-streams the whole per-group
# weight once per 16 rows (~800 MB at 768 rows): 23-126 us at T = 96-768. hipBLASLt's
# strided-batched bf16 GEMM, written directly into the [T, G, R] layout, is 17-31 us
# over the same range (job 34984, cold weights), and the consumer's fp8-grid rounding
# is the same per-32 fake-quant applied once afterwards. Deterministic per shape.
y = out.view(T, G, R)
torch.bmm(x.transpose(0, 1), w.transpose(1, 2), out=y.transpose(0, 1))
if fp8_grid:
from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import fake_quant_fp8_activation

out.copy_(fake_quant_fp8_activation(out))
return out
grid = (G, triton.cdiv(T, _BLOCK_M) * triton.cdiv(R, _BLOCK_N))
_batched_gemm_bf16_fp8_grid_kernel[grid](
x,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,25 +33,30 @@
"gfx950:1152:5120:4096": "hipblaslt_bf16",
"gfx950:1152:5120:64": "hipblaslt_bf16",
"gfx950:1152:5120:8192": "hipblaslt_bf16",
"gfx950:1792:5120:1024": "ds:128,64,256,4,1",
"gfx950:1792:5120:4096": "ds:128,64,256,4,1",
"gfx950:1792:5120:512": "ds:64,64,256,4,1",
"gfx950:1856:5120:1024": "ds:128,64,256,4,1",
"gfx950:1856:5120:128": "hipblaslt_bf16",
"gfx950:1856:5120:16384": "hipblaslt_bf16",
"gfx950:1856:5120:256": "hipblaslt_bf16",
"gfx950:1856:5120:4096": "ds:256,128,128,8,1",
"gfx950:1856:5120:64": "hipblaslt_bf16",
"gfx950:1856:5120:8192": "hipblaslt_bf16",
"gfx950:5120:2048:1024": "hipblaslt_bf16",
"gfx950:5120:2048:128": "hipblaslt_bf16",
"gfx950:5120:2048:1024": "ds:128,128,256,8,1",
"gfx950:5120:2048:128": "ds:64,64,256,4,1",
"gfx950:5120:2048:16384": "hipblaslt_bf16",
"gfx950:5120:2048:256": "hipblaslt_bf16",
"gfx950:5120:2048:256": "ds:64,64,256,4,1",
"gfx950:5120:2048:4096": "ds:128,128,128,4,1",
"gfx950:5120:2048:512": "ds:128,64,256,4,1",
"gfx950:5120:2048:64": "ds:64,64,256,4,1",
"gfx950:5120:2048:8192": "ds:128,256,128,8,1",
"gfx950:8192:1280:1024": "ds:256,128,128,8,1",
"gfx950:8192:1280:1024": "ds:128,128,128,8,1",
"gfx950:8192:1280:128": "ds:64,64,256,4,1",
"gfx950:8192:1280:16384": "ds:128,256,128,8,1",
"gfx950:8192:1280:256": "ds:128,64,256,4,1",
"gfx950:8192:1280:4096": "hipblaslt_bf16",
"gfx950:8192:1280:512": "ds:128,128,256,8,1",
"gfx950:8192:1280:64": "ds:64,64,256,4,1",
"gfx950:8192:1280:8192": "hipblaslt_bf16"
},
Expand All @@ -63,25 +68,30 @@
"gfx950:1152:5120:4096": "ds:128,256,128,8,1",
"gfx950:1152:5120:64": "hipblaslt_bf16",
"gfx950:1152:5120:8192": "hipblaslt_bf16",
"gfx950:1792:5120:1024": "ds:128,64,256,4,1",
"gfx950:1792:5120:4096": "ds:128,64,256,4,1",
"gfx950:1792:5120:512": "ds:64,64,256,4,1",
"gfx950:1856:5120:1024": "ds:128,64,256,4,1",
"gfx950:1856:5120:128": "hipblaslt_bf16",
"gfx950:1856:5120:16384": "ds:128,256,128,8,1",
"gfx950:1856:5120:256": "ds:128,64,256,4,4",
"gfx950:1856:5120:4096": "ds:256,128,128,8,1",
"gfx950:1856:5120:64": "hipblaslt_bf16",
"gfx950:1856:5120:8192": "ds:256,128,128,8,1",
"gfx950:5120:2048:1024": "ds:256,128,128,8,1",
"gfx950:5120:2048:1024": "ds:128,128,256,8,1",
"gfx950:5120:2048:128": "ds:64,64,256,4,1",
"gfx950:5120:2048:16384": "ds:128,256,128,8,1",
"gfx950:5120:2048:256": "ds:128,64,256,4,1",
"gfx950:5120:2048:4096": "ds:128,128,128,4,1",
"gfx950:5120:2048:512": "ds:128,64,256,4,1",
"gfx950:5120:2048:64": "ds:64,64,256,4,1",
"gfx950:5120:2048:8192": "ds:128,256,128,8,1",
"gfx950:8192:1280:1024": "ds:256,128,128,8,1",
"gfx950:8192:1280:1024": "ds:128,128,128,8,1",
"gfx950:8192:1280:128": "ds:64,64,256,4,1",
"gfx950:8192:1280:16384": "ds:128,256,128,8,1",
"gfx950:8192:1280:256": "ds:128,64,256,4,1",
"gfx950:8192:1280:4096": "hipblaslt_bf16",
"gfx950:8192:1280:512": "ds:128,128,256,8,1",
"gfx950:8192:1280:64": "ds:64,64,256,4,1",
"gfx950:8192:1280:8192": "ds:128,256,128,8,1"
},
Expand All @@ -92,4 +102,4 @@
"large_m_picked_by": "hipblaslt_bf16 unless the dot_scaled tile plus the Triton activation quant of a bf16 input is faster (HIP-graph replay)",
"picked_by": "fastest HIP-graph replay time with the bf16 (fp8-grid) activation input"
}
}
}
31 changes: 23 additions & 8 deletions python/sglang/kernels/ops/quantization/mxfp8_native_amd_gfx95.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,7 +189,7 @@ def mxfp8_gemv(


# M > 32: hipBLASLt bf16 on the bf16 copy or the Triton dot_scaled tile, per the table's "large_m"
LARGE_M_BUCKETS = (64, 128, 256, 1024, 4096, 8192, 16384)
LARGE_M_BUCKETS = (64, 128, 256, 512, 1024, 4096, 8192, 16384)
HIPBLASLT_BF16 = "hipblaslt_bf16"


Expand Down Expand Up @@ -281,6 +281,8 @@ def _mxfp8_shuffled_gemm_kernel(
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
OUT_F32: tl.constexpr,
EVEN_M: tl.constexpr = False,
EVEN_N: tl.constexpr = False,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
Expand Down Expand Up @@ -313,13 +315,23 @@ def _mxfp8_shuffled_gemm_kernel(

acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for _ in range(0, k_per_split // BLOCK_K):
x = tl.load(x_ptrs, mask=m_mask[:, None], other=0)
w_raw = tl.load(w_ptrs, mask=blk_mask[:, None], other=0) # [T * S, 2048]
# The masks are loop-invariant: when the tile divides the problem they are all-true, and Triton
# still emits a v_cndmask per element per iteration for them (26 of them in a loop with 8 MFMA).
if EVEN_M:
x = tl.load(x_ptrs)
xs = tl.load(xs_ptrs)
else:
x = tl.load(x_ptrs, mask=m_mask[:, None], other=0)
xs = tl.load(xs_ptrs, mask=m_mask[:, None], other=127)
if EVEN_N:
w_raw = tl.load(w_ptrs) # [T * S, 2048]
ws = tl.load(ws_ptrs)
else:
w_raw = tl.load(w_ptrs, mask=blk_mask[:, None], other=0)
ws = tl.load(ws_ptrs, mask=n_mask[:, None], other=127)
w7 = tl.reshape(w_raw, (T, S, 2, 2, 16, 2, 16)) # (t, s, g1, g0, row, half, e)
w7 = tl.permute(w7, (0, 4, 1, 5, 2, 3, 6)) # (t, row, s, half, g1, g0, e)
w = tl.reshape(w7, (BLOCK_N, BLOCK_K))
xs = tl.load(xs_ptrs, mask=m_mask[:, None], other=127)
ws = tl.load(ws_ptrs, mask=n_mask[:, None], other=127)
acc = tl.dot_scaled(x, xs, "e4m3", w.T, ws, "e4m3", acc)
x_ptrs += BLOCK_K
xs_ptrs += BLOCK_K // 32
Expand All @@ -332,10 +344,11 @@ def _mxfp8_shuffled_gemm_kernel(
+ offs_m[:, None] * stride_om
+ offs_n[None, :]
)
if OUT_F32:
tl.store(o_ptrs, acc, mask=m_mask[:, None] & n_mask[None, :])
o_val = acc if OUT_F32 else acc.to(tl.bfloat16)
if EVEN_M and EVEN_N:
tl.store(o_ptrs, o_val)
else:
tl.store(o_ptrs, acc.to(tl.bfloat16), mask=m_mask[:, None] & n_mask[None, :])
tl.store(o_ptrs, o_val, mask=m_mask[:, None] & n_mask[None, :])


def mxfp8_shuffled_gemm(
Expand Down Expand Up @@ -376,6 +389,8 @@ def mxfp8_shuffled_gemm(
BLOCK_N=bn,
BLOCK_K=bk,
OUT_F32=split_k > 1,
EVEN_M=(m % bm == 0),
EVEN_N=(n % bn == 0),
num_warps=warps,
num_stages=2,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -40,8 +40,12 @@

# aiter batched GEMM fork with wo_b's fp8-grid rounding in its epilogue; None keeps the aiter kernel
_wo_a_fp8_grid_gemm = None
import os as _os

_LARGE_M_EMIT = _os.environ.get("SGLANG_OPT_WO_A_LARGE_M_EMIT", "grid")
if _use_aiter and _is_gfx95_supported and envs.SGLANG_OPT_USE_AITER_BATCHED_GEMM.get():
from sglang.kernels.ops.gemm.gfx95_batched_gemm_bf16_fp8_grid import (
_split_k_applies,
batched_gemm_bf16_fp8_grid as _wo_a_fp8_grid_gemm,
)

Expand Down Expand Up @@ -216,6 +220,23 @@ def wo_a_fp8_grid_matmul(o: torch.Tensor, wo_a: torch.Tensor, fp8_grid: bool):
# the split-K regime ends at 64 rows, so a request's verify and decode rows would
# take different reduction orders; deterministic inference keeps the single chain
split_k = False if get_exec().deterministic.enable_deterministic_inference else None
if (
fp8_grid
and _LARGE_M_EMIT == "mx"
and split_k is not False
and not _split_k_applies(o.shape[0], o.shape[2], wo_a.shape[1])
):
# Above the split-K ceiling the GEMM runs on hipBLASLt and the fp8-grid rounding is a
# separate pass; wo_b's dot_scaled route would quantize that rounded tensor again
# (idempotent). Quantize the raw result once and hand wo_b the (q, scale) pair instead:
# one pass fewer, same values (job in night/woa_mx.sbatch).
from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import (
Mxfp8Activation,
mxfp8_e4m3_quantize,
)

y = _wo_a_fp8_grid_gemm(o, wo_a, fp8_grid=False, split_k=split_k)
return Mxfp8Activation(*mxfp8_e4m3_quantize(y))
y = _wo_a_fp8_grid_gemm(o, wo_a, fp8_grid=fp8_grid, split_k=split_k)
if fp8_grid:
return Fp8GridActivation(y)
Expand Down
Loading