[codex] Reland Wan2.2 ModelOpt CI checkpoints (#25857)

This commit is contained in:
Xiaoyu Zhang
2026-05-20 22:15:25 +08:00
committed by GitHub
parent 55ba03db6a
commit ccbbae00ea
10 changed files with 92 additions and 61 deletions
+30 -32
View File
@@ -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.
<table style={{width: "100%", borderCollapse: "collapse", tableLayout: "fixed"}}>
<colgroup>
@@ -158,10 +159,10 @@ official `black-forest-labs/FLUX.2-dev-NVFP4` repo.
<tr>
<td><code>FP8</code></td>
<td><code>Wan-AI/Wan2.2-T2V-A14B-Diffusers</code></td>
<td><code>--transformer-path</code></td>
<td><code>lmsys/wan22-t2v-a14b-modelopt-fp8-sglang-transformer</code></td>
<td>primary <code>transformer</code> quantized, <code>transformer_2</code> kept BF16</td>
<td>primary-transformer-only path; keep <code>transformer_2</code> on the base checkpoint, and do not describe this as dual-transformer full-model FP8 unless that path is validated separately</td>
<td><code>--model-path</code></td>
<td><code>nvidia/Wan2.2-T2V-A14B-Diffusers-FP8</code></td>
<td>full Diffusers repo with ModelOpt FP8 Wan2.2 components</td>
<td>validated through direct <code>--model-path</code> loading</td>
</tr>
<tr>
<td><code>FP8</code></td>
@@ -206,24 +207,25 @@ official `black-forest-labs/FLUX.2-dev-NVFP4` repo.
<tr>
<td><code>NVFP4</code></td>
<td><code>Wan-AI/Wan2.2-T2V-A14B-Diffusers</code></td>
<td><code>--transformer-path</code></td>
<td><code>lmsys/wan22-t2v-a14b-modelopt-nvfp4-sglang-transformer</code></td>
<td>primary <code>transformer</code> quantized with ModelOpt NVFP4, <code>transformer_2</code> kept BF16</td>
<td>primary-transformer-only path; keep <code>transformer_2</code> on the base checkpoint, and current B200/Blackwell bring-up uses <code>SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=cudnn</code></td>
<td><code>--model-path</code></td>
<td><code>nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4</code></td>
<td>full Diffusers repo with ModelOpt NVFP4 Wan2.2 components</td>
<td>current B200/Blackwell bring-up uses <code>SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND=trtllm</code></td>
</tr>
</tbody>
</table>
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
@@ -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(
+4 -1
View File
@@ -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,
@@ -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
)
@@ -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)
@@ -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,
@@ -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,
),
]
@@ -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,
)
@@ -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}/"
@@ -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(