Diffusers 中的 LatteTransformer3DModel:面向视频生成的 3D 扩散 Transformer 架构与源码解析
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
本篇技术指南以 LatteTransformer3DModel 官方 API 文档 为主线,深入解析 🤗 Diffusers 中用于视频类数据(text-to-video)的 3D Diffusion Transformer 骨干模型:它的模块划分、全部配置参数、前向传播中的数据流,以及如何通过 LattePipeline 完成端到端的文生视频推理。读完本文,你将掌握 LatteTransformer3DModel 的架构原理、参数调优要点,以及加载、推理、量化和内存优化的完整实战方案。
一、模型定位:Latte 与 LatteTransformer3DModel
Latte(Latent Diffusion Transformer)是一种面向视频生成的潜空间扩散 Transformer 模型,其论文为《Latte: Latent Diffusion Transformer for Video Generation》(Monash University、上海 AI 实验室、南京大学与南洋理工大学联合提出)。与直接在像素空间建模的扩散模型不同,Latte 先从输入视频中提取时空 token,再通过一系列 Transformer 块在潜空间中建模视频分布;在视频 token 数量庞大的前提下,Latte 通过"将视频的空间维与时间维解耦"的方式设计出多种高效变体,并从视频片段 patch 嵌入、模型变体、timestep 与类别信息注入、时间位置编码、学习策略等角度进行了系统性的实验验证。
在 Diffusers 仓库中,Latte 的核心模型实现位于 latte_transformer_3d.py,对外暴露的类即为LatteTransformer3DModel;配套的推理管线 pipeline_latte.py 中的LattePipeline组合了该 Transformer 与 VAE、T5 文本编码器、调度器,完成"文本 → 视频"的完整链路。
需要说明的是,
LatteTransformer3DModel本质上是一个视频类数据的 3D Transformer 模型(3D 指"帧数 × 高 × 宽"的时空三维),其实现大量借鉴了 PixArt-α 的模块化设计(如AdaLayerNormSingle、PixArtAlphaTextProjection),理解这一点有助于阅读后续的源码细节。
二、架构总览:五个构建阶段
从 latte_transformer_3d.py 的__init__可以看出,LatteTransformer3DModel的构建逻辑被明确划分为五个阶段:
- 输入层(Patch 嵌入 + 空间位置编码):
self.pos_embed = PatchEmbed(...),将视频每一帧的 latent 切成 patch 并嵌入,同时注入二维正弦余弦位置编码; - 空间 Transformer 块:
self.transformer_blocks(nn.ModuleList包裹的BasicTransformerBlock列表),负责对每一帧内部的空间 token 做自注意力与(可选的)文本交叉注意力; - 时间 Transformer 块:
self.temporal_transformer_blocks,同样由BasicTransformerBlock组成,但cross_attention_dim=None(纯自注意力),负责建模帧与帧之间的时序关系; - 输出层:
norm_out(LayerNorm)+scale_shift_table(AdaLN 调制)+proj_out(把 patch 特征线性映射回像素通道),最后通过unpatchify恢复出视频 latent 张量; - Latte 专属辅助模块:
adaln_single(AdaLayerNormSingle,用于注入 timestep 条件)与caption_projection(PixArtAlphaTextProjection,用于把 T5 文本嵌入投影到 Transformer 隐层维度),以及注册为 buffer 的时间位置嵌入temp_pos_embed(由get_1d_sincos_pos_embed_from_grid生成,非持久化,persistent=False)。
这种"空间块 + 时间块交替堆叠"的结构,正是 Latte 在潜空间中对视频时空维度解耦建模的核心思想。
三、完整配置参数与默认值
LatteTransformer3DModel的所有构造参数都通过@register_to_config注册到config中,因此既可以在实例化时传入,也可以通过from_pretrained从 checkpoint 的config.json加载。下表汇总了 源码 docstring 与签名 中的全部参数:
| 参数 | 类型 / 默认值 | 含义与影响 |
|---|---|---|
num_attention_heads | int,默认16 | 多头注意力的头数,inner_dim = num_attention_heads * attention_head_dim |
attention_head_dim | int,默认88 | 每个注意力头的通道数 |
in_channels | int,可选 | 输入 latent 的通道数(对 Latte-1 通常为 4,对应 VAE 的 latent 通道) |
out_channels | int,可选 | 输出通道数;为None时默认取in_channels |
num_layers | int,默认1 | 空间/时间 Transformer 块的层数(两层各堆叠num_layers个块) |
dropout | float,默认0.0 | Transformer 块内的 dropout 概率 |
cross_attention_dim | int,可选 | 文本encoder_hidden_states的维度,用于空间块的交叉注意力 |
attention_bias | bool,默认False | Transformer 块注意力是否包含 bias 参数 |
sample_size | int,默认64 | latent 图的宽/高(正方形),用于学习位置嵌入,训练时固定 |
patch_size | int,可选 | patch 嵌入层中的 patch 尺寸 |
activation_fn | str,默认"geglu" | 前馈网络使用的激活函数(测试中亦使用"gelu-approximate") |
num_embeds_ada_norm | int,可选 | 训练时使用的扩散步数,用于学习 AdaLN 的时间嵌入数量;推理时最多去噪不超过该步数 |
norm_type | str,默认"layer_norm" | 归一化类型,可选"layer_norm"或"ada_layer_norm"(Latte-1 实际使用"ada_norm_single") |
norm_elementwise_affine | bool,默认True | 归一化层是否使用逐元素仿射参数 |
norm_eps | float,默认1e-5 | 归一化层的 epsilon |
caption_channels | int,可选 | 文本(caption)嵌入的通道数,对应 T5 的输出维度 |
video_length | int,默认16 | 视频帧数,用于生成时间位置编码 |
其中几个参数需要特别说明其设计意图:
sample_size与patch_size:PatchEmbed需要根据sample_size生成固定尺寸的二维正弦余弦位置编码,因此训练阶段必须固定;interpolation_scale = max(sample_size // 64, 1)用于在更高分辨率下对位置编码做插值(见 embeddings.py 中PatchEmbed的实现)。num_embeds_ada_norm:与AdaLayerNormSingle配合,把扩散 timestep 编码为可学习的嵌入,通过 AdaLN 调制注入网络;训练步数决定了可去噪的最大步数上限。video_length:通过get_1d_sincos_pos_embed_from_grid生成一维正弦余弦时间位置编码temp_pos_embed(维度为inner_dim),并在每个时间块之前加到 token 上。
作为参考,测试用例 中构造了一个最小模型,其参数组合为:sample_size=8, patch_size=2, attention_head_dim=8, num_attention_heads=3, caption_channels=32, in_channels=4, cross_attention_dim=24, out_channels=8, attention_bias=True, activation_fn="gelu-approximate", num_embeds_ada_norm=1000, norm_type="ada_norm_single", norm_elementwise_affine=False, norm_eps=1e-6——这是快速验证前向传播与梯度的小型化配置。
四、前向传播:时空数据流详解
forward方法的完整实现在 latte_transformer_3d.py,其输入输出约定如下:
输入参数
hidden_states:形状为(batch size, channel, num_frame, height, width)的视频 latent 张量;timestep(torch.LongTensor,可选):去噪步数,经AdaLayerNormSingle编码后注入;encoder_hidden_states(形状(batch size, sequence len, embed dims),可选):用于交叉注意力的文本条件嵌入;若不提供,交叉注意力退化为自注意力;encoder_attention_mask(可选):支持两种格式——二维 mask(batch, sequence_length)(True=保留,False=丢弃)或三维 bias(batch, 1, sequence_length)(0=保留,-10000=丢弃),二维 mask 会自动转换为 bias 并加到交叉注意力分数上;enable_temporal_attentions(bool,默认True):是否启用时间注意力,关闭后模型仅执行空间建模;return_dict(bool,默认True):为True时返回Transformer2DModelOutput,否则返回裸 tuple。
逐步数据流
- 张量重排:
(B, C, F, H, W)→ 置换为(B, F, C, H, W)并展平为(B*F, C, H, W),把每一帧当作独立的二维图; - Patch 嵌入:
pos_embed将每帧切成 patch 并加上二维位置编码,得到(B*F, num_patches, inner_dim)的 token 序列; - 时间条件注入:
adaln_single(timestep, ...)生成timestep与embedded_timestep(后者用于输出端的调制);这里added_cond_kwargs被置为{"resolution": None, "aspect_ratio": None}; - 文本投影与广播:
caption_projection把 T5 文本嵌入投影到inner_dim,并通过repeat_interleave(num_frame, ...)把文本 token 复制到每一帧上(encoder_hidden_states_spatial);同理,timestep也被按帧数、按 patch 数复制为timestep_spatial与timestep_temp,分别喂给空间块和时间块; - 空间块:对每一帧 token 执行自注意力 + 文本交叉注意力 + 前馈(
transformer_blocks),支持梯度检查点(gradient_checkpointing)路径; - 时间块(
enable_temporal_attentions=True时):先把(B*F, num_patches, D)重排为(B*num_patches, F, D)——即每个空间位置的所有帧聚在一起,第一层时间块之前还会加上temp_pos_embed时间位置编码,然后由temporal_transformer_blocks做纯自注意力(无交叉注意力),最后再重排回(B*F, num_patches, D)供下一层空间块使用; - 输出调制与 unpatchify:
embedded_timestep复制到每帧后,与scale_shift_table相加并chunk(2)拆出 scale/shift,对norm_out后的隐状态做hidden * (1 + scale) + shift调制,经proj_out映射为(patch_size * patch_size * out_channels)的 patch 预测;最后reshape+einsum("nhwpqc->nchpwq")重排为(N, C, H_p*patch, W_p*patch),再 reshape 回(B, C, F, H, W)的视频 latent。
值得注意的是,空间块与时间块的交叉注意力注入位置不同:文本条件只作用于空间块,而时间块仅建模帧间关系——这正是 Latte 将时空解耦的体现。此外,从第 240 行的zip(self.transformer_blocks, self.temporal_transformer_blocks)可以看出,模型按"空间块 → 时间块"交替执行,且两层块的数量都由num_layers控制。
五、端到端实战:文生视频推理
5.1 基础推理
LatteTransformer3DModel通常不作为独立模型使用,而是作为 LattePipeline 的transformer组件。管线还包含AutoencoderKL(VAE)、T5EncoderModel(Latte 使用 t5-v1_1-xxl 变体)与T5Tokenizer、以及调度器(官方使用DDIMScheduler)。
最小推理示例(与 pipeline_latte.py 中的示例 一致):
import torch from diffusers import LattePipeline from diffusers.utils import export_to_gif # 也可以用 "maxin-cn/Latte-1" 替换 checkpoint id pipe = LattePipeline.from_pretrained("maxin-cn/Latte-1", torch_dtype=torch.float16) # 开启内存优化:文本编码器 -> Transformer -> VAE 依次在 CPU/GPU 间搬运 pipe.enable_model_cpu_offload() prompt = "A small cactus with a happy face in the Sahara desert." videos = pipe(prompt).frames[0] export_to_gif(videos, "latte.gif")其中LattePipeline的模块卸载顺序model_cpu_offload_seq = "text_encoder->transformer->vae"(见 pipeline_latte.py),并且tokenizer与text_encoder被声明为_optional_components——这意味着你可以提前用encode_prompt把 prompt 编码为prompt_embeds,之后移除这两个组件、仅凭嵌入直接驱动管线(测试test_save_load_optional_components验证了这一行为,见 test_latte.py)。
5.2 使用 torch.compile 加速
在管线内部,transformer以hidden_states=latent_model_input, encoder_hidden_states=prompt_embeds, timestep=current_timestep, enable_temporal_attentions=..., return_dict=False的方式被调用(见 pipeline_latte.py),因此可以安全地对 transformer 与 VAE 的解码器做编译:
import torch from diffusers import LattePipeline pipeline = LattePipeline.from_pretrained("maxin-cn/Latte-1", torch_dtype=torch.float16).to("cuda") # 将 transformer 与 vae 切换为 channels_last 内存布局 pipeline.transformer.to(memory_format=torch.channels_last) pipeline.vae.to(memory_format=torch.channels_last) # 编译组件 pipeline.transformer = torch.compile(pipeline.transformer) pipeline.vae.decode = torch.compile(pipeline.vae.decode) video = pipeline(prompt="A dog wearing sunglasses floating in space, surreal, nebulae in background").frames[0]5.3 bitsandbytes 8-bit 量化推理
对于显存受限的场景,可以对文本编码器与 Transformer 分别量化后再组装管线(量化后端与选择方式可参考仓库 Quantization 总览):
import torch from diffusers import BitsAndBytesConfig as DiffusersBitsAndBytesConfig, LatteTransformer3DModel, LattePipeline from diffusers.utils import export_to_gif from transformers import BitsAndBytesConfig, T5EncoderModel quant_config = BitsAndBytesConfig(load_in_8bit=True) text_encoder_8bit = T5EncoderModel.from_pretrained( "maxin-cn/Latte-1", subfolder="text_encoder", quantization_config=quant_config, torch_dtype=torch.float16, ) quant_config = DiffusersBitsAndBytesConfig(load_in_8bit=True) transformer_8bit = LatteTransformer3DModel.from_pretrained( "maxin-cn/Latte-1", subfolder="transformer", quantization_config=quant_config, torch_dtype=torch.float16, ) pipeline = LattePipeline.from_pretrained( "maxin-cn/Latte-1", text_encoder=text_encoder_8bit, transformer=transformer_8bit, torch_dtype=torch.float16, device_map="balanced", ) prompt = "A small cactus with a happy face in the Sahara desert." video = pipeline(prompt).frames[0] export_to_gif(video, "latte.gif")这里LatteTransformer3DModel直接通过from_pretrained("maxin-cn/Latte-1", subfolder="transformer", ...)加载官方权重,印证了该模型类与ModelMixin/ConfigMixin的完整集成(保存、加载、量化均开箱可用)。
5.4 直接实例化与独立前向
若需在自定义脚本中独立使用该模型(例如做模型研究或接入自定义管线),可参考测试中的构造方式:
import torch from diffusers import LatteTransformer3DModel transformer = LatteTransformer3DModel( sample_size=8, num_layers=1, patch_size=2, attention_head_dim=8, num_attention_heads=3, caption_channels=32, in_channels=4, cross_attention_dim=24, out_channels=8, attention_bias=True, activation_fn="gelu-approximate", num_embeds_ada_norm=1000, norm_type="ada_norm_single", norm_elementwise_affine=False, norm_eps=1e-6, ) latents = torch.randn(1, 4, 4, 8, 8) # (B, C, F, H, W) timestep = torch.tensor([500]) # 1D 去噪步数 prompt_embeds = torch.randn(1, 16, 24) # (B, seq_len, cross_attention_dim) out = transformer( hidden_states=latents, timestep=timestep, encoder_hidden_states=prompt_embeds, enable_temporal_attentions=True, ) print(out.sample.shape) # 形状与输入 latent 一致六、内存与注意力优化特性
LatteTransformer3DModel继承自ModelMixin, ConfigMixin, CacheMixin,并声明了_supports_gradient_checkpointing = True,因此在训练/微调时可开启梯度检查点以大幅降低显存;_skip_layerwise_casting_patterns = ["pos_embed", "norm"]则表明在做逐层精度转换(layerwise casting)时,位置编码与归一化层会被跳过,以保持数值稳定性。
此外,从测试套件可以确认该模型/管线与 Diffusers 的注意力级加速方案深度兼容(见 test_latte.py):
- Pyramid Attention Broadcast(PAB):通过
spatial_attention_block_skip_range=2、temporal_attention_block_skip_range=2、cross_attention_block_skip_range=2等配置,在指定 timestep 区间(如空间注意力(100, 700)、时间注意力(100, 800))内跳过部分注意力计算并复用相邻步的注意力图,以提速;其块标识符明确指向transformer_blocks(空间/交叉)与temporal_transformer_blocks(时间),印证了本文对模块命名的解读; - FasterCache:基于块级与时间步级跳过(skip_range=2)以及无条件批次跳过,配合注意力权重回调实现缓存式加速;
- MemoryTesterMixin:覆盖 CPU offload、group offload 与 layerwise casting 等内存优化路径的回归测试。
这些测试的存在,说明LatteTransformer3DModel的空间块、时间块与交叉注意力路径均可被 PAB/FasterCache 这类"跳过 + 复用"策略精确识别与调度,是理解其注意力结构命名(transformer_blocks/temporal_transformer_blocks)的重要佐证。
七、从源码验证到实践要点小结
回顾整个分析过程,可以沉淀出以下可直接用于工程的要点:
- 输入形状约定:模型接受 5 维视频 latent
(B, C, F, H, W),输出同形状预测;sample_size决定 patch 位置编码的固定分辨率,推理分辨率变化依赖PatchEmbed的插值机制; - 时空解耦建模:空间块承担文本交叉注意力,时间块为纯自注意力并依赖一维 sincos 时间位置编码;
enable_temporal_attentions=False可关闭时间建模(对应纯图像/逐帧生成场景); - 条件注入:timestep 经
AdaLayerNormSingle编码后既进入 AdaLN 调制(输出端 scale/shift),也作为各 Transformer 块的 AdaLN 输入;文本嵌入经PixArtAlphaTextProjection投影并逐帧广播; - 与管线协作:作为
LattePipeline的可量化、可编译、可 offload 的transformer组件,支持 8-bit 量化、torch.compile、channels_last、CPU offload 等一揽子加速/省显存手段; - 可验证性:所有上述结论均可通过 源码实现、管线调用点 与 测试用例 逐一核对,是学习"如何在 Diffusers 中新增视频扩散 Transformer 骨干"的极佳范本。
【免费下载链接】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),仅供参考