Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
21 commits
Select commit Hold shift + click to select a range
2ffe0a8
feat: add flashinfer's rmsnorm + quant fusion support
Jul 31, 2026
541692b
test: loosen TestGemmaRMSNorm bf16 tolerance to 1e-2
Jul 31, 2026
a04437f
kernel: avoid repeat scalars for _scaled_mm_fp8 cutlass kernel
Jul 31, 2026
367ace8
kernel: scalar scale A support for SM100+
Jul 31, 2026
cdbfb2f
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Jul 31, 2026
f5a1470
fix: merge down fixes
Jul 31, 2026
06a2deb
Merge branch 'main' into 'dev/dlal/norm-quant-fusion'
Aug 1, 2026
c0307ba
review: remove server arg for enabling norm quant fusion
Aug 2, 2026
d9de615
review: add registered tests for fp8 scalar change and norm quant fusion
Aug 2, 2026
4df0fcd
revert: 541692b293e5787d4807595d17db442b91a703ab
Aug 2, 2026
cea4974
review: update registered tests estimated time
Aug 2, 2026
52e61e7
Merge branch 'main' into dev/dlal/norm-quant-fusion
Aug 2, 2026
0a6d274
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 2, 2026
3d14448
Merge branch 'main' into dev/dlal/norm-quant-fusion
BBuf Aug 2, 2026
7a23a40
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 2, 2026
7f12b9f
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 2, 2026
da5c268
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 3, 2026
8de8e53
Merge branch 'main' into dev/dlal/norm-quant-fusion
BBuf Aug 3, 2026
03ec037
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 3, 2026
73c8da9
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 3, 2026
abd5860
Merge branch 'main' into dev/dlal/norm-quant-fusion
DevashishLal-CB Aug 3, 2026
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
185 changes: 185 additions & 0 deletions benchmark/kernels/bench_fused_rmsnorm_fp8_quant.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
"""Microbenchmark: fused RMSNorm + static per-tensor FP8 quant, comparing the
flashinfer default kernels against the CuTe-DSL kernels and the unfused
baseline (RMSNorm followed by a separate static FP8 quant).

Providers:
unfused RMSNorm.forward_cuda + static_quant_fp8
fused flashinfer rmsnorm_quant / fused_add_rmsnorm_quant (default)
fused_cute flashinfer rmsnorm_quant_cute / fused_add_rmsnorm_quant_cute

All fused providers produce an ``(fp8, scale)`` activation (and updated residual
when a residual is supplied), matching what a downstream FP8 static-per-tensor
linear consumes. Covers the no-residual and residual (fused-add) cases across a
few hidden sizes so you can pick the fastest kernel per shape.

Run:
python benchmark/kernels/bench_fused_rmsnorm_fp8_quant.py
"""

import itertools

import numpy as np
import torch
import triton
from flashinfer.norm import fused_add_rmsnorm_quant, rmsnorm_quant
from flashinfer.testing import bench_gpu_time

from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8
from sglang.srt.layers.layernorm import RMSNorm, _flashinfer_rmsnorm_quant_available

if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark")
if not _flashinfer_rmsnorm_quant_available:
raise RuntimeError(
"flashinfer rmsnorm_quant / fused_add_rmsnorm_quant is not available; "
"install flashinfer to benchmark the fused path"
)

try:
from flashinfer.norm import fused_add_rmsnorm_quant_cute, rmsnorm_quant_cute

_CUTE_AVAILABLE = True
except ImportError:
_CUTE_AVAILABLE = False

DEVICE = "cuda"
DTYPE = torch.bfloat16
FP8_DTYPE = torch.float8_e4m3fn
HIDDEN_SIZES = [4096, 8192]
# Per-tensor reciprocal scale (q = normed / scale); 0.05 keeps normed/scale well
# within the e4m3 range for unit-scale activations.
SCALE_VALUE = 0.05


def make_layer(hidden_size):
layer = RMSNorm(hidden_size).to(device=DEVICE, dtype=DTYPE)
layer.weight.data.normal_(mean=1.0, std=0.1)
return layer


