From ccbbae00eaeaff7f7c5cf3ded5877831f0db8239 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Wed, 20 May 2026 22:15:25 +0800 Subject: [PATCH] [codex] Reland Wan2.2 ModelOpt CI checkpoints (#25857) --- .../docs/sglang-diffusion/quantization.mdx | 62 +++++++++---------- .../test_diffusion_nvfp4_scaled_mm.py | 6 +- python/sglang/multimodal_gen/registry.py | 5 +- .../layers/quantization/modelopt_quant.py | 27 +++++--- .../runtime/loader/transformer_load_utils.py | 6 +- .../test/server/consistency_threshold.json | 6 ++ .../multimodal_gen/test/server/gpu_cases.py | 21 ++++--- .../test/server/testcase_configs.py | 12 ++-- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../tools/build_modelopt_nvfp4_transformer.py | 6 +- 10 files changed, 92 insertions(+), 61 deletions(-) diff --git a/docs_new/docs/sglang-diffusion/quantization.mdx b/docs_new/docs/sglang-diffusion/quantization.mdx index 81e3e1ed4..043ef28ce 100644 --- a/docs_new/docs/sglang-diffusion/quantization.mdx +++ b/docs_new/docs/sglang-diffusion/quantization.mdx @@ -110,14 +110,15 @@ backend. ## Validated ModelOpt Checkpoints This section is the canonical support matrix for the nine diffusion ModelOpt -checkpoints currently wired up in SGLang docs and B200 CI coverage. +checkpoints currently wired up in SGLang docs and validation 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`. -Eight of the nine repos live under `lmsys/*`. The FLUX.2 NVFP4 entry keeps the -official `black-forest-labs/FLUX.2-dev-NVFP4` repo. +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. @@ -158,10 +159,10 @@ official `black-forest-labs/FLUX.2-dev-NVFP4` repo. - - - - + + + + @@ -206,24 +207,25 @@ official `black-forest-labs/FLUX.2-dev-NVFP4` repo. - - - - + + + +
FP8 Wan-AI/Wan2.2-T2V-A14B-Diffusers--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--model-pathnvidia/Wan2.2-T2V-A14B-Diffusers-FP8full Diffusers repo with ModelOpt FP8 Wan2.2 componentsvalidated through direct --model-path loading
FP8
NVFP4 Wan-AI/Wan2.2-T2V-A14B-Diffusers--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--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
-These nine checkpoints are also the intended case set for the B200 diffusion CI -job (`multimodal-gen-test-1-b200`). +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`). ## ModelOpt FP8 ### Usage Examples -Converted ModelOpt FP8 checkpoints should be loaded as transformer component -overrides. If the repo or local directory already contains `config.json`, use -`--transformer-path`. +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`. ```bash sglang generate \ @@ -235,8 +237,7 @@ sglang generate \ ```bash sglang generate \ - --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ - --transformer-path lmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformer \ + --model-path nvidia/Wan2.2-T2V-A14B-Diffusers-FP8 \ --prompt "a fox walking through neon rain" \ --save-output ``` @@ -323,14 +324,12 @@ sglang generate \ --save-output ``` -For a dual-transformer Wan2.2 export where only the primary `transformer` -was quantized: +For Wan2.2 NVFP4: ```bash -SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=cudnn \ +SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=trtllm \ sglang generate \ - --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ - --transformer-path lmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer \ + --model-path nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4 \ --prompt "a fox walking through neon rain" \ --save-output ``` @@ -341,17 +340,16 @@ 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 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`. +- 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`. - On Blackwell, the validated Wan2.2 ModelOpt NVFP4 path currently prefers FlashInfer FP4 GEMM via - `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. + `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. - 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/jit_kernel/tests/diffusion/test_diffusion_nvfp4_scaled_mm.py b/python/sglang/jit_kernel/tests/diffusion/test_diffusion_nvfp4_scaled_mm.py index 2122601cb..1214497d7 100644 --- a/python/sglang/jit_kernel/tests/diffusion/test_diffusion_nvfp4_scaled_mm.py +++ b/python/sglang/jit_kernel/tests/diffusion/test_diffusion_nvfp4_scaled_mm.py @@ -140,7 +140,11 @@ def _build_layer( output_size, input_size_half = weight_fp4.shape input_size = input_size_half * 2 method = ModelOptFp4LinearMethod( - ModelOptFp4Config(is_checkpoint_nvfp4_serialized=True, group_size=BLOCK_SIZE) + ModelOptFp4Config( + is_checkpoint_nvfp4_serialized=True, + group_size=BLOCK_SIZE, + swap_weight_nibbles=True, + ) ) layer = torch.nn.Module() method.create_weights( diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index 232384e01..484219238 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -739,7 +739,10 @@ 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"], + hf_model_paths=[ + "Wan-AI/Wan2.2-T2V-A14B-Diffusers", + "nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4", + ], ) 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 c9286aba5..2f3d7c33e 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 = True, + swap_weight_nibbles: bool = False, ) -> 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 = True + swap_weight_nibbles = False # Flat format (config.json quantization_config) quant_method = config.get("quant_algo") @@ -273,7 +273,10 @@ 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", True) + swap_weight_nibbles = config.get( + "swap_weight_nibbles", + config.get("checkpoint_uses_packed_qkv", False), + ) else: # Nested format (hf_quant_config.json) try: @@ -283,7 +286,10 @@ class ModelOptFp4Config(ModelOptQuantConfig): exclude_modules = quant_config.get("exclude_modules", []) swap_weight_nibbles = quant_config.get( "swap_weight_nibbles", - config.get("swap_weight_nibbles", True), + config.get( + "swap_weight_nibbles", + config.get("checkpoint_uses_packed_qkv", False), + ), ) except (ValueError, KeyError): raise ValueError("Cannot find 'quant_algo' in quantization config.") @@ -494,7 +500,9 @@ class ModelOptFp4LinearMethod(LinearMethodBase): w = layer.weight.data w_swapped = _prepare_nvfp4_weight_bytes( w, - swap_weight_nibbles=getattr(self.quant_config, "swap_weight_nibbles", True), + swap_weight_nibbles=getattr( + self.quant_config, "swap_weight_nibbles", False + ), ) _, flashinfer_backend = _get_fp4_gemm_op() @@ -554,8 +562,13 @@ class ModelOptFp4LinearMethod(LinearMethodBase): padded_scales[:B, :M, :K] = scales _, flashinfer_backend = _get_fp4_gemm_op() - if flashinfer_backend is None: - # CUTLASS (sgl_kernel) path: blockwise interleave to TMA layout + 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. 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 dd037034a..204608c24 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 e67ae720c..66512f1c3 100644 --- a/python/sglang/multimodal_gen/test/server/consistency_threshold.json +++ b/python/sglang/multimodal_gen/test/server/consistency_threshold.json @@ -13,6 +13,12 @@ "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 7ac8274d9..0d3a8e44a 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -11,8 +11,9 @@ 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_TRANSFORMER, - MODELOPT_WAN22_NVFP4_TRANSFORMER, + MODELOPT_WAN22_FP8_MODEL, + MODELOPT_WAN22_NVFP4_B200_ENV_VARS, + MODELOPT_WAN22_NVFP4_MODEL, T2V_PROMPT, DiffusionSamplingParams, DiffusionServerArgs, @@ -419,6 +420,7 @@ 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", @@ -426,13 +428,15 @@ 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=DEFAULT_WAN_2_2_T2V_A14B_MODEL_NAME_FOR_TEST, + model_path=MODELOPT_WAN22_FP8_MODEL, modality="video", sampling_params=MODELOPT_T2V_CI_sampling_params, - extras=["--transformer-path", MODELOPT_WAN22_FP8_TRANSFORMER], + extras=[], + run_consistency_check=True, ), _make_modelopt_ci_case( "hunyuanvideo_modelopt_fp8_t2v", @@ -469,6 +473,7 @@ 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", @@ -477,14 +482,16 @@ 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=DEFAULT_WAN_2_2_T2V_A14B_MODEL_NAME_FOR_TEST, + model_path=MODELOPT_WAN22_NVFP4_MODEL, modality="video", sampling_params=MODELOPT_T2V_CI_sampling_params, - extras=["--transformer-path", MODELOPT_WAN22_NVFP4_TRANSFORMER], - env_vars=MODELOPT_NVFP4_B200_ENV_VARS, + extras=[], + env_vars=MODELOPT_WAN22_NVFP4_B200_ENV_VARS, + run_consistency_check=True, ), ] diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 401f8a633..a8ff76c3e 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_TRANSFORMER = "lmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformer" +MODELOPT_WAN22_FP8_MODEL = "nvidia/Wan2.2-T2V-A14B-Diffusers-FP8" MODELOPT_HUNYUANVIDEO_FP8_TRANSFORMER = ( "lmsys/hunyuanvideo-modelopt-fp8-sglang-transformer" ) @@ -451,10 +451,11 @@ 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_TRANSFORMER = ( - "lmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer" -) +MODELOPT_WAN22_NVFP4_MODEL = "nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4" 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( @@ -465,6 +466,7 @@ 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, @@ -477,7 +479,7 @@ def _make_modelopt_ci_case( ), sampling_params, run_perf_check=False, - run_consistency_check=False, + run_consistency_check=run_consistency_check, 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 54a3d6c13..29da29308 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 = "c8305f1dd8cc82197f36c17d0f503adc94016cc7" +SGL_TEST_FILES_CI_DATA_REVISION = "94eab4fcca6d4ddc77cdb3622f13033b61e81002" 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 e6f13b306..7bd7b2c81 100644 --- a/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py +++ b/python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py @@ -218,9 +218,7 @@ 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 if pattern_preset == "flux1-nvfp4" else True) + swap_weight_nibbles if swap_weight_nibbles is not None else False ) output_config = _updated_quant_config( _load_config(source_dir), @@ -373,7 +371,7 @@ def _parse_args() -> argparse.Namespace: default=None, help=( "Whether the runtime should swap packed FP4 nibbles before padding. " - "Defaults to false for --pattern-preset flux1-nvfp4 and true otherwise." + "Defaults to false." ), ) parser.add_argument(