Feat/add w4a16 moe support to nemotron (#25655)
This commit is contained in:
@@ -142,10 +142,10 @@ The following table summarizes quantization method support across NVIDIA and AMD
|
|||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>modelopt_fp4</code></td>
|
<td><code>modelopt_fp4</code></td>
|
||||||
<td>Yes (Blackwell/SM100+)</td>
|
<td>Yes (SM80-SM90 via Marlin; SM100+ native FP4)</td>
|
||||||
<td>No</td>
|
<td>No</td>
|
||||||
<td>No</td>
|
<td>No</td>
|
||||||
<td><a href="https://github.com/NVIDIA/Model-Optimizer">NVIDIA ModelOpt</a>; native FP4 on Blackwell (B200, GB200)</td>
|
<td><a href="https://github.com/NVIDIA/Model-Optimizer">NVIDIA ModelOpt</a>; use Marlin W4A16 fallback on Ampere/Hopper and native FP4 backends on Blackwell</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>petit_nvfp4</code></td>
|
<td><code>petit_nvfp4</code></td>
|
||||||
@@ -257,8 +257,8 @@ Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. Wh
|
|||||||
<tbody>
|
<tbody>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>auto</code></td>
|
<td><code>auto</code></td>
|
||||||
<td>SM100/120</td>
|
<td>SM80+</td>
|
||||||
<td>Auto-selects: <code>flashinfer_cudnn</code> on SM120; <code>flashinfer_cutlass</code> on SM100</td>
|
<td>Auto-selects: <code>flashinfer_cudnn</code> on SM120; <code>flashinfer_cutedsl</code> on SM100; <code>marlin</code> on SM80-SM90; <code>flashinfer_cutlass</code> otherwise</td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>cutlass</code></td>
|
<td><code>cutlass</code></td>
|
||||||
@@ -275,15 +275,25 @@ Backend selection is supported only for **blockwise FP8** and **NVFP4** GEMM. Wh
|
|||||||
<td>SM100/120 (CUDA 13+, cuDNN 9.15+)</td>
|
<td>SM100/120 (CUDA 13+, cuDNN 9.15+)</td>
|
||||||
<td>FlashInfer cuDNN backend; used on SM120 for performance</td>
|
<td>FlashInfer cuDNN backend; used on SM120 for performance</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><code>flashinfer_cutedsl</code></td>
|
||||||
|
<td>SM100</td>
|
||||||
|
<td>FlashInfer CuTe DSL backend</td>
|
||||||
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td><code>flashinfer_trtllm</code></td>
|
<td><code>flashinfer_trtllm</code></td>
|
||||||
<td>SM100</td>
|
<td>SM100</td>
|
||||||
<td>FlashInfer TensorRT-LLM backend</td>
|
<td>FlashInfer TensorRT-LLM backend</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td><code>marlin</code></td>
|
||||||
|
<td>SM80-SM90</td>
|
||||||
|
<td>Weight-only W4A16 fallback for NVFP4 checkpoints</td>
|
||||||
|
</tr>
|
||||||
</tbody>
|
</tbody>
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
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
|
## Offline Quantization
|
||||||
|
|
||||||
|
|||||||
@@ -1220,9 +1220,9 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
|||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--fp4-gemm-backend`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--fp4-gemm-backend`</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Choose the runner backend for NVFP4 GEMM operations. Options: 'auto' (default; selects <code>flashinfer_cudnn</code> on SM120, <code>flashinfer_cutedsl</code> on SM100, <code>flashinfer_cutlass</code> 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.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Choose the runner backend for NVFP4 GEMM operations. Options: 'auto' (default; selects <code>flashinfer_cudnn</code> on SM120, <code>flashinfer_cutedsl</code> on SM100, <code>marlin</code> on SM80-SM90, <code>flashinfer_cutlass</code> 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.</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>auto</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>auto</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>cutlass</code>, <code>flashinfer_cudnn</code>, <code>flashinfer_cutedsl</code>, <code>flashinfer_cutlass</code>, <code>flashinfer_trtllm</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>cutlass</code>, <code>flashinfer_cudnn</code>, <code>flashinfer_cutedsl</code>, <code>flashinfer_cutlass</code>, <code>flashinfer_trtllm</code>, <code>marlin</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--disable-flashinfer-autotune`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--disable-flashinfer-autotune`</td>
|
||||||
@@ -1452,7 +1452,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
|||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --moe-runner-backend`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --moe-runner-backend`</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Choose the runner backend for MoE.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Choose the runner backend for MoE.</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>` auto`</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>` auto`</td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>deep_gemm</code>, <code>triton</code>, <code>triton_kernel</code>, <code>flashinfer_trtllm</code>, <code>flashinfer_trtllm_routed</code>, <code>flashinfer_cutlass</code>, <code>flashinfer_mxfp4</code>, <code>flashinfer_cutedsl</code>, <code>cutlass</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>deep_gemm</code>, <code>triton</code>, <code>triton_kernel</code>, <code>flashinfer_trtllm</code>, <code>flashinfer_trtllm_routed</code>, <code>flashinfer_cutlass</code>, <code>flashinfer_mxfp4</code>, <code>flashinfer_cutedsl</code>, <code>cutlass</code>, <code>marlin</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --flashinfer-mxfp4-moe-precision`</td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>` --flashinfer-mxfp4-moe-precision`</td>
|
||||||
|
|||||||
@@ -483,12 +483,15 @@ __global__ void Marlin(
|
|||||||
constexpr int b_sh_stage = b_sh_stride * thread_k_blocks;
|
constexpr int b_sh_stage = b_sh_stride * thread_k_blocks;
|
||||||
constexpr int b_sh_wr_iters = b_sh_stage / b_sh_wr_delta;
|
constexpr int b_sh_wr_iters = b_sh_stage / b_sh_wr_delta;
|
||||||
|
|
||||||
// Scale sizes/strides without act_order
|
// Scale sizes/strides without act_order.
|
||||||
int s_gl_stride = prob_n / 8;
|
// NVFP4 packs FP8 (8-bit) scales into shared/global memory at twice the
|
||||||
constexpr int s_sh_stride = 16 * thread_n_blocks / 8;
|
// density of the half-precision scale path, so the strides scale with the
|
||||||
constexpr int s_tb_groups = !has_act_order && group_blocks != -1 && group_blocks < thread_k_blocks
|
// element size.
|
||||||
? thread_k_blocks / group_blocks / (w_type == host::kFE2M1f ? 2 : 1)
|
constexpr bool is_8bit_scale = w_type == host::kFE2M1f;
|
||||||
: 1;
|
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;
|
constexpr int s_sh_stage = s_tb_groups * s_sh_stride;
|
||||||
int s_gl_rd_delta = s_gl_stride;
|
int s_gl_rd_delta = s_gl_stride;
|
||||||
|
|
||||||
@@ -540,8 +543,7 @@ __global__ void Marlin(
|
|||||||
if constexpr (group_blocks == -1) {
|
if constexpr (group_blocks == -1) {
|
||||||
s_gl_rd = s_sh_stride * slice_col + threadIdx.x;
|
s_gl_rd = s_sh_stride * slice_col + threadIdx.x;
|
||||||
} else {
|
} else {
|
||||||
s_gl_rd = s_gl_stride * ((thread_k_blocks * slice_row) / group_blocks) / (w_type == host::kFE2M1f ? 2 : 1) +
|
s_gl_rd = s_gl_stride * ((thread_k_blocks * slice_row) / group_blocks) + s_sh_stride * slice_col + threadIdx.x;
|
||||||
s_sh_stride * slice_col + threadIdx.x;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
auto s_sh_wr = 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
|
// we scale a `half2` tile in column-major layout in the former and in
|
||||||
// row-major in the latter case.
|
// row-major in the latter case.
|
||||||
int s_sh_rd;
|
int s_sh_rd;
|
||||||
if constexpr (group_blocks != -1 && w_type == host::kFE2M1f) {
|
if constexpr (group_blocks != -1)
|
||||||
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)
|
|
||||||
s_sh_rd = 8 * ((threadIdx.x / 32) % (thread_n_blocks / 4)) + (threadIdx.x % 32) / 4;
|
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)))
|
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;
|
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);
|
cur_k += k_iter_size * (k % b_sh_wr_iters);
|
||||||
|
|
||||||
int k_blocks = cur_k / 16;
|
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;
|
int4* sh_s_stage = sh_s + s_sh_stage * pipe;
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,27 @@
|
|||||||
import sys
|
import sys
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
from sgl_kernel.scalar_type import scalar_types
|
from sgl_kernel.scalar_type import scalar_types
|
||||||
|
|
||||||
from sglang.jit_kernel.gptq_marlin import gptq_marlin_gemm
|
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.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=13, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
|
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
|
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__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import itertools
|
import itertools
|
||||||
import sys
|
import sys
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
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.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 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.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=10, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
|
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)
|
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__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
|
|||||||
@@ -1,10 +1,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import functools
|
import functools
|
||||||
|
import hashlib
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import pathlib
|
import pathlib
|
||||||
|
import re
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import (
|
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}));"
|
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
|
@cache_once
|
||||||
def _resolve_kernel_path() -> pathlib.Path:
|
def _resolve_kernel_path() -> pathlib.Path:
|
||||||
cur_dir = pathlib.Path(__file__).parent.resolve()
|
cur_dir = pathlib.Path(__file__).parent.resolve()
|
||||||
@@ -201,6 +233,8 @@ def load_jit(
|
|||||||
extra_include_paths += _REGISTERED_DEPENDENCIES[dep]()
|
extra_include_paths += _REGISTERED_DEPENDENCIES[dep]()
|
||||||
|
|
||||||
module_name = "sgl_kernel_jit_" + "_".join(str(arg) for arg in args)
|
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:
|
if header_only:
|
||||||
cpp_wrappers = cpp_wrappers or []
|
cpp_wrappers = cpp_wrappers or []
|
||||||
cuda_wrappers = cuda_wrappers or []
|
cuda_wrappers = cuda_wrappers or []
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
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 "
|
"Use flashinfer_trtllm as MoE runner backend on sm100 for "
|
||||||
f"{model_arch}"
|
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:
|
else:
|
||||||
server_args.moe_runner_backend = "flashinfer_cutlass"
|
server_args.moe_runner_backend = "flashinfer_cutlass"
|
||||||
|
|
||||||
|
|||||||
@@ -15,14 +15,19 @@ if _is_cuda:
|
|||||||
from sglang.jit_kernel.moe_wna16_marlin import moe_wna16_marlin_gemm
|
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
|
from sgl_kernel.scalar_type import scalar_types
|
||||||
|
|
||||||
if (
|
if (
|
||||||
not has_zp
|
not has_zp
|
||||||
and num_bits == 4
|
and num_bits == 4
|
||||||
and scales is not None
|
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
|
return scalar_types.float4_e2m1f
|
||||||
if has_zp:
|
if has_zp:
|
||||||
@@ -66,12 +71,16 @@ def fused_marlin_moe(
|
|||||||
sort_indices2: Optional[torch.Tensor] = None,
|
sort_indices2: Optional[torch.Tensor] = None,
|
||||||
w1_zeros: Optional[torch.Tensor] = None,
|
w1_zeros: Optional[torch.Tensor] = None,
|
||||||
w2_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,
|
workspace: Optional[torch.Tensor] = None,
|
||||||
num_bits: int = 8,
|
num_bits: int = 8,
|
||||||
is_k_full: bool = True,
|
is_k_full: bool = True,
|
||||||
inplace: bool = False,
|
inplace: bool = False,
|
||||||
routed_scaling_factor: Optional[float] = None,
|
routed_scaling_factor: Optional[float] = None,
|
||||||
clamp_limit: Optional[float] = None,
|
clamp_limit: Optional[float] = None,
|
||||||
|
activation: str = "silu",
|
||||||
|
is_gated: bool = True,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
This function computes a Mixture of Experts (MoE) layer using two sets of
|
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 w1_scale.dtype == torch.float8_e8m0fnu
|
||||||
and w2_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:
|
if is_mxfp4_marlin:
|
||||||
assert hidden_states.dtype == torch.bfloat16, (
|
assert hidden_states.dtype == torch.bfloat16, (
|
||||||
"MXFP4 Marlin with E8M0 scales is only instantiated for bfloat16 "
|
"MXFP4 Marlin with E8M0 scales is only instantiated for bfloat16 "
|
||||||
f"activations, got {hidden_states.dtype}"
|
f"activations, got {hidden_states.dtype}"
|
||||||
)
|
)
|
||||||
else:
|
elif not is_nvfp4_marlin:
|
||||||
assert (
|
assert (
|
||||||
hidden_states.dtype == w1_scale.dtype
|
hidden_states.dtype == w1_scale.dtype
|
||||||
), f"moe_wna16_marlin_gemm assumes hidden_states.dtype ({hidden_states.dtype}) == w1_scale.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]
|
E = w1.shape[0]
|
||||||
N = w2.shape[1] * 16
|
N = w2.shape[1] * 16
|
||||||
topk = topk_ids.shape[1]
|
topk = topk_ids.shape[1]
|
||||||
|
gemm1_n = 2 * N if is_gated else N
|
||||||
|
|
||||||
# M block size selection logic
|
# M block size selection logic
|
||||||
# TODO: tune this further for specific models
|
# 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
|
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_type1 = get_scalar_type(
|
||||||
scalar_type2 = get_scalar_type(num_bits, w2_zeros is not None, w2_scale)
|
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(
|
intermediate_cache2 = torch.empty(
|
||||||
(M * topk_ids.shape[1], N),
|
(M * topk_ids.shape[1], N),
|
||||||
@@ -169,12 +190,12 @@ def fused_marlin_moe(
|
|||||||
dtype=hidden_states.dtype,
|
dtype=hidden_states.dtype,
|
||||||
)
|
)
|
||||||
intermediate_cache13 = torch.empty(
|
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,
|
device=hidden_states.device,
|
||||||
dtype=hidden_states.dtype,
|
dtype=hidden_states.dtype,
|
||||||
)
|
)
|
||||||
intermediate_cache1 = intermediate_cache13[: M * topk_ids.shape[1] * 2 * N]
|
intermediate_cache1 = intermediate_cache13[: M * topk_ids.shape[1] * gemm1_n]
|
||||||
intermediate_cache1 = intermediate_cache1.view(-1, 2 * N)
|
intermediate_cache1 = intermediate_cache1.view(-1, gemm1_n)
|
||||||
intermediate_cache3 = intermediate_cache13[: M * topk_ids.shape[1] * K]
|
intermediate_cache3 = intermediate_cache13[: M * topk_ids.shape[1] * K]
|
||||||
intermediate_cache3 = intermediate_cache3.view(-1, K)
|
intermediate_cache3 = intermediate_cache3.view(-1, K)
|
||||||
|
|
||||||
@@ -189,7 +210,7 @@ def fused_marlin_moe(
|
|||||||
w1,
|
w1,
|
||||||
None, # b_bias_or_none
|
None, # b_bias_or_none
|
||||||
w1_scale,
|
w1_scale,
|
||||||
None, # global_scale_or_none
|
w1_global_scale,
|
||||||
w1_zeros,
|
w1_zeros,
|
||||||
g_idx1,
|
g_idx1,
|
||||||
sort_indices1,
|
sort_indices1,
|
||||||
@@ -204,7 +225,7 @@ def fused_marlin_moe(
|
|||||||
is_ep=expert_map is not None,
|
is_ep=expert_map is not None,
|
||||||
b_q_type=scalar_type1,
|
b_q_type=scalar_type1,
|
||||||
size_m=M,
|
size_m=M,
|
||||||
size_n=2 * N,
|
size_n=gemm1_n,
|
||||||
size_k=K,
|
size_k=K,
|
||||||
is_k_full=is_k_full,
|
is_k_full=is_k_full,
|
||||||
use_atomic_add=use_atomic_add,
|
use_atomic_add=use_atomic_add,
|
||||||
@@ -212,14 +233,20 @@ def fused_marlin_moe(
|
|||||||
is_zp_float=False,
|
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(
|
swiglu_limit_func(
|
||||||
intermediate_cache2,
|
intermediate_cache2,
|
||||||
intermediate_cache1.view(-1, 2 * N),
|
intermediate_cache1.view(-1, gemm1_n),
|
||||||
clamp_limit,
|
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:
|
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:
|
if expert_map is not None:
|
||||||
intermediate_cache3.zero_()
|
intermediate_cache3.zero_()
|
||||||
@@ -230,7 +257,7 @@ def fused_marlin_moe(
|
|||||||
w2,
|
w2,
|
||||||
None, # b_bias_or_none
|
None, # b_bias_or_none
|
||||||
w2_scale,
|
w2_scale,
|
||||||
None, # global_scale_or_none
|
w2_global_scale,
|
||||||
w2_zeros,
|
w2_zeros,
|
||||||
g_idx2,
|
g_idx2,
|
||||||
sort_indices2,
|
sort_indices2,
|
||||||
|
|||||||
@@ -71,6 +71,9 @@ class MarlinMoeQuantInfo(MoeQuantInfo):
|
|||||||
|
|
||||||
# Optional
|
# Optional
|
||||||
expert_map: Optional[torch.Tensor] = None
|
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")
|
@register_fused_func("none", "marlin")
|
||||||
@@ -87,7 +90,12 @@ def fused_experts_none_to_marlin(
|
|||||||
hidden_states = dispatch_output.hidden_states
|
hidden_states = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
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 (
|
if (
|
||||||
MARLIN_MOE_WORKSPACE is None
|
MARLIN_MOE_WORKSPACE is None
|
||||||
@@ -124,6 +132,7 @@ def fused_experts_none_to_marlin(
|
|||||||
gating_output=topk_output.router_logits,
|
gating_output=topk_output.router_logits,
|
||||||
topk_weights=topk_output.topk_weights,
|
topk_weights=topk_output.topk_weights,
|
||||||
topk_ids=topk_output.topk_ids,
|
topk_ids=topk_output.topk_ids,
|
||||||
|
global_num_experts=quant_info.global_num_experts,
|
||||||
expert_map=quant_info.expert_map,
|
expert_map=quant_info.expert_map,
|
||||||
g_idx1=quant_info.w13_g_idx,
|
g_idx1=quant_info.w13_g_idx,
|
||||||
g_idx2=quant_info.w2_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,
|
sort_indices2=quant_info.w2_g_idx_sort_indices,
|
||||||
w1_zeros=quant_info.w13_qzeros,
|
w1_zeros=quant_info.w13_qzeros,
|
||||||
w2_zeros=quant_info.w2_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,
|
workspace=MARLIN_MOE_WORKSPACE,
|
||||||
num_bits=quant_info.weight_bits,
|
num_bits=quant_info.weight_bits,
|
||||||
is_k_full=quant_info.is_k_full,
|
is_k_full=quant_info.is_k_full,
|
||||||
inplace=marlin_inplace,
|
inplace=marlin_inplace,
|
||||||
routed_scaling_factor=runner_config.routed_scaling_factor,
|
routed_scaling_factor=runner_config.routed_scaling_factor,
|
||||||
clamp_limit=runner_config.swiglu_limit,
|
clamp_limit=runner_config.swiglu_limit,
|
||||||
|
activation=runner_config.activation,
|
||||||
|
is_gated=runner_config.is_gated,
|
||||||
).to(hidden_states.dtype)
|
).to(hidden_states.dtype)
|
||||||
|
|
||||||
return StandardCombineInput(
|
return StandardCombineInput(
|
||||||
|
|||||||
@@ -6,7 +6,12 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import torch
|
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
|
from sglang.srt.utils.custom_op import register_custom_op_from_extern
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -95,6 +100,7 @@ class Fp4GemmRunnerBackend(Enum):
|
|||||||
FLASHINFER_CUTEDSL = "flashinfer_cutedsl"
|
FLASHINFER_CUTEDSL = "flashinfer_cutedsl"
|
||||||
FLASHINFER_CUTLASS = "flashinfer_cutlass"
|
FLASHINFER_CUTLASS = "flashinfer_cutlass"
|
||||||
FLASHINFER_TRTLLM = "flashinfer_trtllm"
|
FLASHINFER_TRTLLM = "flashinfer_trtllm"
|
||||||
|
MARLIN = "marlin"
|
||||||
|
|
||||||
def is_auto(self) -> bool:
|
def is_auto(self) -> bool:
|
||||||
return self == Fp4GemmRunnerBackend.AUTO
|
return self == Fp4GemmRunnerBackend.AUTO
|
||||||
@@ -114,6 +120,9 @@ class Fp4GemmRunnerBackend(Enum):
|
|||||||
def is_flashinfer_cutedsl(self) -> bool:
|
def is_flashinfer_cutedsl(self) -> bool:
|
||||||
return self == Fp4GemmRunnerBackend.FLASHINFER_CUTEDSL
|
return self == Fp4GemmRunnerBackend.FLASHINFER_CUTEDSL
|
||||||
|
|
||||||
|
def is_marlin(self) -> bool:
|
||||||
|
return self == Fp4GemmRunnerBackend.MARLIN
|
||||||
|
|
||||||
def is_flashinfer(self) -> bool:
|
def is_flashinfer(self) -> bool:
|
||||||
return self.value.startswith("flashinfer_")
|
return self.value.startswith("flashinfer_")
|
||||||
|
|
||||||
@@ -151,6 +160,8 @@ def initialize_fp4_gemm_config(server_args: ServerArgs) -> None:
|
|||||||
backend = "flashinfer_cudnn"
|
backend = "flashinfer_cudnn"
|
||||||
elif is_sm100_supported():
|
elif is_sm100_supported():
|
||||||
backend = "flashinfer_cutedsl"
|
backend = "flashinfer_cutedsl"
|
||||||
|
elif is_cuda() and (10, 0) > get_device_capability() >= (8, 0):
|
||||||
|
backend = "marlin"
|
||||||
else:
|
else:
|
||||||
backend = "flashinfer_cutlass"
|
backend = "flashinfer_cutlass"
|
||||||
|
|
||||||
|
|||||||
@@ -56,6 +56,8 @@ GPTQ_MARLIN_MIN_THREAD_K = 128
|
|||||||
GPTQ_MARLIN_MAX_PARALLEL = 16
|
GPTQ_MARLIN_MAX_PARALLEL = 16
|
||||||
|
|
||||||
MARLIN_SUPPORTED_GROUP_SIZES = [-1, 32, 64, 128]
|
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
|
# 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
|
# 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"are supported (for group_size = {group_size}, "
|
||||||
f"device_capability = {device_capability}, zp = {has_zp}).",
|
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 (
|
return (
|
||||||
False,
|
False,
|
||||||
f"Marlin does not support group_size = {group_size}. "
|
f"Marlin does not support group_size = {group_size} for "
|
||||||
f"Only group_sizes = {MARLIN_SUPPORTED_GROUP_SIZES} "
|
f"quant_type = {quant_type}. Only group_sizes = {allowed_group_sizes} "
|
||||||
"are supported.",
|
"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
|
intermediate_size_per_partition = layer.intermediate_size_per_partition
|
||||||
# apply_router_weight_on_input is not supported for moe marlin
|
# apply_router_weight_on_input is not supported for moe marlin
|
||||||
supports_router_weight = not layer.moe_runner_config.apply_router_weight_on_input
|
supports_router_weight = not layer.moe_runner_config.apply_router_weight_on_input
|
||||||
# moe marlin requires the activation to be silu
|
if layer.moe_runner_config.is_gated:
|
||||||
supports_activation = layer.moe_runner_config.activation == "silu"
|
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)
|
# gate-up: (n, k) = (intermediate_size_per_partition * 2, hidden_size)
|
||||||
# down: (n, k) = (hidden_size, intermediate_size_per_partition)
|
# down: (n, k) = (hidden_size, intermediate_size_per_partition)
|
||||||
|
|||||||
@@ -3,17 +3,182 @@ from __future__ import annotations
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.quantization.marlin_utils import (
|
from sglang.srt.layers.quantization.marlin_utils import (
|
||||||
|
USE_FP32_REDUCE_DEFAULT,
|
||||||
marlin_make_workspace,
|
marlin_make_workspace,
|
||||||
marlin_permute_bias,
|
marlin_permute_bias,
|
||||||
marlin_permute_scales,
|
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 import is_cuda
|
||||||
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
|
||||||
if _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
|
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(
|
def mxfp4_marlin_process_scales(
|
||||||
marlin_scales: torch.Tensor,
|
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(
|
layer.w2_weight_bias = torch.nn.Parameter(
|
||||||
_permute_bias(w2_bias_data), requires_grad=False
|
_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)
|
||||||
|
|||||||
@@ -48,6 +48,11 @@ from sglang.srt.layers.quantization.fp8_utils import (
|
|||||||
is_blackwell_supported,
|
is_blackwell_supported,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
|
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.unquant import UnquantizedLinearMethod
|
||||||
from sglang.srt.layers.quantization.utils import (
|
from sglang.srt.layers.quantization.utils import (
|
||||||
convert_to_channelwise,
|
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.radix_attention import RadixAttention
|
||||||
from sglang.srt.layers.utils import alias_or_bind_derived_param, copy_or_rebind_param
|
from sglang.srt.layers.utils import alias_or_bind_derived_param, copy_or_rebind_param
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
|
get_device_capability,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_sm120_supported,
|
is_sm120_supported,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
@@ -1190,7 +1196,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_min_capability(cls) -> int:
|
def get_min_capability(cls) -> int:
|
||||||
return 100
|
return 80
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def common_group_size(cfg: dict) -> int:
|
def common_group_size(cfg: dict) -> int:
|
||||||
@@ -1367,6 +1373,8 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
|
|
||||||
layer.input_size_per_partition = input_size_per_partition
|
layer.input_size_per_partition = input_size_per_partition
|
||||||
layer.output_size_per_partition = output_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:
|
if input_size_per_partition % 16 != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Unsupported model when in features size is not multiple of 16"
|
"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
|
# Store original output size before any padding
|
||||||
layer.output_size_per_partition = layer.weight.shape[0]
|
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():
|
if get_fp4_gemm_runner_backend().is_flashinfer_trtllm():
|
||||||
# FlashInfer TRTLLM FP4 GEMM requires a different weight layout.
|
# FlashInfer TRTLLM FP4 GEMM requires a different weight layout.
|
||||||
# FlashInfer provides nvfp4_quantize to quantize + shuffle the
|
# FlashInfer provides nvfp4_quantize to quantize + shuffle the
|
||||||
@@ -1579,6 +1604,18 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> 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
|
# `_accepts_prequantized_fp4` is the explicit opt-in so an accidental
|
||||||
# tuple from unrelated code can't silently bypass quantization.
|
# tuple from unrelated code can't silently bypass quantization.
|
||||||
if getattr(layer, "_accepts_prequantized_fp4", False) and isinstance(x, tuple):
|
if getattr(layer, "_accepts_prequantized_fp4", False) and isinstance(x, tuple):
|
||||||
@@ -1639,11 +1676,17 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
|
|
||||||
def __init__(self, quant_config: ModelOptFp4Config):
|
def __init__(self, quant_config: ModelOptFp4Config):
|
||||||
self.quant_config = quant_config
|
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(
|
raise ValueError(
|
||||||
"Current platform does not support NVFP4"
|
"Current platform does not support NVFP4"
|
||||||
" quantization. Please use Blackwell and"
|
" quantization with the selected MoE backend. Please use "
|
||||||
" above."
|
"Blackwell and above, or use moe_runner_backend=marlin on SM80+."
|
||||||
)
|
)
|
||||||
self.enable_flashinfer_trtllm_moe = (
|
self.enable_flashinfer_trtllm_moe = (
|
||||||
get_moe_runner_backend().is_flashinfer_trtllm()
|
get_moe_runner_backend().is_flashinfer_trtllm()
|
||||||
@@ -1844,6 +1887,18 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
else:
|
else:
|
||||||
w13_weight_scale_2 = layer.w13_weight_scale_2[:]
|
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
|
# Calculate input scales based on strategy
|
||||||
if self.enable_flashinfer_cutlass_moe or self.enable_flashinfer_trtllm_moe:
|
if self.enable_flashinfer_cutlass_moe or self.enable_flashinfer_trtllm_moe:
|
||||||
w13_input_scale = layer.w13_input_scale.max().to(torch.float32)
|
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()
|
moe_runner_backend = get_moe_runner_backend()
|
||||||
|
|
||||||
if moe_runner_backend.is_auto():
|
if moe_runner_backend.is_auto():
|
||||||
# TRTLLM is currently the most performant and tested FP4 MoE
|
if is_cuda() and (8, 0) <= get_device_capability() < (10, 0):
|
||||||
# backend, so use it as the default.
|
moe_runner_backend = MoeRunnerBackend.MARLIN
|
||||||
moe_runner_backend = MoeRunnerBackend.FLASHINFER_TRTLLM
|
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():
|
if moe_runner_backend.is_flashinfer_cutedsl():
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl # noqa: F401 – triggers @register_fused_func
|
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
|
# tuple). Defer per-attribute access to the branches that actually
|
||||||
# consume them.
|
# consume them.
|
||||||
activation = self.moe_runner_config.activation
|
activation = self.moe_runner_config.activation
|
||||||
|
moe_runner_backend = getattr(
|
||||||
|
self, "_moe_runner_backend", get_moe_runner_backend()
|
||||||
|
)
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
activation in _SUPPORTED_ACT_STRS
|
activation in _SUPPORTED_ACT_STRS
|
||||||
), f"{activation=} not in supported {_SUPPORTED_ACT_STRS}"
|
), f"{activation=} not in supported {_SUPPORTED_ACT_STRS}"
|
||||||
moe_runner_config = self.moe_runner_config
|
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
|
# FlashInfer TRTLLM FP4 path
|
||||||
if self.enable_flashinfer_trtllm_moe and hasattr(layer, "g1_scale_c"):
|
if self.enable_flashinfer_trtllm_moe and hasattr(layer, "g1_scale_c"):
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
||||||
|
|||||||
@@ -255,6 +255,7 @@ FP4_GEMM_RUNNER_BACKEND_CHOICES = [
|
|||||||
"flashinfer_cutedsl",
|
"flashinfer_cutedsl",
|
||||||
"flashinfer_cutlass",
|
"flashinfer_cutlass",
|
||||||
"flashinfer_trtllm",
|
"flashinfer_trtllm",
|
||||||
|
"marlin",
|
||||||
]
|
]
|
||||||
|
|
||||||
RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu", "slru", "priority"]
|
RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu", "slru", "priority"]
|
||||||
@@ -5688,12 +5689,13 @@ class ServerArgs:
|
|||||||
default=ServerArgs.fp4_gemm_runner_backend,
|
default=ServerArgs.fp4_gemm_runner_backend,
|
||||||
dest="fp4_gemm_runner_backend",
|
dest="fp4_gemm_runner_backend",
|
||||||
help="Choose the runner backend for NVFP4 GEMM operations. "
|
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), "
|
"'cutlass' (SGLang CUTLASS kernel), "
|
||||||
"'flashinfer_cutlass' (FlashInfer CUTLASS backend), "
|
"'flashinfer_cutlass' (FlashInfer CUTLASS backend), "
|
||||||
"'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), "
|
"'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), "
|
||||||
"'flashinfer_cutedsl' (FlashInfer CuTe DSL backend), "
|
"'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(
|
parser.add_argument(
|
||||||
"--disable-flashinfer-autotune",
|
"--disable-flashinfer-autotune",
|
||||||
|
|||||||
@@ -283,6 +283,11 @@ is_sm100_supported = lru_cache(maxsize=1)(
|
|||||||
_check_cuda_device_version, device_capability_majors=[10], cuda_version=(12, 8)
|
_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)(
|
is_sm90_supported = lru_cache(maxsize=1)(
|
||||||
partial(
|
partial(
|
||||||
_check_cuda_device_version, device_capability_majors=[9], cuda_version=(12, 3)
|
_check_cuda_device_version, device_capability_majors=[9], cuda_version=(12, 3)
|
||||||
|
|||||||
@@ -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)
|
res_list[i] = res_list[i].to(w.device)
|
||||||
|
|
||||||
return res_list
|
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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -7,6 +7,7 @@ Run with `python3 test/manual/models/test_nvidia_nemotron_3_nano_archived.py`.
|
|||||||
|
|
||||||
import unittest
|
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.kits.lm_eval_kit import LMEvalMixin
|
||||||
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
|
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
|
||||||
|
|
||||||
@@ -43,5 +44,27 @@ class TestNvidiaNemotron3Nano30BBF16FlashInfer(LMEvalMixin, DefaultServerBase):
|
|||||||
] + NEMOTRON_3_NANO_THINKING_ARGS
|
] + 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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.srt.configs.load_config import LoadConfig
|
|||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
|
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
|
||||||
from sglang.srt.layers.quantization.modelopt_quant import (
|
from sglang.srt.layers.quantization.modelopt_quant import (
|
||||||
|
ModelOptFp4Config,
|
||||||
ModelOptMixedPrecisionConfig,
|
ModelOptMixedPrecisionConfig,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.loader import ModelOptModelLoader
|
from sglang.srt.model_loader.loader import ModelOptModelLoader
|
||||||
@@ -646,7 +647,10 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def test_mixed_precision_uses_nvfp4_min_capability(self):
|
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):
|
def test_mixed_precision_quant_layer_resolution_after_mapping(self):
|
||||||
quant_config = ModelOptMixedPrecisionConfig.from_config(
|
quant_config = ModelOptMixedPrecisionConfig.from_config(
|
||||||
|
|||||||
Reference in New Issue
Block a user