From 5c3775823e81d2a7a6430d2250690f0c6e2ce5ba Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 25 May 2026 00:25:34 +0800 Subject: [PATCH] [diffusion] chore: use model-aware vae channels_last_3d policy (#26214) --- python/sglang/multimodal_gen/envs.py | 6 +- .../loader/component_loaders/vae_loader.py | 17 ++- .../server/test_component_accuracy_1_gpu.py | 4 +- .../server/test_component_accuracy_2_gpu.py | 4 +- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../test/unit/test_vae_loader.py | 123 ++++++++++++++++++ 6 files changed, 147 insertions(+), 9 deletions(-) diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index 76f8af71c..febd85b60 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -56,7 +56,7 @@ if TYPE_CHECKING: # model loading SGLANG_USE_RUNAI_MODEL_STREAMER: bool = True SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None - SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: bool = True + SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: str = "auto" SGLANG_USE_CUDA_HUNYUANVIDEO_GROUP_NORM_SILU: bool = False SGLANG_USE_ROCM_VAE: bool = False SGLANG_USE_ROCM_CUDNN_BENCHMARK: bool = False @@ -250,8 +250,8 @@ environment_variables: dict[str, Callable[[], Any]] = { # If set, sgl_diffusion will enable stage logging, which will print the time # taken for each stage "SGLANG_DIFFUSION_STAGE_LOGGING": _lazy_bool("SGLANG_DIFFUSION_STAGE_LOGGING"), - "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": _lazy_bool( - "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", "true" + "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": _lazy_str( + "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", "auto" ), # ================== cache-dit Env Vars ================== # Enable cache-dit acceleration for DiT inference diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py index 2f63437bd..9c21d44d1 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py @@ -60,7 +60,9 @@ def _convert_conv3d_weights_to_channels_last_3d(module: nn.Module) -> int: return num_converted -def _should_use_channels_last_3d(server_args: ServerArgs, component_name: str) -> bool: +def _should_use_channels_last_3d( + server_args: ServerArgs | None, component_name: str +) -> bool: if component_name not in ( "vae", "video_vae", @@ -68,9 +70,18 @@ def _should_use_channels_last_3d(server_args: ServerArgs, component_name: str) - return False override = os.getenv(VAE_CHANNELS_LAST_3D_ENV) - if override is None or override.strip().lower() == "auto": + if override is not None and override.strip().lower() != "auto": + return get_bool_env_var(VAE_CHANNELS_LAST_3D_ENV) + + if server_args is None: + return False + + pipeline_name = server_args.pipeline_config.__class__.__name__ + if pipeline_name.startswith("QwenImage"): return True - return get_bool_env_var(VAE_CHANNELS_LAST_3D_ENV) + if "Wan" in pipeline_name and server_args.num_gpus == 1: + return True + return False class VAELoader(ComponentLoader): diff --git a/python/sglang/multimodal_gen/test/server/test_component_accuracy_1_gpu.py b/python/sglang/multimodal_gen/test/server/test_component_accuracy_1_gpu.py index 716ce0196..67741bb6f 100644 --- a/python/sglang/multimodal_gen/test/server/test_component_accuracy_1_gpu.py +++ b/python/sglang/multimodal_gen/test/server/test_component_accuracy_1_gpu.py @@ -15,7 +15,9 @@ from sglang.multimodal_gen.test.server.accuracy_utils import ( ) from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine -VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS = {"wan2_1_t2v_1.3b"} +VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS = { + "wan2_1_t2v_1.3b", +} VAE_CHANNELS_LAST_3D_PARITY_CASES = [ case for case in ACCURACY_ONE_GPU_CASES diff --git a/python/sglang/multimodal_gen/test/server/test_component_accuracy_2_gpu.py b/python/sglang/multimodal_gen/test/server/test_component_accuracy_2_gpu.py index 87b96fab5..bb5c0a54a 100644 --- a/python/sglang/multimodal_gen/test/server/test_component_accuracy_2_gpu.py +++ b/python/sglang/multimodal_gen/test/server/test_component_accuracy_2_gpu.py @@ -15,7 +15,9 @@ from sglang.multimodal_gen.test.server.accuracy_utils import ( ) from sglang.multimodal_gen.test.server.component_accuracy import AccuracyEngine -VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS = {"wan2_2_i2v_a14b_2gpu"} +VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS = { + "wan2_2_i2v_a14b_2gpu", +} VAE_CHANNELS_LAST_3D_PARITY_CASES = [ case for case in ACCURACY_TWO_GPU_CASES diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 7c9a440b4..0e115293b 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 = "5b728ad6dee869c0c720ef3b668ac7a0d98b0f9a" +SGL_TEST_FILES_CI_DATA_REVISION = "b7455318873fc5af399c8447b3bb0d9471a5084c" 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/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index 6198b22b9..df0e9761a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -3,12 +3,40 @@ from unittest.mock import patch import torch +from sglang.multimodal_gen.runtime.loader.component_loaders import vae_loader from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import ( _backfill_ltx2_audio_vae_latent_stats, + _should_use_channels_last_3d, ) from sglang.multimodal_gen.runtime.models.vaes.parallel import wan_common_utils +class _FakeServerArgs: + def __init__(self, pipeline_config, num_gpus=1): + self.pipeline_config = pipeline_config + self.num_gpus = num_gpus + + +class QwenImagePipelineConfig: + pass + + +class WanT2V480PConfig: + pass + + +class FastWan2_2_TI2V_5B_Config: + pass + + +class Wan2_2_I2V_A14B_Config: + pass + + +class LTX2PipelineConfig: + pass + + class TestVAELoader(unittest.TestCase): def test_backfill_ltx2_audio_vae_latent_stats_maps_official_keys(self): loaded = { @@ -45,6 +73,101 @@ class TestVAELoader(unittest.TestCase): self.assertNotIn("latents_mean", loaded) self.assertNotIn("latents_std", loaded) + def test_channels_last_3d_defaults_true_for_qwen_image_on_cuda(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + self.assertTrue(_should_use_channels_last_3d(server_args, "vae")) + + def test_channels_last_3d_defaults_true_for_single_gpu_wan_on_cuda(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(WanT2V480PConfig(), num_gpus=1) + self.assertTrue(_should_use_channels_last_3d(server_args, "video_vae")) + + def test_channels_last_3d_defaults_true_for_single_gpu_fast_wan_on_cuda(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(FastWan2_2_TI2V_5B_Config(), num_gpus=1) + self.assertTrue(_should_use_channels_last_3d(server_args, "video_vae")) + + def test_channels_last_3d_defaults_false_for_multi_gpu_wan_on_cuda(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(Wan2_2_I2V_A14B_Config(), num_gpus=2) + self.assertFalse(_should_use_channels_last_3d(server_args, "video_vae")) + + def test_channels_last_3d_defaults_false_for_ltx_on_cuda(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(LTX2PipelineConfig(), num_gpus=2) + self.assertFalse(_should_use_channels_last_3d(server_args, "video_vae")) + + def test_channels_last_3d_can_be_disabled_by_env(self): + with ( + patch.dict( + "os.environ", {"SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": "false"} + ), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + self.assertFalse(_should_use_channels_last_3d(server_args, "vae")) + + def test_channels_last_3d_can_be_enabled_by_env(self): + with ( + patch.dict("os.environ", {"SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": "true"}), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(LTX2PipelineConfig(), num_gpus=2) + self.assertTrue(_should_use_channels_last_3d(server_args, "video_vae")) + + def test_channels_last_3d_auto_uses_model_policy(self): + with ( + patch.dict("os.environ", {"SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": "auto"}), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + wan_args = _FakeServerArgs(WanT2V480PConfig(), num_gpus=1) + ltx_args = _FakeServerArgs(LTX2PipelineConfig(), num_gpus=2) + + self.assertTrue(_should_use_channels_last_3d(wan_args, "video_vae")) + self.assertFalse(_should_use_channels_last_3d(ltx_args, "video_vae")) + + def test_channels_last_3d_skips_non_video_vae_components(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=True), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + self.assertFalse(_should_use_channels_last_3d(server_args, "audio_vae")) + + def test_channels_last_3d_skips_unsupported_platforms(self): + with ( + patch.dict("os.environ", {}, clear=True), + patch.object(vae_loader.current_platform, "is_cuda", return_value=False), + patch.object(vae_loader.current_platform, "is_rocm", return_value=False), + ): + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + self.assertFalse(_should_use_channels_last_3d(server_args, "vae")) + @unittest.skipUnless( hasattr(torch, "channels_last_3d"), "channels_last_3d is unavailable" )