perf(diffusion): decode Wan VAE in BF16 (#32697)
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user