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
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,13 @@ def _unsupported_derived_weight_cache_error(
"online weight updates."
)

if model is not None:
# Model-owned caches can publish the same rank-uniform constraint
# without importing individual model implementations in the updater.
for module in model.modules():
reason = getattr(module, "_derived_weight_cache_error", None)
if reason is not None:
return reason
from sglang.kernels.ops.attention.dsv4.gemm import hpc_bf16xfp32_gemm_enabled

if hpc_bf16xfp32_gemm_enabled():
Expand Down
156 changes: 151 additions & 5 deletions python/sglang/srt/models/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,8 +133,12 @@
_is_gfx95 = is_gfx95_supported()
_is_hip = is_hip()
_QWEN3_5_MOE_TEXT_MODEL_TYPES = ("qwen3_5_moe_text", "qwen4_exp_text")
# qwen4_exp shares these classes, but the ROCm packed path is Qwen3.5-only.
_QWEN3_5_ROCM_PACKED_MODEL_TYPES = ("qwen3_5_text", "qwen3_5_moe_text")
_is_xpu = is_xpu()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
from aiter.tuned_gemm import tgemm
_hip_use_alt_stream = get_bool_env_var("SGLANG_ALT_STREAM") and _is_hip
_gdn_use_alt_stream = _is_cuda or (
get_bool_env_var("SGLANG_GDN_QKVZ_BA_ALT_STREAM", "False") and _hip_use_alt_stream
Expand Down Expand Up @@ -394,6 +398,10 @@ def __init__(
self._bind_packed_weight_loaders(self.in_proj_ba)
self._fused_in_proj_weight: Optional[torch.Tensor] = None
self._fused_in_proj_qkvz_width = 0
self._fused_in_proj_ba_width = 0
self._fused_in_proj_scale: Optional[torch.Tensor] = None
self._fused_in_proj_sources = None
self._derived_weight_cache_error = None
self._fused_input_proj_cpu_enabled = LazyValue(
lambda: (
_is_cpu
Expand Down Expand Up @@ -653,16 +661,27 @@ def fix_query_key_value_ordering(
return query, key, value, z, b, a

def finalize_fused_in_proj(self) -> None:
"""Stack in_proj_qkvz + in_proj_ba into one GEMM weight;
the module weights become row views of it,
so weight reload and dtype checks still see them."""
if not _is_cuda or self._fused_in_proj_weight is not None:
"""Prepare one BF16 or FP8 GEMM for both input projections.

BF16 parameters alias the packed rows. FP8 keeps the original layouts
for the separate path and publishes its derived-cache update constraint.
"""
if not (_is_cuda or _use_aiter) or self._fused_in_proj_weight is not None:
return
if (
_use_aiter
and self.config.model_type not in _QWEN3_5_ROCM_PACKED_MODEL_TYPES
):
return
if get_lora().enable_lora or get_lora().lora_paths:
# LoRA wraps the individual Linear modules; the fused GEMM would
# bypass their adapters.
return
if _use_aiter and not get_bool_env_var("SGLANG_QWEN35_PACKED_IN_PROJ", "True"):
return
qkvz, ba = self.in_proj_qkvz, self.in_proj_ba
if _use_aiter and self._finalize_fused_fp8_in_proj():
return
if not (
isinstance(qkvz.quant_method, UnquantizedLinearMethod)
and isinstance(ba.quant_method, UnquantizedLinearMethod)
Expand All @@ -678,7 +697,124 @@ def finalize_fused_in_proj(self) -> None:
ba.weight.data = fused[self._fused_in_proj_qkvz_width :]
self._fused_in_proj_weight = fused

def _finalize_fused_fp8_in_proj(self) -> bool:
"""Pack loaded Quark per-channel FP8 projections without requantizing."""
from aiter.ops.shuffle import shuffle_weight

from sglang.srt.layers.quantization.fp8_utils import use_aiter_bpreshuffle_gemm
from sglang.srt.layers.quantization.quark.schemes.quark_w8a8_fp8 import (
QuarkW8A8Fp8,
)

projections = (self.in_proj_qkvz, self.in_proj_ba)
if not all(
type(getattr(proj, "scheme", None)) is QuarkW8A8Fp8
and proj.scheme.weight_qscheme == "per_channel"
and proj.scheme.per_token
and proj.input_scale is None
and proj.bias is None
and proj.weight.dtype == torch.float8_e4m3fn
and proj.weight.is_cuda
and proj.weight.t().is_contiguous()
and proj.weight.shape[1] % 16 == 0
and proj.weight_scale.numel() == proj.weight.shape[1]
for proj in projections
):
return False
# Keep the derived-cache update constraint uniform across ranks. With
# PP, an attention-only stage may have no GDN cache to publish it from.
if get_parallel().pp_size != 1:
return False
widths = [proj.weight.shape[1] for proj in projections]
if projections[0].weight.shape[0] != projections[1].weight.shape[0]:
return False
# CK's bpreshuffle GEMM requires N % 64 == 0. TP4 has BA=32,
# so pad only the packed result, preserving the original Linear sizes.
padded_width = (sum(widths) + 63) // 64 * 64
weight = torch.zeros(
(padded_width, projections[0].weight.shape[0]),
dtype=projections[0].weight.dtype,
device=projections[0].weight.device,
)
scale = torch.ones((padded_width, 1), dtype=torch.float32, device=weight.device)
offset = 0
for proj, width in zip(projections, widths):
source = proj.weight.t()
if not use_aiter_bpreshuffle_gemm(width):
source = shuffle_weight(source, (16, 16))
# AITER's shuffle is local to 16 output rows; concatenating aligned
# shuffled blocks is identical to shuffling their concatenation.
weight[offset : offset + width].copy_(source)
scale[offset : offset + width].copy_(proj.weight_scale.view(-1, 1))
offset += width
self._fused_in_proj_weight = weight
self._fused_in_proj_scale = scale
self._fused_in_proj_qkvz_width, self._fused_in_proj_ba_width = widths
self._fused_in_proj_sources = tuple(
(value, value.data_ptr(), None if value.is_inference() else value._version)
for proj in projections
for value in (proj.weight, proj.weight_scale)
)
self._derived_weight_cache_error = (
"Online weight updates are not supported with packed FP8 GDN input "
"projections: captured graphs retain a derived weight/scale buffer. "
"Restart with SGLANG_QWEN35_PACKED_IN_PROJ=0 before updating weights."
)
return True

def _fused_fp8_in_proj_sources_valid(self) -> bool:
current = tuple(
value
for proj in (self.in_proj_qkvz, self.in_proj_ba)
for value in (proj.weight, proj.weight_scale)
)
return all(
value is source
and value.data_ptr() == pointer
and (version is None or value._version == version)
for value, (source, pointer, version) in zip(
current, self._fused_in_proj_sources
)
)

def _forward_input_proj(self, hidden_states: torch.Tensor):
if _use_aiter and self._fused_in_proj_weight is not None:
# Unquantized BF16 projections consume the bf16 side of the fused
# AR+RMSNorm tuple; one aiter GEMM replaces the two separate
# projections. Measured on MI355X the packed GEMM only pays off for
# decode/verify-sized batches, so larger batches keep the separate
# projections below.
x = hidden_states[0] if isinstance(hidden_states, tuple) else hidden_states
if (
x.dtype == torch.bfloat16
and 0 < x.shape[0] <= 64
and not torch.compiler.is_compiling()
):
if self._fused_in_proj_scale is not None:
if self._fused_fp8_in_proj_sources_valid():
from sglang.srt.layers.quantization.fp8_utils import (
apply_fp8_linear,
)

fused_out = apply_fp8_linear(
x,
self._fused_in_proj_weight.t(),
self._fused_in_proj_scale,
use_per_token_if_dynamic=True,
)
split = self._fused_in_proj_qkvz_width
return (
fused_out[:, :split],
fused_out[:, split : split + self._fused_in_proj_ba_width],
)
else:
fused_out = tgemm.mm(
x, self._fused_in_proj_weight, None, otype=x.dtype
)
return (
fused_out[:, : self._fused_in_proj_qkvz_width],
fused_out[:, self._fused_in_proj_qkvz_width :],
)
# AMD/aiter fused AR+RMSNorm+per-group-quant path ships a
# ``(bf16, fp8, scale)`` 3-tuple so the FP8 ``in_proj_qkvz`` can
# consume ``(fp8, scale)`` (skipping its internal quant) while the
Expand All @@ -688,7 +824,8 @@ def _forward_input_proj(self, hidden_states: torch.Tensor):
return self._forward_input_proj_fused_quant_amd(hidden_states)

if (
self._fused_in_proj_weight is not None
not _use_aiter
and self._fused_in_proj_weight is not None
and hidden_states.dtype == torch.bfloat16
# Measured on cuBLAS above ~1k rows:
# the merged (m, 4120) GEMM is ~10% slower than the two separate GEMMs.
Expand Down Expand Up @@ -1724,6 +1861,15 @@ def get_input_embeddings(self):
return self.embed_tokens

def prepare_before_cuda_graph_capture(self, model_runner) -> None:
if _use_aiter and self.config.model_type in _QWEN3_5_ROCM_PACKED_MODEL_TYPES:
packed = 0
for module in self.modules():
if isinstance(module, Qwen3_5GatedDeltaNet):
module.finalize_fused_in_proj()
packed += int(module._fused_in_proj_weight is not None)
logger.info(
"Packed BF16/FP8 GDN input projection enabled for %d layers", packed
)
if self.flashinfer_mnnvl_cutedsl_fusion is None:
return
from sglang.srt.layers.moe.qwen35_flashinfer_fusion import (
Expand Down
142 changes: 142 additions & 0 deletions test/registered/amd/test_qwen35_gdn_packed_fp8_in_proj.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
"""Quark FP8 GDN packing: actual TP widths, scales, padding and dispatch."""

import os
import types
import unittest
from unittest.mock import patch

import torch

from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase

register_amd_ci(est_time=90, suite="stage-b-test-1-gpu-small-amd-mi35x")


class _FP8Linear(torch.nn.Module):
def __init__(self, rows):
super().__init__()
from sglang.srt.layers.quantization.quark.schemes.quark_w8a8_fp8 import (
QuarkW8A8Fp8,
)

self.weight = torch.nn.Parameter(
torch.randn(rows, 4096, device="cuda").to(torch.float8_e4m3fn),
requires_grad=False,
)
self.weight_scale = torch.nn.Parameter(
torch.linspace(0.001, 0.1, rows, device="cuda", dtype=torch.float32),
requires_grad=False,
)
self.scheme = QuarkW8A8Fp8(
{"qscheme": "per_channel"},
{"qscheme": "per_channel", "is_dynamic": True},
)
self.scheme.process_weights_after_loading(self)
self.bias = None
self.quant_method = object()

def forward(self, x):
return self.scheme.apply_weights(self, x), None


@unittest.skipUnless(torch.cuda.is_available() and torch.version.hip, "ROCm FP8 GEMM")
class TestQwen35GDNPackedFP8InProj(CustomTestCase):
@classmethod
def setUpClass(cls):
os.environ.setdefault("SGLANG_USE_AITER", "1")
from sglang.srt.models import qwen3_5

if not qwen3_5._use_aiter:
raise unittest.SkipTest("SGLANG_USE_AITER was disabled at import")
cls.module = qwen3_5
cls.gdn_cls = qwen3_5.Qwen3_5GatedDeltaNet

def make_gdn(self, tp, pp_size=1, model_type="qwen3_5_moe_text"):
gdn = types.SimpleNamespace(
config=types.SimpleNamespace(model_type=model_type),
in_proj_qkvz=_FP8Linear(20480 // tp),
in_proj_ba=_FP8Linear(128 // tp),
_fused_in_proj_weight=None,
_fused_in_proj_scale=None,
_fused_in_proj_sources=None,
alt_stream=None,
_fused_input_proj_cpu_enabled=types.SimpleNamespace(value=False),
)
for name in (
"_finalize_fused_fp8_in_proj",
"_fused_fp8_in_proj_sources_valid",
"_forward_input_proj_fused_quant_amd",
):
setattr(gdn, name, types.MethodType(getattr(self.gdn_cls, name), gdn))
with (
patch.object(
self.module,
"get_parallel",
return_value=types.SimpleNamespace(pp_size=pp_size),
),
patch.object(
self.module,
"get_lora",
return_value=types.SimpleNamespace(enable_lora=False, lora_paths=None),
),
patch.dict(os.environ, {"SGLANG_QWEN35_PACKED_IN_PROJ": "1"}),
):
self.gdn_cls.finalize_fused_in_proj(gdn)
return gdn

def test_fp8_dispatch_matches_separate_projections(self):
torch.manual_seed(42)
for tp in (2, 4):
gdn = self.make_gdn(tp)
self.assertIsNotNone(gdn._fused_in_proj_scale)
self.assertEqual(gdn._fused_in_proj_weight.shape[0] % 64, 0)
self.assertTrue(gdn._fused_fp8_in_proj_sources_valid())
for m in (1, 4, 33, 64, 65, 300):
with self.subTest(tp=tp, tokens=m):
x = torch.randn(m, 4096, device="cuda", dtype=torch.bfloat16)
expected = (gdn.in_proj_qkvz(x)[0], gdn.in_proj_ba(x)[0])
for hidden in (x, (x, None, None)):
got = self.gdn_cls._forward_input_proj(gdn, hidden)
for actual, reference in zip(got, expected):
self.assertEqual(actual.shape, reference.shape)
error = (actual.float() - reference.float()).norm()
self.assertLess(
(error / reference.float().norm()).item(), 0.008
)
if m > 64:
self.assertTrue(
all(torch.equal(a, b) for a, b in zip(got, expected))
)

def test_replaced_parameters_disable_cached_projection(self):
gdn = self.make_gdn(4)
gdn.in_proj_ba.weight = torch.nn.Parameter(
gdn.in_proj_ba.weight.clone(), requires_grad=False
)
self.assertFalse(gdn._fused_fp8_in_proj_sources_valid())
x = torch.randn(4, 4096, device="cuda", dtype=torch.bfloat16)
got = self.gdn_cls._forward_input_proj(gdn, x)
expected = (gdn.in_proj_qkvz(x)[0], gdn.in_proj_ba(x)[0])
self.assertTrue(all(torch.equal(a, b) for a, b in zip(got, expected)))

def test_pipeline_parallel_keeps_original_projections(self):
gdn = self.make_gdn(4, pp_size=2)
self.assertIsNone(gdn._fused_in_proj_weight)

def test_fp8_packing_is_limited_to_qwen35(self):
for model_type, supported in (
("qwen3_5_text", True),
("qwen3_5_moe_text", True),
("qwen4_exp_text", False),
("other_text_model", False),
):
with self.subTest(model_type=model_type):
gdn = self.make_gdn(4, model_type=model_type)
self.assertEqual(gdn._fused_in_proj_weight is not None, supported)
self.assertEqual(gdn._fused_in_proj_scale is not None, supported)
self.assertEqual(hasattr(gdn, "_derived_weight_cache_error"), supported)


if __name__ == "__main__":
unittest.main()
Loading
Loading