[AMD] fix: use the hardware fp8 e4m3 convert on gfx950 (#37140)

Signed-off-by: amd-danli103 <danli103@amd.com>
This commit is contained in:
amd-danli103
2026-09-08 03:01:53 -07:00
committed by GitHub
parent 775f17b07c
commit 141febf329
6 changed files with 260 additions and 7 deletions
@@ -0,0 +1,116 @@
"""The DSv4 fp8 e4m3 conversion on its own, byte for byte against torch.
Every fp8 store in the DSv4 tree quantizes through ``pack_fp8``, and none of the
callers can see when it is wrong: the value has already been divided by a
quantization scale, so a bad rounding or saturation boundary comes back as fp8 being
lossier than it should be, not as a failure. ``pack_fp8`` on ROCm used to be a
hand-written bit twiddle, and it had two: the whole top exponent segment saturated
to the max normal, and the binade under the min subnormal flushed to zero instead of
rounding up to it. Both are pinned below.
gfx950 only. gfx942 still runs the software cast, bugs and all, because the
instruction cannot produce the fnuz bytes that arch writes; that one needs its own fix.
The reference is torch's own cast.
"""
import math
import unittest
import torch
from sglang.kernels.ops.attention.dsv4.fp8_cvt import cvt_fp8_e4m3
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, fp8_max
from sglang.srt.utils import is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase
# the default amd runner is mi300, where pack_fp8 still takes the software path this
# does not cover
register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd-mi35x")
DEVICE = torch.device("cuda")
# start of the top exponent segment, i.e. the largest power of two the format holds
TOP_BINADE = 2.0 ** math.floor(math.log2(fp8_max))
def _representable():
vals = torch.arange(256, dtype=torch.uint8).view(fp8_dtype).float()
return vals[vals.isfinite()]
def _domain():
"""every representable value, every midpoint between two of them, every bf16"""
vals = _representable()
mids = ((vals[:, None] + vals[None, :]) / 2).flatten()
cases = torch.cat(
[
vals,
mids,
torch.linspace(-fp8_max, fp8_max, 100003),
torch.arange(1 << 16, dtype=torch.int32).view(torch.bfloat16).float(),
]
)
# stay inside the range: past the max the two casts are allowed to disagree on
# whether to clamp or produce NaN, which is not what this is testing
cases = cases[cases.isfinite() & (cases.abs() <= fp8_max)].unique()
if cases.numel() % 2:
cases = cases[:-1]
return cases.contiguous()
def _as_bytes(x):
# the conversion runs two values at a time, so the length has to stay even --
# e4m3fnuz has an odd number of representable values (only 0x80 is NaN)
if x.numel() % 2:
x = x[:-1]
x = x.contiguous().to(DEVICE)
return cvt_fp8_e4m3(x), x.to(fp8_dtype).view(torch.uint8)
@unittest.skipUnless(
torch.cuda.is_available() and is_hip() and is_gfx95_supported(),
"the gfx95 path of pack_fp8 is what this pins",
)
class TestDsv4Fp8Cast(CustomTestCase):
def test_matches_torch_over_the_whole_range(self):
cases = _domain()
got, want = _as_bytes(cases)
bad = got != want
if bool(bad.any()):
v = cases.to(DEVICE)[bad]
worst = v.abs().argmax()
self.fail(
f"{int(bad.sum())} of {cases.numel()} bytes differ, "
f"|v| in [{v.abs().min():.4e}, {v.abs().max():.4e}]; e.g. "
f"{v[worst]:.6g} -> {got[bad][worst].item():#04x} "
f"(torch {want[bad][worst].item():#04x})"
)
def test_top_exponent_segment_is_not_saturated(self):
# testing the exponent alone here used to write every value from TOP_BINADE up
# to the max out as the max normal
vals = _representable()
# pack_fp8 clips to fp8_max, so anything the format holds above it never comes
# out of the conversion
top = vals[(vals.abs() >= TOP_BINADE) & (vals.abs() <= fp8_max)]
self.assertGreater(top.numel(), 2)
got, want = _as_bytes(top)
self.assertTrue(torch.equal(got, want))
# and they really are distinct values, not all the same byte
self.assertGreater(int(got.unique().numel()), 2)
def test_binade_below_the_min_subnormal_rounds_up(self):
min_subnormal = _representable().abs()
min_subnormal = min_subnormal[min_subnormal > 0].min().item()
# (midpoint, min subnormal): rounds up. the midpoint itself is a tie and goes
# to even, i.e. to zero
band = torch.linspace(min_subnormal / 2, min_subnormal, 2049)[1:-1]
band = torch.cat([band, -band])
got, want = _as_bytes(band.contiguous())
self.assertTrue(torch.equal(got, want))
self.assertTrue(bool((got & 0x7F).all()))
if __name__ == "__main__":
unittest.main()
@@ -17,19 +17,22 @@ branches introduced by the scheduling optimization.
from __future__ import annotations
import pytest
import sgl_kernel # noqa: F401 the ROCm path dispatches to torch.ops.sgl_kernel
import torch
from sglang.kernels.ops.attention.dsv4 import (
fused_q_indexer_rope_first_quant,
fused_q_indexer_rope_hadamard_quant,
)
from sglang.srt.utils import is_hip
from sglang.srt.utils import is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
_is_hip = is_hip()
register_cuda_ci(est_time=13, stage="base-b", runner_config="1-gpu-large")
register_amd_ci(est_time=45, suite="jit-kernel-unit-test-amd")
# the mi35x suite rather than the default AMD one: the ROCm case below is gfx95-only, and
# everything else in here skips on HIP, so the mi300 registration only ever produced skips
register_amd_ci(est_time=45, suite="stage-b-test-1-gpu-small-amd-mi35x")
HEAD_DIM = 128
ROPE_DIM = 64
@@ -45,10 +48,10 @@ N_HEADS = 64
BATCHES = [1, 8, 64, 256, 512, 2048]
def _skip_if_unavailable():
def _skip_if_unavailable(hip_ok=False):
if not torch.cuda.is_available():
pytest.skip("CUDA required")
if _is_hip:
if _is_hip and not hip_ok:
pytest.skip("Indexer fused Q kernel is CUDA-specific")
@@ -73,7 +76,14 @@ def _fp8_dequant_ok(q_fp8, ref, scale):
@pytest.mark.parametrize("pos_dtype", [torch.int32, torch.int64])
@pytest.mark.parametrize("batch", BATCHES)
def test_v4_rope_hadamard_quant_matches_reference(batch, pos_dtype):
_skip_if_unavailable()
# runs on gfx95 too: elementwise.py routes this one to the AOT op there, and that op
# carries its own copy of the cast, so this is the only coverage it gets
_skip_if_unavailable(hip_ok=True)
if _is_hip:
if not is_gfx95_supported():
pytest.skip("gfx942 keeps the software cast in the AOT copy")
if pos_dtype is torch.int64:
pytest.skip("the ROCm AOT op takes int32 positions only")
dev = "cuda"
g = torch.Generator(device=dev).manual_seed(0)
q = torch.randn(
@@ -157,6 +167,9 @@ def test_v32_rope_first_quant_matches_reference(batch):
# Strided weight (the non-contiguous wk slice) matches contiguous (V4 path).
# ----------------------------------------------------------------------------
def test_v4_strided_weight_matches_contiguous():
# stays CUDA-only. the ROCm op reads the weight linearly, so a non-contiguous slice
# comes out wrong there -- unrelated to the cast, and latent, since the indexer hands
# it the contiguous weights_proj output
_skip_if_unavailable()
dev = "cuda"
B = 512 # grid-stride regime