[codex] Add flashinfer TRTLLM backend for diffusion NVFP4 (#22717)
This commit is contained in:
@@ -187,6 +187,9 @@ sglang generate \
|
|||||||
over the compatibility `--model-path` flow.
|
over the compatibility `--model-path` flow.
|
||||||
- For local directories, SGLang first looks for `*-mixed.safetensors`, then
|
- For local directories, SGLang first looks for `*-mixed.safetensors`, then
|
||||||
falls back to loading from the directory.
|
falls back to loading from the directory.
|
||||||
|
- To force the generic diffusion ModelOpt FP4 path onto a specific FlashInfer
|
||||||
|
backend, set `SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND`. Supported values
|
||||||
|
include `flashinfer_cudnn`, `flashinfer_cutlass`, and `flashinfer_trtllm`.
|
||||||
- On disk, the quantization config stays `quant_method=modelopt` with
|
- On disk, the quantization config stays `quant_method=modelopt` with
|
||||||
`quant_algo=NVFP4`; the `modelopt-nvfp4` label here is again a documentation
|
`quant_algo=NVFP4`; the `modelopt-nvfp4` label here is again a documentation
|
||||||
family name rather than a serialized config key.
|
family name rather than a serialized config key.
|
||||||
|
|||||||
@@ -5,10 +5,14 @@ import pytest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.jit_kernel.nvfp4 import cutlass_scaled_fp4_mm, scaled_fp4_quant
|
from sglang.jit_kernel.nvfp4 import cutlass_scaled_fp4_mm, scaled_fp4_quant
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization import (
|
||||||
|
modelopt_quant as diffusion_modelopt_quant,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
||||||
ModelOptFp4Config,
|
ModelOptFp4Config,
|
||||||
ModelOptFp4LinearMethod,
|
ModelOptFp4LinearMethod,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.srt.layers.quantization.modelopt_quant import pad_nvfp4_weight
|
from sglang.srt.layers.quantization.modelopt_quant import pad_nvfp4_weight
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
@@ -110,12 +114,29 @@ def _quantize_weight_for_checkpoint(
|
|||||||
return weight_fp4, weight_scale_linear.contiguous()
|
return weight_fp4, weight_scale_linear.contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
def _set_diffusion_fp4_backend(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, backend: str | None
|
||||||
|
) -> None:
|
||||||
|
if backend is None:
|
||||||
|
monkeypatch.delenv(
|
||||||
|
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND", raising=False
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
monkeypatch.setenv("SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND", backend)
|
||||||
|
|
||||||
|
current_platform.__class__.get_modelopt_flashinfer_fp4_backend.cache_clear()
|
||||||
|
current_platform.__class__.get_modelopt_fp4_gemm_op.cache_clear()
|
||||||
|
diffusion_modelopt_quant._get_fp4_gemm_op.cache_clear()
|
||||||
|
|
||||||
|
|
||||||
def _build_layer(
|
def _build_layer(
|
||||||
weight_fp4: torch.Tensor,
|
weight_fp4: torch.Tensor,
|
||||||
weight_scale_linear: torch.Tensor,
|
weight_scale_linear: torch.Tensor,
|
||||||
input_global_scale: torch.Tensor,
|
input_global_scale: torch.Tensor,
|
||||||
weight_global_scale: torch.Tensor,
|
weight_global_scale: torch.Tensor,
|
||||||
) -> None:
|
*,
|
||||||
|
weight_scale_device: torch.device | str | None = None,
|
||||||
|
) -> tuple[ModelOptFp4LinearMethod, torch.nn.Module]:
|
||||||
output_size, input_size_half = weight_fp4.shape
|
output_size, input_size_half = weight_fp4.shape
|
||||||
input_size = input_size_half * 2
|
input_size = input_size_half * 2
|
||||||
method = ModelOptFp4LinearMethod(
|
method = ModelOptFp4LinearMethod(
|
||||||
@@ -142,19 +163,62 @@ def _build_layer(
|
|||||||
(1.0 / weight_global_scale).reshape_as(layer.weight_scale_2)
|
(1.0 / weight_global_scale).reshape_as(layer.weight_scale_2)
|
||||||
)
|
)
|
||||||
layer.weight_scale.data.copy_(weight_scale_linear)
|
layer.weight_scale.data.copy_(weight_scale_linear)
|
||||||
|
if weight_scale_device is not None:
|
||||||
|
layer.weight_scale = torch.nn.Parameter(
|
||||||
|
layer.weight_scale.detach().to(weight_scale_device), requires_grad=False
|
||||||
|
)
|
||||||
|
|
||||||
method.process_weights_after_loading(layer)
|
method.process_weights_after_loading(layer)
|
||||||
|
|
||||||
expected_weight, expected_padding_cols = pad_nvfp4_weight(weight_fp4)
|
_, flashinfer_backend = current_platform.get_modelopt_fp4_gemm_op()
|
||||||
expected_scale_shape = (
|
if flashinfer_backend == "trtllm":
|
||||||
((output_size + 128 - 1) // 128) * 128,
|
expected_weight, _ = pad_nvfp4_weight(
|
||||||
(((input_size // BLOCK_SIZE) + 4 - 1) // 4) * 4,
|
weight_fp4, n_alignment=128, k_alignment=0
|
||||||
)
|
)
|
||||||
|
expected_scale = weight_scale_linear
|
||||||
|
if expected_scale.shape[0] != expected_weight.shape[0]:
|
||||||
|
pad_n = expected_weight.shape[0] - expected_scale.shape[0]
|
||||||
|
expected_scale = torch.nn.functional.pad(expected_scale, (0, 0, 0, pad_n))
|
||||||
|
|
||||||
assert torch.equal(layer.weight, expected_weight)
|
expected_padding_cols = 0
|
||||||
assert layer.weight_scale_interleaved.shape == expected_scale_shape
|
if expected_scale.shape[1] % 4 != 0:
|
||||||
assert layer.weight_scale_interleaved.dtype == torch.float8_e4m3fn
|
padded_scale_k = ((expected_scale.shape[1] + 4 - 1) // 4) * 4
|
||||||
assert layer.weights_padding_cols == expected_padding_cols
|
pad_scale_k = padded_scale_k - expected_scale.shape[1]
|
||||||
|
expected_scale = torch.nn.functional.pad(
|
||||||
|
expected_scale, (0, pad_scale_k, 0, 0)
|
||||||
|
)
|
||||||
|
pad_weight_k = pad_scale_k * 8
|
||||||
|
expected_weight = torch.nn.functional.pad(
|
||||||
|
expected_weight, (0, pad_weight_k, 0, 0)
|
||||||
|
)
|
||||||
|
expected_padding_cols = pad_weight_k
|
||||||
|
|
||||||
|
expected_weight = flashinfer.shuffle_matrix_a(
|
||||||
|
expected_weight.view(torch.uint8), 128
|
||||||
|
)
|
||||||
|
expected_scale = (
|
||||||
|
flashinfer.shuffle_matrix_sf_a(expected_scale.view(torch.uint8), 128)
|
||||||
|
.reshape(expected_scale.shape)
|
||||||
|
.view(torch.float8_e4m3fn)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(layer.weight, expected_weight)
|
||||||
|
assert torch.equal(
|
||||||
|
layer.weight_scale_interleaved.view(torch.uint8),
|
||||||
|
expected_scale.view(torch.uint8),
|
||||||
|
)
|
||||||
|
assert layer.weights_padding_cols == expected_padding_cols
|
||||||
|
else:
|
||||||
|
expected_weight, expected_padding_cols = pad_nvfp4_weight(weight_fp4)
|
||||||
|
expected_scale_shape = (
|
||||||
|
((output_size + 128 - 1) // 128) * 128,
|
||||||
|
(((input_size // BLOCK_SIZE) + 4 - 1) // 4) * 4,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(layer.weight, expected_weight)
|
||||||
|
assert layer.weight_scale_interleaved.shape == expected_scale_shape
|
||||||
|
assert layer.weight_scale_interleaved.dtype == torch.float8_e4m3fn
|
||||||
|
assert layer.weights_padding_cols == expected_padding_cols
|
||||||
torch.testing.assert_close(
|
torch.testing.assert_close(
|
||||||
layer.alpha,
|
layer.alpha,
|
||||||
(1.0 / (input_global_scale * weight_global_scale)).to(torch.float32),
|
(1.0 / (input_global_scale * weight_global_scale)).to(torch.float32),
|
||||||
@@ -163,6 +227,7 @@ def _build_layer(
|
|||||||
layer.input_scale_inv,
|
layer.input_scale_inv,
|
||||||
input_global_scale.to(torch.float32),
|
input_global_scale.to(torch.float32),
|
||||||
)
|
)
|
||||||
|
return method, layer
|
||||||
|
|
||||||
|
|
||||||
def _resolve_mode(mode: str):
|
def _resolve_mode(mode: str):
|
||||||
@@ -170,6 +235,8 @@ def _resolve_mode(mode: str):
|
|||||||
return scaled_fp4_quant, cutlass_scaled_fp4_mm, None
|
return scaled_fp4_quant, cutlass_scaled_fp4_mm, None
|
||||||
if mode == "flashinfer2":
|
if mode == "flashinfer2":
|
||||||
return flashinfer.fp4_quantize, flashinfer.mm_fp4, "cudnn"
|
return flashinfer.fp4_quantize, flashinfer.mm_fp4, "cudnn"
|
||||||
|
if mode == "flashinfer_trtllm":
|
||||||
|
return flashinfer.fp4_quantize, flashinfer.mm_fp4, "trtllm"
|
||||||
raise ValueError(f"Unknown mode: {mode}")
|
raise ValueError(f"Unknown mode: {mode}")
|
||||||
|
|
||||||
|
|
||||||
@@ -177,8 +244,14 @@ def _resolve_mode(mode: str):
|
|||||||
not _nvfp4_supported(),
|
not _nvfp4_supported(),
|
||||||
reason="Diffusion NVFP4 scaled mm correctness requires Blackwell GPUs",
|
reason="Diffusion NVFP4 scaled mm correctness requires Blackwell GPUs",
|
||||||
)
|
)
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"backend", [None, "flashinfer_trtllm"], ids=["default", "flashinfer_trtllm"]
|
||||||
|
)
|
||||||
@pytest.mark.parametrize("m,n,k", TEST_CASES)
|
@pytest.mark.parametrize("m,n,k", TEST_CASES)
|
||||||
def test_checkpoint_processing(m: int, n: int, k: int) -> None:
|
def test_checkpoint_processing(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, backend: str | None, m: int, n: int, k: int
|
||||||
|
) -> None:
|
||||||
|
_set_diffusion_fp4_backend(monkeypatch, backend)
|
||||||
generator = torch.Generator(device=DEVICE)
|
generator = torch.Generator(device=DEVICE)
|
||||||
generator.manual_seed(20260404 + m + n + k)
|
generator.manual_seed(20260404 + m + n + k)
|
||||||
|
|
||||||
@@ -247,5 +320,80 @@ def test_flux2_shape_correctness(mode: str) -> None:
|
|||||||
assert diff < DEEPGEMM_FP4_MAX_DIFF, f"{mode=}, {m=}, {n=}, {k=}, {diff=:.6f}"
|
assert diff < DEEPGEMM_FP4_MAX_DIFF, f"{mode=}, {m=}, {n=}, {k=}, {diff=:.6f}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
not _nvfp4_supported(),
|
||||||
|
reason="Diffusion NVFP4 scaled mm correctness requires Blackwell GPUs",
|
||||||
|
)
|
||||||
|
def test_flux2_shape_correctness_flashinfer_trtllm(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
_set_diffusion_fp4_backend(monkeypatch, "flashinfer_trtllm")
|
||||||
|
|
||||||
|
m, n, k = FLUX2_PROJECTION_SHAPE
|
||||||
|
generator = torch.Generator(device=DEVICE)
|
||||||
|
generator.manual_seed(20260404 + m + n + k + 17)
|
||||||
|
|
||||||
|
x = torch.randn((m, k), device=DEVICE, dtype=DTYPE, generator=generator)
|
||||||
|
weight = torch.randn((n, k), device=DEVICE, dtype=DTYPE, generator=generator)
|
||||||
|
input_global_scale = _make_global_scale(x)
|
||||||
|
weight_global_scale = _make_global_scale(weight)
|
||||||
|
weight_fp4, weight_scale_linear = _quantize_weight_for_checkpoint(
|
||||||
|
weight, weight_global_scale
|
||||||
|
)
|
||||||
|
|
||||||
|
method, layer = _build_layer(
|
||||||
|
weight_fp4, weight_scale_linear, input_global_scale, weight_global_scale
|
||||||
|
)
|
||||||
|
actual = method.apply(layer, x)
|
||||||
|
|
||||||
|
x_fp4, x_scale_swizzled = flashinfer.fp4_quantize(x, input_global_scale)
|
||||||
|
weight_fp4_ref, weight_scale_swizzled = flashinfer.fp4_quantize(
|
||||||
|
weight, weight_global_scale
|
||||||
|
)
|
||||||
|
if x_scale_swizzled.dtype == torch.uint8:
|
||||||
|
x_scale_swizzled = x_scale_swizzled.view(torch.float8_e4m3fn)
|
||||||
|
if weight_scale_swizzled.dtype == torch.uint8:
|
||||||
|
weight_scale_swizzled = weight_scale_swizzled.view(torch.float8_e4m3fn)
|
||||||
|
|
||||||
|
expected = torch.matmul(
|
||||||
|
_dequantize_nvfp4(x_fp4, x_scale_swizzled, input_global_scale),
|
||||||
|
_dequantize_nvfp4(
|
||||||
|
weight_fp4_ref, weight_scale_swizzled, weight_global_scale
|
||||||
|
).t(),
|
||||||
|
)
|
||||||
|
|
||||||
|
diff = _calc_diff(actual, expected.to(dtype=DTYPE))
|
||||||
|
assert diff < DEEPGEMM_FP4_MAX_DIFF, f"{m=}, {n=}, {k=}, {diff=:.6f}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
not _nvfp4_supported(),
|
||||||
|
reason="Diffusion NVFP4 scaled mm correctness requires Blackwell GPUs",
|
||||||
|
)
|
||||||
|
def test_checkpoint_processing_flashinfer_trtllm_cpu_weight_scale(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
_set_diffusion_fp4_backend(monkeypatch, "flashinfer_trtllm")
|
||||||
|
|
||||||
|
m, n, k = FLUX2_PROJECTION_SHAPE
|
||||||
|
generator = torch.Generator(device=DEVICE)
|
||||||
|
generator.manual_seed(20260413 + m + n + k)
|
||||||
|
|
||||||
|
weight = torch.randn((n, k), device=DEVICE, dtype=DTYPE, generator=generator)
|
||||||
|
input_global_scale = torch.tensor(512.0, device=DEVICE, dtype=torch.float32)
|
||||||
|
weight_global_scale = _make_global_scale(weight)
|
||||||
|
weight_fp4, weight_scale_linear = _quantize_weight_for_checkpoint(
|
||||||
|
weight, weight_global_scale
|
||||||
|
)
|
||||||
|
|
||||||
|
_build_layer(
|
||||||
|
weight_fp4,
|
||||||
|
weight_scale_linear,
|
||||||
|
input_global_scale,
|
||||||
|
weight_global_scale,
|
||||||
|
weight_scale_device="cpu",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ if TYPE_CHECKING:
|
|||||||
SGLANG_CACHE_DIT_SECONDARY_TS_ORDER: int = 1
|
SGLANG_CACHE_DIT_SECONDARY_TS_ORDER: int = 1
|
||||||
# model loading
|
# model loading
|
||||||
SGLANG_USE_RUNAI_MODEL_STREAMER: bool = True
|
SGLANG_USE_RUNAI_MODEL_STREAMER: bool = True
|
||||||
|
SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None
|
||||||
SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: bool = False
|
SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: bool = False
|
||||||
SGLANG_USE_ROCM_VAE: bool = False
|
SGLANG_USE_ROCM_VAE: bool = False
|
||||||
SGLANG_USE_ROCM_CUDNN_BENCHMARK: bool = False
|
SGLANG_USE_ROCM_CUDNN_BENCHMARK: bool = False
|
||||||
@@ -278,7 +279,13 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
|||||||
"SGLANG_USE_RUNAI_MODEL_STREAMER": _lazy_bool(
|
"SGLANG_USE_RUNAI_MODEL_STREAMER": _lazy_bool(
|
||||||
"SGLANG_USE_RUNAI_MODEL_STREAMER", "true"
|
"SGLANG_USE_RUNAI_MODEL_STREAMER", "true"
|
||||||
),
|
),
|
||||||
# FlashInfer FP4 GEMM backend for the generic diffusion NVFP4 fallback.
|
# FlashInfer FP4 GEMM backend override for diffusion NVFP4.
|
||||||
|
# Supported values:
|
||||||
|
# - auto
|
||||||
|
# - flashinfer_cudnn
|
||||||
|
# - flashinfer_cutlass
|
||||||
|
# - flashinfer_trtllm
|
||||||
|
# Legacy aliases `cudnn` and `trtllm` are also accepted.
|
||||||
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": _lazy_str(
|
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": _lazy_str(
|
||||||
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND"
|
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND"
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -37,10 +37,15 @@ from sglang.srt.layers.quantization.utils import (
|
|||||||
requantize_with_max_scale,
|
requantize_with_max_scale,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.common import copy_or_rebind_param
|
from sglang.srt.layers.utils.common import copy_or_rebind_param
|
||||||
from sglang.srt.utils.common import round_up
|
from sglang.srt.utils.common import is_flashinfer_available, round_up
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if is_flashinfer_available():
|
||||||
|
import flashinfer
|
||||||
|
else:
|
||||||
|
flashinfer = None
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def _get_fp4_quantize_op():
|
def _get_fp4_quantize_op():
|
||||||
@@ -61,6 +66,14 @@ def _prepare_nvfp4_weight_bytes(
|
|||||||
return ((weight >> 4) | (weight << 4)).contiguous()
|
return ((weight >> 4) | (weight << 4)).contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
def _require_flashinfer():
|
||||||
|
if flashinfer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"flashinfer is required for the diffusion NVFP4 FlashInfer path."
|
||||||
|
)
|
||||||
|
return flashinfer
|
||||||
|
|
||||||
|
|
||||||
class ModelOptQuantConfig(QuantizationConfig):
|
class ModelOptQuantConfig(QuantizationConfig):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -479,8 +492,51 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
|
|
||||||
w = layer.weight.data
|
w = layer.weight.data
|
||||||
w_swapped = _prepare_nvfp4_weight_bytes(
|
w_swapped = _prepare_nvfp4_weight_bytes(
|
||||||
w, swap_weight_nibbles=self.quant_config.swap_weight_nibbles
|
w,
|
||||||
|
swap_weight_nibbles=getattr(self.quant_config, "swap_weight_nibbles", True),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_, flashinfer_backend = _get_fp4_gemm_op()
|
||||||
|
if flashinfer_backend == "trtllm":
|
||||||
|
flashinfer_ops = _require_flashinfer()
|
||||||
|
|
||||||
|
weight, _ = pad_nvfp4_weight(w_swapped, n_alignment=128, k_alignment=0)
|
||||||
|
scales = layer.weight_scale
|
||||||
|
if scales.shape[0] != weight.shape[0]:
|
||||||
|
pad_n = weight.shape[0] - scales.shape[0]
|
||||||
|
scales = torch.nn.functional.pad(scales, (0, 0, 0, pad_n))
|
||||||
|
|
||||||
|
scale_k = scales.shape[1]
|
||||||
|
weights_padding_cols = 0
|
||||||
|
if scale_k % 4 != 0:
|
||||||
|
padded_scale_k = round_up(scale_k, 4)
|
||||||
|
pad_scale_k = padded_scale_k - scale_k
|
||||||
|
scales = torch.nn.functional.pad(scales, (0, pad_scale_k, 0, 0))
|
||||||
|
pad_weight_k = pad_scale_k * 8
|
||||||
|
weight = torch.nn.functional.pad(weight, (0, pad_weight_k, 0, 0))
|
||||||
|
weights_padding_cols = pad_weight_k
|
||||||
|
|
||||||
|
epilogue_tile_m = 128
|
||||||
|
shuffled_scale_shape = scales.shape
|
||||||
|
if not weight.is_cuda:
|
||||||
|
weight = weight.cuda()
|
||||||
|
if scales.device != weight.device:
|
||||||
|
scales = scales.to(device=weight.device)
|
||||||
|
weight = flashinfer_ops.shuffle_matrix_a(
|
||||||
|
weight.view(torch.uint8), epilogue_tile_m
|
||||||
|
)
|
||||||
|
scales = (
|
||||||
|
flashinfer_ops.shuffle_matrix_sf_a(
|
||||||
|
scales.view(torch.uint8), epilogue_tile_m
|
||||||
|
)
|
||||||
|
.reshape(shuffled_scale_shape)
|
||||||
|
.view(torch.float8_e4m3fn)
|
||||||
|
)
|
||||||
|
|
||||||
|
layer.weights_padding_cols = weights_padding_cols
|
||||||
|
copy_or_rebind_param(layer, "weight", weight)
|
||||||
|
copy_or_rebind_param(layer, "weight_scale_interleaved", scales)
|
||||||
|
return
|
||||||
weight, weights_padding_cols = pad_nvfp4_weight(w_swapped)
|
weight, weights_padding_cols = pad_nvfp4_weight(w_swapped)
|
||||||
layer.weights_padding_cols = weights_padding_cols
|
layer.weights_padding_cols = weights_padding_cols
|
||||||
copy_or_rebind_param(layer, "weight", weight)
|
copy_or_rebind_param(layer, "weight", weight)
|
||||||
|
|||||||
@@ -124,30 +124,38 @@ class CudaPlatformBase(Platform):
|
|||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def get_modelopt_flashinfer_fp4_backend(cls) -> str:
|
def get_modelopt_flashinfer_fp4_backend(cls) -> str:
|
||||||
backend = envs.SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND
|
backend = envs.SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND
|
||||||
|
default_backend = "cudnn" if cls.is_blackwell() else "auto"
|
||||||
if backend is None:
|
if backend is None:
|
||||||
return "cudnn" if cls.is_blackwell() else "auto"
|
return default_backend
|
||||||
|
|
||||||
backend = backend.lower()
|
backend = backend.lower()
|
||||||
if backend not in {"auto", "cudnn"}:
|
backend = {
|
||||||
|
"flashinfer_cudnn": "cudnn",
|
||||||
|
"flashinfer_cutlass": "cutlass",
|
||||||
|
"flashinfer_trtllm": "trtllm",
|
||||||
|
"trtllm": "trtllm",
|
||||||
|
"cudnn": "cudnn",
|
||||||
|
"auto": "auto",
|
||||||
|
}.get(backend, backend)
|
||||||
|
if backend not in {"auto", "cudnn", "cutlass", "trtllm"}:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Unsupported SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=%r. "
|
"Unsupported SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=%r. "
|
||||||
"Falling back to %r.",
|
"Falling back to %r.",
|
||||||
backend,
|
backend,
|
||||||
"cudnn" if cls.is_blackwell() else "auto",
|
default_backend,
|
||||||
)
|
)
|
||||||
return "cudnn" if cls.is_blackwell() else "auto"
|
return default_backend
|
||||||
return backend
|
return backend
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def get_modelopt_fp4_gemm_op(cls) -> tuple[Callable | None, str | None]:
|
def get_modelopt_fp4_gemm_op(cls) -> tuple[Callable | None, str | None]:
|
||||||
|
requested_backend = envs.SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND
|
||||||
|
prefer_flashinfer = requested_backend is not None
|
||||||
|
|
||||||
# TODO: Remove this explicit FlashInfer preference once the sm100 CUTLASS
|
# TODO: Remove this explicit FlashInfer preference once the sm100 CUTLASS
|
||||||
# LargeM dispatch grows a validated fallback for Blackwell NVFP4 shapes
|
# LargeM dispatch grows a validated fallback for Blackwell NVFP4 shapes
|
||||||
# such as Wan2.2's large-M attention projections.
|
# such as Wan2.2's large-M attention projections.
|
||||||
prefer_flashinfer = (
|
|
||||||
envs.SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND is not None
|
|
||||||
)
|
|
||||||
|
|
||||||
if prefer_flashinfer:
|
if prefer_flashinfer:
|
||||||
try:
|
try:
|
||||||
from flashinfer import mm_fp4 as flashinfer_mm_fp4
|
from flashinfer import mm_fp4 as flashinfer_mm_fp4
|
||||||
@@ -155,8 +163,10 @@ class CudaPlatformBase(Platform):
|
|||||||
return flashinfer_mm_fp4, cls.get_modelopt_flashinfer_fp4_backend()
|
return flashinfer_mm_fp4, cls.get_modelopt_flashinfer_fp4_backend()
|
||||||
except ImportError:
|
except ImportError:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND is set, "
|
"Requested SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=%r "
|
||||||
"but flashinfer.mm_fp4 is unavailable. Falling back to cutlass."
|
"but flashinfer.mm_fp4 is unavailable. Falling back to "
|
||||||
|
"cutlass.",
|
||||||
|
requested_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -23,8 +23,10 @@ Example:
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import contextlib
|
||||||
import json
|
import json
|
||||||
import math
|
import math
|
||||||
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Sequence
|
from typing import Any, Sequence
|
||||||
|
|
||||||
@@ -299,18 +301,7 @@ def build_sampling_kwargs(
|
|||||||
return kwargs
|
return kwargs
|
||||||
|
|
||||||
|
|
||||||
def run_variant(
|
def _normalize_single_result(result: Any):
|
||||||
*,
|
|
||||||
server_kwargs: dict[str, Any],
|
|
||||||
sampling_kwargs: dict[str, Any],
|
|
||||||
):
|
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import (
|
|
||||||
DiffGenerator,
|
|
||||||
)
|
|
||||||
|
|
||||||
with DiffGenerator.from_pretrained(local_mode=True, **server_kwargs) as generator:
|
|
||||||
result = generator.generate(sampling_params_kwargs=sampling_kwargs)
|
|
||||||
|
|
||||||
if isinstance(result, list):
|
if isinstance(result, list):
|
||||||
if len(result) != 1:
|
if len(result) != 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -322,6 +313,114 @@ def run_variant(
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _clear_diffusion_fp4_backend_caches() -> None:
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization import (
|
||||||
|
modelopt_quant as diffusion_modelopt_quant,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
|
||||||
|
diffusion_modelopt_quant._get_fp4_gemm_op.cache_clear()
|
||||||
|
current_platform.__class__.get_modelopt_fp4_gemm_op.cache_clear()
|
||||||
|
current_platform.__class__.get_modelopt_flashinfer_fp4_backend.cache_clear()
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def override_diffusion_fp4_backend(backend: str | None):
|
||||||
|
env_name = "SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND"
|
||||||
|
previous = os.environ.get(env_name)
|
||||||
|
|
||||||
|
if backend is None:
|
||||||
|
os.environ.pop(env_name, None)
|
||||||
|
else:
|
||||||
|
os.environ[env_name] = backend
|
||||||
|
|
||||||
|
_clear_diffusion_fp4_backend_caches()
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
if previous is None:
|
||||||
|
os.environ.pop(env_name, None)
|
||||||
|
else:
|
||||||
|
os.environ[env_name] = previous
|
||||||
|
_clear_diffusion_fp4_backend_caches()
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_total_duration_ms(result: Any) -> float | None:
|
||||||
|
metrics = getattr(result, "metrics", None)
|
||||||
|
if not isinstance(metrics, dict):
|
||||||
|
return None
|
||||||
|
total_duration_ms = metrics.get("total_duration_ms")
|
||||||
|
if total_duration_ms is None:
|
||||||
|
return None
|
||||||
|
return float(total_duration_ms)
|
||||||
|
|
||||||
|
|
||||||
|
def run_variant(
|
||||||
|
*,
|
||||||
|
server_kwargs: dict[str, Any],
|
||||||
|
sampling_kwargs: dict[str, Any],
|
||||||
|
fp4_gemm_backend: str | None,
|
||||||
|
warmup_runs: int,
|
||||||
|
measure_runs: int,
|
||||||
|
):
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import (
|
||||||
|
DiffGenerator,
|
||||||
|
)
|
||||||
|
|
||||||
|
if warmup_runs < 0:
|
||||||
|
raise ValueError("warmup_runs must be >= 0.")
|
||||||
|
if measure_runs <= 0:
|
||||||
|
raise ValueError("measure_runs must be >= 1.")
|
||||||
|
|
||||||
|
with override_diffusion_fp4_backend(fp4_gemm_backend):
|
||||||
|
with DiffGenerator.from_pretrained(
|
||||||
|
local_mode=True, **server_kwargs
|
||||||
|
) as generator:
|
||||||
|
for _ in range(warmup_runs):
|
||||||
|
_normalize_single_result(
|
||||||
|
generator.generate(sampling_params_kwargs=sampling_kwargs)
|
||||||
|
)
|
||||||
|
|
||||||
|
measured_results = []
|
||||||
|
for _ in range(measure_runs):
|
||||||
|
measured_results.append(
|
||||||
|
_normalize_single_result(
|
||||||
|
generator.generate(sampling_params_kwargs=sampling_kwargs)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
final_result = measured_results[-1]
|
||||||
|
generation_times = [float(result.generation_time) for result in measured_results]
|
||||||
|
peak_memories = [float(result.peak_memory_mb) for result in measured_results]
|
||||||
|
total_duration_ms = [
|
||||||
|
duration
|
||||||
|
for duration in (
|
||||||
|
_extract_total_duration_ms(result) for result in measured_results
|
||||||
|
)
|
||||||
|
if duration is not None
|
||||||
|
]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"result": final_result,
|
||||||
|
"fp4_gemm_backend": fp4_gemm_backend or "default",
|
||||||
|
"warmup_runs": warmup_runs,
|
||||||
|
"measure_runs": measure_runs,
|
||||||
|
"generation_time_s": generation_times[-1],
|
||||||
|
"avg_generation_time_s": sum(generation_times) / len(generation_times),
|
||||||
|
"per_run_generation_time_s": generation_times,
|
||||||
|
"peak_memory_mb": peak_memories[-1],
|
||||||
|
"max_peak_memory_mb": max(peak_memories) if peak_memories else 0.0,
|
||||||
|
"per_run_peak_memory_mb": peak_memories,
|
||||||
|
"total_duration_ms": total_duration_ms[-1] if total_duration_ms else None,
|
||||||
|
"avg_total_duration_ms": (
|
||||||
|
sum(total_duration_ms) / len(total_duration_ms)
|
||||||
|
if total_duration_ms
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
"per_run_total_duration_ms": total_duration_ms,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def _to_jsonable(result: dict[str, Any]) -> dict[str, Any]:
|
def _to_jsonable(result: dict[str, Any]) -> dict[str, Any]:
|
||||||
return json.loads(json.dumps(result, allow_nan=True))
|
return json.loads(json.dumps(result, allow_nan=True))
|
||||||
|
|
||||||
@@ -353,6 +452,22 @@ def main() -> None:
|
|||||||
parser.add_argument("--trajectory-step-index", type=int, default=-1)
|
parser.add_argument("--trajectory-step-index", type=int, default=-1)
|
||||||
parser.add_argument("--reference-transformer-path")
|
parser.add_argument("--reference-transformer-path")
|
||||||
parser.add_argument("--candidate-transformer-path")
|
parser.add_argument("--candidate-transformer-path")
|
||||||
|
parser.add_argument(
|
||||||
|
"--reference-fp4-gemm-backend",
|
||||||
|
help=(
|
||||||
|
"Optional NVFP4 GEMM backend override for the reference run, e.g. "
|
||||||
|
"'flashinfer_trtllm'."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--candidate-fp4-gemm-backend",
|
||||||
|
help=(
|
||||||
|
"Optional NVFP4 GEMM backend override for the candidate run, e.g. "
|
||||||
|
"'flashinfer_trtllm'."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument("--warmup-runs", type=int, default=0)
|
||||||
|
parser.add_argument("--measure-runs", type=int, default=1)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--reference-component-path",
|
"--reference-component-path",
|
||||||
action="append",
|
action="append",
|
||||||
@@ -423,23 +538,37 @@ def main() -> None:
|
|||||||
output_dir=str(save_root / "candidate") if save_root else None,
|
output_dir=str(save_root / "candidate") if save_root else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
reference = run_variant(
|
reference_run = run_variant(
|
||||||
server_kwargs=ref_server_kwargs,
|
server_kwargs=ref_server_kwargs,
|
||||||
sampling_kwargs=ref_sampling_kwargs,
|
sampling_kwargs=ref_sampling_kwargs,
|
||||||
|
fp4_gemm_backend=args.reference_fp4_gemm_backend,
|
||||||
|
warmup_runs=args.warmup_runs,
|
||||||
|
measure_runs=args.measure_runs,
|
||||||
)
|
)
|
||||||
candidate = run_variant(
|
candidate_run = run_variant(
|
||||||
server_kwargs=cand_server_kwargs,
|
server_kwargs=cand_server_kwargs,
|
||||||
sampling_kwargs=cand_sampling_kwargs,
|
sampling_kwargs=cand_sampling_kwargs,
|
||||||
|
fp4_gemm_backend=args.candidate_fp4_gemm_backend,
|
||||||
|
warmup_runs=args.warmup_runs,
|
||||||
|
measure_runs=args.measure_runs,
|
||||||
)
|
)
|
||||||
|
reference = reference_run["result"]
|
||||||
|
candidate = candidate_run["result"]
|
||||||
|
|
||||||
result = {
|
result = {
|
||||||
"model_path": args.model_path,
|
"model_path": args.model_path,
|
||||||
"prompt": args.prompt,
|
"prompt": args.prompt,
|
||||||
"seed": args.seed,
|
"seed": args.seed,
|
||||||
|
"warmup_runs": args.warmup_runs,
|
||||||
|
"measure_runs": args.measure_runs,
|
||||||
"server_kwargs": {
|
"server_kwargs": {
|
||||||
"reference": ref_server_kwargs,
|
"reference": ref_server_kwargs,
|
||||||
"candidate": cand_server_kwargs,
|
"candidate": cand_server_kwargs,
|
||||||
},
|
},
|
||||||
|
"backend_overrides": {
|
||||||
|
"reference_fp4_gemm_backend": reference_run["fp4_gemm_backend"],
|
||||||
|
"candidate_fp4_gemm_backend": candidate_run["fp4_gemm_backend"],
|
||||||
|
},
|
||||||
"sampling_kwargs": {
|
"sampling_kwargs": {
|
||||||
"width": args.width,
|
"width": args.width,
|
||||||
"height": args.height,
|
"height": args.height,
|
||||||
@@ -449,15 +578,13 @@ def main() -> None:
|
|||||||
"guidance_scale_2": args.guidance_scale_2,
|
"guidance_scale_2": args.guidance_scale_2,
|
||||||
},
|
},
|
||||||
"reference_generation": {
|
"reference_generation": {
|
||||||
"generation_time_s": reference.generation_time,
|
key: value for key, value in reference_run.items() if key != "result"
|
||||||
"peak_memory_mb": reference.peak_memory_mb,
|
}
|
||||||
"output_file_path": reference.output_file_path,
|
| {"output_file_path": reference.output_file_path},
|
||||||
},
|
|
||||||
"candidate_generation": {
|
"candidate_generation": {
|
||||||
"generation_time_s": candidate.generation_time,
|
key: value for key, value in candidate_run.items() if key != "result"
|
||||||
"peak_memory_mb": candidate.peak_memory_mb,
|
}
|
||||||
"output_file_path": candidate.output_file_path,
|
| {"output_file_path": candidate.output_file_path},
|
||||||
},
|
|
||||||
"trajectory_metrics": summarize_trajectory_metrics(
|
"trajectory_metrics": summarize_trajectory_metrics(
|
||||||
reference.trajectory_latents,
|
reference.trajectory_latents,
|
||||||
candidate.trajectory_latents,
|
candidate.trajectory_latents,
|
||||||
@@ -484,6 +611,12 @@ def main() -> None:
|
|||||||
"trajectory_selected_step": result["trajectory_metrics"][
|
"trajectory_selected_step": result["trajectory_metrics"][
|
||||||
"selected_step_index"
|
"selected_step_index"
|
||||||
],
|
],
|
||||||
|
"reference_avg_generation_time_s": result["reference_generation"][
|
||||||
|
"avg_generation_time_s"
|
||||||
|
],
|
||||||
|
"candidate_avg_generation_time_s": result["candidate_generation"][
|
||||||
|
"avg_generation_time_s"
|
||||||
|
],
|
||||||
"trajectory_cosine": selected["cosine_similarity"],
|
"trajectory_cosine": selected["cosine_similarity"],
|
||||||
"trajectory_mae": selected["mae"],
|
"trajectory_mae": selected["mae"],
|
||||||
"frame0_psnr_db": frame0["psnr_db"],
|
"frame0_psnr_db": frame0["psnr_db"],
|
||||||
|
|||||||
Reference in New Issue
Block a user