diff --git a/docs_new/docs/advanced_features/quantization.mdx b/docs_new/docs/advanced_features/quantization.mdx
index bd59a0497..935651e9c 100644
--- a/docs_new/docs/advanced_features/quantization.mdx
+++ b/docs_new/docs/advanced_features/quantization.mdx
@@ -142,10 +142,10 @@ The following table summarizes quantization method support across NVIDIA and AMD
modelopt_fp4 |
- Yes (Blackwell/SM100+) |
+ Yes (SM80-SM90 via Marlin; SM100+ native FP4) |
No |
No |
- NVIDIA ModelOpt; native FP4 on Blackwell (B200, GB200) |
+ NVIDIA ModelOpt; use Marlin W4A16 fallback on Ampere/Hopper and native FP4 backends on Blackwell |
petit_nvfp4 |
@@ -257,8 +257,8 @@ Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. Wh
auto |
- SM100/120 |
- Auto-selects: flashinfer_cudnn on SM120; flashinfer_cutlass on SM100 |
+ SM80+ |
+ Auto-selects: flashinfer_cudnn on SM120; flashinfer_cutedsl on SM100; marlin on SM80-SM90; flashinfer_cutlass otherwise |
cutlass |
@@ -275,15 +275,25 @@ Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. Wh
SM100/120 (CUDA 13+, cuDNN 9.15+) |
FlashInfer cuDNN backend; used on SM120 for performance |
+
+ flashinfer_cutedsl |
+ SM100 |
+ FlashInfer CuTe DSL backend |
+
flashinfer_trtllm |
SM100 |
FlashInfer TensorRT-LLM backend |
+
+ marlin |
+ SM80-SM90 |
+ Weight-only W4A16 fallback for NVFP4 checkpoints |
+
-When FlashInfer is unavailable for NVFP4, the SGLang CUTLASS kernel is used as an automatic fallback.
+On Blackwell, when FlashInfer is unavailable for NVFP4, the SGLang CUTLASS kernel is used as an automatic fallback. On SM80-SM90, `auto` selects Marlin for NVFP4.
## Offline Quantization
diff --git a/docs_new/docs/advanced_features/server_arguments.mdx b/docs_new/docs/advanced_features/server_arguments.mdx
index c2a9f282e..6cf6d6343 100644
--- a/docs_new/docs/advanced_features/server_arguments.mdx
+++ b/docs_new/docs/advanced_features/server_arguments.mdx
@@ -1220,9 +1220,9 @@ Please consult the documentation below and [server_args.py](https://github.com/s
| `--fp4-gemm-backend` |
- Choose the runner backend for NVFP4 GEMM operations. Options: 'auto' (default; selects flashinfer_cudnn on SM120, flashinfer_cutedsl on SM100, flashinfer_cutlass otherwise), 'cutlass' (SGLang CUTLASS kernel), 'flashinfer_cutlass' (FlashInfer CUTLASS backend), 'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), 'flashinfer_cutedsl' (FlashInfer CuTe DSL backend), 'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling). All FlashInfer backends fall back to sgl-kernel CUTLASS when FlashInfer is unavailable. |
+ Choose the runner backend for NVFP4 GEMM operations. Options: 'auto' (default; selects flashinfer_cudnn on SM120, flashinfer_cutedsl on SM100, marlin on SM80-SM90, flashinfer_cutlass otherwise), 'cutlass' (SGLang CUTLASS kernel), 'flashinfer_cutlass' (FlashInfer CUTLASS backend), 'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), 'flashinfer_cutedsl' (FlashInfer CuTe DSL backend), 'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling), 'marlin' (weight-only W4A16 fallback for SM80-SM90). All FlashInfer backends fall back to sgl-kernel CUTLASS when FlashInfer is unavailable. |
auto |
- auto, cutlass, flashinfer_cudnn, flashinfer_cutedsl, flashinfer_cutlass, flashinfer_trtllm |
+ auto, cutlass, flashinfer_cudnn, flashinfer_cutedsl, flashinfer_cutlass, flashinfer_trtllm, marlin |
| `--disable-flashinfer-autotune` |
@@ -1452,7 +1452,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
` --moe-runner-backend` |
Choose the runner backend for MoE. |
` auto` |
- auto, deep_gemm, triton, triton_kernel, flashinfer_trtllm, flashinfer_trtllm_routed, flashinfer_cutlass, flashinfer_mxfp4, flashinfer_cutedsl, cutlass |
+ auto, deep_gemm, triton, triton_kernel, flashinfer_trtllm, flashinfer_trtllm_routed, flashinfer_cutlass, flashinfer_mxfp4, flashinfer_cutedsl, cutlass, marlin |
| ` --flashinfer-mxfp4-moe-precision` |
diff --git a/python/sglang/jit_kernel/csrc/gemm/marlin/marlin_template.h b/python/sglang/jit_kernel/csrc/gemm/marlin/marlin_template.h
index 6c4112e63..ca41d44d5 100644
--- a/python/sglang/jit_kernel/csrc/gemm/marlin/marlin_template.h
+++ b/python/sglang/jit_kernel/csrc/gemm/marlin/marlin_template.h
@@ -483,12 +483,15 @@ __global__ void Marlin(
constexpr int b_sh_stage = b_sh_stride * thread_k_blocks;
constexpr int b_sh_wr_iters = b_sh_stage / b_sh_wr_delta;
- // Scale sizes/strides without act_order
- int s_gl_stride = prob_n / 8;
- constexpr int s_sh_stride = 16 * thread_n_blocks / 8;
- constexpr int s_tb_groups = !has_act_order && group_blocks != -1 && group_blocks < thread_k_blocks
- ? thread_k_blocks / group_blocks / (w_type == host::kFE2M1f ? 2 : 1)
- : 1;
+ // Scale sizes/strides without act_order.
+ // NVFP4 packs FP8 (8-bit) scales into shared/global memory at twice the
+ // density of the half-precision scale path, so the strides scale with the
+ // element size.
+ constexpr bool is_8bit_scale = w_type == host::kFE2M1f;
+ int s_gl_stride = prob_n / (is_8bit_scale ? 16 : 8);
+ constexpr int s_sh_stride = 16 * thread_n_blocks / (is_8bit_scale ? 16 : 8);
+ constexpr int s_tb_groups =
+ !has_act_order && group_blocks != -1 && group_blocks < thread_k_blocks ? thread_k_blocks / group_blocks : 1;
constexpr int s_sh_stage = s_tb_groups * s_sh_stride;
int s_gl_rd_delta = s_gl_stride;
@@ -540,8 +543,7 @@ __global__ void Marlin(
if constexpr (group_blocks == -1) {
s_gl_rd = s_sh_stride * slice_col + threadIdx.x;
} else {
- s_gl_rd = s_gl_stride * ((thread_k_blocks * slice_row) / group_blocks) / (w_type == host::kFE2M1f ? 2 : 1) +
- s_sh_stride * slice_col + threadIdx.x;
+ s_gl_rd = s_gl_stride * ((thread_k_blocks * slice_row) / group_blocks) + s_sh_stride * slice_col + threadIdx.x;
}
}
auto s_sh_wr = threadIdx.x;
@@ -563,15 +565,7 @@ __global__ void Marlin(
// we scale a `half2` tile in column-major layout in the former and in
// row-major in the latter case.
int s_sh_rd;
- if constexpr (group_blocks != -1 && w_type == host::kFE2M1f) {
- auto warp_id = threadIdx.x / 32;
- int n_warps = thread_n_blocks / 4;
- int warp_row = warp_id / n_warps;
-
- s_sh_rd = 8 * ((threadIdx.x / 32) % (thread_n_blocks / 4)) + (threadIdx.x % 32) / 4;
- s_sh_rd = s_sh_rd * 2 + warp_row % 2;
-
- } else if constexpr (group_blocks != -1)
+ if constexpr (group_blocks != -1)
s_sh_rd = 8 * ((threadIdx.x / 32) % (thread_n_blocks / 4)) + (threadIdx.x % 32) / 4;
else if constexpr (group_blocks == -1 && (m_block_size_8 || (has_zp && !dequant_skip_flop)))
s_sh_rd = 8 * ((threadIdx.x / 32) % (thread_n_blocks / 4)) + (threadIdx.x % 32) / 8;
@@ -876,7 +870,7 @@ __global__ void Marlin(
cur_k += k_iter_size * (k % b_sh_wr_iters);
int k_blocks = cur_k / 16;
- int cur_group_id = k_blocks / (group_blocks * (w_type == host::kFE2M1f ? 2 : 1));
+ int cur_group_id = k_blocks / group_blocks;
int4* sh_s_stage = sh_s + s_sh_stage * pipe;
diff --git a/python/sglang/jit_kernel/tests/test_gptq_marlin.py b/python/sglang/jit_kernel/tests/test_gptq_marlin.py
index 379d1132a..495bfcf14 100644
--- a/python/sglang/jit_kernel/tests/test_gptq_marlin.py
+++ b/python/sglang/jit_kernel/tests/test_gptq_marlin.py
@@ -1,13 +1,27 @@
import sys
+from types import SimpleNamespace
import pytest
import torch
from sgl_kernel.scalar_type import scalar_types
from sglang.jit_kernel.gptq_marlin import gptq_marlin_gemm
-from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace
+from sglang.srt.layers.quantization.marlin_utils import (
+ check_marlin_supported,
+ marlin_make_workspace,
+)
+from sglang.srt.layers.quantization.marlin_utils_fp4 import (
+ apply_fp4_marlin_linear,
+ nvfp4_marlin_process_global_scale,
+ prepare_nvfp4_layer_for_marlin,
+)
+from sglang.srt.utils.common import is_sm80_supported, is_sm90_supported
from sglang.test.ci.ci_register import register_cuda_ci
-from sglang.test.test_marlin_utils import awq_marlin_quantize, marlin_quantize
+from sglang.test.test_marlin_utils import (
+ awq_marlin_quantize,
+ make_nvfp4_weight_and_ref,
+ marlin_quantize,
+)
register_cuda_ci(est_time=13, suite="base-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
@@ -101,5 +115,80 @@ def test_gptq_marlin_gemm(
assert max_diff < 0.04
+@pytest.mark.skip(reason="Skip, test pass locally but compiling takes too long in CI")
+@pytest.mark.skipif(
+ not (is_sm80_supported() or is_sm90_supported()),
+ reason="NVFP4 Marlin fallback tests require CUDA SM8X/SM9X",
+)
+@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
+def test_nvfp4_marlin_support_and_scale_transforms_sm80_sm90(dtype):
+ major, minor = torch.cuda.get_device_capability()
+ capability = major * 10 + minor
+ assert check_marlin_supported(
+ scalar_types.float4_e2m1f,
+ group_size=16,
+ has_zp=False,
+ device_capability=capability,
+ )
+
+ global_scale = torch.tensor(1.0, dtype=dtype, device="cuda")
+ actual_global_scale = nvfp4_marlin_process_global_scale(global_scale)
+ assert actual_global_scale.is_cuda
+ assert actual_global_scale.ndim == 1
+ assert actual_global_scale.numel() == 1
+ if dtype == torch.float16:
+ assert actual_global_scale.item() == 128.0
+ else:
+ assert actual_global_scale.item() == 2.0**119
+
+
+@pytest.mark.skip(reason="Skip, test pass locally but compiling takes too long in CI")
+@pytest.mark.skipif(
+ not (is_sm80_supported() or is_sm90_supported()),
+ reason="NVFP4 Marlin dense numeric test requires CUDA SM80, SM86, or SM90",
+)
+@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
+def test_nvfp4_marlin_dense_matches_dequant_reference(dtype):
+ torch.manual_seed(0)
+
+ size_m = 17
+ size_k = 256
+ size_n = 192
+ group_size = 16
+
+ a_input = torch.randn((size_m, size_k), dtype=dtype, device="cuda") / 10
+ fp4_weight, scales, global_scale, weight_ref = make_nvfp4_weight_and_ref(
+ size_n, size_k, dtype, group_size=group_size
+ )
+
+ layer = torch.nn.Module()
+ layer.quant_config = SimpleNamespace(group_size=group_size)
+ layer.output_size_per_partition = size_n
+ layer.input_size_per_partition = size_k
+ layer.params_dtype = dtype
+ layer.weight = torch.nn.Parameter(fp4_weight, requires_grad=False)
+ layer.weight_scale = torch.nn.Parameter(scales, requires_grad=False)
+ layer.weight_global_scale = torch.nn.Parameter(
+ global_scale.reshape(1), requires_grad=False
+ )
+ prepare_nvfp4_layer_for_marlin(layer)
+
+ output = apply_fp4_marlin_linear(
+ a_input,
+ layer.weight,
+ layer.weight_scale,
+ layer.weight_global_scale,
+ layer.workspace,
+ size_n,
+ size_k,
+ use_fp32_reduce=True,
+ )
+
+ output_ref = torch.matmul(a_input, weight_ref.T)
+ torch.cuda.synchronize()
+
+ torch.testing.assert_close(output, output_ref, rtol=0.04, atol=0.04)
+
+
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
diff --git a/python/sglang/jit_kernel/tests/test_moe_wna16_marlin.py b/python/sglang/jit_kernel/tests/test_moe_wna16_marlin.py
index 54852f1a9..96e6a0ba3 100644
--- a/python/sglang/jit_kernel/tests/test_moe_wna16_marlin.py
+++ b/python/sglang/jit_kernel/tests/test_moe_wna16_marlin.py
@@ -1,5 +1,6 @@
import itertools
import sys
+from types import SimpleNamespace
import pytest
import torch
@@ -7,8 +8,17 @@ from sgl_kernel.scalar_type import scalar_types
from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm
from sglang.srt.layers.moe.fused_moe_triton import moe_align_block_size
+from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import fused_marlin_moe
+from sglang.srt.layers.quantization.marlin_utils_fp4 import (
+ prepare_moe_nvfp4_layer_for_marlin,
+)
+from sglang.srt.utils.common import is_sm80_supported, is_sm90_supported
from sglang.test.ci.ci_register import register_cuda_ci
-from sglang.test.test_marlin_utils import awq_marlin_quantize, marlin_quantize
+from sglang.test.test_marlin_utils import (
+ awq_marlin_quantize,
+ make_nvfp4_weight_and_ref,
+ marlin_quantize,
+)
register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
@@ -339,5 +349,267 @@ def test_moe_wna16_marlin_gemm(
torch.testing.assert_close(c_jit, c_aot, rtol=0, atol=0)
+@pytest.mark.skip(reason="Skip, test pass locally but compiling takes too long in CI")
+@pytest.mark.skipif(
+ not (is_sm80_supported() or is_sm90_supported()),
+ reason="Non-gated NVFP4 Marlin fallback test requires CUDA SM8X/SM9X",
+)
+def test_fused_marlin_moe_non_gated_relu2():
+ torch.manual_seed(0)
+
+ m = 17
+ n = 128
+ k = 256
+ e = 4
+ topk = 2
+ dtype = torch.float16
+ group_size = 128
+ quant_type = scalar_types.uint4b8
+
+ hidden_states = torch.randn((m, k), device="cuda", dtype=dtype) / 10
+ w_ref1, qweight1, scales1, zeros1, g_idx1, sort_indices1 = _setup_moe_weights(
+ e, n, k, quant_type, group_size, False, dtype
+ )
+ w_ref2, qweight2, scales2, zeros2, g_idx2, sort_indices2 = _setup_moe_weights(
+ e, k, n, quant_type, group_size, False, dtype
+ )
+
+ router_logits = torch.randn((m, e), device="cuda", dtype=dtype)
+ score_softmax = torch.softmax(router_logits, dim=-1, dtype=torch.float32)
+ topk_weights, topk_ids = torch.topk(score_softmax, topk)
+
+ output = fused_marlin_moe(
+ hidden_states=hidden_states,
+ w1=qweight1,
+ w2=qweight2,
+ w1_scale=scales1,
+ w2_scale=scales2,
+ gating_output=router_logits,
+ topk_weights=topk_weights,
+ topk_ids=topk_ids,
+ g_idx1=g_idx1,
+ g_idx2=g_idx2,
+ sort_indices1=sort_indices1,
+ sort_indices2=sort_indices2,
+ w1_zeros=zeros1,
+ w2_zeros=zeros2,
+ num_bits=4,
+ is_k_full=True,
+ routed_scaling_factor=1.0,
+ activation="relu2",
+ is_gated=False,
+ )
+
+ output_ref = torch.zeros_like(hidden_states)
+ for token_idx in range(m):
+ for route_idx in range(topk):
+ expert_id = topk_ids[token_idx, route_idx]
+ intermediate = hidden_states[token_idx] @ w_ref1[expert_id].T
+ intermediate = torch.square(torch.relu(intermediate))
+ routed = intermediate @ w_ref2[expert_id].T
+ output_ref[token_idx] += routed * topk_weights[token_idx, route_idx]
+
+ torch.cuda.synchronize()
+ torch.testing.assert_close(output, output_ref, rtol=0.04, atol=0.04)
+
+
+@pytest.mark.skip(reason="Skip, test pass locally but compiling takes too long in CI")
+@pytest.mark.skipif(
+ not (is_sm80_supported() or is_sm90_supported()),
+ reason="NVFP4 Marlin MoE padding test requires CUDA SM8X/SM9X",
+)
+def test_fused_marlin_moe_nvfp4_non_gated_padded_intermediate_launches():
+ torch.manual_seed(0)
+
+ m = 17
+ intermediate_size = 192
+ hidden_size = 256
+ e = 4
+ topk = 2
+ dtype = torch.bfloat16
+ nvfp4_group_size = 16
+
+ layer = torch.nn.Module()
+ layer.quant_config = SimpleNamespace(group_size=nvfp4_group_size)
+ layer.moe_runner_config = SimpleNamespace(is_gated=False)
+ layer.params_dtype = dtype
+ layer.intermediate_size_per_partition = intermediate_size
+ layer.w13_weight = torch.nn.Parameter(
+ torch.randint(
+ 0,
+ 256,
+ (e, intermediate_size, hidden_size // 2),
+ device="cuda",
+ dtype=torch.uint8,
+ ),
+ requires_grad=False,
+ )
+ layer.w2_weight = torch.nn.Parameter(
+ torch.randint(
+ 0,
+ 256,
+ (e, hidden_size, intermediate_size // 2),
+ device="cuda",
+ dtype=torch.uint8,
+ ),
+ requires_grad=False,
+ )
+ layer.w13_weight_scale = torch.nn.Parameter(
+ torch.rand(
+ (e, intermediate_size, hidden_size // nvfp4_group_size),
+ device="cuda",
+ dtype=dtype,
+ ),
+ requires_grad=False,
+ )
+ layer.w2_weight_scale = torch.nn.Parameter(
+ torch.rand(
+ (e, hidden_size, intermediate_size // nvfp4_group_size),
+ device="cuda",
+ dtype=dtype,
+ ),
+ requires_grad=False,
+ )
+ layer.w13_weight_scale_2 = torch.nn.Parameter(
+ torch.ones((e,), device="cuda", dtype=dtype), requires_grad=False
+ )
+ layer.w2_weight_scale_2 = torch.nn.Parameter(
+ torch.ones((e,), device="cuda", dtype=dtype), requires_grad=False
+ )
+ prepare_moe_nvfp4_layer_for_marlin(layer)
+
+ assert layer.w13_weight.shape[1] * 16 == 256
+ assert layer.w2_weight.shape[1] * 16 == 256
+
+ hidden_states = torch.randn((m, hidden_size), device="cuda", dtype=dtype) / 10
+
+ score = torch.randn((m, e), device="cuda", dtype=dtype)
+ score_softmax = torch.softmax(score, dim=-1, dtype=torch.float32)
+ topk_weights, topk_ids = torch.topk(score_softmax, topk)
+
+ out = fused_marlin_moe(
+ hidden_states=hidden_states,
+ w1=layer.w13_weight,
+ w2=layer.w2_weight,
+ w1_scale=layer.w13_weight_scale,
+ w2_scale=layer.w2_weight_scale,
+ gating_output=score,
+ topk_weights=topk_weights,
+ topk_ids=topk_ids,
+ w1_global_scale=layer.w13_weight_scale_2,
+ w2_global_scale=layer.w2_weight_scale_2,
+ workspace=layer.workspace,
+ num_bits=4,
+ is_k_full=True,
+ routed_scaling_factor=1.0,
+ activation="relu2",
+ is_gated=False,
+ )
+
+ torch.cuda.synchronize()
+ assert out.shape == (m, hidden_size)
+
+
+@pytest.mark.skip(reason="Skip, test pass locally but compiling takes too long in CI")
+@pytest.mark.skipif(
+ not (is_sm80_supported() or is_sm90_supported()),
+ reason="NVFP4 Marlin MoE numeric test requires CUDA SM80, SM86, or SM90",
+)
+def test_fused_marlin_moe_nvfp4_non_gated_matches_dequant_reference():
+ torch.manual_seed(0)
+
+ m = 17
+ intermediate_size = 192
+ hidden_size = 256
+ e = 4
+ topk = 2
+ dtype = torch.bfloat16
+ group_size = 16
+ routed_scaling_factor = 1.0
+
+ w13_packed_l, w13_scales_l, w13_gscale_l, w13_ref_l = [], [], [], []
+ w2_packed_l, w2_scales_l, w2_gscale_l, w2_ref_l = [], [], [], []
+ for _ in range(e):
+ packed, scales, gscale, ref = make_nvfp4_weight_and_ref(
+ intermediate_size, hidden_size, dtype, group_size=group_size
+ )
+ w13_packed_l.append(packed)
+ w13_scales_l.append(scales)
+ w13_gscale_l.append(gscale)
+ w13_ref_l.append(ref)
+
+ packed, scales, gscale, ref = make_nvfp4_weight_and_ref(
+ hidden_size, intermediate_size, dtype, group_size=group_size
+ )
+ w2_packed_l.append(packed)
+ w2_scales_l.append(scales)
+ w2_gscale_l.append(gscale)
+ w2_ref_l.append(ref)
+
+ layer = torch.nn.Module()
+ layer.quant_config = SimpleNamespace(group_size=group_size)
+ layer.moe_runner_config = SimpleNamespace(is_gated=False)
+ layer.params_dtype = dtype
+ layer.intermediate_size_per_partition = intermediate_size
+ layer.w13_weight = torch.nn.Parameter(
+ torch.stack(w13_packed_l), requires_grad=False
+ )
+ layer.w2_weight = torch.nn.Parameter(torch.stack(w2_packed_l), requires_grad=False)
+ layer.w13_weight_scale = torch.nn.Parameter(
+ torch.stack(w13_scales_l), requires_grad=False
+ )
+ layer.w2_weight_scale = torch.nn.Parameter(
+ torch.stack(w2_scales_l), requires_grad=False
+ )
+ layer.w13_weight_scale_2 = torch.nn.Parameter(
+ torch.stack(w13_gscale_l), requires_grad=False
+ )
+ layer.w2_weight_scale_2 = torch.nn.Parameter(
+ torch.stack(w2_gscale_l), requires_grad=False
+ )
+ prepare_moe_nvfp4_layer_for_marlin(layer)
+
+ # Scale activations down so relu² doesn't blow up intermediate magnitudes;
+ # this keeps output values small so tighter element-wise tolerance is realistic.
+ hidden_states = torch.randn((m, hidden_size), device="cuda", dtype=dtype) / 20
+ router_logits = torch.randn((m, e), device="cuda", dtype=dtype)
+ score_softmax = torch.softmax(router_logits, dim=-1, dtype=torch.float32)
+ topk_weights, topk_ids = torch.topk(score_softmax, topk)
+
+ output = fused_marlin_moe(
+ hidden_states=hidden_states,
+ w1=layer.w13_weight,
+ w2=layer.w2_weight,
+ w1_scale=layer.w13_weight_scale,
+ w2_scale=layer.w2_weight_scale,
+ gating_output=router_logits,
+ topk_weights=topk_weights,
+ topk_ids=topk_ids,
+ w1_global_scale=layer.w13_weight_scale_2,
+ w2_global_scale=layer.w2_weight_scale_2,
+ workspace=layer.workspace,
+ num_bits=4,
+ is_k_full=True,
+ routed_scaling_factor=routed_scaling_factor,
+ activation="relu2",
+ is_gated=False,
+ )
+
+ w13_ref = torch.stack(w13_ref_l)
+ w2_ref = torch.stack(w2_ref_l)
+ output_ref = torch.zeros_like(hidden_states)
+ for token_idx in range(m):
+ for route_idx in range(topk):
+ expert_id = topk_ids[token_idx, route_idx]
+ intermediate = hidden_states[token_idx] @ w13_ref[expert_id].T
+ intermediate = torch.square(torch.relu(intermediate))
+ routed = intermediate @ w2_ref[expert_id].T
+ output_ref[token_idx] += routed * topk_weights[token_idx, route_idx]
+ output_ref *= routed_scaling_factor
+
+ torch.cuda.synchronize()
+ torch.testing.assert_close(output, output_ref, rtol=0.05, atol=0.25)
+
+
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
diff --git a/python/sglang/jit_kernel/utils.py b/python/sglang/jit_kernel/utils.py
index 1f6783ad0..4b04ce1c5 100644
--- a/python/sglang/jit_kernel/utils.py
+++ b/python/sglang/jit_kernel/utils.py
@@ -1,10 +1,12 @@
from __future__ import annotations
import functools
+import hashlib
import importlib.util
import logging
import os
import pathlib
+import re
from contextlib import contextmanager
from dataclasses import dataclass
from typing import (
@@ -65,6 +67,36 @@ def _make_wrapper(tup: Tuple[str, str]) -> str:
return f"TVM_FFI_DLL_EXPORT_TYPED_FUNC({export_name}, ({kernel_name}));"
+_LOCAL_INCLUDE_RE = re.compile(r'^\s*#\s*include\s+"([^"]+)"', re.MULTILINE)
+
+
+def _local_jit_source_hash(source_files: List[str]) -> str:
+ """Hash JIT source contents so TVM-FFI cache keys track included headers."""
+ digest = hashlib.sha256()
+ seen: set[pathlib.Path] = set()
+ stack = [pathlib.Path(path).resolve() for path in source_files]
+
+ while stack:
+ path = stack.pop()
+ if path in seen or not path.is_file():
+ continue
+ seen.add(path)
+
+ data = path.read_bytes()
+ digest.update(str(path).encode())
+ digest.update(b"\0")
+ digest.update(data)
+ digest.update(b"\0")
+
+ text = data.decode("utf-8", errors="ignore")
+ for include in _LOCAL_INCLUDE_RE.findall(text):
+ include_path = (path.parent / include).resolve()
+ if include_path.is_file():
+ stack.append(include_path)
+
+ return digest.hexdigest()[:16]
+
+
@cache_once
def _resolve_kernel_path() -> pathlib.Path:
cur_dir = pathlib.Path(__file__).parent.resolve()
@@ -201,6 +233,8 @@ def load_jit(
extra_include_paths += _REGISTERED_DEPENDENCIES[dep]()
module_name = "sgl_kernel_jit_" + "_".join(str(arg) for arg in args)
+ if cpp_files or cuda_files:
+ module_name += "_" + _local_jit_source_hash(cpp_files + cuda_files)
if header_only:
cpp_wrappers = cpp_wrappers or []
cuda_wrappers = cuda_wrappers or []
diff --git a/python/sglang/srt/arg_groups/nemotron_h_hook.py b/python/sglang/srt/arg_groups/nemotron_h_hook.py
index 87b568f3c..8a31d04a2 100644
--- a/python/sglang/srt/arg_groups/nemotron_h_hook.py
+++ b/python/sglang/srt/arg_groups/nemotron_h_hook.py
@@ -1,7 +1,7 @@
import logging
from typing import TYPE_CHECKING
-from sglang.srt.utils.common import is_sm100_supported
+from sglang.srt.utils.common import get_device_capability, is_cuda, is_sm100_supported
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -36,6 +36,19 @@ def apply_nemotron_h_defaults(server_args: "ServerArgs", model_arch: str) -> Non
"Use flashinfer_trtllm as MoE runner backend on sm100 for "
f"{model_arch}"
)
+ elif (
+ (
+ model_config.quantization in ("modelopt_fp4", "modelopt_mixed")
+ or server_args.quantization == "modelopt_fp4"
+ )
+ and is_cuda()
+ and (8, 0) <= get_device_capability() < (10, 0)
+ ):
+ server_args.moe_runner_backend = "marlin"
+ logger.info(
+ "Use marlin as MoE runner backend on SM80-SM90 for "
+ f"{model_arch} {model_config.quantization}"
+ )
else:
server_args.moe_runner_backend = "flashinfer_cutlass"
diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
index f62fe7278..d48c7f675 100644
--- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_marlin_moe.py
@@ -15,14 +15,19 @@ if _is_cuda:
from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm
-def get_scalar_type(num_bits: int, has_zp: bool, scales: Optional[torch.Tensor] = None):
+def get_scalar_type(
+ num_bits: int,
+ has_zp: bool,
+ scales: Optional[torch.Tensor] = None,
+ global_scale: Optional[torch.Tensor] = None,
+):
from sgl_kernel.scalar_type import scalar_types
if (
not has_zp
and num_bits == 4
and scales is not None
- and scales.dtype == torch.float8_e8m0fnu
+ and (scales.dtype == torch.float8_e8m0fnu or global_scale is not None)
):
return scalar_types.float4_e2m1f
if has_zp:
@@ -66,12 +71,16 @@ def fused_marlin_moe(
sort_indices2: Optional[torch.Tensor] = None,
w1_zeros: Optional[torch.Tensor] = None,
w2_zeros: Optional[torch.Tensor] = None,
+ w1_global_scale: Optional[torch.Tensor] = None,
+ w2_global_scale: Optional[torch.Tensor] = None,
workspace: Optional[torch.Tensor] = None,
num_bits: int = 8,
is_k_full: bool = True,
inplace: bool = False,
routed_scaling_factor: Optional[float] = None,
clamp_limit: Optional[float] = None,
+ activation: str = "silu",
+ is_gated: bool = True,
) -> torch.Tensor:
"""
This function computes a Mixture of Experts (MoE) layer using two sets of
@@ -118,12 +127,19 @@ def fused_marlin_moe(
and w1_scale.dtype == torch.float8_e8m0fnu
and w2_scale.dtype == torch.float8_e8m0fnu
)
+ is_nvfp4_marlin = (
+ num_bits == 4
+ and w1_zeros is None
+ and w2_zeros is None
+ and w1_global_scale is not None
+ and w2_global_scale is not None
+ )
if is_mxfp4_marlin:
assert hidden_states.dtype == torch.bfloat16, (
"MXFP4 Marlin with E8M0 scales is only instantiated for bfloat16 "
f"activations, got {hidden_states.dtype}"
)
- else:
+ elif not is_nvfp4_marlin:
assert (
hidden_states.dtype == w1_scale.dtype
), f"moe_wna16_marlin_gemm assumes hidden_states.dtype ({hidden_states.dtype}) == w1_scale.dtype ({w1_scale.dtype})"
@@ -136,6 +152,7 @@ def fused_marlin_moe(
E = w1.shape[0]
N = w2.shape[1] * 16
topk = topk_ids.shape[1]
+ gemm1_n = 2 * N if is_gated else N
# M block size selection logic
# TODO: tune this further for specific models
@@ -160,8 +177,12 @@ def fused_marlin_moe(
max_workspace_size, dtype=torch.int, device=device, requires_grad=False
)
- scalar_type1 = get_scalar_type(num_bits, w1_zeros is not None, w1_scale)
- scalar_type2 = get_scalar_type(num_bits, w2_zeros is not None, w2_scale)
+ scalar_type1 = get_scalar_type(
+ num_bits, w1_zeros is not None, w1_scale, w1_global_scale
+ )
+ scalar_type2 = get_scalar_type(
+ num_bits, w2_zeros is not None, w2_scale, w2_global_scale
+ )
intermediate_cache2 = torch.empty(
(M * topk_ids.shape[1], N),
@@ -169,12 +190,12 @@ def fused_marlin_moe(
dtype=hidden_states.dtype,
)
intermediate_cache13 = torch.empty(
- (M * topk_ids.shape[1] * max(2 * N, K),),
+ (M * topk_ids.shape[1] * max(gemm1_n, K),),
device=hidden_states.device,
dtype=hidden_states.dtype,
)
- intermediate_cache1 = intermediate_cache13[: M * topk_ids.shape[1] * 2 * N]
- intermediate_cache1 = intermediate_cache1.view(-1, 2 * N)
+ intermediate_cache1 = intermediate_cache13[: M * topk_ids.shape[1] * gemm1_n]
+ intermediate_cache1 = intermediate_cache1.view(-1, gemm1_n)
intermediate_cache3 = intermediate_cache13[: M * topk_ids.shape[1] * K]
intermediate_cache3 = intermediate_cache3.view(-1, K)
@@ -189,7 +210,7 @@ def fused_marlin_moe(
w1,
None, # b_bias_or_none
w1_scale,
- None, # global_scale_or_none
+ w1_global_scale,
w1_zeros,
g_idx1,
sort_indices1,
@@ -204,7 +225,7 @@ def fused_marlin_moe(
is_ep=expert_map is not None,
b_q_type=scalar_type1,
size_m=M,
- size_n=2 * N,
+ size_n=gemm1_n,
size_k=K,
is_k_full=is_k_full,
use_atomic_add=use_atomic_add,
@@ -212,14 +233,20 @@ def fused_marlin_moe(
is_zp_float=False,
)
- if clamp_limit is not None:
+ if activation == "silu" and is_gated and clamp_limit is not None:
swiglu_limit_func(
intermediate_cache2,
- intermediate_cache1.view(-1, 2 * N),
+ intermediate_cache1.view(-1, gemm1_n),
clamp_limit,
)
+ elif activation == "silu" and is_gated:
+ silu_and_mul(intermediate_cache1.view(-1, gemm1_n), intermediate_cache2)
+ elif activation == "silu" and not is_gated:
+ intermediate_cache2 = F.silu(intermediate_cache1.view(-1, N))
+ elif activation == "relu2" and not is_gated:
+ intermediate_cache2 = torch.square(F.relu(intermediate_cache1.view(-1, N)))
else:
- silu_and_mul(intermediate_cache1.view(-1, 2 * N), intermediate_cache2)
+ raise ValueError(f"Unsupported activation: {activation=}, with {is_gated=}")
if expert_map is not None:
intermediate_cache3.zero_()
@@ -230,7 +257,7 @@ def fused_marlin_moe(
w2,
None, # b_bias_or_none
w2_scale,
- None, # global_scale_or_none
+ w2_global_scale,
w2_zeros,
g_idx2,
sort_indices2,
diff --git a/python/sglang/srt/layers/moe/moe_runner/marlin.py b/python/sglang/srt/layers/moe/moe_runner/marlin.py
index 4e335f330..bedf48ffb 100644
--- a/python/sglang/srt/layers/moe/moe_runner/marlin.py
+++ b/python/sglang/srt/layers/moe/moe_runner/marlin.py
@@ -71,6 +71,9 @@ class MarlinMoeQuantInfo(MoeQuantInfo):
# Optional
expert_map: Optional[torch.Tensor] = None
+ global_num_experts: int = -1
+ w13_global_scale: Optional[torch.Tensor] = None
+ w2_global_scale: Optional[torch.Tensor] = None
@register_fused_func("none", "marlin")
@@ -87,7 +90,12 @@ def fused_experts_none_to_marlin(
hidden_states = dispatch_output.hidden_states
topk_output = dispatch_output.topk_output
- assert runner_config.activation == "silu", "Only SiLU activation is supported."
+ if runner_config.is_gated:
+ assert runner_config.activation == "silu", "Only gated SiLU is supported."
+ elif runner_config.activation not in {"silu", "relu2"}:
+ raise ValueError(
+ f"Unsupported Marlin MoE activation: {runner_config.activation}"
+ )
if (
MARLIN_MOE_WORKSPACE is None
@@ -124,6 +132,7 @@ def fused_experts_none_to_marlin(
gating_output=topk_output.router_logits,
topk_weights=topk_output.topk_weights,
topk_ids=topk_output.topk_ids,
+ global_num_experts=quant_info.global_num_experts,
expert_map=quant_info.expert_map,
g_idx1=quant_info.w13_g_idx,
g_idx2=quant_info.w2_g_idx,
@@ -131,12 +140,16 @@ def fused_experts_none_to_marlin(
sort_indices2=quant_info.w2_g_idx_sort_indices,
w1_zeros=quant_info.w13_qzeros,
w2_zeros=quant_info.w2_qzeros,
+ w1_global_scale=quant_info.w13_global_scale,
+ w2_global_scale=quant_info.w2_global_scale,
workspace=MARLIN_MOE_WORKSPACE,
num_bits=quant_info.weight_bits,
is_k_full=quant_info.is_k_full,
inplace=marlin_inplace,
routed_scaling_factor=runner_config.routed_scaling_factor,
clamp_limit=runner_config.swiglu_limit,
+ activation=runner_config.activation,
+ is_gated=runner_config.is_gated,
).to(hidden_states.dtype)
return StandardCombineInput(
diff --git a/python/sglang/srt/layers/quantization/fp4_utils.py b/python/sglang/srt/layers/quantization/fp4_utils.py
index a7a64f25e..ad8ee9058 100644
--- a/python/sglang/srt/layers/quantization/fp4_utils.py
+++ b/python/sglang/srt/layers/quantization/fp4_utils.py
@@ -6,7 +6,12 @@ from typing import TYPE_CHECKING, Optional
import torch
-from sglang.srt.utils.common import is_sm100_supported, is_sm120_supported
+from sglang.srt.utils.common import (
+ get_device_capability,
+ is_cuda,
+ is_sm100_supported,
+ is_sm120_supported,
+)
from sglang.srt.utils.custom_op import register_custom_op_from_extern
if TYPE_CHECKING:
@@ -95,6 +100,7 @@ class Fp4GemmRunnerBackend(Enum):
FLASHINFER_CUTEDSL = "flashinfer_cutedsl"
FLASHINFER_CUTLASS = "flashinfer_cutlass"
FLASHINFER_TRTLLM = "flashinfer_trtllm"
+ MARLIN = "marlin"
def is_auto(self) -> bool:
return self == Fp4GemmRunnerBackend.AUTO
@@ -114,6 +120,9 @@ class Fp4GemmRunnerBackend(Enum):
def is_flashinfer_cutedsl(self) -> bool:
return self == Fp4GemmRunnerBackend.FLASHINFER_CUTEDSL
+ def is_marlin(self) -> bool:
+ return self == Fp4GemmRunnerBackend.MARLIN
+
def is_flashinfer(self) -> bool:
return self.value.startswith("flashinfer_")
@@ -151,6 +160,8 @@ def initialize_fp4_gemm_config(server_args: ServerArgs) -> None:
backend = "flashinfer_cudnn"
elif is_sm100_supported():
backend = "flashinfer_cutedsl"
+ elif is_cuda() and (10, 0) > get_device_capability() >= (8, 0):
+ backend = "marlin"
else:
backend = "flashinfer_cutlass"
diff --git a/python/sglang/srt/layers/quantization/marlin_utils.py b/python/sglang/srt/layers/quantization/marlin_utils.py
index 0c754b9b2..0cfc0298d 100644
--- a/python/sglang/srt/layers/quantization/marlin_utils.py
+++ b/python/sglang/srt/layers/quantization/marlin_utils.py
@@ -56,6 +56,8 @@ GPTQ_MARLIN_MIN_THREAD_K = 128
GPTQ_MARLIN_MAX_PARALLEL = 16
MARLIN_SUPPORTED_GROUP_SIZES = [-1, 32, 64, 128]
+# NVFP SUPPORT 16, while MXFP4 supports 32 and 16.
+FP4_MARLIN_SUPPORTED_GROUP_SIZES = [16, 32]
# In case there is a performance issue with Marlin, the variable below can be
# changed to False, which allows Marlin to perform global reductions in fp16
@@ -137,11 +139,15 @@ def _check_marlin_supported(
f"are supported (for group_size = {group_size}, "
f"device_capability = {device_capability}, zp = {has_zp}).",
)
- if group_size is None or group_size not in MARLIN_SUPPORTED_GROUP_SIZES:
+ if quant_type == scalar_types.float4_e2m1f:
+ allowed_group_sizes = FP4_MARLIN_SUPPORTED_GROUP_SIZES
+ else:
+ allowed_group_sizes = MARLIN_SUPPORTED_GROUP_SIZES
+ if group_size is None or group_size not in allowed_group_sizes:
return (
False,
- f"Marlin does not support group_size = {group_size}. "
- f"Only group_sizes = {MARLIN_SUPPORTED_GROUP_SIZES} "
+ f"Marlin does not support group_size = {group_size} for "
+ f"quant_type = {quant_type}. Only group_sizes = {allowed_group_sizes} "
"are supported.",
)
@@ -239,8 +245,13 @@ def check_moe_marlin_supports_layer(layer: FusedMoE, group_size: int) -> bool:
intermediate_size_per_partition = layer.intermediate_size_per_partition
# apply_router_weight_on_input is not supported for moe marlin
supports_router_weight = not layer.moe_runner_config.apply_router_weight_on_input
- # moe marlin requires the activation to be silu
- supports_activation = layer.moe_runner_config.activation == "silu"
+ if layer.moe_runner_config.is_gated:
+ supports_activation = layer.moe_runner_config.activation == "silu"
+ else:
+ supports_activation = layer.moe_runner_config.activation in {
+ "silu",
+ "relu2",
+ }
# gate-up: (n, k) = (intermediate_size_per_partition * 2, hidden_size)
# down: (n, k) = (hidden_size, intermediate_size_per_partition)
diff --git a/python/sglang/srt/layers/quantization/marlin_utils_fp4.py b/python/sglang/srt/layers/quantization/marlin_utils_fp4.py
index 4ebca81d9..c0adbb49f 100644
--- a/python/sglang/srt/layers/quantization/marlin_utils_fp4.py
+++ b/python/sglang/srt/layers/quantization/marlin_utils_fp4.py
@@ -3,17 +3,182 @@ from __future__ import annotations
import torch
from sglang.srt.layers.quantization.marlin_utils import (
+ USE_FP32_REDUCE_DEFAULT,
marlin_make_workspace,
marlin_permute_bias,
marlin_permute_scales,
+ should_use_atomic_add_reduce,
)
+from sglang.srt.layers.quantization.utils import get_scalar_types
from sglang.srt.utils import is_cuda
+from sglang.srt.utils.custom_op import register_custom_op
_is_cuda = is_cuda()
if _is_cuda:
+ from sglang.jit_kernel.gptq_marlin import gptq_marlin_gemm
from sglang.jit_kernel.gptq_marlin_repack import gptq_marlin_repack
+ScalarType, scalar_types = get_scalar_types()
+
+
+def nvfp4_marlin_process_scales(marlin_scales: torch.Tensor) -> torch.Tensor:
+ if not (marlin_scales >= 0).all():
+ # NVFP4 ModelOpt scales are expected to be non-negative. Keep this as
+ # a warning so unusual checkpoints can still load for diagnosis.
+ import logging
+
+ logging.getLogger(__name__).warning_once(
+ "NVFP4 Marlin assumes non-negative scales, but negative scales "
+ "were found. Accuracy may be degraded."
+ )
+
+ marlin_scales = marlin_scales.to(torch.half)
+ marlin_scales = marlin_scales.view(-1, 4)[:, [0, 2, 1, 3]].view(
+ marlin_scales.size(0), -1
+ )
+ marlin_scales = (marlin_scales * (2**7)).view(torch.int16) << 1
+ marlin_scales = marlin_scales.view(torch.float8_e4m3fn)
+ return marlin_scales[:, 1::2].contiguous()
+
+
+def nvfp4_marlin_process_global_scale(global_scale: torch.Tensor) -> torch.Tensor:
+ assert global_scale.dtype in [torch.half, torch.bfloat16]
+ global_scale_shape = global_scale.shape
+ fp4_exponent = 2
+ if global_scale.dtype == torch.half:
+ target_exponent = 5
+ elif global_scale.dtype == torch.bfloat16:
+ target_exponent = 8
+ exponent_bias = 2 ** (target_exponent - 1) - 2 ** (fp4_exponent - 1)
+ global_scale = global_scale * (2.0 ** (exponent_bias - 7))
+ if global_scale_shape == torch.Size([]):
+ global_scale = global_scale.reshape(1)
+ return global_scale
+
+
+def fake_apply_fp4_marlin_linear(
+ input: torch.Tensor,
+ weight: torch.Tensor,
+ weight_scale: torch.Tensor,
+ weight_global_scale: torch.Tensor,
+ workspace: torch.Tensor,
+ size_n: int,
+ size_k: int,
+ bias: torch.Tensor | None = None,
+ use_fp32_reduce: bool = USE_FP32_REDUCE_DEFAULT,
+) -> torch.Tensor:
+ del weight, weight_scale, weight_global_scale, workspace, size_k, bias
+ out_shape = input.shape[:-1] + (size_n,)
+ return input.new_empty(out_shape)
+
+
+@register_custom_op(fake_impl=fake_apply_fp4_marlin_linear)
+def apply_fp4_marlin_linear(
+ input: torch.Tensor,
+ weight: torch.Tensor,
+ weight_scale: torch.Tensor,
+ weight_global_scale: torch.Tensor,
+ workspace: torch.Tensor,
+ size_n: int,
+ size_k: int,
+ bias: torch.Tensor | None = None,
+ use_fp32_reduce: bool = USE_FP32_REDUCE_DEFAULT,
+) -> torch.Tensor:
+ if input.dtype not in (torch.float16, torch.bfloat16):
+ raise RuntimeError("NVFP4 Marlin requires FP16 or BF16 activations.")
+
+ reshaped_x = input.reshape(-1, input.shape[-1])
+ out_shape = input.shape[:-1] + (size_n,)
+
+ use_atomic_add = should_use_atomic_add_reduce(
+ m=reshaped_x.size(0),
+ n=size_n,
+ k=size_k,
+ device=input.device,
+ dtype=input.dtype,
+ )
+
+ output = gptq_marlin_gemm(
+ a=reshaped_x,
+ c=None,
+ b_q_weight=weight,
+ b_scales=weight_scale,
+ global_scale=weight_global_scale,
+ b_zeros=None,
+ g_idx=None,
+ perm=None,
+ workspace=workspace,
+ b_q_type=scalar_types.float4_e2m1f,
+ size_m=reshaped_x.size(0),
+ size_n=size_n,
+ size_k=size_k,
+ is_k_full=True,
+ use_atomic_add=use_atomic_add,
+ use_fp32_reduce=use_fp32_reduce,
+ )
+
+ if bias is not None:
+ output.add_(bias)
+
+ return output.reshape(out_shape)
+
+
+def prepare_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
+ if getattr(layer, "quant_config", None) is not None:
+ group_size = layer.quant_config.group_size
+ if group_size != 16:
+ raise ValueError(f"NVFP4 Marlin requires group_size=16, got {group_size}.")
+
+ part_size_n = layer.output_size_per_partition
+ part_size_k = layer.input_size_per_partition
+ param_dtype = getattr(layer, "params_dtype", getattr(layer, "orig_dtype", None))
+ if param_dtype not in (torch.float16, torch.bfloat16):
+ raise RuntimeError("NVFP4 Marlin requires FP16 or BF16 activation dtype.")
+
+ assert layer.weight.shape == (part_size_n, part_size_k // 2)
+
+ if part_size_n % 64 != 0:
+ raise ValueError(
+ f"NVFP4 Marlin requires output_size_per_partition to be a multiple of 64, "
+ f"got {part_size_n}."
+ )
+
+ device = layer.weight.device
+ layer.workspace = marlin_make_workspace(device)
+
+ perm = torch.empty(0, dtype=torch.int, device=device)
+ qweight = layer.weight.view(torch.int32).T.contiguous()
+ marlin_qweight = gptq_marlin_repack(
+ b_q_weight=qweight,
+ perm=perm,
+ size_k=part_size_k,
+ size_n=part_size_n,
+ num_bits=4,
+ )
+ layer.weight = torch.nn.Parameter(marlin_qweight, requires_grad=False)
+
+ weight_scale = layer.weight_scale.T.contiguous().to(param_dtype)
+ weight_scale = marlin_permute_scales(
+ s=weight_scale,
+ size_k=part_size_k,
+ size_n=part_size_n,
+ group_size=16,
+ )
+ weight_scale = nvfp4_marlin_process_scales(weight_scale)
+ layer.weight_scale = torch.nn.Parameter(weight_scale, requires_grad=False)
+
+ weight_global_scale = layer.weight_global_scale.to(param_dtype)
+ weight_global_scale = nvfp4_marlin_process_global_scale(weight_global_scale)
+ layer.weight_global_scale = torch.nn.Parameter(
+ weight_global_scale, requires_grad=False
+ )
+
+ if hasattr(layer, "bias") and layer.bias is not None:
+ assert layer.bias.shape == (part_size_n,)
+ bias = marlin_permute_bias(layer.bias)
+ layer.bias = torch.nn.Parameter(bias, requires_grad=False)
+
def mxfp4_marlin_process_scales(
marlin_scales: torch.Tensor,
@@ -161,3 +326,121 @@ def prepare_moe_mxfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
layer.w2_weight_bias = torch.nn.Parameter(
_permute_bias(w2_bias_data), requires_grad=False
)
+
+
+def prepare_moe_nvfp4_layer_for_marlin(layer: torch.nn.Module) -> None:
+ if layer.quant_config.group_size != 16:
+ raise ValueError(
+ f"NVFP4 Marlin MoE requires group_size=16, got {layer.quant_config.group_size}."
+ )
+
+ w13 = layer.w13_weight.data
+ w2 = layer.w2_weight.data
+ w13_scale = layer.w13_weight_scale.data
+ w2_scale = layer.w2_weight_scale.data
+ w13_global_scale = layer.w13_weight_scale_2.data
+ w2_global_scale = layer.w2_weight_scale_2.data
+ w13_bias = getattr(layer, "w13_bias", None)
+ w2_bias = getattr(layer, "w2_bias", None)
+
+ num_experts = w13.shape[0]
+ num_shards = 2 if layer.moe_runner_config.is_gated else 1
+ intermediate_size = layer.intermediate_size_per_partition
+ hidden_size = w13.shape[2] * 2
+ param_dtype = layer.params_dtype
+ if param_dtype not in (torch.float16, torch.bfloat16):
+ raise RuntimeError("NVFP4 Marlin MoE requires FP16 or BF16 activations.")
+
+ device = w13.device
+ layer.workspace = marlin_make_workspace(device, 4)
+ perm = torch.empty(0, dtype=torch.int, device=device)
+
+ if not layer.moe_runner_config.is_gated:
+ padded_intermediate_size = ((intermediate_size + 127) // 128) * 128
+ intermediate_size_pad = padded_intermediate_size - intermediate_size
+ if intermediate_size_pad:
+ w13 = torch.nn.functional.pad(w13, (0, 0, 0, intermediate_size_pad))
+ w13_scale = torch.nn.functional.pad(
+ w13_scale, (0, 0, 0, intermediate_size_pad)
+ )
+ w2 = torch.nn.functional.pad(w2, (0, intermediate_size_pad // 2, 0, 0))
+ w2_scale = torch.nn.functional.pad(
+ w2_scale, (0, intermediate_size_pad // 16)
+ )
+ if w13_bias is not None:
+ w13_bias = torch.nn.functional.pad(w13_bias, (0, intermediate_size_pad))
+ intermediate_size = padded_intermediate_size
+
+ def _repack_weight(weight: torch.Tensor, is_w13: bool) -> torch.Tensor:
+ if is_w13:
+ size_n, size_k = intermediate_size * num_shards, hidden_size
+ else:
+ size_n, size_k = hidden_size, intermediate_size
+ assert weight.shape == (num_experts, size_n, size_k // 2)
+
+ tensor_list = []
+ for i in range(num_experts):
+ qweight = weight[i].view(torch.int32).T.contiguous()
+ marlin_qweight = gptq_marlin_repack(
+ b_q_weight=qweight,
+ perm=perm,
+ size_k=size_k,
+ size_n=size_n,
+ num_bits=4,
+ )
+ tensor_list.append(marlin_qweight)
+ return torch.stack(tensor_list)
+
+ def _permute_scales(scales: torch.Tensor, is_w13: bool) -> torch.Tensor:
+ scales = scales.to(param_dtype)
+ if is_w13:
+ size_n, size_k = intermediate_size * num_shards, hidden_size
+ else:
+ size_n, size_k = hidden_size, intermediate_size
+
+ tensor_list = []
+ for i in range(num_experts):
+ scale = scales[i].T.contiguous()
+ marlin_scales = marlin_permute_scales(
+ s=scale,
+ size_k=size_k,
+ size_n=size_n,
+ group_size=16,
+ )
+ tensor_list.append(nvfp4_marlin_process_scales(marlin_scales))
+ return torch.stack(tensor_list)
+
+ def _process_global_scale(global_scale: torch.Tensor) -> torch.Tensor:
+ return nvfp4_marlin_process_global_scale(global_scale.to(param_dtype))
+
+ def _permute_bias(bias: torch.Tensor | None) -> torch.Tensor | None:
+ if bias is None:
+ return None
+ tensor_list = []
+ for i in range(num_experts):
+ tensor_list.append(marlin_permute_bias(bias[i].to(param_dtype)))
+ return torch.stack(tensor_list)
+
+ layer.w13_weight = torch.nn.Parameter(
+ _repack_weight(w13, True), requires_grad=False
+ )
+ layer.w2_weight = torch.nn.Parameter(_repack_weight(w2, False), requires_grad=False)
+ layer.w13_weight_scale = torch.nn.Parameter(
+ _permute_scales(w13_scale, True), requires_grad=False
+ )
+ layer.w2_weight_scale = torch.nn.Parameter(
+ _permute_scales(w2_scale, False), requires_grad=False
+ )
+ layer.w13_weight_scale_2 = torch.nn.Parameter(
+ _process_global_scale(w13_global_scale), requires_grad=False
+ )
+ layer.w2_weight_scale_2 = torch.nn.Parameter(
+ _process_global_scale(w2_global_scale), requires_grad=False
+ )
+
+ if w13_bias is not None:
+ layer.w13_bias = torch.nn.Parameter(
+ _permute_bias(w13_bias), requires_grad=False
+ )
+ if w2_bias is not None:
+ layer.w2_bias = torch.nn.Parameter(_permute_bias(w2_bias), requires_grad=False)
diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py
index c91643e83..c27efda65 100755
--- a/python/sglang/srt/layers/quantization/modelopt_quant.py
+++ b/python/sglang/srt/layers/quantization/modelopt_quant.py
@@ -48,6 +48,11 @@ from sglang.srt.layers.quantization.fp8_utils import (
is_blackwell_supported,
)
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
+from sglang.srt.layers.quantization.marlin_utils_fp4 import (
+ apply_fp4_marlin_linear,
+ prepare_moe_nvfp4_layer_for_marlin,
+ prepare_nvfp4_layer_for_marlin,
+)
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.layers.quantization.utils import (
convert_to_channelwise,
@@ -59,6 +64,7 @@ from sglang.srt.layers.quantization.utils import (
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.utils import alias_or_bind_derived_param, copy_or_rebind_param
from sglang.srt.utils.common import (
+ get_device_capability,
is_cuda,
is_sm120_supported,
next_power_of_2,
@@ -1190,7 +1196,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
@classmethod
def get_min_capability(cls) -> int:
- return 100
+ return 80
@staticmethod
def common_group_size(cfg: dict) -> int:
@@ -1367,6 +1373,8 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
+ layer.params_dtype = params_dtype
+ layer.quant_config = self.quant_config
if input_size_per_partition % 16 != 0:
raise ValueError(
"Unsupported model when in features size is not multiple of 16"
@@ -1434,6 +1442,23 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
# Store original output size before any padding
layer.output_size_per_partition = layer.weight.shape[0]
+ if get_fp4_gemm_runner_backend().is_marlin():
+ if self.quant_config.group_size != 16:
+ raise ValueError(
+ f"NVFP4 Marlin requires group_size=16, got {self.quant_config.group_size}."
+ )
+ copy_or_rebind_param(layer, "input_global_scale", input_scale_2)
+ copy_or_rebind_param(layer, "weight_global_scale", weight_scale_2)
+ prepare_nvfp4_layer_for_marlin(layer)
+ layer.weights_padding_cols = 0
+ return
+
+ if not is_blackwell_supported():
+ raise ValueError(
+ "ModelOpt NVFP4 native dense GEMM backends require SM100+. "
+ "Use --fp4-gemm-backend marlin on SM80-SM90."
+ )
+
if get_fp4_gemm_runner_backend().is_flashinfer_trtllm():
# FlashInfer TRTLLM FP4 GEMM requires a different weight layout.
# FlashInfer provides nvfp4_quantize to quantize + shuffle the
@@ -1579,6 +1604,18 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
+ if get_fp4_gemm_runner_backend().is_marlin():
+ return apply_fp4_marlin_linear(
+ input=x,
+ weight=layer.weight,
+ weight_scale=layer.weight_scale,
+ weight_global_scale=layer.weight_global_scale,
+ workspace=layer.workspace,
+ size_n=layer.output_size_per_partition,
+ size_k=layer.input_size_per_partition,
+ bias=bias,
+ )
+
# `_accepts_prequantized_fp4` is the explicit opt-in so an accidental
# tuple from unrelated code can't silently bypass quantization.
if getattr(layer, "_accepts_prequantized_fp4", False) and isinstance(x, tuple):
@@ -1639,11 +1676,17 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
def __init__(self, quant_config: ModelOptFp4Config):
self.quant_config = quant_config
- if not is_blackwell_supported():
+ moe_runner_backend = get_moe_runner_backend()
+ if moe_runner_backend.is_auto() and is_cuda():
+ capability = get_device_capability()
+ use_marlin_fallback = (8, 0) <= capability < (10, 0)
+ else:
+ use_marlin_fallback = moe_runner_backend.is_marlin()
+ if not is_blackwell_supported() and not use_marlin_fallback:
raise ValueError(
"Current platform does not support NVFP4"
- " quantization. Please use Blackwell and"
- " above."
+ " quantization with the selected MoE backend. Please use "
+ "Blackwell and above, or use moe_runner_backend=marlin on SM80+."
)
self.enable_flashinfer_trtllm_moe = (
get_moe_runner_backend().is_flashinfer_trtllm()
@@ -1844,6 +1887,18 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
else:
w13_weight_scale_2 = layer.w13_weight_scale_2[:]
+ moe_runner_backend = getattr(
+ self, "_moe_runner_backend", get_moe_runner_backend()
+ )
+ if moe_runner_backend.is_marlin():
+ copy_or_rebind_param(
+ layer,
+ "w13_weight_scale_2",
+ w13_weight_scale_2.contiguous(),
+ )
+ prepare_moe_nvfp4_layer_for_marlin(layer)
+ return
+
# Calculate input scales based on strategy
if self.enable_flashinfer_cutlass_moe or self.enable_flashinfer_trtllm_moe:
w13_input_scale = layer.w13_input_scale.max().to(torch.float32)
@@ -2115,9 +2170,14 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
moe_runner_backend = get_moe_runner_backend()
if moe_runner_backend.is_auto():
- # TRTLLM is currently the most performant and tested FP4 MoE
- # backend, so use it as the default.
- moe_runner_backend = MoeRunnerBackend.FLASHINFER_TRTLLM
+ if is_cuda() and (8, 0) <= get_device_capability() < (10, 0):
+ moe_runner_backend = MoeRunnerBackend.MARLIN
+ else:
+ # TRTLLM is currently the most performant and tested FP4 MoE
+ # backend, so use it as the default.
+ moe_runner_backend = MoeRunnerBackend.FLASHINFER_TRTLLM
+
+ self._moe_runner_backend = moe_runner_backend
if moe_runner_backend.is_flashinfer_cutedsl():
import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl # noqa: F401 – triggers @register_fused_func
@@ -2137,12 +2197,42 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
# tuple). Defer per-attribute access to the branches that actually
# consume them.
activation = self.moe_runner_config.activation
+ moe_runner_backend = getattr(
+ self, "_moe_runner_backend", get_moe_runner_backend()
+ )
assert (
activation in _SUPPORTED_ACT_STRS
), f"{activation=} not in supported {_SUPPORTED_ACT_STRS}"
moe_runner_config = self.moe_runner_config
+ if moe_runner_backend.is_marlin():
+ from sglang.srt.layers.moe.moe_runner.marlin import MarlinMoeQuantInfo
+
+ expert_map = None
+ global_num_experts = -1
+ if hasattr(layer, "dispatcher") and hasattr(
+ layer.dispatcher, "local_expert_mapping"
+ ):
+ expert_map = layer.dispatcher.local_expert_mapping
+ if expert_map is not None:
+ global_num_experts = self.moe_runner_config.num_experts
+
+ quant_info = MarlinMoeQuantInfo(
+ w13_qweight=layer.w13_weight,
+ w2_qweight=layer.w2_weight,
+ w13_scales=layer.w13_weight_scale,
+ w2_scales=layer.w2_weight_scale,
+ w13_g_idx_sort_indices=None,
+ w2_g_idx_sort_indices=None,
+ weight_bits=4,
+ w13_global_scale=layer.w13_weight_scale_2,
+ w2_global_scale=layer.w2_weight_scale_2,
+ expert_map=expert_map,
+ global_num_experts=global_num_experts,
+ )
+ return self.runner.run(dispatch_output, quant_info)
+
# FlashInfer TRTLLM FP4 path
if self.enable_flashinfer_trtllm_moe and hasattr(layer, "g1_scale_c"):
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index 6cf1314e5..a0d707fa3 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -255,6 +255,7 @@ FP4_GEMM_RUNNER_BACKEND_CHOICES = [
"flashinfer_cutedsl",
"flashinfer_cutlass",
"flashinfer_trtllm",
+ "marlin",
]
RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu", "slru", "priority"]
@@ -5688,12 +5689,13 @@ class ServerArgs:
default=ServerArgs.fp4_gemm_runner_backend,
dest="fp4_gemm_runner_backend",
help="Choose the runner backend for NVFP4 GEMM operations. "
- "Options: 'auto' (default; selects flashinfer_cudnn on SM120, flashinfer_cutedsl on SM100, flashinfer_cutlass otherwise), "
+ "Options: 'auto' (default; selects flashinfer_cudnn on SM120, flashinfer_cutedsl on SM100, marlin on SM80-SM90, flashinfer_cutlass otherwise), "
"'cutlass' (SGLang CUTLASS kernel), "
"'flashinfer_cutlass' (FlashInfer CUTLASS backend), "
"'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), "
"'flashinfer_cutedsl' (FlashInfer CuTe DSL backend), "
- "'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling). ",
+ "'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling), "
+ "'marlin' (weight-only W4A16 fallback for SM80+). ",
)
parser.add_argument(
"--disable-flashinfer-autotune",
diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py
index ec428fb24..bdf496316 100644
--- a/python/sglang/srt/utils/common.py
+++ b/python/sglang/srt/utils/common.py
@@ -283,6 +283,11 @@ is_sm100_supported = lru_cache(maxsize=1)(
_check_cuda_device_version, device_capability_majors=[10], cuda_version=(12, 8)
)
)
+is_sm80_supported = lru_cache(maxsize=1)(
+ partial(
+ _check_cuda_device_version, device_capability_majors=[8], cuda_version=(11, 0)
+ )
+)
is_sm90_supported = lru_cache(maxsize=1)(
partial(
_check_cuda_device_version, device_capability_majors=[9], cuda_version=(12, 3)
diff --git a/python/sglang/test/test_marlin_utils.py b/python/sglang/test/test_marlin_utils.py
index 834f677e8..bd53c8332 100644
--- a/python/sglang/test/test_marlin_utils.py
+++ b/python/sglang/test/test_marlin_utils.py
@@ -170,3 +170,47 @@ def awq_marlin_quantize(w: torch.Tensor, quant_type: ScalarType, group_size: int
res_list[i] = res_list[i].to(w.device)
return res_list
+
+
+def make_nvfp4_weight_and_ref(
+ size_n: int,
+ size_k: int,
+ dtype: torch.dtype,
+ group_size: int = 16,
+ device: str = "cuda",
+):
+ """Build a random NVFP4-quantized weight and its FP dequantized reference.
+
+ Returns:
+ fp4_weight: (size_n, size_k // 2) uint8, two packed FP4 (E2M1) values per byte
+ scales: (size_n, size_k // group_size) FP8 E4M3 per-group scales
+ global_scale: scalar in `dtype`, the FP16/BF16 outer scale
+ weight_ref: (size_n, size_k) tensor in `dtype` = dequantized weight
+ """
+ fp4_weight = torch.randint(
+ 0, 256, (size_n, size_k // 2), dtype=torch.uint8, device=device
+ )
+ scale_source = torch.randn((size_n, size_k), dtype=dtype, device=device)
+ # /6 = FP4 (E2M1) max; /448 = FP8 (E4M3) max — sets each level to its dtype's full range.
+ scales = scale_source.view(size_n, -1, group_size).abs().max(-1)[0] / 6
+ global_scale = scales.max() / 448
+ scales = (scales / global_scale).to(torch.float8_e4m3fn)
+
+ def _unpack(byte_view: torch.Tensor) -> torch.Tensor:
+ # Convert 4-bit E2M1 nibble (in upper bits of a uint8) to FP8 E4M3 bit pattern.
+ unpacked = (byte_view & 0b10000000) | ((byte_view & 0b01110000) >> 2)
+ return unpacked.view(torch.float8_e4m3fn).to(dtype) * (2**6)
+
+ part_low = _unpack(fp4_weight)
+ part_high = _unpack(fp4_weight << 4)
+
+ weight_ref = torch.cat([part_high.unsqueeze(2), part_low.unsqueeze(2)], 2).view(
+ size_n, size_k
+ )
+ weight_ref = (
+ weight_ref
+ * global_scale.to(dtype)
+ * scales.repeat_interleave(group_size, 1).to(dtype)
+ )
+
+ return fp4_weight, scales, global_scale, weight_ref
diff --git a/test/lm_eval_configs/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.yaml b/test/lm_eval_configs/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.yaml
new file mode 100644
index 000000000..0d8a8dfa2
--- /dev/null
+++ b/test/lm_eval_configs/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.yaml
@@ -0,0 +1,13 @@
+model_name: "nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4"
+tasks:
+- name: "gsm8k"
+ metrics:
+ - name: "exact_match,strict-match"
+ value: 0.847
+ - name: "exact_match,flexible-extract"
+ value: 0.444
+limit: 1319
+num_concurrent: 128
+num_fewshot: 5
+apply_chat_template: false
+fewshot_as_multiturn: true
diff --git a/test/manual/models/test_nvidia_nemotron_3_nano_archived.py b/test/manual/models/test_nvidia_nemotron_3_nano_archived.py
index 6d9ffe49c..edb9f6c54 100644
--- a/test/manual/models/test_nvidia_nemotron_3_nano_archived.py
+++ b/test/manual/models/test_nvidia_nemotron_3_nano_archived.py
@@ -7,6 +7,7 @@ Run with `python3 test/manual/models/test_nvidia_nemotron_3_nano_archived.py`.
import unittest
+from sglang.srt.utils import is_sm80_supported, is_sm90_supported
from sglang.test.kits.lm_eval_kit import LMEvalMixin
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
@@ -43,5 +44,27 @@ class TestNvidiaNemotron3Nano30BBF16FlashInfer(LMEvalMixin, DefaultServerBase):
] + NEMOTRON_3_NANO_THINKING_ARGS
+@unittest.skip("Skip, test pass locally but compiling takes too long in CI")
+@unittest.skipIf(
+ not (is_sm80_supported() or is_sm90_supported()),
+ "NVFP4 Marlin fallback test requires CUDA SM8X/SM9X",
+)
+class TestNvidiaNemotron3Nano30BNVFP4Marlin(LMEvalMixin, DefaultServerBase):
+ """Test Nemotron-3-Nano-30B NVFP4 model with the Marlin path."""
+
+ model = "nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4"
+ model_config_name = "lm_eval_configs/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.yaml"
+ other_args = [
+ "--tp-size",
+ "1",
+ "--quantization",
+ "modelopt_fp4",
+ "--fp4-gemm-backend",
+ "marlin",
+ "--moe-runner-backend",
+ "marlin",
+ ] + NEMOTRON_3_NANO_THINKING_ARGS
+
+
if __name__ == "__main__":
unittest.main()
diff --git a/test/registered/unit/model_loader/test_modelopt_loader.py b/test/registered/unit/model_loader/test_modelopt_loader.py
index fe3724b67..35ea5609b 100644
--- a/test/registered/unit/model_loader/test_modelopt_loader.py
+++ b/test/registered/unit/model_loader/test_modelopt_loader.py
@@ -15,6 +15,7 @@ from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
from sglang.srt.layers.quantization.modelopt_quant import (
+ ModelOptFp4Config,
ModelOptMixedPrecisionConfig,
)
from sglang.srt.model_loader.loader import ModelOptModelLoader
@@ -646,7 +647,10 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
)
def test_mixed_precision_uses_nvfp4_min_capability(self):
- self.assertEqual(ModelOptMixedPrecisionConfig.get_min_capability(), 100)
+ self.assertEqual(
+ ModelOptMixedPrecisionConfig.get_min_capability(),
+ ModelOptFp4Config.get_min_capability(),
+ )
def test_mixed_precision_quant_layer_resolution_after_mapping(self):
quant_config = ModelOptMixedPrecisionConfig.from_config(