def make_inputs(num_tokens, hidden_size, add_residual):
x = torch.randn(num_tokens, hidden_size, device=DEVICE, dtype=DTYPE)
residual = torch.randn_like(x) if add_residual else None
scale = torch.tensor([SCALE_VALUE], device=DEVICE, dtype=torch.float32)
return x, residual, scale


def run_unfused(layer, x, residual, scale):
out = layer(x, residual)
if residual is not None:
normed, residual_out = out
q, q_scale = static_quant_fp8(normed, scale)
return (q, q_scale), residual_out
q, q_scale = static_quant_fp8(out, scale)
return q, q_scale


def _run_fused(kernel, add_kernel, layer, x, residual, scale):
out = torch.empty_like(x, dtype=FP8_DTYPE)
if residual is not None:
# In-place: residual += x, then out = quant(rmsnorm(residual) * w).
add_kernel(out, x, residual, layer.weight.data, scale, layer.variance_epsilon)
return (out, scale), residual
kernel(out, x, layer.weight.data, scale, layer.variance_epsilon)
return out, scale


def run_fused_default(layer, x, residual, scale):
return _run_fused(rmsnorm_quant, fused_add_rmsnorm_quant, layer, x, residual, scale)


def run_fused_cute(layer, x, residual, scale):
return _run_fused(
rmsnorm_quant_cute, fused_add_rmsnorm_quant_cute, layer, x, residual, scale
)


RUNNERS = {
"unfused": run_unfused,
"fused": run_fused_default,
"fused_cute": run_fused_cute,
}

# (provider key, plot label, style)
_PROVIDERS = [
("unfused", "rmsnorm + static_quant_fp8 (unfused)", ("blue", "-")),
("fused", "rmsnorm_quant (fused, default)", ("green", "-")),
]
if _CUTE_AVAILABLE:
_PROVIDERS.append(
("fused_cute", "rmsnorm_quant_cute (fused, cute-dsl)", ("red", "-"))
)


def _bench_ms(fn, args, quantiles=(0.5, 0.2, 0.8)):
# Pass the GPU tensors as input_args so flashinfer's cold_l2_cache flush can
# find them; a zero-arg callable trips its "no GPU tensors found" warning and
# silently disables cold-L2 timing.
times = bench_gpu_time(
fn=fn,
input_args=args,
use_cuda_graph=True,
dry_run_time_ms=25,
repeat_time_ms=100,
)
return tuple(float(np.percentile(times, q * 100)) for q in quantiles)


def _check_correctness():
"""One-shot sanity check that every fused provider agrees with the unfused
baseline within FP8 precision."""
fused_providers = [p for p in RUNNERS if p != "unfused"]
for hidden_size, add_residual in itertools.product(HIDDEN_SIZES, [False, True]):
layer = make_layer(hidden_size)
x, residual, scale = make_inputs(64, hidden_size, add_residual)
with torch.inference_mode():
ref = run_unfused(
layer, x.clone(), residual.clone() if add_residual else None, scale
)
(uq, _), _ = ref if add_residual else (ref, None)
ref_deq = uq.float() * scale
for provider in fused_providers:
if provider == "fused_cute" and not _CUTE_AVAILABLE:
continue
with torch.inference_mode():
out = RUNNERS[provider](
layer, x.clone(), residual.clone() if add_residual else None, scale
)
(q, _), _ = out if add_residual else (out, None)
cos = torch.nn.functional.cosine_similarity(
(q.float() * scale).flatten(), ref_deq.flatten(), dim=0
).item()
assert (
cos > 0.99
), f"{provider} h={hidden_size} residual={add_residual} cos={cos:.4f}"
print("correctness check passed (all fused providers vs unfused within FP8)")


