[diffusion] Default Hunyuan VAE to tiled decode (#36012)
This commit is contained in:
@@ -93,7 +93,9 @@ class HunyuanConfig(PipelineConfig):
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=HunyuanVideoConfig)
|
||||
# VAE
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
vae_config: VAEConfig = field(
|
||||
default_factory=lambda: HunyuanVAEConfig(parallel_decode_mode="tiled")
|
||||
)
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
|
||||
FastHunyuanConfig,
|
||||
HunyuanConfig,
|
||||
)
|
||||
|
||||
|
||||
def test_hunyuan_configs_default_to_parallel_tiled_vae_decode():
|
||||
for config_cls in (HunyuanConfig, FastHunyuanConfig):
|
||||
assert config_cls().vae_config.parallel_decode_mode == "tiled"
|
||||
|
||||
|
||||
def test_hunyuan_parallel_decode_mode_can_be_overridden():
|
||||
for config_cls in (HunyuanConfig, FastHunyuanConfig):
|
||||
config = config_cls()
|
||||
|
||||
config.update_config_from_dict(
|
||||
{"vae_config.parallel_decode_mode": "spatial_shard"}
|
||||
)
|
||||
|
||||
assert config.vae_config.parallel_decode_mode == "spatial_shard"
|
||||
Reference in New Issue
Block a user