ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Diffusers 中的 WanTransformer3DModel:Wan 2.1/2.2 视频生成 3D Diffusion Transformer 架构与实战解析

2026/9/10 14:03:47 拓冰建站 浏览量
Diffusers 中的 WanTransformer3DModel:Wan 2.1/2.2 视频生成 3D Diffusion Transformer 架构与实战解析 Diffusers 中的 WanTransformer3DModelWan 2.1/2.2 视频生成 3D Diffusion Transformer 架构与实战解析【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusersWanTransformer3DModel 是 Diffusers 为阿里 Wan 团队开源的 Wan 2.1 / Wan 2.2 系列视频模型实现的 3D Diffusion Transformer 主干网络负责在去噪过程中对视频潜变量进行逐帧、逐位置的联合建模是文生视频T2V、图生视频I2V与视频生视频V2V管线的核心去噪引擎。读完本文你将掌握该模型的完整类接口、全部可配置参数与默认值、前向推理的输入输出契约以及如何独立加载、单文件转换并在 Wan 全系管线中实际使用它。模型定位与适用场景WanTransformer3DModel 在仓库中的完整实现位于 transformer_wan.py它继承自ModelMixin、ConfigMixin、PeftAdapterMixin、FromOriginalModelMixin、CacheMixin与AttentionMixin六个基类这意味着它同时具备标准模型的保存/加载能力、register_to_config配置化构造能力、LoRA 适配器挂载能力、从原始 Wan 官方权重.safetensors单文件直接转换加载的能力以及缓存与注意力处理器扩展能力。从源码结构看该模型服务于四条 Wan 管线pipeline_wan.py — Wan 文生视频T2Vpipeline_wan_i2v.py — Wan 图生视频I2Vpipeline_wan_video2video.py — Wan 视频生视频V2Vpipeline_wan_22.py含WanImageToVideoPipeline的 2.2 变体见 pipeline_wan_22_image_to_video.py— Wan 2.2 两阶段去噪其中transformer负责高噪声阶段、transformer_2负责低噪声阶段。模型的输入是 VAE 编码后的 5D 视频潜变量(batch_size, num_channels, num_frames, height, width)输出为同形状的 5D 去噪结果因此被命名为“3D Transformer”——其注意力既沿空间高、宽展开也沿时间维帧展开从而对视频类数据建立时空联合建模。快速加载示例官方 API 文档wan_transformer_3d.md给出的标准加载方式如下from diffusers import WanTransformer3DModel transformer WanTransformer3DModel.from_pretrained( Wan-AI/Wan2.1-T2V-1.3B-Diffusers, subfoldertransformer, dtypetorch.bfloat16 )关键点解读subfoldertransformerWan 官方 Diffusers 仓库采用多组件目录结构vae、text_encoder、transformer、image_encoder等Transformer 权重固定存放在transformer子目录下。这一点在仓库多个位置得到印证WanPipeline初始化时以transformer作为模型组件名注册single_file_model.py 中也将WanTransformer3DModel的default_subfolder配置为transformer。dtypetorch.bfloat16Wan 系列模型官方推荐使用 bfloat16 精度加载与推理在 pipeline_wan.py 的示例中文本编码器与 VAE 以torch.float32加载、Transformer 以torch.bfloat16加载二者在推理时由管线内部完成 dtype 对齐prompt_embeds.to(transformer_dtype)。类签名与核心参数详解WanTransformer3DModel.__init__的全部参数、默认值与语义如下与源码 docstring 及 transformer_wan.py 中的实际签名一致参数默认值说明patch_size(1, 2, 2)视频嵌入的 3D patch 尺寸(t_patch, h_patch, w_patch)由 3D 卷积完成 patchifynum_attention_heads40注意力头数量attention_head_dim128每个注意力头的通道数in_channels16输入通道数即 VAE 潜变量通道数out_channels16输出通道数缺省时取in_channelstext_dim4096文本嵌入UMT5-XXL的输入维度freq_dim256时间步正弦嵌入维度ffn_dim13824前馈网络中间维度num_layers40Transformer Block 层数cross_attn_normTrue是否启用交叉注意力归一化qk_normrms_norm_across_heads是否启用 Q/K 归一化跨头 RMSNormeps1e-6归一化层 epsilonimage_dimNone图像嵌入维度I2V 模型为 1280对应 CLIP vision encoder 输出added_kv_proj_dimNone附加 KV 投影的通道数为None时不使用rope_max_seq_len1024旋转位置编码RoPE最大序列长度pos_embed_seq_lenNone图像嵌入的可学习位置编码长度说明类 docstring 中的“文本嵌入固定长度 512”指的是 UMT5 文本编码器在 Wan 2.1 中默认生成 512 个 token 的文本条件而text_dim4096是每个 token 的嵌入维度。二者对应关系可参见 pipeline_wan.py 中max_sequence_length512的默认值以及WanAttnProcessor中“512 is the context length of the text encoder”的源码注释。与测试配置的对应关系tests/models/transformers/test_models_transformer_wan.py 中的微型模型配置验证了上述参数的真实语义patch_size(1, 2, 2)、num_attention_heads2、attention_head_dim12、in_channels4、text_dim16、ffn_dim32、num_layers2、qk_normrms_norm_across_heads、rope_max_seq_len32。测试用 dummy 输入形状为(1, 4, 2, 16, 16)即 1 个 batch、4 通道、2 帧、16×16 空间分辨率。同时测试还给出了真实模型的维度画像Wan 2.2 I2V 模型为in_channels3616 个视频潜变量通道 20 个 mask 通道、text_dim4096这正是WanTransformer3DModel在 I2V 任务中的标准配置见 test_models_transformer_wan.py。模型内部架构WanTransformer3DModel的前向流程transformer_wan.py可划分为五个阶段1. 3D patchify 与 RoPE 位置编码输入(B, C, F, H, W)首先经过WanRotaryPosEmbed将注意力头维度按 6:2:2 比例拆分为时间维t_dim、高度维h_dim、宽度维w_dim分别通过get_1d_rotary_pos_embed生成 1D 旋转频率再广播扩展为 3D 网格频率。freqs_dtype在 Apple SiliconMPS上使用float32其余平台使用float64。patch_embeddingnn.Conv3d(in_channels, inner_dim, kernel_sizepatch_size, stridepatch_size)即用(1, 2, 2)的 3D 卷积核完成 patchify将(F, H, W)压缩为(F/1, H/2, W/2)的 token 网格随后flatten(2).transpose(1, 2)并contiguous()转换为序列形式。2. 条件嵌入时间 文本 图像WanTimeTextImageEmbedding内部包含三个子模块timesteps_projTimesteps正弦频率编码→time_embedderTimestepEmbeddingMLP→ SiLU 激活 →time_proj线性层输出timestep_proj维度为inner_dim * 6text_embedderPixArtAlphaTextProjection(text_embed_dim4096, diminner_dim, act_fngelu_tanh)将 UMT5 文本嵌入投影到模型宽度image_embedder可选I2V 专用WanImageEmbedding由FP32LayerNorm → FeedForward → FP32LayerNorm组成并为pos_embed_seq_len注册可学习位置编码。I2V 时图像嵌入与文本嵌入在序列维度拼接torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim1)。3. Transformer Block 堆叠模型由num_layers默认 40个WanTransformerBlock堆叠而成每个 Block 采用 AdaLN自适应层归一化 门控残差的结构内部包含自注意力attn1 交叉注意力attn2对文本条件——交叉注意力层在added_kv_proj_dim非空时I2V 模型额外附加图像 KV 投影前馈网络ffnFeedForward激活函数gelu-approximatescale_shift_table形状(1, 6, dim)将timestep_proj拆分为 6 个 AdaLN 系数shift_msa / scale_msa / gate_msa控制自注意力前的归一化与残差门控c_shift_msa / c_scale_msa / c_gate_msa控制 FFN 前的归一化与残差门控。当temb.ndim 4Wan 2.2 ti2v 逐 token 时间步时AdaLN 系数按序列维度展开否则按 batch 维度展开Wan 2.1 / Wan 2.2 14B。4. 注意力处理器WanAttention的默认处理器为WanAttnProcessor其实现transformer_wan.py依赖 PyTorch 2.0 的scaled_dot_product_attention构造时若无该 API 会抛出 ImportError。核心行为Q/K 在进入注意力前分别经norm_q/norm_kRMSNorm归一化支持 RoPE 旋转嵌入的apply_rotary_embI2V 任务中将encoder_hidden_states按“前 512 token 文本、其余 token 图像”切分源码硬编码文本上下文长度 512图像 token 走add_k_proj/add_v_proj投影并单独计算一次注意力其结果与主注意力输出相加支持fuse_projections()将 QKV或 KV投影融合为单个线性层以减少内核调用。历史上存在WanAttnProcessor2_0现已被声明为弃用并直接转发到WanAttnProcessortransformer_wan.py。5. 输出归一化与 unpatchify最后一层为norm_outFP32LayerNormproj_out线性层输出out_channels * prod(patch_size)通道 输出侧scale_shift_table形状(1, 2, inner_dim)提供 shift/scale最终通过reshape → permute → flatten的逆 patchify 操作还原为(B, C, F, H, W)形状。FP32 稳定性设计模型类级属性transformer_wan.py定义了多处与精度相关的行为_keep_in_fp32_modulesrope、time_embedder、scale_shift_table、norm1/norm2/norm3保持 FP32 计算_skip_layerwise_casting_patternspatch_embedding、condition_embedder、norm跳过逐层精度转换_supports_gradient_checkpointing True支持梯度检查点训练时以时间换显存_no_split_modules [WanTransformerBlock]与accelerate的 device_map 自动切分配合_repeated_blocks [WanTransformerBlock]标识可复用/重复的 Block 结构。归一化层全部使用FP32LayerNorm且前向中多处显式.float()运算后再.type_as(hidden_states)回退精度这是 Wan 在 bfloat16 推理下保持数值稳定的关键实现细节。forward 方法输入输出契约forward方法transformer_wan.py签名如下def forward( self, hidden_states: torch.Tensor, # (B, C, F, H, W) 视频潜变量 timestep: torch.LongTensor, # 去噪时间步shape 为 (B,) 或 (B, seq_len) encoder_hidden_states: torch.Tensor, # (B, seq_len, text_dim) 文本/条件嵌入 encoder_hidden_states_image: torch.Tensor | None None, # I2V 图像条件嵌入 return_dict: bool True, attention_kwargs: dict[str, Any] | None None, ) - torch.Tensor | dict[str, torch.Tensor]:行为要点hidden_states形状必须为(batch_size, num_channels, num_frames, height, width)且帧数、高、宽需能被patch_size整除timestep支持两种形态标量时间步(B,)Wan 2.1 / 2.2 14B以及按序列展开的(B, seq_len)Wan 2.2 ti2v 逐 token 时间步expand_timestepsTrue时由管线构造见 pipeline_wan.pyreturn_dictTrue时返回Transformer2DModelOutput其sample字段为去噪后的 5D 张量return_dictFalse时返回(sample,)元组attention_kwargs会传递给注意力处理器如用于 LoRA 缩放或注意力控制。在 Wan 管线中的实际调用方式Wan 系列管线将WanTransformer3DModel作为去噪主干使用。以 pipeline_wan.py 中的去噪循环为例with current_model.cache_context(cond): noise_pred current_model( hidden_stateslatent_model_input, timesteptimestep, encoder_hidden_statesprompt_embeds, attention_kwargsattention_kwargs, return_dictFalse, )[0] if self.do_classifier_free_guidance: with current_model.cache_context(uncond): noise_uncond current_model( hidden_stateslatent_model_input, timesteptimestep, encoder_hidden_statesnegative_prompt_embeds, attention_kwargsattention_kwargs, return_dictFalse, )[0] noise_pred noise_uncond guidance_scale * (noise_pred - noise_uncond)无分类器引导CFG时模型被调用两次条件 无条件输出按noise_uncond w * (noise_pred - noise_uncond)融合cache_context(cond/uncond)是CacheMixin提供的缓存上下文用于跨步缓存如 MagCache管线会自动根据 Transformer 的patch_size与 VAE 缩放因子校正输出分辨率h_multiple_of vae_scale_factor_spatial * patch_size[1]、w_multiple_of vae_scale_factor_spatial * patch_size[2]pipeline_wan.py即默认(1, 2, 2)patch 下宽高必须是 16 的倍数Wan 2.2 两阶段去噪时transformer处理t boundary_ratio * num_train_timesteps的高噪声步transformer_2处理低噪声步pipeline_wan.py。在 I2V 管线中WanImageToVideoPipeline通过 CLIP vision encoder 提取图像嵌入并传入encoder_hidden_states_image见 pipeline_wan_i2v.py 的encode_image取hidden_states[-2]层输出从而驱动 Transformer 的附加图像注意力路径。单文件权重加载与 LoRA借助FromOriginalModelMixinWanTransformer3DModel支持直接从原始 Wan 官方单文件权重加载。仓库注册的转换入口single_file_model.pyWanTransformer3DModel: { checkpoint_mapping_fn: convert_wan_transformer_to_diffusers, default_subfolder: transformer, }即使用from_single_file加载 ComfyUI 重新打包的wan2.1_t2v_1.3B_bf16.safetensors或wan2.1_i2v_480p_14B_fp8_e4m3fn.safetensors等权重时会通过convert_wan_transformer_to_diffusers完成键名映射。对应测试见 test_model_wan_transformer3d_single_file.py。同时类支持 PeftAdapterMixin 提供的 LoRA 适配能力WanLoraLoaderMixin供管线使用与模型级apply_lora_scale(attention_kwargs)装饰器transformer_wan.py保证了 LoRA 权重在注意力层的正确注入与缩放tests/lora/test_lora_layers_wan.py中亦有对应覆盖。从零构建一个最小可运行的模型实例如需脱离预训练权重验证模型结构可按测试配置构造微型实例参考 test_models_transformer_wan.pyimport torch from diffusers import WanTransformer3DModel transformer WanTransformer3DModel( patch_size(1, 2, 2), num_attention_heads2, attention_head_dim12, in_channels4, out_channels4, text_dim16, freq_dim256, ffn_dim32, num_layers2, cross_attn_normTrue, qk_normrms_norm_across_heads, rope_max_seq_len32, ) hidden_states torch.randn(1, 4, 2, 16, 16) # (B, C, F, H, W) encoder_hidden_states torch.randn(1, 12, 16) # (B, seq_len, text_dim) timestep torch.randint(0, 1000, (1,)) output transformer( hidden_stateshidden_states, timesteptimestep, encoder_hidden_statesencoder_hidden_states, return_dictFalse, )[0] print(output.shape) # (1, 4, 2, 16, 16)注意输入形状约束num_frames2需能被patch_size[0]1整除height/width16需能被patch_size[1:]2整除。此外从测试基类可知该模型已通过梯度检查点、TorchCompile、bitsandbytes / TorchAO / GGUF 量化、LoRA、内存优化与训练态等全套模型测试test_models_transformer_wan.py可在上述优化路径下稳定运行。性能优化要点结合类级属性与管线用法落地 Wan 视频生成时值得关注的优化手段包括精度策略Transformer 使用 bfloat16VAE 与文本编码器使用 float32参考 pipeline_wan.py 官方示例关键模块自动保持 FP32 计算无需手动干预。分辨率与 patch 对齐生成前确保宽高为vae_scale_factor_spatial * patch_size[1]默认 16的整数倍管线会自动向下取整校正。模型卸载与并行model_cpu_offload_seq text_encoder-transformer-transformer_2-vae定义了 CPU 卸载顺序_cp_plan属性transformer_wan.py提供了上下文并行Context Parallel的切分计划其中特意禁用了encoder_hidden_states的切分因为 I2V 图像编码器固定输出 257 个 token导致拼接后的 769 token 无法被设备数整除——这是多卡推理时需要注意的约束。投影融合调用WanAttention.fuse_projections()可将 QKV/KV/附加 KV 投影融合为单个线性层减少 kernel 启动开销。总结WanTransformer3DModel 是 Diffusers 中 Wan 2.1 / Wan 2.2 视频生成体系的架构核心以 3D patchify RoPE 时空位置编码 AdaLN 门控 Transformer Block 附加图像注意力为骨架支撑 T2V、I2V、V2V 三条管线以及 Wan 2.2 的两阶段去噪方案。其完整实现、默认配置与测试覆盖均可直接在 transformer_wan.py 与 test_models_transformer_wan.py 中查阅与复现是理解视频 DiT 架构及在此基础上二次开发LoRA、量化、并行化的理想起点。【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考