[diffusion] optimization: keep vae decoder weights in their decode dtype from load (#35967)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-22 21:31:06 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 3c69a4c744
commit 61981e1fcd
3 changed files with 131 additions and 1 deletions
+7
View File
@@ -21,6 +21,7 @@ if TYPE_CHECKING:
SGLANG_DIFFUSION_LOGGING_LEVEL: str = "INFO"
SGLANG_DIFFUSION_LOGGING_PREFIX: str = ""
SGLANG_DIFFUSION_TRACE_FUNCTION: int = 0
SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST: bool = False
SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork"
SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda"
SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = ""
@@ -265,6 +266,12 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": _lazy_str(
"SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", "auto"
),
# Kill-switch: keep VAE decoder weights in their checkpoint dtype at load
# instead of the decode compute dtype the decode stage would round them to
# on first use anyway.
"SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST": _lazy_bool(
"SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST"
),
# ================== cache-dit Env Vars ==================
# Enable cache-dit acceleration for DiT inference
# CUDA-IPC transport for 2-rank Ulysses all-to-all (NVLink same-node)
@@ -5,6 +5,7 @@ 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.pipeline_configs.ltx_2 import LTX2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImagePipelineConfig,
@@ -30,7 +31,11 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
from sglang.multimodal_gen.runtime.utils.precision import (
autocast_enabled,
resolve_component_precision,
resolve_decode_precision,
)
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
from sglang.srt.model_loader.checkpoint_quantization import (
resolve_checkpoint_quant_spec,
@@ -126,6 +131,41 @@ def _should_use_channels_last_3d(
return False
def _hold_decoder_weights_in_decode_dtype(
vae, server_args: ServerArgs, component_name: str
) -> None:
"""Round decoder weights to their decode compute dtype at load.
The decode stage persists these frozen weights in the autocast dtype on
first use (``prepare_autocast_linear_weights``), so the rounding itself is
already part of the output. Doing it at load makes residency plans, host
pins, and every host-to-device copy carry the halved size: MiniMax-H3's
video decoder drops from 9.7 to ~4.9 GiB, which is the difference between
restreaming a third of it per tile and holding all 36 blocks on a 12 GiB
card for the decode.
"""
if component_name not in ("vae", "video_vae"):
return
if envs.SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST:
return
prepare = getattr(vae, "prepare_decoder_autocast_weights", None)
if prepare is None:
return
dtype = resolve_decode_precision(server_args, component_name)
if dtype == torch.float32:
return
if not autocast_enabled(dtype, server_args.disable_autocast):
return
converted = prepare(dtype)
if converted:
logger.info(
"VAE: %s holds %d decoder weights in %s from load",
component_name,
converted,
dtype,
)
def _match_checkpoint_dtypes(loaded: dict, target_state: dict) -> dict:
"""Convert checkpoint tensors whose dtype differs from their parameter's.
@@ -235,6 +275,7 @@ class VAELoader(ComponentLoader):
logger.info(
"VAE: converted %d Conv3d weights to channels_last_3d", n
)
_hold_decoder_weights_in_decode_dtype(vae, server_args, component_name)
vae = current_platform.optimize_vae(vae)
return vae
@@ -308,5 +349,6 @@ class VAELoader(ComponentLoader):
if n > 0:
logger.info("VAE: converted %d Conv3d weights to channels_last_3d", n)
_hold_decoder_weights_in_decode_dtype(vae, server_args, component_name)
vae = current_platform.optimize_vae(vae)
return vae
@@ -0,0 +1,81 @@
"""The loader rounds VAE decoder weights to the decode dtype it will compute in.
The decode stage persists these frozen weights in the autocast dtype on first
use, so the rounding is already part of every output; what the tests pin down
is when the loader is allowed to do it early and when it must leave the
checkpoint dtype alone.
"""
from types import SimpleNamespace
import pytest
import torch
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
_hold_decoder_weights_in_decode_dtype,
)
class _RecordingVAE:
def __init__(self):
self.prepared_with = None
def prepare_decoder_autocast_weights(self, dtype) -> int:
self.prepared_with = dtype
return 144
def _server_args(decode_precision="fp16", disable_autocast=False):
return SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_decode_precision=decode_precision,
vae_precision="fp32",
),
disable_autocast=disable_autocast,
)
@pytest.fixture(autouse=True)
def _amp_supported(monkeypatch):
from sglang.multimodal_gen.runtime.utils import precision
monkeypatch.setattr(precision.current_platform, "is_amp_supported", lambda: True)
def test_the_decoder_is_rounded_to_the_decode_dtype_at_load():
vae = _RecordingVAE()
_hold_decoder_weights_in_decode_dtype(vae, _server_args(), "video_vae")
assert vae.prepared_with == torch.float16
def test_disabling_autocast_keeps_the_checkpoint_dtype():
vae = _RecordingVAE()
_hold_decoder_weights_in_decode_dtype(
vae, _server_args(disable_autocast=True), "video_vae"
)
assert vae.prepared_with is None
def test_the_kill_switch_keeps_the_checkpoint_dtype(monkeypatch):
monkeypatch.setenv("SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST", "1")
vae = _RecordingVAE()
_hold_decoder_weights_in_decode_dtype(vae, _server_args(), "video_vae")
assert vae.prepared_with is None
def test_an_fp32_decode_precision_is_left_alone():
vae = _RecordingVAE()
_hold_decoder_weights_in_decode_dtype(
vae, _server_args(decode_precision="fp32"), "video_vae"
)
assert vae.prepared_with is None
def test_the_audio_vae_is_not_touched():
vae = _RecordingVAE()
_hold_decoder_weights_in_decode_dtype(vae, _server_args(), "audio_vae")
assert vae.prepared_with is None
def test_a_vae_without_the_hook_is_skipped():
_hold_decoder_weights_in_decode_dtype(object(), _server_args(), "video_vae")