configs = [
triton.testing.Benchmark(
x_names=["num_tokens"],
x_vals=[512, 1024, 2048, 4096, 8192, 16384],
x_log=False,
line_arg="provider",
line_vals=[p[0] for p in _PROVIDERS],
line_names=[p[1] for p in _PROVIDERS],
styles=[p[2] for p in _PROVIDERS],
ylabel="latency (ms)",
plot_name=f"rmsnorm_fp8_quant_h{hidden_size}_residual{add_residual}",
args={"hidden_size": hidden_size, "add_residual": add_residual},
)
for hidden_size, add_residual in itertools.product(HIDDEN_SIZES, [False, True])
]


@triton.testing.perf_report(configs)
def benchmark(num_tokens, hidden_size, add_residual, provider):
layer = make_layer(hidden_size)
x, residual, scale = make_inputs(num_tokens, hidden_size, add_residual)
return _bench_ms(RUNNERS[provider], (layer, x, residual, scale))


if __name__ == "__main__":
torch.manual_seed(0)
_check_correctness()
benchmark.run(print_data=True, show_plots=False)
30 changes: 22 additions & 8 deletions python/sglang/kernels/aot/benchmark/bench_fp8_gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,7 @@
import triton
from sgl_kernel import fp8_scaled_mm as sgl_scaled_mm

from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
per_tensor_quant_fp8,
)
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import per_tensor_quant_fp8
from sglang.utils import is_in_ci

# Optional vLLM import
Expand Down Expand Up @@ -106,7 +104,7 @@ def sglang_scaled_fp8_quant(
if IS_CI:
batch_sizes = [1] # Single batch size for CI
else:
batch_sizes = [1, 16, 64, 128, 256, 512, 1024, 2048]
batch_sizes = [1, 2, 8, 16, 64, 128, 256, 512, 1024, 2048]

# Filter line_vals based on vLLM availability
if VLLM_AVAILABLE:
Expand All @@ -115,24 +113,39 @@ def sglang_scaled_fp8_quant(
"vllm-fp8-bf16",
"sglang-fp8-fp16",
"sglang-fp8-bf16",
"sglang-scalar-a-fp8-fp16",
"sglang-scalar-a-fp8-bf16",
]
line_names = [
"vllm-fp8-fp16",
"vllm-fp8-bf16",
"sglang-fp8-fp16",
"sglang-fp8-bf16",
"sglang-scalar-a-fp8-fp16",
"sglang-scalar-a-fp8-bf16",
]
styles = [
("green", "-"),
("green", "--"),
("blue", "-"),
("blue", "--"),
("red", "-"),
("red", "--"),
]
styles = [("green", "-"), ("green", "--"), ("blue", "-"), ("blue", "--")]
else:
line_vals = [
"sglang-fp8-fp16",
"sglang-fp8-bf16",
"sglang-scalar-a-fp8-fp16",
"sglang-scalar-a-fp8-bf16",
]
line_names = [
"sglang-fp8-fp16",
"sglang-fp8-bf16",
"sglang-scalar-a-fp8-fp16",
"sglang-scalar-a-fp8-bf16",
]
styles = [("blue", "-"), ("blue", "--")]
styles = [("blue", "-"), ("blue", "--"), ("red", "-"), ("red", "--")]


@triton.testing.perf_report(
Expand Down Expand Up @@ -174,8 +187,9 @@ def benchmark(batch_size, provider, N, K):
lambda: vllm_scaled_mm(a_fp8, b_fp8, scale_a_fp8, scale_b_fp8, dtype),
quantiles=quantiles,
)
elif "sglang-fp8" in provider:
a_fp8, scale_a_fp8 = sglang_scaled_fp8_quant(a, scale_a)
elif "sglang" in provider:
a_scale = scale_a_scalar if "scalar-a" in provider else scale_a
a_fp8, scale_a_fp8 = sglang_scaled_fp8_quant(a, a_scale)
b_fp8, scale_b_fp8 = sglang_scaled_fp8_quant(b, scale_b)
b_fp8 = b_fp8.t()
ms, min_ms, max_ms = triton.testing.do_bench_cudagraph(
Expand Down
Loading
Loading