From c8cea6d4aad92a6aa2c2821f73f6a152cee77a2b Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 23 May 2026 10:20:30 +0800 Subject: [PATCH] [diffusion] feat: auto-select vae channels_last_3d (#26121) --- .../loader/component_loaders/vae_loader.py | 28 +-- .../models/vaes/parallel/wan_common_utils.py | 8 +- .../test/server/accuracy_hooks.py | 33 +++- .../test/server/accuracy_utils.py | 4 +- .../test/server/component_accuracy.py | 175 ++++++++++++++++++ .../server/test_component_accuracy_1_gpu.py | 20 ++ .../server/test_component_accuracy_2_gpu.py | 20 ++ .../test/unit/test_vae_loader.py | 44 +++++ 8 files changed, 311 insertions(+), 21 deletions(-) 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 2ff38095e..2f63437bd 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 @@ -5,7 +5,6 @@ import torch import torch.nn as nn from safetensors.torch import load_file as safetensors_load_file -from sglang.multimodal_gen import envs from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentLoader, @@ -18,6 +17,7 @@ from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.models.registry import ModelRegistry from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( get_diffusers_component_config, ) @@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.utils import PRECISION_TO_TYPE logger = init_logger(__name__) +VAE_CHANNELS_LAST_3D_ENV = "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D" def _backfill_ltx2_audio_vae_latent_stats( @@ -59,6 +60,19 @@ 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: + if component_name not in ( + "vae", + "video_vae", + ) or not (current_platform.is_cuda() or current_platform.is_rocm()): + return False + + override = os.getenv(VAE_CHANNELS_LAST_3D_ENV) + if override is None or override.strip().lower() == "auto": + return True + return get_bool_env_var(VAE_CHANNELS_LAST_3D_ENV) + + class VAELoader(ComponentLoader): """Shared loader for (video/audio) VAE modules.""" @@ -120,11 +134,7 @@ class VAELoader(ComponentLoader): trust_remote_code=server_args.trust_remote_code, ) vae = vae.to(device=target_device, dtype=vae_dtype) - if ( - component_name in ("vae", "video_vae") - and torch.cuda.is_available() - and getattr(envs, "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", False) - ): + if _should_use_channels_last_3d(server_args, component_name): n = _convert_conv3d_weights_to_channels_last_3d(vae) if n > 0: logger.info( @@ -167,11 +177,7 @@ class VAELoader(ComponentLoader): if unexpected_keys: logger.warning("VAE unexpected keys: %s", unexpected_keys) - if ( - component_name in ("vae", "video_vae") - and torch.cuda.is_available() - and getattr(envs, "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", False) - ): + if _should_use_channels_last_3d(server_args, component_name): n = _convert_conv3d_weights_to_channels_last_3d(vae) if n > 0: logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py index eb1c5ece0..4d4fbde43 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/parallel/wan_common_utils.py @@ -7,10 +7,16 @@ import torch.nn.functional as F from sglang.multimodal_gen.runtime.platforms import current_platform +def _channels_last_3d_supported_by_platform() -> bool: + return hasattr(torch, "channels_last_3d") and ( + current_platform.is_cuda() or current_platform.is_rocm() + ) + + def _conv3d_weight_is_channels_last_3d(weight: torch.Tensor) -> bool: return ( weight.dim() == 5 - and hasattr(torch, "channels_last_3d") + and _channels_last_3d_supported_by_platform() and weight.is_contiguous(memory_format=torch.channels_last_3d) ) diff --git a/python/sglang/multimodal_gen/test/server/accuracy_hooks.py b/python/sglang/multimodal_gen/test/server/accuracy_hooks.py index a495904d8..09ca8f17d 100644 --- a/python/sglang/multimodal_gen/test/server/accuracy_hooks.py +++ b/python/sglang/multimodal_gen/test/server/accuracy_hooks.py @@ -38,6 +38,7 @@ DEFAULT_TRANSFORMER_TEXT_CHANNELS = 4096 DEFAULT_TRANSFORMER_POOLED_CHANNELS = 768 DEFAULT_VAE_LATENT_CHANNELS = 16 DEFAULT_VAE_LATENT_SPATIAL_SIZE = 32 +DEFAULT_VAE_VIDEO_LATENT_FRAMES = 3 LARGE_CHANNEL_LAYOUT_THRESHOLD = 128 @@ -610,17 +611,35 @@ def _infer_vae_latent_channels(model: nn.Module) -> int: def _build_vae_hook_inputs( case: Any, model: nn.Module, device: str, ref_model: Optional[nn.Module] = None ) -> Inputs: - del case, ref_model + del ref_model latent_channels = _infer_vae_latent_channels(model) + model_path = getattr(getattr(case, "server_args", None), "model_path", "").lower() + modality = getattr(getattr(case, "server_args", None), "modality", None) + use_wan_video_latent = ( + modality == "video" + and "wan" in model_path + and any(isinstance(module, nn.Conv3d) for module in model.modules()) + ) + shape = ( + ( + 1, + latent_channels, + DEFAULT_VAE_VIDEO_LATENT_FRAMES, + DEFAULT_VAE_LATENT_SPATIAL_SIZE, + DEFAULT_VAE_LATENT_SPATIAL_SIZE, + ) + if use_wan_video_latent + else ( + 1, + latent_channels, + DEFAULT_VAE_LATENT_SPATIAL_SIZE, + DEFAULT_VAE_LATENT_SPATIAL_SIZE, + ) + ) rng = _DeterministicRNG() return { "z": rng.randn( - ( - 1, - latent_channels, - DEFAULT_VAE_LATENT_SPATIAL_SIZE, - DEFAULT_VAE_LATENT_SPATIAL_SIZE, - ), + shape, device, torch.bfloat16, ) diff --git a/python/sglang/multimodal_gen/test/server/accuracy_utils.py b/python/sglang/multimodal_gen/test/server/accuracy_utils.py index cd544e6b9..6a2dcb6bf 100644 --- a/python/sglang/multimodal_gen/test/server/accuracy_utils.py +++ b/python/sglang/multimodal_gen/test/server/accuracy_utils.py @@ -760,12 +760,12 @@ def _run_staged_native_component_accuracy_case( ref = ref.to(device=device, dtype=torch.bfloat16).eval() if component == ComponentType.VAE: - from sglang.multimodal_gen import envs from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import ( _convert_conv3d_weights_to_channels_last_3d, + _should_use_channels_last_3d, ) - if torch.cuda.is_available() and envs.SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: + if _should_use_channels_last_3d(runtime_server_args, "vae"): _convert_conv3d_weights_to_channels_last_3d(ref) ref_call = profile.prepare_reference_call(ref, inputs) ref_autocast = ( diff --git a/python/sglang/multimodal_gen/test/server/component_accuracy.py b/python/sglang/multimodal_gen/test/server/component_accuracy.py index 1e97beb76..8f6a26561 100644 --- a/python/sglang/multimodal_gen/test/server/component_accuracy.py +++ b/python/sglang/multimodal_gen/test/server/component_accuracy.py @@ -2,6 +2,7 @@ from __future__ import annotations import gc import os +from contextlib import contextmanager from dataclasses import dataclass from typing import Any, Dict, List, Optional, Tuple @@ -74,6 +75,69 @@ from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase logger = init_logger(__name__) MIN_MATCH_RATIO = float(os.getenv("SGLANG_DIFFUSION_WEIGHT_MATCH_RATIO", "0.98")) +VAE_CHANNELS_LAST_3D_ENV = "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D" +VAE_CHANNELS_LAST_3D_PARITY_THRESHOLD = float( + os.getenv("SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D_PARITY_THRESHOLD", "0.999") +) + + +@contextmanager +def _temporary_vae_channels_last_3d(enabled: bool): + previous = os.environ.get(VAE_CHANNELS_LAST_3D_ENV) + os.environ[VAE_CHANNELS_LAST_3D_ENV] = "true" if enabled else "false" + try: + yield + finally: + if previous is None: + os.environ.pop(VAE_CHANNELS_LAST_3D_ENV, None) + else: + os.environ[VAE_CHANNELS_LAST_3D_ENV] = previous + + +@dataclass +class Conv3dLayoutStats: + calls: int = 0 + channels_last_input_calls: int = 0 + channels_last_weight_calls: int = 0 + mixed_layout_calls: int = 0 + + +@contextmanager +def _record_conv3d_layouts(): + stats = Conv3dLayoutStats() + original_conv3d = torch.nn.functional.conv3d + + def wrapped_conv3d(input, weight, *args, **kwargs): + if ( + isinstance(input, torch.Tensor) + and isinstance(weight, torch.Tensor) + and input.dim() == 5 + and weight.dim() == 5 + and hasattr(torch, "channels_last_3d") + ): + input_channels_last = input.is_contiguous( + memory_format=torch.channels_last_3d + ) + weight_channels_last = weight.is_contiguous( + memory_format=torch.channels_last_3d + ) + else: + input_channels_last = False + weight_channels_last = False + + stats.calls += 1 + stats.channels_last_input_calls += int(input_channels_last) + stats.channels_last_weight_calls += int(weight_channels_last) + stats.mixed_layout_calls += int( + weight_channels_last and not input_channels_last + ) + return original_conv3d(input, weight, *args, **kwargs) + + torch.nn.functional.conv3d = wrapped_conv3d + try: + yield stats + finally: + torch.nn.functional.conv3d = original_conv3d @dataclass(frozen=True) @@ -584,3 +648,114 @@ class AccuracyEngine: ) return sgl_component.eval(), ref_component.eval(), str(device) + + @staticmethod + def run_vae_channels_last_3d_parity( + case: DiffusionTestCase, + num_gpus: int, + ) -> None: + component = ComponentType.VAE + spec = COMPONENT_SPECS[component] + hub_id = case.server_args.model_path + component_selection = select_component_source( + hub_id, + case.server_args.extras, + component, + spec.model_index_keys, + ) + sgl_args = build_accuracy_server_args( + component_selection.base_model_id, + component_selection.base_model_root, + case, + component, + num_gpus, + component_selection.component_paths, + ) + + baseline_vae = None + channels_last_vae = None + try: + initialize_parallel_runtime(sgl_args) + set_global_server_args(sgl_args) + device = get_local_torch_device() + + with _temporary_vae_channels_last_3d(False): + baseline_vae = _load_sglang_component( + component_selection.source_path, + sgl_args, + component, + spec.reference_library, + ).to(device=device, dtype=torch.bfloat16) + + with _temporary_vae_channels_last_3d(True): + channels_last_vae = _load_sglang_component( + component_selection.source_path, + sgl_args, + component, + spec.reference_library, + ).to(device=device, dtype=torch.bfloat16) + + baseline_vae.eval() + channels_last_vae.eval() + + profile = resolve_component_native_profile(component) + inputs = profile.build_inputs( + case, baseline_vae, str(device), channels_last_vae + ) + baseline_call = profile.prepare_sglang_call(baseline_vae, inputs) + channels_last_call = profile.prepare_sglang_call(channels_last_vae, inputs) + + with torch.no_grad(): + with _record_conv3d_layouts() as baseline_layout: + baseline_raw = AccuracyEngine._execute_with_native_hook( + baseline_call + ) + with _record_conv3d_layouts() as channels_last_layout: + channels_last_raw = AccuracyEngine._execute_with_native_hook( + channels_last_call + ) + + baseline_out = profile.normalize_sglang_output(baseline_raw) + channels_last_out = profile.normalize_sglang_output(channels_last_raw) + + AccuracyEngine.check_accuracy( + channels_last_out, + baseline_out, + f"{case.id}_vae_channels_last_3d", + VAE_CHANNELS_LAST_3D_PARITY_THRESHOLD, + ) + + logger.info( + "[%s_vae_channels_last_3d] Conv3d layout baseline: calls=%d, " + "input_cl3d=%d, weight_cl3d=%d, mixed=%d | channels_last: " + "calls=%d, input_cl3d=%d, weight_cl3d=%d, mixed=%d", + case.id, + baseline_layout.calls, + baseline_layout.channels_last_input_calls, + baseline_layout.channels_last_weight_calls, + baseline_layout.mixed_layout_calls, + channels_last_layout.calls, + channels_last_layout.channels_last_input_calls, + channels_last_layout.channels_last_weight_calls, + channels_last_layout.mixed_layout_calls, + ) + if channels_last_layout.calls == 0: + raise RuntimeError( + f"{case.id}: VAE channels_last_3d guard did not execute Conv3d" + ) + if channels_last_layout.channels_last_weight_calls == 0: + raise RuntimeError( + f"{case.id}: VAE channels_last_3d guard did not see channels_last_3d Conv3d weights" + ) + if channels_last_layout.mixed_layout_calls: + raise RuntimeError( + f"{case.id}: {channels_last_layout.mixed_layout_calls} Conv3d calls used " + "channels_last_3d weights with non-channels_last_3d inputs" + ) + finally: + if baseline_vae is not None: + del baseline_vae + if channels_last_vae is not None: + del channels_last_vae + AccuracyEngine.reset_parallel_runtime() + AccuracyEngine.clear_memory() 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 1d41f9b0b..716ce0196 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,6 +15,13 @@ 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_CASES = [ + case + for case in ACCURACY_ONE_GPU_CASES + if case.id in VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS +] + @pytest.mark.parametrize("case", ACCURACY_ONE_GPU_CASES, ids=lambda case: case.id) class TestComponentAccuracy1GPU: @@ -63,3 +70,16 @@ class TestComponentAccuracy1GPU: case, case.server_args.num_gpus, ) + + +@pytest.mark.parametrize( + "case", VAE_CHANNELS_LAST_3D_PARITY_CASES, ids=lambda case: case.id +) +class TestVAEChannelsLast3DParity1GPU: + """1-GPU VAE guard for channels_last_3d drift.""" + + def test_vae_channels_last_3d_parity(self, case): + AccuracyEngine.run_vae_channels_last_3d_parity( + case, + case.server_args.num_gpus, + ) 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 8e8d3dfaa..87b96fab5 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,6 +15,13 @@ 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_CASES = [ + case + for case in ACCURACY_TWO_GPU_CASES + if case.id in VAE_CHANNELS_LAST_3D_PARITY_CASE_IDS +] + @pytest.mark.parametrize("case", ACCURACY_TWO_GPU_CASES, ids=lambda case: case.id) class TestComponentAccuracy2GPU: @@ -63,3 +70,16 @@ class TestComponentAccuracy2GPU: case, case.server_args.num_gpus, ) + + +@pytest.mark.parametrize( + "case", VAE_CHANNELS_LAST_3D_PARITY_CASES, ids=lambda case: case.id +) +class TestVAEChannelsLast3DParity2GPU: + """2-GPU VAE guard for channels_last_3d drift.""" + + def test_vae_channels_last_3d_parity(self, case): + AccuracyEngine.run_vae_channels_last_3d_parity( + case, + case.server_args.num_gpus, + ) 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 fd9b52f56..6198b22b9 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -1,10 +1,12 @@ import unittest +from unittest.mock import patch import torch from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import ( _backfill_ltx2_audio_vae_latent_stats, ) +from sglang.multimodal_gen.runtime.models.vaes.parallel import wan_common_utils class TestVAELoader(unittest.TestCase): @@ -43,6 +45,48 @@ class TestVAELoader(unittest.TestCase): self.assertNotIn("latents_mean", loaded) self.assertNotIn("latents_std", loaded) + @unittest.skipUnless( + hasattr(torch, "channels_last_3d"), "channels_last_3d is unavailable" + ) + def test_match_conv3d_input_format_skips_non_cuda_platforms(self): + x = torch.randn(1, 3, 2, 4, 4) + weight = torch.randn(3, 3, 1, 1, 1).contiguous( + memory_format=torch.channels_last_3d + ) + + with ( + patch.object( + wan_common_utils.current_platform, "is_cuda", return_value=False + ), + patch.object( + wan_common_utils.current_platform, "is_rocm", return_value=False + ), + ): + out = wan_common_utils.match_conv3d_input_format(x, weight) + + self.assertIs(out, x) + + @unittest.skipUnless( + hasattr(torch, "channels_last_3d"), "channels_last_3d is unavailable" + ) + def test_match_conv3d_input_format_uses_channels_last_3d_on_cuda(self): + x = torch.randn(1, 3, 2, 4, 4) + weight = torch.randn(3, 3, 1, 1, 1).contiguous( + memory_format=torch.channels_last_3d + ) + + with ( + patch.object( + wan_common_utils.current_platform, "is_cuda", return_value=True + ), + patch.object( + wan_common_utils.current_platform, "is_rocm", return_value=False + ), + ): + out = wan_common_utils.match_conv3d_input_format(x, weight) + + self.assertTrue(out.is_contiguous(memory_format=torch.channels_last_3d)) + if __name__ == "__main__": unittest.main()