vllm-omni 扩散模型张量并行(Tensor Parallel)接入指南:以 Z-Image 为参考的完整实现
【免费下载链接】vllm-omniA framework for efficient model inference with omni-modality models项目地址: https://gitcode.com/GitHub_Trending/vl/vllm-omni
本指南面向希望在 vllm-omni 中为扩散 Transformer(DiT)模型添加张量并行(Tensor Parallel,TP)支持的研究与工程人员,完整梳理从识别可切分线性层、替换为 vLLM 并行层、校验维度约束,到端到端测试与问题排查的全过程。读完本文,你将掌握以 Z-Image 为参考的标准化 TP 改造范式,并能复用到 FLUX、Qwen-Image 等其他扩散模型上。
一、背景:为什么扩散 Transformer 需要 Tensor Parallel
Tensor Parallel(TP)是一种模型并行技术,其核心思想是将模型权重按张量维度切分到多张 GPU 上:每张 GPU 只持有模型参数的一部分,也只计算每一层输出的一个分片。对于包含大规模 Attention 与 MLP 层的扩散 Transformer 而言,沿模型维度(dim)切分可以带来两个直接收益:
- 显存可承载:更大模型得以在单卡放不下的情况下分布式运行;
- 接近线性的加速:切分后的每张卡计算量约降为原来的 1/N,配合通信即可逼近线性加速比。
vllm-omni 的张量并行实现建立在 vLLM 的 Parallel Layers 之上,这些并行层位于vllm.model_executor.layers.linear,其类型与职责如下:
| 层类型 | 用途 | 权重切分方式 |
|---|---|---|
ColumnParallelLinear | FFN 第一层、独立的 QKV 投影 | 按列切分(输出维度) |
RowParallelLinear | FFN 第二层、Attention 输出投影 | 按行切分(输入维度) |
QKVParallelLinear | 多头 / 分组查询注意力(GQA)的 QKV 投影 | 自动处理 head 复制,按 head 切分 |
ReplicatedLinear | 不应切分的层(如时间步嵌入、最终输出投影) | 不切分(全卡复制) |
说明:本文所依据的设计文档为 docs/design/feature/tensor_parallel.md,其中 Z-Image 作为参考实现;本节关于 vLLM Parallel Layers 的对应关系来自 vLLM 上游的模型接入文档,本文不再赘述外部链接。
二、改造前的第一步:识别需要切分的 Linear 层
在动手改代码之前,需要先盘点模型中所有nn.Linear,并回答两个关键问题:
- 哪些层应该按列并行(权重按列切分)?—— 典型答案是 FFN 的第一个投影、Attention 的 QKV 投影;
- 哪些层应该按行并行(权重按行切分)?—— 典型答案是 FFN 的第二个投影、Attention 的输出投影。
一个实用的判断基准:一个"列并行 → 激活函数/注意力 → 行并行"的配对构成一个完整的切分闭环。列并行层输出的分片张量(shape 为[batch, seq, hidden/N])经过无需跨卡通信的激活或注意力计算后,由行并行层通过 All-Reduce 归约为完整维度,从而保证下一层拿到的是与单卡完全一致的输入。
此外,模型中还有一些不能切分的层,例如:
- 时间步嵌入 MLP(输出参与 adaLN 调制,对精度敏感);
- patch embedding、caption embedding 等入口层;
- 最终输出投影层(直接产生输出 latent)。
这些层在 Z-Image 参考实现中统一使用ReplicatedLinear,且很多刻意保持全精度(quant_config=None),详见源码注释(见 z_image_transformer.py 中 TimestepEmbedder 的说明)。
三、改造核心:用并行层替换nn.Linear
3.1 MLP 块(Up-Down 模式)
扩散模型 FFN 的标准形态是"升维再降维"。参考实现(见 FeedForward)中,Z-Image 使用MergedColumnParallelLinear(将 w1/w3 两个门控投影合并)配合SiluAndMul激活:
class FeedForward(nn.Module): def __init__(self, dim: int, hidden_dim: int): super().__init__() # 列并行:权重按列切分,输出为 [hidden_dim/N] self.w13 = MergedColumnParallelLinear( dim, [hidden_dim] * 2, # w1 与 w3 合并,各输出 hidden_dim bias=False, return_bias=False, ) self.act = SiluAndMul() # 作用于切分后的张量,无需通信 # 行并行:输入来自列并行层的输出(已切分) self.w2 = RowParallelLinear( hidden_dim, dim, bias=False, input_is_parallel=True, # 关键:输入已被 w13 切分 return_bias=False, ) def forward(self, x): # x: [batch, seq, dim](每张卡上完整复制) # w13 输出切分为 [batch, seq, hidden_dim/N] hidden_states = self.w13(x) hidden_states = self.act(hidden_states) # w2 通过 All-Reduce 恢复完整 dim hidden_states = self.w2(hidden_states) return hidden_states对照设计文档中的示意代码(ColumnParallelLinear→RowParallelLinear),实际源码采用了MergedColumnParallelLinear合并 w1/w3 的写法,两者在 TP 语义上完全一致,合并写法还减少了 Kernel 调用次数。注意 Z-Image 的 FFN 隐藏维度取int(dim / 3 * 8),在dim=3840时为10240。
3.2 Attention 块(QKV-Out 模式)
Attention 的 QKV 投影使用QKVParallelLinear,它会把总 head 数平均切分到各卡,每卡持有num_heads/N个头,并自动处理 GQA 下的 KV head 复制。参考实现见 ZImageAttention:
from vllm_omni.diffusion.attention.layer import Attention class ZImageAttention(nn.Module): def __init__(self, dim: int, num_heads: int, num_kv_heads: int): super().__init__() self.head_dim = dim // num_heads # 列并行:QKV 权重按 head 切分,每卡获得 num_heads/N 个头 self.to_qkv = QKVParallelLinear( hidden_size=dim, head_size=self.head_dim, total_num_heads=num_heads, # 模型总 head 数 total_num_kv_heads=num_kv_heads, # 模型总 KV head 数 bias=False, return_bias=False, ) # 行并行:Attention 输出按行切分,All-Reduce 恢复完整维度 self.to_out = RowParallelLinear( dim, dim, bias=False, input_is_parallel=True, # 输入来自注意力输出(按 head 切分) return_bias=False, ) # 关键:Attention 使用“本地 head 数”,而非模型总 head 数 self.attn = Attention( num_heads=self.to_qkv.num_heads, # 本卡实际持有的 head 数 head_size=self.head_dim, softmax_scale=1.0 / (self.head_dim**0.5), causal=False, num_kv_heads=self.to_qkv.num_kv_heads, # 本卡持有的 KV head 数 ) def forward(self, x): qkv, _ = self.to_qkv(x) # [batch, seq, (q+k+v)*head_dim/N] q_size = self.to_qkv.num_heads * self.head_dim # 本地 q 尺寸 kv_size = self.to_qkv.num_kv_heads * self.head_dim # 本地 kv 尺寸 q, k, v = qkv.split([q_size, kv_size, kv_size], dim=-1) # 每张卡独立计算注意力,无需跨卡通信 out = self.attn(q, k, v) out = self.to_out(out) # All-Reduce 回完整 dim return out三个必须遵守的要点:
ColumnParallelLinear/QKVParallelLinear与RowParallelLinear成对出现,这是 TP 的标准配对;RowParallelLinear必须设置input_is_parallel=True,因为其输入来自列并行层的切分输出;- Attention 使用本地 head 数(
self.to_qkv.num_heads),不能用模型总 head 数去切分 QKV,否则会出现维度不匹配。
3.3 权重加载与切分
并行层替换完成后,还需要让权重加载逻辑适配"合并参数"的形态。Z-Image 在load_weights中声明了stacked_params_mapping(见 z_image_transformer.py):
stacked_params_mapping = [ (".to_qkv.", ".to_q.", "q"), # to_q/to_k/to_v 合并进 to_qkv (".to_qkv.", ".to_k.", "k"), (".to_qkv.", ".to_v.", "v"), (".w13", ".w1", 0), # w1/w3 合并进 w13 (".w13", ".w3", 1), ]加载权重时,切分后的每个并行层参数通过其weight_loader将完整权重按 TP 维度切片到对应卡上。类属性packed_modules_mapping同时服务于量化 checkpoint 适配器与 LoRA 对融合投影的处理。这也解释了为什么"只替换nn.Linear、不动权重加载"会导致权重形状对不上——两者必须同步改造。
四、第三步:校验 TP 约束(可整除性)
TP 正确运行的硬性前提是:所有会被切分的维度必须能被tensor_parallel_size整除。设计文档给出的约束表如下:
| 维度 | 原因 | 错误示例 |
|---|---|---|
num_heads | head 数由 QKVParallelLinear 按列切分 | num_heads=30, tp=4❌(30 % 4 ≠ 0) |
num_kv_heads | KV head 数由 QKVParallelLinear 切分 | num_kv_heads=30, tp=4❌(30 % 4 ≠ 0) |
Z-Image 在此基础上将校验落到了代码中:validate_zimage_tp_constraints(见 z_image_transformer.py)会在模型初始化时对dim、n_heads、n_kv_heads、ffn_hidden_dim(=dim/3*8)、最终输出维度final_out_dims(=patch_size² × f_patch_size × in_channels)逐一检查可整除性,不满足时抛出带"支持的 TP 候选值"提示的ValueError。TP size 从 forward context 中读取(_get_tensor_parallel_size_from_context),这意味着同一个模型代码在单卡推理时(tp=1)与多卡 TP 推理时都能复用。
对应的单元测试(test_zimage_tp_constraints.py)验证了典型场景:
dim=3840, n_heads=30, n_kv_heads=30, tp=2→ 通过,ffn_hidden_dim=10240,final_out_dims=[64],支持的 TP 候选为[1, 2];tp=4→ 因n_heads % tp != 0抛错(30 无法被 4 整除);tp=3→ 因ffn_hidden_dim % tp != 0抛错(10240 无法被 3 整除)。
在规划 TP 规模时,应先用这些约束倒推候选值:对 Z-Image 而言,30 个 head 与 10240 的 FFN 隐藏维度共同决定了实际可用 TP 规模只有 1 和 2。
五、端到端验证:如何测试 TP 是否生效
5.1 Python API 方式
按设计文档,TP 通过DiffusionParallelConfig(tensor_parallel_size=N)开启:
from vllm_omni import Omni from vllm_omni.diffusion.data import DiffusionParallelConfig from vllm_omni.inputs.data import OmniDiffusionSamplingParams parallel_config = DiffusionParallelConfig(tensor_parallel_size=2) omni = Omni(model="your-model-name", parallel_config=parallel_config) output = omni.generate( "a cup of coffee on the table", OmniDiffusionSamplingParams(num_inference_steps=50), )DiffusionParallelConfig(见 vllm_omni/diffusion/data.py)是扩散模型分布式执行的统一配置,除tensor_parallel_size(默认 1)外,还包含pipeline_parallel_size、data_parallel_size、sequence_parallel_size(Ulysses/Ring/AllGather-KV)、cfg_parallel_size、vae_patch_parallel_size、text_encoder_tp_size等,可按需与 TP 组合使用。
5.2 命令行方式
官方离线推理示例(examples/offline_inference/text_to_image)提供了现成的--tensor-parallel-size参数(见 text_to_image.py):
cd examples/offline_inference/text_to_image python text_to_image.py \ --model Your-org/your-model \ --prompt "a cup of coffee on the table" \ --negative-prompt "ugly, unclear" \ --cfg-scale 4.0 \ --num-inference-steps 50 \ --output "tp_enabled.png" \ --tensor-parallel-size 2启动日志中会打印并行配置,例如Parallel configuration: tensor_parallel_size=2, ...(见 text_to_image.py),可用于确认 TP 确实被传入。
5.3 仓库自带的自动化验证
仓库在 tests/diffusion/distributed/test_tensor_parallel.py 中提供了完整的 TP 端到端(E2E)回归测试,以Tongyi-MAI/Z-Image-Turbo为基准模型,对比TP=1与TP=2:
- 正确性:同一 seed 下分别用 TP=1 与 TP=2 生成 512×512 图像,计算两图的 mean 绝对差与 P99 绝对差,断言
mean_abs_diff <= 3e-2且p99_abs_diff <= 2.5e-1,从像素级验证切分不改变生成语义; - 性能:使用
DeviceMemoryMonitor以 20ms 间隔采样显存,断言 TP=2 的每请求中位耗时低于 TP=1(ROCm 平台除外); - 显存:断言 TP=2 的峰值显存低于 TP=1,直接印证"权重切分降低单卡显存"的目标。
该测试标记为hardware_test,需要至少 2 张 CUDA 卡(或 2 张 ROCm 卡),NPU 平台当前会跳过(测试文件中注明 TP parity 测试目前仅支持 CUDA 与 ROCm)。
5.4 人工验证清单
设计文档给出的手工验证步骤:
- 查看日志中的
e2e_time_ms,确认是否有加速; - 对比 TP 关闭与开启时生成图像的质量;
- 确认显存占用按比例下降;
- 将对比结果记录在 PR 中(含正确性、速度、显存三项数据)。
六、常见问题排查
问题 1:TP 未生效(仍跑在单卡上,无显存节省与加速)
症状:模型在单张 GPU 上运行,日志显示tensor_parallel_size=1。
根因与解法:
- 仍然使用
nn.Linear:并行层未被替换。解法:替换为并行等价层:# ❌ 错误:普通线性层无法切分 self.proj = nn.Linear(dim, dim) # ✅ 正确:行并行层 self.proj = RowParallelLinear(dim, dim, input_is_parallel=True) - 未正确传入并行配置:确认
DiffusionParallelConfig(tensor_parallel_size=N)已传入,且启动日志中的并行配置显示tensor_parallel_size=N(而非默认的deploy/default)。
问题 2:前向过程报RuntimeError: shape mismatch
症状:forward 中张量形状对不上。
根因与解法:
- 缺少
input_is_parallel=True:RowParallelLinear期望接收已切分的输入,却收到了完整张量。解法:当输入来自列并行层时务必显式开启:# ✅ 正确配对 self.w1 = ColumnParallelLinear(dim, hidden_dim, return_bias=False) self.w2 = RowParallelLinear( hidden_dim, dim, input_is_parallel=True, # 输入已被 w1 切分 return_bias=False, ) - QKV 切分尺寸用了总 head 数:切分后的 QKV 必须按本地(local)head 数计算:
# ❌ 错误:用模型总 head 数 q_size = self.total_num_heads * self.head_dim # ✅ 正确:用并行层暴露的本地 head 数 q_size = self.to_qkv.num_heads * self.head_dim
问题 3:初始化报"维度不可整除"
症状:加载模型时抛出ValueError: ... requires n_heads % tensor_parallel_size == 0之类错误,并附带Supported tp candidates列表。
根因与解法:所选 TP 规模与模型结构不匹配。解法:按错误提示中给出的候选值调整tensor_parallel_size,或修改模型结构使相关维度可被目标 TP 规模整除(后者一般不建议,属于架构级改动)。
七、仓库内参考实现一览
以下为改造时的对照范本与测试锚点:
| 模型 / 组件 | 路径 | 模式 | 备注 |
|---|---|---|---|
| Z-Image | z_image_transformer.py | 标准 TP | 完整实现,含约束校验(validate_zimage_tp_constraints)与序列并行_sp_plan |
| FLUX | flux_transformer.py | 双流(Dual-stream) | 图像流与文本流分别切分 |
| Qwen-Image | qwen_image_transformer.py | 标准 TP + RoPE | 带旋转位置编码的切分处理 |
| TP E2E 测试 | test_tensor_parallel.py | 端到端 | TP=1 vs TP=2 的正确性、性能、显存三方对比 |
| 约束单元测试 | test_zimage_tp_constraints.py | 单元测试 | 验证可整除性校验逻辑 |
Z-Image 除 TP 外还叠加了序列并行(SP,对应 diffusers 中的上下文并行 CP):_sp_plan指定了在unified_prepare模块边界对 unified 序列、RoPE cos/sin、注意力掩码按序列维切分,并在all_final_layer处聚拢输出(见 z_image_transformer.py)。实际落地时,TP 解决"权重放不下",SP 解决"序列太长",两者可组合使用。
八、总结
为扩散 Transformer 接入 Tensor Parallel 支持,遵循四步流程即可:
- ✅识别线性层——判断哪些层该切分(列并行/行并行),哪些层必须复制(
ReplicatedLinear); - ✅替换为并行层——Attention 用
QKVParallelLinear+RowParallelLinear,FFN 用ColumnParallelLinear/MergedColumnParallelLinear+RowParallelLinear,并同步改造权重加载的合并映射; - ✅校验 TP 约束——确保
num_heads、num_kv_heads、dim、FFN 隐藏维度、最终输出维度均可被tensor_parallel_size整除; - ✅测试——以
tensor_parallel_size=N运行,核对e2e_time_ms、峰值显存与生成图像质量,并以 TP=1 为基线做正确性对比。
以 Z-Image 为参考的这一套"识别 → 替换 → 校验 → 测试"方法论,可以直接迁移到仓库中的其他扩散模型(FLUX、Qwen-Image 等),是 vllm-omni 中为任意 DiT 模型落地多卡推理的标准路径。
【免费下载链接】vllm-omniA framework for efficient model inference with omni-modality models项目地址: https://gitcode.com/GitHub_Trending/vl/vllm-omni
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考