[diffusion] refactor: gate fast vae paths by quality (#33849)

This commit is contained in:
Mick
2026-08-07 12:39:55 +08:00
committed by GitHub
parent 914644e81c
commit c2657cc4bf
10 changed files with 105 additions and 37 deletions
@@ -76,6 +76,11 @@ def _get_request_field_or_extra(request, field_name):
return _get_extra_field(request, field_name)
def _runtime_sampling_quality(quality: str | None) -> str | None:
"""Keep OpenAI's automatic default out of SGLang's sampling contract."""
return None if quality in (None, "auto") else quality
def _parse_extra_container(value: Any) -> dict[str, Any]:
if isinstance(value, str):
try:
@@ -309,6 +314,7 @@ async def generations(
use_system_prompt=_get_extra_field(request, "use_system_prompt"),
use_guardrails=_get_extra_field(request, "use_guardrails"),
enable_teacache=request.enable_teacache,
quality=_runtime_sampling_quality(request.quality),
output_compression=request.output_compression,
output_quality=request.output_quality,
diffusers_kwargs=request.diffusers_kwargs,
@@ -418,6 +424,7 @@ async def edits(
guidance_scale: Optional[float] = Form(None),
true_cfg_scale: Optional[float] = Form(None),
num_inference_steps: Optional[int] = Form(None),
quality: Optional[str] = Form(None),
output_quality: Optional[str] = Form("default"),
output_compression: Optional[int] = Form(None),
enable_teacache: Optional[bool] = Form(False),
@@ -478,6 +485,7 @@ async def edits(
num_inference_steps=num_inference_steps,
enable_teacache=enable_teacache,
num_frames=num_frames,
quality=_runtime_sampling_quality(quality),
output_compression=output_compression,
output_quality=output_quality,
enable_upscaling=enable_upscaling,
@@ -104,6 +104,7 @@ _MULTIPART_EXTRA_FORM_FIELDS = (
"action_normalization",
"condition_frame_indexes_vision",
"condition_video_keep",
"quality",
)
@@ -371,6 +372,7 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque
"upscaling_model_path": request.upscaling_model_path,
"upscaling_scale": request.upscaling_scale,
"output_path": request.output_path,
"quality": _extra_value(request, "quality"),
"output_compression": request.output_compression,
"output_quality": request.output_quality,
"perf_dump_path": request.perf_dump_path,
@@ -0,0 +1,41 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared decode-scoped gate for optional VAE fast paths."""
from contextlib import contextmanager
from weakref import WeakKeyDictionary
import torch.nn as nn
class VaeFastPathGate:
"""Mutable flag shared by the wrappers installed on one VAE."""
__slots__ = ("enabled",)
def __init__(self) -> None:
self.enabled = False
_VAE_FAST_PATH_GATES: WeakKeyDictionary[nn.Module, VaeFastPathGate] = (
WeakKeyDictionary()
)
def register_vae_fast_path_gate(vae: nn.Module, gate: VaeFastPathGate) -> None:
_VAE_FAST_PATH_GATES[vae] = gate
@contextmanager
def use_vae_fast_path(vae: nn.Module, enabled: bool):
"""Enable an installed VAE fast path for one decode and always reset it."""
gate = _VAE_FAST_PATH_GATES.get(vae)
if gate is None:
yield
return
previous_enabled = gate.enabled
gate.enabled = enabled
try:
yield
finally:
gate.enabled = previous_enabled
@@ -8,9 +8,8 @@ decoder module family (``ResnetBlock2D`` GroupNorm+SiLU chains,
All rewrites are mathematically exact re-associations of the original
operators. Wrappers are installed once at VAE load and dispatch on a
request-scoped :class:`VaeFastPathGate` (published as
``_sgl_vae_fast_path_gate``): ``quality == "high"`` runs the fast paths, the
``"lossless"`` default runs the original module path bit-for-bit.
decode-scoped :class:`VaeFastPathGate`: ``quality == "high"`` runs the fast
paths, the ``"lossless"`` default runs the original module path bit-for-bit.
- channels_last: run the decoder in NHWC so cuDNN convs skip the transpose
kernels; parameter layout is swapped at decode entry to match the gate.
@@ -28,6 +27,10 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import (
VaeFastPathGate,
register_vae_fast_path_gate,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -43,19 +46,6 @@ except ImportError: # pragma: no cover
_HAS_TRITON = False
class VaeFastPathGate:
"""Mutable fast-path flag shared by every wrapper of one VAE; enabled by
``DecodingStage`` while decoding a ``quality == "high"`` request."""
__slots__ = ("enabled",)
def __init__(self) -> None:
self.enabled = False
GATE_ATTR = "_sgl_vae_fast_path_gate"
# ---------------------------------------------------------------------------
# Fuse A: two-pass GroupNorm(+SiLU) fusion (channels_last Triton kernel)
# ---------------------------------------------------------------------------
@@ -330,7 +320,7 @@ def _decoder_layout_forward(self, *args, **kwargs):
memory_format=(torch.channels_last if want_cl else torch.contiguous_format)
)
self._sgl_channels_last = want_cl
logger.info(
logger.debug(
"%s: decoder switched to %s layout.",
self._sgl_label,
"channels_last (NHWC)" if want_cl else "contiguous (NCHW)",
@@ -396,7 +386,7 @@ def _install_decoder_fast_paths(vae: nn.Module, label: str) -> nn.Module:
m._sgl_folded_v = None
m.forward = MethodType(_attn_fast_forward, m)
n_norm = _install_norm_silu(decoder, ResnetBlock2D, gate)
setattr(vae, GATE_ATTR, gate)
register_vae_fast_path_gate(vae, gate)
logger.info(
"%s: installed quality-gated decoder fast paths "
"(channels_last dispatch, %d fused upsamplers, %d fast attention "
@@ -3,20 +3,19 @@
Fuses every decoder ``WanRMS_norm -> SiLU`` chain into one Triton kernel on
the channels_last_3d layout. Wrappers are installed once at VAE load and
dispatch on a request-scoped :class:`VaeFastPathGate` (published as
``_sgl_vae_fast_path_gate``): ``quality == "high"`` runs the fused kernel
(not bitwise-identical to aten, hence gated), the ``"lossless"`` default
runs the original module path bit-for-bit. Install is all-or-nothing and
fail-closed.
dispatch on a decode-scoped :class:`VaeFastPathGate`: ``quality == "high"``
runs the fused kernel (not bitwise-identical to aten, hence gated), the
``"lossless"`` default runs the original module path bit-for-bit. Install is
all-or-nothing and fail-closed.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.multimodal_gen.runtime.models.vaes.flux2_vae_cuda_opt import (
GATE_ATTR,
from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import (
VaeFastPathGate,
register_vae_fast_path_gate,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -133,7 +132,7 @@ def maybe_optimize_wan_vae(vae: nn.Module) -> nn.Module:
n_norm = _install_norm_silu(decoder, gate)
if n_norm is None:
return vae
setattr(vae, GATE_ATTR, gate)
register_vae_fast_path_gate(vae, gate)
logger.info(
"Wan VAE: installed quality-gated fast path (%d RMSNorm+SiLU fusions).",
n_norm,
@@ -20,6 +20,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im
ComponentUse,
)
from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE
from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import (
use_vae_fast_path,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
PipelineStage,
@@ -312,14 +315,7 @@ class DecodingStage(PipelineStage):
assert vae is not None
self.vae = vae
# Request-scoped VAE fast-path gate (see flux2_vae_cuda_opt):
# quality == "high" opts this decode into the near-lossless fast
# paths; the "lossless" default keeps the bit-exact original
# module path. VAEs without installed wrappers have no gate.
gate = getattr(vae, "_sgl_vae_fast_path_gate", None)
if gate is not None:
gate.enabled = getattr(batch.sampling_params, "quality", None) == "high"
try:
with use_vae_fast_path(vae, batch.sampling_params.quality == "high"):
frames = self.decode(batch.latents, server_args, vae_dtype=vae_dtype)
# decode trajectory latents if needed
@@ -348,9 +344,6 @@ class DecodingStage(PipelineStage):
trajectory_decoded = [decoded_tensor[:, i] for i in range(T)]
else:
trajectory_decoded = None
finally:
if gate is not None:
gate.enabled = False
frames = server_args.pipeline_config.post_decoding(frames, server_args)
@@ -10,6 +10,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import (
_fallback_image_urls,
_get_response_resize,
_raise_if_image_variant_not_found,
_runtime_sampling_quality,
_select_image_variant_cloud_url,
_select_image_variant_path,
)
@@ -40,6 +41,13 @@ def test_url_response_returns_one_item_per_output_path():
]
def test_runtime_sampling_quality_preserves_the_openai_default():
assert _runtime_sampling_quality(None) is None
assert _runtime_sampling_quality("auto") is None
assert _runtime_sampling_quality("lossless") == "lossless"
assert _runtime_sampling_quality("high") == "high"
def test_image_response_includes_resize_for_every_output():
response = _build_image_response_kwargs(
["first.png", "second.png"],
@@ -0,0 +1,21 @@
import torch.nn as nn
from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import (
VaeFastPathGate,
register_vae_fast_path_gate,
use_vae_fast_path,
)
def test_vae_fast_path_gate_is_decode_scoped_and_nestable():
vae = nn.Module()
gate = VaeFastPathGate()
register_vae_fast_path_gate(vae, gate)
with use_vae_fast_path(vae, True):
assert gate.enabled
with use_vae_fast_path(vae, False):
assert not gate.enabled
assert gate.enabled
assert not gate.enabled
@@ -22,6 +22,7 @@ def test_video_api_forwards_profiling_options():
profile=True,
num_profiled_timesteps=3,
profile_all_stages=False,
quality="high",
)
server_args = SimpleNamespace(
backend="auto",
@@ -48,3 +49,4 @@ def test_video_api_forwards_profiling_options():
assert kwargs["profile"] is True
assert kwargs["num_profiled_timesteps"] == 3
assert kwargs["profile_all_stages"] is False
assert kwargs["quality"] == "high"