vLLM-Omni VAE 并行详解:Patch/Tile 并行与 Wan 空间分片解码的实现与配置指南
【免费下载链接】vllm-omniA framework for efficient model inference with omni-modality models项目地址: https://gitcode.com/GitHub_Trending/vl/vllm-omni
在 vLLM-Omni 中,VAE(Variational AutoEncoder)解码往往是高分辨率图像生成和长视频生成的显存瓶颈。本篇指南围绕docs/user_guide/diffusion/parallelism/vae_parallelism.md展开,讲清 VAE 并行的三种解码策略(tile 分块、patch 分块、Wan 空间分片)各自的工作原理与选型依据,并给出从 Python API 到vllm serve服务端的完整配置方法、参数约束与源码级排查路径。读完后你可以:正确启用 VAE patch 并行以降低 VAE decode 峰值显存、理解它与 DiT 并行组(DiT process group)的共享关系、以及在遇到"配置被静默忽略"类问题时能快速定位根因。
一、总览:VAE 并行在 vLLM-Omni 中的定位
VAE parallelism 将 VAE 的 decode/encode 工作分摊到多张 GPU 上。当前仓库实现了两条路线:
- VAE patch/tile parallelism:把 latent 空间切成空间上的 tile 或 patch,各 rank 解码一部分,再由 rank 0 拼接成完整结果;
- Wan spatial-shard decode:针对 Wan VAE,沿高度或宽度方向分片 decoder 特征图,并在空间卷积处交换 halo 行/列。
适用场景(来自原指南):
- 高分辨率图像生成:VAE decode 成为显存瓶颈时;
- 显存受限环境:VAE decode 激活峰值超出可用 VRAM;
- 多 GPU 环境:希望把 VAE 阶段也放到分布式资源上利用起来。
各模型的支持情况见 Supported Models 中的VAE-Patch-Parallel列。
两种策略对照表(原文完整继承)
VAE patch parallelism 依据图像尺寸自动选择两种策略:
| 策略 | 适用场景 | 工作方式 | 重叠区域处理 | 输出质量 |
|---|---|---|---|---|
| Tiled Decode | 大图像(触发 VAE tiling) | 把既有的 VAE tiling 计算分摊到各 rank,每个 rank 解码一组重叠 tile | 复用 VAE 原生的blend_v与blend_h函数无缝合并重叠区域 | Bit-identical(与单卡 tiling 逻辑相同) |
| Patch Decode | 小图像(不触发 VAE tiling) | 把 latent 切成带 halo 的空间 patch,每个 rank 解码一个 patch 及其边界上下文 | halo 区域提供边缘上下文;核心区域直接拼接,不做 blending | 近似一致(diff < 0.5%,视觉不可感知) |
从源码看,策略选择逻辑位于 VaePatchParallelism.decode:它先判断 latent 是否满足 diffusers 的 tiling 触发条件(z.shape[-1] > tile_latent_min_size or z.shape[-2] > tile_latent_min_size),满足则走_distributed_tiled_decode,否则走_distributed_patch_decode,与上表描述一致。
与 DiT 并行组的关系
VAE patch parallelism复用 DiT 的 process group(dit_group),不会单独初始化新的 ProcessGroup。这意味着:
- 共享 ranks:VAE patch 并行使用与 DiT 并行(Tensor Parallel、Sequence Parallel 等)相同的 GPU ranks;
- 组合使用:VAE patch 并行通常与其他并行方式一起使用;
- 配置对齐:
vae_patch_parallel_size不应大于 DiT process group 的大小。
实现上,DistributedVaeExecutor 在初始化时直接取get_world_group().device_group(worker 全 WORLD 范围)作为通信组,world_size/rank均取自该组;实际参与的 rank 数为min(vae_patch_parallel_size, world_size)。
二、快速上手(Quick Start)
基本用法
最简可运行示例:
from vllm_omni import Omni from vllm_omni.inputs.data import OmniDiffusionSamplingParams from vllm_omni.diffusion.data import DiffusionParallelConfig # TP=2 for DiT, VAE patch parallel also uses these 2 GPUs omni = Omni( model="Tongyi-MAI/Z-Image-Turbo", parallel_config=DiffusionParallelConfig( tensor_parallel_size=2, # Enable tensor parallelism for DiT vae_patch_parallel_size=2, # Enable VAE patch parallelism ), vae_use_tiling=True, # Required for VAE patch parallelism ) outputs = omni.generate( "a futuristic city at sunset, high resolution, 8k", OmniDiffusionSamplingParams( num_inference_steps=9, height=1152, # High resolution benefits from VAE patch parallel width=1152, ), )要点:
DiffusionParallelConfig定义于 vllm_omni/diffusion/data.py,其中vae_patch_parallel_size默认 1(即关闭),vae_parallel_mode默认"tile";vae_use_tiling在OmniDiffusionConfig层默认False(见 data.py)。使用 VAE patch 并行时需要开启它,但即使忘了开,注册器也会在启动时自动补上(见下文"配置校验"一节)。
三、示例脚本
离线推理
使用examples/offline_inference/text_to_image/下的 text_to_image.py:
# Text-to-Image with Z-Image python examples/offline_inference/text_to_image/text_to_image.py \ --model Tongyi-MAI/Z-Image-Turbo \ --prompt "a futuristic city at sunset" \ --height 1152 \ --width 1152 \ --tensor-parallel-size 2 \ --vae-patch-parallel-size 2 \ --vae-use-tiling在线服务(Online Serving)
在线服务通过--vae-patch-parallel-size启用(该 CLI 参数在 vllm_omni/entrypoints/cli/serve.py 中定义):
# Text-to-Image with Z-Image, TP=2 + VAE patch parallel=2 vllm serve Tongyi-MAI/Z-Image-Turbo --omni --port 8091 \ --tensor-parallel-size 2 \ --vae-patch-parallel-size 2 \ --vae-use-tiling四、配置参数详解
DiffusionParallelConfig 中的参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
vae_patch_parallel_size | int | 1 | VAE patch/tile 并行使用的 GPU 数。设为 2 或更大即启用。应与tensor_parallel_size一致,因为二者共享同一个 process group。 |
vae_parallel_mode | str | "tile" | VAE 并行解码策略:"tile"(默认 tile/patch 并行解码)、"spatial_shard_height"、"spatial_shard_width"(空间分片解码,仅 Wan 支持)。见下文 Wan 空间分片解码。 |
源码中该字段及约束(vllm_omni/diffusion/data.py):
vae_patch_parallel_size:"Number of ranks used for VAE patch/tile parallelism (decode/encode)";vae_parallel_mode:"spatial_shard_*" 模式是 decode-only,且要求vae_patch_parallel_size与 DiT group size 匹配,否则运行时回退到 tile 并行解码;- 配置校验器
_validate_parallel_config会断言vae_patch_parallel_size > 0且vae_parallel_mode ∈ {"tile", "spatial_shard_height", "spatial_shard_width"},不满足会在构造DiffusionParallelConfig时直接报错(data.py)。
附加要求
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
vae_use_tiling | bool | False | 使用 VAE patch 并行时必须设为True。 |
!!! note "自动开启 VAE Tiling" 当vae_patch_parallel_size > 1且模型具备分布式 VAE(DistributedVaeMixin)时,若vae_use_tiling尚未开启,系统会自动将其置为True。
这个自动行为在注册器中有明确实现,vllm_omni/diffusion/registry.py 中:
vae_pp_size = od_config.parallel_config.vae_patch_parallel_size is_distributed_vae = hasattr(model, "vae") and isinstance(model.vae, DistributedVaeMixin) if vae_pp_size > 1 and not is_distributed_vae: logger.warning( "vae_patch_parallel_size=%d is set but VAE patch parallelism is NOT enabled for %s; ignoring.", vae_pp_size, od_config.model_class_name, ) if vae_pp_size > 1 and is_distributed_vae and not od_config.vae_use_tiling: logger.info( "vae_patch_parallel_size=%d requires vae_use_tiling; automatically enabling it.", vae_pp_size, ) od_config.vae_use_tiling = True ... if is_distributed_vae: model.vae.set_parallel_size(vae_pp_size, mode=od_config.parallel_config.vae_parallel_mode)即:注册器在加载 pipeline 后统一判断 VAE 是否为DistributedVaeMixin实例,不支持则打警告并忽略配置;支持则自动开启 tiling,并把parallel_size与vae_parallel_mode一并传给 VAE(DistributedVaeMixin.set_parallel_size)。
五、源码级原理:tile 并行与 patch 并行如何工作
分布式执行器(DistributedVaeExecutor)
通用执行框架在 distributed_vae_executor.py,其execute主流程为:
- 切分:
operator.split(z)把输入 latent 切成一组TileTask(带tile_id、网格坐标与 workload); - 负载均衡:
_balance_tasks按 workload 降序做贪心分配(总是把下一个最大任务分给当前累计负载最小的 rank),避免各 rank 忙闲不均; - 本地解码:每个 rank 只解码分配给自己的 tile(
operator.exec); - 形状协商:
_compute_global_padding_shape通过all_reduce(MAX)求出所有 tile 的最大尺寸,保证 gather 时张量形状统一; - 打包与收集:每个 rank 把本地 tile 与元信息(tile_id、H/W)打包成定长张量,
all_gather到所有 rank; - 拼接:rank 0 调用
operator.merge重建完整张量,非 rank 0 返回空张量; - 结果同步:
_sync_final_result先广播形状、再广播内容,使所有 rank 拿到一致的完整输出。
Tiled Decode:复用 diffusers 原生 tiling 逻辑
_distributed_tiled_decode 是 diffusersAutoencoderKL.tiled_decode的分布式版本(仅 decode 路径):
- 按照
overlap_size = tile_latent_min_size * (1 - tile_overlap_factor)的步长在 latent 上生成 tile 网格; - tile 到 rank 的分配使用
tile_rank = (tile_id + 1) % pp_size,偏移 1 位是为了让 rank 0 避开最大的边界 tile(tile_id=0),让最重的一块落到其他 rank 上; - 各 rank 解码后
gather到 rank 0,rank 0 用 VAE 原生的blend_v/blend_h在blend_extent重叠区做线性混合,再裁剪拼接——这正是"bit-identical"质量承诺的来源,因为它执行的是与单卡 tiling 完全相同的混合与拼接代码。
Patch Decode:小图像也能受益
_distributed_patch_decode 针对"单卡本来不会触发 tiling"的中小尺寸输入:
- 网格切分:
_factor_pp_grid为pp_size选一个接近正方形的 (rows, cols) 因子分解,每个 rank 负责一个 patch; - halo 计算:
halo = max(halo_base, min(core_h, core_w) // 2),其中halo_base来自 tile overlap 参数。每个 rank 解码"核心区 + halo 边界上下文",然后只裁出核心区(ch0:ch1, cw0:cw1),halo 的贡献被丢弃,直接拼接核心块——因此质量是"近似一致"而非逐 bit 相同; - 拼接同样发生在 rank 0:先 gather 各 rank 的核心 RGB 块(不足部分补零),再按原 latent 网格坐标填回输出张量。
回退与容错
VaePatchParallelism.decode中有多重保护(vae_patch_parallel.py):
- latent 非 4D、
vae_patch_parallel_size <= 1、分布式未初始化、VAE 未开启use_tiling、取不到 process group 等情况,一律回退到原始vae.decode; - 若 rank 0 的并行解码产出为空,打印
VAE patch parallel decode produced empty output on rank0; falling back to vae.decode.并回退; - 最终通过"广播形状 + 广播张量"让所有 rank 持有同一份完整输出,pipeline 下游无需感知并行细节。
另外还有一条更轻量的挂接路径 maybe_wrap_vae_decode_with_patch_parallelism:它以实例级覆写的方式包装vae.decode,通过能力检查(有decode/decoder属性)而非严格的 diffusers 类型检查来支持自定义 VAE,并用_vllm_vae_patch_parallel_installed标志防止重复安装。
哪些 VAE 实现了 DistributedVaeMixin?
仓库中实现 DistributedVaeMixin 的 VAE 分布在vllm_omni/diffusion/distributed/autoencoders/与部分模型目录下,例如:
- autoencoder_kl_wan.py(Wan,额外支持空间分片)
- autoencoder_kl_qwenimage.py(Qwen-Image)
- autoencoder_kl_hunyuan_video_15.py、autoencoder_kl_hunyuan.py
- autoencoder_kl_ltx2.py 及 ltx2/vae/distributed.py
- 模型内定制:bagel/autoencoder.py、magi2/turbo_vae.py、minimax_h3/vae.py 等
is_distributed_enabled()的判定条件是:parallel_size > 1、分布式已初始化、use_tiling为 True,且min(parallel_size, world_size) > 1;若parallel_size超过 WORLD size,会打印vae_patch_parallel_size=... is greater than WORLD=...; using WORLD size=...警告并按 WORLD size 截断——这正是下一节"配置超过 DiT group 大小"问题的运行时表现。
六、Wan 空间分片解码(Spatially-Sharded Decode)
默认的vae_parallel_mode="tile"把整块 tile 分给各 rank。针对WanVAE 还有备选策略——空间分片解码,通过vae_parallel_mode="spatial_shard_height"或"spatial_shard_width"选择。
它不向各 rank 分派独立 tile,而是把 decoder 特征图沿高度(spatial_shard_height)或宽度(spatial_shard_width)方向切分,并在空间卷积(spatial convolutions)前后于相邻 rank 间交换 halo 行/列。这样跨分片边界处的感受野保持正确,结果与单卡解码在数值误差范围内一致。
Python API
from vllm_omni import Omni from vllm_omni.diffusion.data import DiffusionParallelConfig omni = Omni( model="Wan-AI/Wan2.1-T2V-1.3B-Diffusers", parallel_config=DiffusionParallelConfig( tensor_parallel_size=2, vae_patch_parallel_size=2, # must match the DiT group size vae_parallel_mode="spatial_shard_width", # or "spatial_shard_height" ), )CLI / 服务端
vllm serve Wan-AI/Wan2.1-T2V-1.3B-Diffusers --omni \ --tensor-parallel-size 2 \ --vae-patch-parallel-size 2 \ --vae-parallel-mode spatial_shard_width约束与行为
- 空间分片解码是decode-only,目前仅对WanVAE 实现,其他模型会忽略
spatial_shard_*模式; - 要求
vae_patch_parallel_size与 DiT process group 大小匹配,不匹配时 VAE 记录警告并在运行时回退到 tile 并行解码; - 对同一个 VAE 实例,
spatial_shard_height与spatial_shard_width互斥(decoder 就地为单一 split 维度打补丁)。
源码印证:autoencoder_kl_wan.py 中_spatial_shard_decode_enabled会检查is_distributed_enabled()与 split 维度配置;命中后调用 wan_spatial_shard.spatial_shard_decode。后者内部通过install_wan_spatial_shard_decode(wan_spatial_shard.py)对 decoder 的就地打补丁——补丁时校验已安装的 split 维度,若再次以不同维度安装会复用/保持单一维度,与"互斥"的约束一致。
端到端基准
评估端到端时延/吞吐时,以期望的vae_parallel_mode启动服务后,可直接复用仓库现成的 diffusion serving benchmark:
python3 benchmarks/diffusion/diffusion_benchmark_serving.py \ --endpoint /v1/videos --dataset random --task t2v --num-prompts 1 \ --height 480 --width 832 --num-frames 17 --max-concurrency 1七、最佳实践
何时使用
适合:
- 高分辨率图像生成与长视频生成;
- VAE decode 导致 OOM 的显存受限部署;
- 多 GPU 环境。
不适合:
- VAE decode 本就不是瓶颈的低分辨率图像/视频;
- 单 GPU 环境——单卡应使用 vae tiling decode,而不是并行 vae tiling decode;
- 不支持 VAE patch parallel 的模型。
八、常见问题排查(Troubleshooting)
问题 1:模型不支持 VAE Patch Parallel
现象:
WARNING: vae_patch_parallel_size=2 is set but VAE patch parallelism is NOT enabled for xxxPipeline; ignoring.根因:VAE Patch Parallelism 要求模型的 VAE 实现DistributedVaeMixin。启动时 vllm_omni/diffusion/registry.py 检查实例化后的 pipeline 是否具有.vae属性且为DistributedVaeMixin实例;若不是,该配置被静默忽略(仅打警告):
vae_pp_size = od_config.parallel_config.vae_patch_parallel_size is_distributed_vae = hasattr(model, "vae") and isinstance(model.vae, DistributedVaeMixin) if vae_pp_size > 1 and not is_distributed_vae: logger.warning( "vae_patch_parallel_size=%d is set but VAE patch parallelism is NOT enabled for %s; ignoring.", vae_pp_size, od_config.model_class_name, )解决方案:
- 改用受支持的模型(推荐):查看 Supported Models 的 VAE-Patch-Parallel 列;
- 为新模型添加支持:在其 VAE 类上实现
DistributedVaeMixin(欢迎贡献)。
问题 2:vae_patch_parallel_size超过 DiT Process Group 大小
现象:出现警告信息,且 VAE patch parallel size 被调整为 DiT process group size。
根因:VAE Patch Parallelism 复用 DiT process group(pp_size = min(vae_patch_parallel_size, world_size),见 DistributedVaeMixin.is_distributed_enabled 与 VaePatchParallelism.decode 中的截断逻辑)。
建议:始终把vae_patch_parallel_size设为不大于 DiT process group size 的值。
注意 DiT process group size 等于:
dit_parallel_size = data_parallel_size × cfg_parallel_size × sequence_parallel_size × pipeline_parallel_size × tensor_parallel_size其中 sequence_parallel_size = ulysses_degree × ring_degree。
九、总结
- 启用 VAE Patch Parallelism:在
DiffusionParallelConfig中设置vae_patch_parallel_size,vae_use_tiling=True,以降低 VAE decode 峰值显存; - 利用长序列收益:VAE patch 并行的收益在长序列解码(高分辨率、长视频)场景下最明显;
- 组合其他并行方式:建议与 Tensor Parallel 或 CFG-Parallel 一起使用,以获得最大显存节省。
配置入口一览:Python API 使用DiffusionParallelConfig(vllm_omni/diffusion/data.py),CLI/serve 使用--vae-patch-parallel-size与--vae-parallel-mode(vllm_omni/entrypoints/cli/serve.py),离线脚本参考 text_to_image.py,模型支持矩阵参考 docs/user_guide/diffusion_features.md。
【免费下载链接】vllm-omniA framework for efficient model inference with omni-modality models项目地址: https://gitcode.com/GitHub_Trending/vl/vllm-omni
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考