diff --git a/docs_new/docs/sglang-diffusion/quantization.mdx b/docs_new/docs/sglang-diffusion/quantization.mdx index 043ef28ce..81e3e1ed4 100644 --- a/docs_new/docs/sglang-diffusion/quantization.mdx +++ b/docs_new/docs/sglang-diffusion/quantization.mdx @@ -110,15 +110,14 @@ backend. ## Validated ModelOpt Checkpoints This section is the canonical support matrix for the nine diffusion ModelOpt -checkpoints currently wired up in SGLang docs and validation coverage. +checkpoints currently wired up in SGLang docs and B200 CI coverage. Published checkpoints keep the serialized quantization config as `quant_method=modelopt`; the FP8 vs NVFP4 split below is a documentation label derived from `quant_algo`. -Six of the nine repos live under `lmsys/*`. The Wan2.2 entries use NVIDIA's -official full Diffusers repos, and the FLUX.2 NVFP4 entry keeps the official -`black-forest-labs/FLUX.2-dev-NVFP4` repo. +Eight of the nine repos live under `lmsys/*`. The FLUX.2 NVFP4 entry keeps the +official `black-forest-labs/FLUX.2-dev-NVFP4` repo. @@ -159,10 +158,10 @@ official full Diffusers repos, and the FLUX.2 NVFP4 entry keeps the official - - - - + + + + @@ -207,25 +206,24 @@ official full Diffusers repos, and the FLUX.2 NVFP4 entry keeps the official - - - - + + + +
FP8 Wan-AI/Wan2.2-T2V-A14B-Diffusers--model-pathnvidia/Wan2.2-T2V-A14B-Diffusers-FP8full Diffusers repo with ModelOpt FP8 Wan2.2 componentsvalidated through direct --model-path loading--transformer-pathlmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformerprimary transformer quantized, transformer_2 kept BF16primary-transformer-only path; keep transformer_2 on the base checkpoint, and do not describe this as dual-transformer full-model FP8 unless that path is validated separately
FP8
NVFP4 Wan-AI/Wan2.2-T2V-A14B-Diffusers--model-pathnvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4full Diffusers repo with ModelOpt NVFP4 Wan2.2 componentscurrent B200/Blackwell bring-up uses SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=trtllm--transformer-pathlmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformerprimary transformer quantized with ModelOpt NVFP4, transformer_2 kept BF16primary-transformer-only path; keep transformer_2 on the base checkpoint, and current B200/Blackwell bring-up uses SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=cudnn
-The FP8 rows run in the regular H100 1-GPU diffusion CI shard; the NVFP4 rows -run in the B200 diffusion CI shard (`multimodal-gen-test-1-b200`). +These nine checkpoints are also the intended case set for the B200 diffusion CI +job (`multimodal-gen-test-1-b200`). ## ModelOpt FP8 ### Usage Examples -Converted ModelOpt FP8 transformer repos should be loaded as transformer -component overrides. If the repo or local directory already contains -`config.json`, use `--transformer-path`. Full Diffusers repos such as the -NVIDIA Wan2.2 FP8 checkpoint can be passed directly with `--model-path`. +Converted ModelOpt FP8 checkpoints should be loaded as transformer component +overrides. If the repo or local directory already contains `config.json`, use +`--transformer-path`. ```bash sglang generate \ @@ -237,7 +235,8 @@ sglang generate \ ```bash sglang generate \ - --model-path nvidia/Wan2.2-T2V-A14B-Diffusers-FP8 \ + --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ + --transformer-path lmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformer \ --prompt "a fox walking through neon rain" \ --save-output ``` @@ -324,12 +323,14 @@ sglang generate \ --save-output ``` -For Wan2.2 NVFP4: +For a dual-transformer Wan2.2 export where only the primary `transformer` +was quantized: ```bash -SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=trtllm \ +SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=cudnn \ sglang generate \ - --model-path nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4 \ + --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ + --transformer-path lmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer \ --prompt "a fox walking through neon rain" \ --save-output ``` @@ -340,16 +341,17 @@ sglang generate \ directories that already include `config.json`. - Use `--transformer-weights-path` for raw NVFP4 exports, individual safetensors files, or repo layouts that should be treated as weights first. -- For legacy mixed Wan2.2 transformer overrides, the primary - `--transformer-path` override targets only `transformer`. Use a per-component - override such as `--transformer-2-path` only when you intentionally want a - non-default `transformer_2`. +- For dual-transformer pipelines such as `Wan2.2-T2V-A14B-Diffusers`, the + primary `--transformer-path` override targets only `transformer`. Use a + per-component override such as `--transformer-2-path` only when you + intentionally want a non-default `transformer_2`. - On Blackwell, the validated Wan2.2 ModelOpt NVFP4 path currently prefers FlashInfer FP4 GEMM via - `SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=trtllm`. -- This environment-variable override selects the validated Wan2.2 NVFP4 - full-repo path on Blackwell while the other NVFP4 CI cases continue to use - the generic `cudnn` backend. + `SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=cudnn`. +- This environment-variable override is a current workaround for NVFP4 cases + where the default sglang JIT/CUTLASS `sm100` path rejects a large-M shape at + `can_implement()`. The intended long-term fix is to add a validated CUTLASS + fallback for those shapes rather than rely on the override. - Direct `--model-path` loading is a compatibility path for FLUX.2 NVFP4-style repos or local directories. - If `--transformer-weights-path` is provided explicitly, it takes precedence diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index 484219238..232384e01 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -739,10 +739,7 @@ def _register_configs(): register_configs( sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam, pipeline_config_cls=Wan2_2_T2V_A14B_Config, - hf_model_paths=[ - "Wan-AI/Wan2.2-T2V-A14B-Diffusers", - "nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4", - ], + hf_model_paths=["Wan-AI/Wan2.2-T2V-A14B-Diffusers"], ) register_configs( sampling_param_cls=Wan2_2_I2V_A14B_SamplingParam, diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py index 2f3d7c33e..c9286aba5 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -202,7 +202,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): exclude_modules: List[str] = None, packed_modules_mapping: Optional[Dict[str, List[str]]] = None, checkpoint_uses_packed_qkv: bool = False, - swap_weight_nibbles: bool = False, + swap_weight_nibbles: bool = True, ) -> None: super().__init__(exclude_modules, packed_modules_mapping) self.is_checkpoint_nvfp4_serialized = is_checkpoint_nvfp4_serialized @@ -261,7 +261,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): def from_config(cls, config: Dict[str, Any]) -> ModelOptFp4Config: group_size = None exclude_modules = [] - swap_weight_nibbles = False + swap_weight_nibbles = True # Flat format (config.json quantization_config) quant_method = config.get("quant_algo") @@ -273,10 +273,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): first_group = next(iter(config_groups.values()), {}) group_size = first_group.get("weights", {}).get("group_size") exclude_modules = config.get("ignore", []) - swap_weight_nibbles = config.get( - "swap_weight_nibbles", - config.get("checkpoint_uses_packed_qkv", False), - ) + swap_weight_nibbles = config.get("swap_weight_nibbles", True) else: # Nested format (hf_quant_config.json) try: @@ -286,10 +283,7 @@ class ModelOptFp4Config(ModelOptQuantConfig): exclude_modules = quant_config.get("exclude_modules", []) swap_weight_nibbles = quant_config.get( "swap_weight_nibbles", - config.get( - "swap_weight_nibbles", - config.get("checkpoint_uses_packed_qkv", False), - ), + config.get("swap_weight_nibbles", True), ) except (ValueError, KeyError): raise ValueError("Cannot find 'quant_algo' in quantization config.") @@ -500,9 +494,7 @@ class ModelOptFp4LinearMethod(LinearMethodBase): w = layer.weight.data w_swapped = _prepare_nvfp4_weight_bytes( w, - swap_weight_nibbles=getattr( - self.quant_config, "swap_weight_nibbles", False - ), + swap_weight_nibbles=getattr(self.quant_config, "swap_weight_nibbles", True), ) _, flashinfer_backend = _get_fp4_gemm_op() @@ -562,13 +554,8 @@ class ModelOptFp4LinearMethod(LinearMethodBase): padded_scales[:B, :M, :K] = scales _, flashinfer_backend = _get_fp4_gemm_op() - uses_flux1_scale_layout = not getattr( - self.quant_config, "checkpoint_uses_packed_qkv", False - ) and getattr(layer, "prefix", "").startswith( - ("transformer_blocks.", "single_transformer_blocks.") - ) - if flashinfer_backend is None or uses_flux1_scale_layout: - # CUTLASS and FLUX.1 CUDNN paths need the TMA scale layout. + if flashinfer_backend is None: + # CUTLASS (sgl_kernel) path: blockwise interleave to TMA layout padded_scales = padded_scales.reshape( B, M_padded // 128, 4, 32, K_padded // 4, 4 ) diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 204608c24..dd037034a 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -88,12 +88,12 @@ def _merge_modelopt_fp4_configs( inferred_config.packed_modules_mapping = getattr( existing_config, "packed_modules_mapping", {} ) + inferred_config.swap_weight_nibbles = getattr( + existing_config, "swap_weight_nibbles", True + ) inferred_config.checkpoint_uses_packed_qkv = getattr( inferred_config, "checkpoint_uses_packed_qkv", False ) or getattr(existing_config, "checkpoint_uses_packed_qkv", False) - inferred_config.swap_weight_nibbles = getattr( - inferred_config, "swap_weight_nibbles", False - ) or getattr(existing_config, "swap_weight_nibbles", False) if getattr(inferred_config, "group_size", None) is None: inferred_config.group_size = getattr(existing_config, "group_size", None) diff --git a/python/sglang/multimodal_gen/test/server/consistency_threshold.json b/python/sglang/multimodal_gen/test/server/consistency_threshold.json index 66512f1c3..e67ae720c 100644 --- a/python/sglang/multimodal_gen/test/server/consistency_threshold.json +++ b/python/sglang/multimodal_gen/test/server/consistency_threshold.json @@ -13,12 +13,6 @@ "psnr_threshold": 24.0, "mean_abs_diff_threshold": 8.0 }, - "flux1_modelopt_fp8_t2i": { - "clip_threshold": 0.92, - "ssim_threshold": 0.94, - "psnr_threshold": 28.0, - "mean_abs_diff_threshold": 8.0 - }, "flux_2_klein_image_t2i": { "clip_threshold": 0.94, "ssim_threshold": 0.78, diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 0d3a8e44a..7ac8274d9 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -11,9 +11,8 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( MODELOPT_NVFP4_B200_ENV_VARS, MODELOPT_QWEN_IMAGE_EDIT_FP8_TRANSFORMER, MODELOPT_QWEN_IMAGE_FP8_TRANSFORMER, - MODELOPT_WAN22_FP8_MODEL, - MODELOPT_WAN22_NVFP4_B200_ENV_VARS, - MODELOPT_WAN22_NVFP4_MODEL, + MODELOPT_WAN22_FP8_TRANSFORMER, + MODELOPT_WAN22_NVFP4_TRANSFORMER, T2V_PROMPT, DiffusionSamplingParams, DiffusionServerArgs, @@ -420,7 +419,6 @@ else: modality="image", sampling_params=MODELOPT_T2I_CI_sampling_params, extras=["--transformer-path", MODELOPT_FLUX1_FP8_TRANSFORMER], - run_consistency_check=True, ), _make_modelopt_ci_case( "flux2_modelopt_fp8_t2i", @@ -428,15 +426,13 @@ else: modality="image", sampling_params=MODELOPT_T2I_CI_sampling_params, extras=["--transformer-path", MODELOPT_FLUX2_FP8_TRANSFORMER], - run_consistency_check=True, ), _make_modelopt_ci_case( "wan22_modelopt_fp8_t2v", - model_path=MODELOPT_WAN22_FP8_MODEL, + model_path=DEFAULT_WAN_2_2_T2V_A14B_MODEL_NAME_FOR_TEST, modality="video", sampling_params=MODELOPT_T2V_CI_sampling_params, - extras=[], - run_consistency_check=True, + extras=["--transformer-path", MODELOPT_WAN22_FP8_TRANSFORMER], ), _make_modelopt_ci_case( "hunyuanvideo_modelopt_fp8_t2v", @@ -473,7 +469,6 @@ else: sampling_params=MODELOPT_T2I_CI_sampling_params, extras=["--transformer-path", MODELOPT_FLUX1_NVFP4_TRANSFORMER], env_vars=MODELOPT_NVFP4_B200_ENV_VARS, - run_consistency_check=True, ), _make_modelopt_ci_case( "flux2_modelopt_nvfp4_t2i", @@ -482,16 +477,14 @@ else: sampling_params=MODELOPT_T2I_CI_sampling_params, extras=["--transformer-weights-path", MODELOPT_FLUX2_NVFP4_WEIGHTS], env_vars=MODELOPT_NVFP4_B200_ENV_VARS, - run_consistency_check=True, ), _make_modelopt_ci_case( "wan22_modelopt_nvfp4_t2v", - model_path=MODELOPT_WAN22_NVFP4_MODEL, + model_path=DEFAULT_WAN_2_2_T2V_A14B_MODEL_NAME_FOR_TEST, modality="video", sampling_params=MODELOPT_T2V_CI_sampling_params, - extras=[], - env_vars=MODELOPT_WAN22_NVFP4_B200_ENV_VARS, - run_consistency_check=True, + extras=["--transformer-path", MODELOPT_WAN22_NVFP4_TRANSFORMER], + env_vars=MODELOPT_NVFP4_B200_ENV_VARS, ), ] diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index a8ff76c3e..401f8a633 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -441,7 +441,7 @@ HUNYUAN3D_SHAPE_sampling_params = DiffusionSamplingParams( MODELOPT_FLUX1_FP8_TRANSFORMER = "lmsys/flux1-dev-modelopt-fp8-sglang-transformer" MODELOPT_FLUX2_FP8_TRANSFORMER = "lmsys/flux2-dev-modelopt-fp8-sglang-transformer" -MODELOPT_WAN22_FP8_MODEL = "nvidia/Wan2.2-T2V-A14B-Diffusers-FP8" +MODELOPT_WAN22_FP8_TRANSFORMER = "lmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformer" MODELOPT_HUNYUANVIDEO_FP8_TRANSFORMER = ( "lmsys/hunyuanvideo-modelopt-fp8-sglang-transformer" ) @@ -451,11 +451,10 @@ MODELOPT_QWEN_IMAGE_EDIT_FP8_TRANSFORMER = ( ) MODELOPT_FLUX1_NVFP4_TRANSFORMER = "lmsys/flux1-dev-modelopt-nvfp4-sglang-transformer" MODELOPT_FLUX2_NVFP4_WEIGHTS = "black-forest-labs/FLUX.2-dev-NVFP4" -MODELOPT_WAN22_NVFP4_MODEL = "nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4" +MODELOPT_WAN22_NVFP4_TRANSFORMER = ( + "lmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer" +) MODELOPT_NVFP4_B200_ENV_VARS = {"SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": "cudnn"} -MODELOPT_WAN22_NVFP4_B200_ENV_VARS = { - "SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": "trtllm" -} def _make_modelopt_ci_case( @@ -466,7 +465,6 @@ def _make_modelopt_ci_case( sampling_params: DiffusionSamplingParams, extras: list[str], env_vars: dict[str, str] | None = None, - run_consistency_check: bool = False, ) -> DiffusionTestCase: return DiffusionTestCase( case_id, @@ -479,7 +477,7 @@ def _make_modelopt_ci_case( ), sampling_params, run_perf_check=False, - run_consistency_check=run_consistency_check, + run_consistency_check=False, run_component_accuracy_check=False, ) diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 29da29308..54a3d6c13 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -33,7 +33,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "94eab4fcca6d4ddc77cdb3622f13033b61e81002" +SGL_TEST_FILES_CI_DATA_REVISION = "c8305f1dd8cc82197f36c17d0f503adc94016cc7" SGL_TEST_FILES_CONSISTENCY_GT_ROOT = ( "https://raw.githubusercontent.com/" f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/" diff --git a/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py b/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py index 7bd7b2c81..e6f13b306 100644 --- a/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py +++ b/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py @@ -218,7 +218,9 @@ def build_modelopt_nvfp4_transformer( patterns.extend(keep_bf16_patterns) resolved_swap_weight_nibbles = ( - swap_weight_nibbles if swap_weight_nibbles is not None else False + swap_weight_nibbles + if swap_weight_nibbles is not None + else (False if pattern_preset == "flux1-nvfp4" else True) ) output_config = _updated_quant_config( _load_config(source_dir), @@ -371,7 +373,7 @@ def _parse_args() -> argparse.Namespace: default=None, help=( "Whether the runtime should swap packed FP4 nibbles before padding. " - "Defaults to false." + "Defaults to false for --pattern-preset flux1-nvfp4 and true otherwise." ), ) parser.add_argument(