Skip to content
Merged
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
11 changes: 11 additions & 0 deletions python/sglang/srt/layers/quantization/fp8_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,17 @@ def materialize_bpreshuffle_fp8_scale(scale: torch.Tensor) -> torch.Tensor:
return scale.t().contiguous().t() if scale.dim() == 2 else scale


def view_aiter_fused_rms_transposed_fp8_scale(scale: torch.Tensor) -> torch.Tensor:
"""Expose AITER fused-RMS ``transpose_scale=True`` storage logically.

The fused-RMS op returns transposed physical bytes through a row-major-looking
view. Restore logical ``[M, G]`` indexing without copying those bytes.
"""
if scale.dim() != 2:
return scale
return torch.as_strided(scale, scale.shape, (1, scale.shape[0]))


def materialize_bpreshuffle_fp8_scale_tuple(
value: Tuple[torch.Tensor, ...],
) -> Tuple[torch.Tensor, ...]:
Expand Down
8 changes: 8 additions & 0 deletions python/sglang/srt/models/deepseek_v4.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,9 @@
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.quantization.fp8_utils import (
view_aiter_fused_rms_transposed_fp8_scale,
)
from sglang.srt.layers.rotary_embedding import get_rope_wrapper
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.utils.cp_utils import (
Expand Down Expand Up @@ -249,6 +252,11 @@ def _fused_rmsnorm_fp8_quant(hidden_states, weight, eps):
output_unquantized_inp1=True,
transpose_scale=_use_aiter_bpreshuffle_gfx95,
)
if _use_aiter_bpreshuffle_gfx95:
x_quant = (
x_quant[0],
view_aiter_fused_rms_transposed_fp8_scale(x_quant[1]),
)
return x_quant, x_bf16


Expand Down
42 changes: 42 additions & 0 deletions test/registered/quant/test_fused_rms_fp8_group_quant.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,10 @@
import torch
import torch.nn.functional as F

from sglang.srt.layers.quantization.fp8_utils import (
materialize_bpreshuffle_fp8_scale,
view_aiter_fused_rms_transposed_fp8_scale,
)
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase

Expand Down Expand Up @@ -145,6 +149,44 @@ def test_fused_rms_fp8_group_quant(self):
with self.subTest(M=M, N1=N1, N2=N2, group_size=g, dtype=dtype, seed=seed):
self._case(M, N1, N2, g, dtype, seed)

def test_transposed_scale_matches_bpreshuffle_layout_contract(self):
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant

common_kwargs = dict(
inp2=None,
inp2_weight=None,
inp2_epsilon=None,
group_size=128,
dtype_quant=torch.float8_e4m3fn,
res1=None,
output_unquantized_inp1=False,
)

for m, k in ((1, 1024), (64, 1024), (1, 4096), (64, 4096)):
with self.subTest(m=m, k=k):
torch.manual_seed(0)
x = torch.randn(m, k, dtype=torch.bfloat16, device="cuda")
weight = torch.ones(k, dtype=torch.float32, device="cuda")

(q_row_major, scale_row_major), *_ = fused_rms_fp8_group_quant(
x, weight, 1e-6, transpose_scale=False, **common_kwargs
)
(q_transposed, scale_transposed), *_ = fused_rms_fp8_group_quant(
x, weight, 1e-6, transpose_scale=True, **common_kwargs
)

repaired = view_aiter_fused_rms_transposed_fp8_scale(scale_transposed)
materialized = materialize_bpreshuffle_fp8_scale(repaired)

torch.testing.assert_close(q_transposed, q_row_major, rtol=0, atol=0)
torch.testing.assert_close(repaired, scale_row_major, rtol=0, atol=0)
torch.testing.assert_close(
materialized, scale_row_major, rtol=0, atol=0
)
self.assertEqual(repaired.stride(), (1, repaired.shape[0]))
self.assertEqual(repaired.data_ptr(), scale_transposed.data_ptr())
self.assertEqual(materialized.data_ptr(), scale_transposed.data_ptr())


if __name__ == "__main__":
unittest.main(verbosity=2)
55 changes: 55 additions & 0 deletions test/registered/unit/layers/test_fp8_bpreshuffle_scale.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,12 @@
import unittest
from unittest.mock import patch

import torch

from sglang.srt.layers.quantization.fp8_utils import (
materialize_bpreshuffle_fp8_scale,
materialize_bpreshuffle_fp8_scale_tuple,
view_aiter_fused_rms_transposed_fp8_scale,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
Expand All @@ -31,6 +33,59 @@ def test_materialization_is_idempotent_for_bpreshuffle_layout(self):

self.assertTrue(torch.equal(rematerialized, scale))
self.assertEqual(rematerialized.stride(), materialized.stride())
self.assertEqual(rematerialized.data_ptr(), materialized.data_ptr())

def test_repairs_aiter_scale_before_downstream_layout_handling(self):
"""AITER-transposed scale bytes must retain their logical indexing.

AITER ``transpose_scale=True`` returns transposed physical storage with
row-major-looking metadata. Treating that metadata as logical layout
permutes the scales during CK materialization.
"""
logical_scale = torch.arange(12, dtype=torch.float32).reshape(3, 4)
aiter_scale = logical_scale.t().contiguous().view(logical_scale.shape)

repaired = view_aiter_fused_rms_transposed_fp8_scale(aiter_scale)
materialized = materialize_bpreshuffle_fp8_scale(repaired)
renormalized = view_aiter_fused_rms_transposed_fp8_scale(repaired)

self.assertTrue(torch.equal(repaired, logical_scale))
self.assertTrue(torch.equal(materialized, logical_scale))
self.assertTrue(torch.equal(renormalized, logical_scale))
self.assertEqual(repaired.stride(), (1, logical_scale.shape[0]))
self.assertEqual(repaired.data_ptr(), aiter_scale.data_ptr())
self.assertEqual(materialized.data_ptr(), aiter_scale.data_ptr())
self.assertEqual(renormalized.stride(), repaired.stride())
self.assertEqual(renormalized.data_ptr(), aiter_scale.data_ptr())

def test_deepseek_v4_repairs_fused_rms_scale_at_producer(self):
"""DeepSeek-V4 must repair fused-RMS scale metadata before CK consumes it."""
from sglang.srt.models import deepseek_v4

q_input = torch.ones((3, 1024), dtype=torch.float32)
x_bf16 = torch.ones((3, 1024), dtype=torch.bfloat16)
logical_scale = torch.arange(24, dtype=torch.float32).reshape(3, 8)
aiter_scale = logical_scale.t().contiguous().view(logical_scale.shape)
fused_output = ((q_input, aiter_scale), x_bf16, None, None)

with (
patch.object(
deepseek_v4,
"fused_rms_fp8_group_quant",
return_value=fused_output,
create=True,
),
patch.object(deepseek_v4, "_use_aiter_bpreshuffle_gfx95", True),
):
x_quant, x_unquantized = deepseek_v4._fused_rmsnorm_fp8_quant(
q_input, torch.ones(1024), 1e-6
)

self.assertIs(x_quant[0], q_input)
self.assertIs(x_unquantized, x_bf16)
self.assertTrue(torch.equal(x_quant[1], logical_scale))
self.assertEqual(x_quant[1].stride(), (1, logical_scale.shape[0]))
self.assertEqual(x_quant[1].data_ptr(), aiter_scale.data_ptr())

def test_tuple_helper_keeps_extra_tuple_payload(self):
q_input = torch.ones((3, 8), dtype=torch.float32)
Expand Down
Loading