perf(diffusion): decode Wan VAE in BF16 (#32697)

This commit is contained in:
Xiaoyu Zhang
2026-07-29 21:57:50 +08:00
committed by GitHub
parent 917e900d4d
commit 4f5b50c576
10 changed files with 99 additions and 12 deletions
@@ -224,6 +224,7 @@ class PipelineConfig:
# VAE configuration # VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig) vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp32" vae_precision: str = "fp32"
vae_decode_precision: str | None = None
vae_tiling: bool = True vae_tiling: bool = True
vae_slicing: bool = False vae_slicing: bool = False
vae_sp: bool = True vae_sp: bool = True
@@ -798,6 +799,17 @@ class PipelineConfig:
choices=["fp32", "fp16", "bf16"], choices=["fp32", "fp16", "bf16"],
help="Precision for VAE", help="Precision for VAE",
) )
parser.add_argument(
f"--{prefix_with_dot}vae-decode-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}vae_decode_precision",
default=PipelineConfig.vae_decode_precision,
choices=["fp32", "fp16", "bf16"],
help=(
"Optional decode-only VAE precision override. "
"Defaults to --vae-precision when unset."
),
)
parser.add_argument( parser.add_argument(
f"--{prefix_with_dot}vae-tiling", f"--{prefix_with_dot}vae-tiling",
action=StoreBoolean, action=StoreBoolean,
@@ -18,6 +18,7 @@ class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config):
is_causal: bool = True is_causal: bool = True
task_type: ModelTaskType = ModelTaskType.TI2V task_type: ModelTaskType = ModelTaskType.TI2V
vae_precision: str = "bf16" vae_precision: str = "bf16"
vae_decode_precision: str = "bf16"
flow_shift: float | None = 5.0 flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field( dmd_denoising_steps: list[int] | None = field(
@@ -89,6 +89,7 @@ class WanT2V480PConfig(PipelineConfig):
# Precision for each component # Precision for each component
precision: str = "bf16" precision: str = "bf16"
vae_precision: str = "fp32" vae_precision: str = "fp32"
vae_decode_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32",)) text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32",))
def __post_init__(self): def __post_init__(self):
@@ -239,6 +240,7 @@ class Wan2_2_TI2V_5B_Config(WanT2V480PConfig, WanI2VCommonConfig):
flow_shift: float | None = 5.0 flow_shift: float | None = 5.0
task_type: ModelTaskType = ModelTaskType.TI2V task_type: ModelTaskType = ModelTaskType.TI2V
expand_timesteps: bool = True expand_timesteps: bool = True
vae_decode_precision: str = "fp32"
# ti2v, 5B # ti2v, 5B
vae_stride = (4, 16, 16) vae_stride = (4, 16, 16)
@@ -269,6 +271,7 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig): class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
flow_shift: float | None = 12.0 flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875 boundary_ratio: float | None = 0.875
vae_decode_precision: str = "fp32"
def __post_init__(self) -> None: def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio self.dit_config.boundary_ratio = self.boundary_ratio
@@ -279,6 +282,7 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
class Wan2_2_I2V_A14B_Config(WanI2V720PConfig): class Wan2_2_I2V_A14B_Config(WanI2V720PConfig):
flow_shift: float | None = 5.0 flow_shift: float | None = 5.0
boundary_ratio: float | None = 0.900 boundary_ratio: float | None = 0.900
vae_decode_precision: str = "fp32"
def __post_init__(self) -> None: def __post_init__(self) -> None:
super().__post_init__() super().__post_init__()
@@ -165,6 +165,7 @@ async def get_models(request: Request):
"task_type": server_args.pipeline_config.task_type.name, "task_type": server_args.pipeline_config.task_type.name,
"dit_precision": server_args.pipeline_config.dit_precision, "dit_precision": server_args.pipeline_config.dit_precision,
"vae_precision": server_args.pipeline_config.vae_precision, "vae_precision": server_args.pipeline_config.vae_precision,
"vae_decode_precision": server_args.pipeline_config.vae_decode_precision,
} }
if model_info: if model_info:
@@ -34,7 +34,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.precision import ( from sglang.multimodal_gen.runtime.utils.precision import (
autocast_context, autocast_context,
autocast_enabled, autocast_enabled,
resolve_precision, resolve_decode_precision,
temporary_module_dtype, temporary_module_dtype,
) )
from sglang.multimodal_gen.runtime.utils.torch_compile import ( from sglang.multimodal_gen.runtime.utils.torch_compile import (
@@ -117,9 +117,7 @@ class DecodingStage(PipelineStage):
def component_uses( def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]: ) -> list[ComponentUse]:
vae_dtype = resolve_precision( vae_dtype = resolve_decode_precision(server_args, self.component_name)
server_args, self.component_name, precision_attr="vae_precision"
)
stage_name = self._component_stage_name(stage_name) stage_name = self._component_stage_name(stage_name)
return [ return [
ComponentUse( ComponentUse(
@@ -204,7 +202,9 @@ class DecodingStage(PipelineStage):
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents) latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
server_args: Configuration containing: server_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False) - disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16") - pipeline_config.vae_decode_precision: optional decode-only
VAE precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_precision: fallback VAE precision
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency - pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns: Returns:
@@ -212,10 +212,8 @@ class DecodingStage(PipelineStage):
normalized to [0, 1] range and moved to CPU as float32 normalized to [0, 1] range and moved to CPU as float32
""" """
latents = latents.to(get_local_torch_device()) latents = latents.to(get_local_torch_device())
# Setup VAE precision from user policy. # The caller resolves the decode-only override before component use so
vae_dtype = resolve_precision( # residency and execution agree on the target dtype.
server_args, self.component_name, precision_attr="vae_precision"
)
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast) vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
# scale and shift # scale and shift
@@ -281,9 +279,7 @@ class DecodingStage(PipelineStage):
# load vae if not already loaded (used for memory constrained devices) # load vae if not already loaded (used for memory constrained devices)
self.load_model() self.load_model()
vae_dtype = resolve_precision( vae_dtype = resolve_decode_precision(server_args, self.component_name)
server_args, self.component_name, precision_attr="vae_precision"
)
with self.use_declared_component( with self.use_declared_component(
component_name=self.component_name, component_name=self.component_name,
module=self.vae, module=self.vae,
@@ -29,6 +29,25 @@ def resolve_precision(
return precision_to_dtype(precision, field_name or precision_attr) return precision_to_dtype(precision, field_name or precision_attr)
def resolve_decode_precision(server_args, component_name: str = "vae") -> torch.dtype:
pipeline_config = server_args.pipeline_config
if component_name in ("audio_vae", "vocoder"):
return resolve_precision(
server_args,
component_name,
precision_attr="audio_vae_precision",
)
decode_precision = getattr(pipeline_config, "vae_decode_precision", None)
if decode_precision is not None:
return precision_to_dtype(decode_precision, "vae_decode_precision")
return resolve_precision(
server_args,
component_name,
precision_attr="vae_precision",
)
def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]: def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]:
pipeline_config = getattr(server_args, "pipeline_config", None) pipeline_config = getattr(server_args, "pipeline_config", None)
if pipeline_config is None: if pipeline_config is None:
@@ -2,6 +2,7 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import torch
import torch.nn as nn import torch.nn as nn
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
@@ -19,6 +20,19 @@ class FakeVAE(nn.Module):
class TestDecodingStageParallelism(unittest.TestCase): class TestDecodingStageParallelism(unittest.TestCase):
def test_component_use_honors_decode_precision_override(self):
stage = DecodingStage(FakeVAE())
server_args = SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_precision="fp32",
vae_decode_precision="bf16",
)
)
[component_use] = stage.component_uses(server_args)
self.assertEqual(component_use.target_dtype, torch.bfloat16)
def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks( def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks(
self, self,
): ):
@@ -17,6 +17,10 @@ class TestLongLive2AdjustNumFrames(unittest.TestCase):
def test_rounds_to_causal_block_aligned_latents(self): def test_rounds_to_causal_block_aligned_latents(self):
self.assertEqual(self.config.adjust_num_frames(65), 61) self.assertEqual(self.config.adjust_num_frames(65), 61)
def test_preserves_bf16_vae_decode_precision(self):
self.assertEqual(self.config.vae_precision, "bf16")
self.assertEqual(self.config.vae_decode_precision, "bf16")
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -70,6 +70,7 @@ autocast_enabled = precision.autocast_enabled
get_module_dtype = precision.get_module_dtype get_module_dtype = precision.get_module_dtype
precision_to_dtype = precision.precision_to_dtype precision_to_dtype = precision.precision_to_dtype
resolve_component_precision = precision.resolve_component_precision resolve_component_precision = precision.resolve_component_precision
resolve_decode_precision = precision.resolve_decode_precision
resolve_precision = precision.resolve_precision resolve_precision = precision.resolve_precision
temporary_module_dtype = precision.temporary_module_dtype temporary_module_dtype = precision.temporary_module_dtype
@@ -104,6 +105,7 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase):
def _server_args(self, **overrides): def _server_args(self, **overrides):
config = { config = {
"vae_precision": "fp16", "vae_precision": "fp16",
"vae_decode_precision": None,
"audio_vae_precision": "bf16", "audio_vae_precision": "bf16",
"dit_precision": "fp32", "dit_precision": "fp32",
"image_encoder_precision": "fp16", "image_encoder_precision": "fp16",
@@ -128,6 +130,18 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "Unsupported custom_precision"): with self.assertRaisesRegex(ValueError, "Unsupported custom_precision"):
precision_to_dtype("fp8", "custom_precision") precision_to_dtype("fp8", "custom_precision")
def test_decode_precision_override_and_fallback(self):
self.assertEqual(
resolve_decode_precision(self._server_args()),
torch.float16,
)
self.assertEqual(
resolve_decode_precision(self._server_args(vae_decode_precision="bf16")),
torch.bfloat16,
)
with self.assertRaisesRegex(ValueError, "Unsupported vae_decode_precision"):
resolve_decode_precision(self._server_args(vae_decode_precision="fp8"))
def test_component_precision_mapping(self): def test_component_precision_mapping(self):
server_args = self._server_args() server_args = self._server_args()
expected = { expected = {
@@ -688,6 +688,28 @@ class TestWarmupImageIsModelValid(unittest.TestCase):
class TestOffloadDefaults(unittest.TestCase): class TestOffloadDefaults(unittest.TestCase):
def test_wan_decode_precision_defaults(self):
for pipeline_config in (
WanT2V480PConfig(),
WanI2V480PConfig(),
):
with self.subTest(pipeline_config=pipeline_config.__class__.__name__):
self.assertEqual(pipeline_config.vae_precision, "fp32")
self.assertEqual(pipeline_config.vae_decode_precision, "bf16")
for pipeline_config in (
FastWan2_2_TI2V_5B_Config(),
Wan2_2_T2V_A14B_Config(),
Wan2_2_I2V_A14B_Config(),
):
with self.subTest(pipeline_config=pipeline_config.__class__.__name__):
self.assertEqual(pipeline_config.vae_precision, "fp32")
self.assertEqual(pipeline_config.vae_decode_precision, "fp32")
generic_config = PipelineConfig()
self.assertEqual(generic_config.vae_precision, "fp32")
self.assertIsNone(generic_config.vae_decode_precision)
def _from_dict_with_pipeline_config( def _from_dict_with_pipeline_config(
self, self,
pipeline_config, pipeline_config,