news 2026/9/10 11:41:11

Diffusers 中的 LatteTransformer3DModel:面向视频生成的 3D 扩散 Transformer 架构与源码解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Diffusers 中的 LatteTransformer3DModel:面向视频生成的 3D 扩散 Transformer 架构与源码解析

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-α 的模块化设计(如AdaLayerNormSinglePixArtAlphaTextProjection),理解这一点有助于阅读后续的源码细节。

二、架构总览:五个构建阶段

从 latte_transformer_3d.py 的__init__可以看出,LatteTransformer3DModel的构建逻辑被明确划分为五个阶段:

  1. 输入层(Patch 嵌入 + 空间位置编码)self.pos_embed = PatchEmbed(...),将视频每一帧的 latent 切成 patch 并嵌入,同时注入二维正弦余弦位置编码;
  2. 空间 Transformer 块self.transformer_blocksnn.ModuleList包裹的BasicTransformerBlock列表),负责对每一帧内部的空间 token 做自注意力与(可选的)文本交叉注意力;
  3. 时间 Transformer 块self.temporal_transformer_blocks,同样由BasicTransformerBlock组成,但cross_attention_dim=None(纯自注意力),负责建模帧与帧之间的时序关系;
  4. 输出层norm_out(LayerNorm)+scale_shift_table(AdaLN 调制)+proj_out(把 patch 特征线性映射回像素通道),最后通过unpatchify恢复出视频 latent 张量;
  5. Latte 专属辅助模块adaln_singleAdaLayerNormSingle,用于注入 timestep 条件)与caption_projectionPixArtAlphaTextProjection,用于把 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_headsint,默认16多头注意力的头数,inner_dim = num_attention_heads * attention_head_dim
attention_head_dimint,默认88每个注意力头的通道数
in_channelsint,可选输入 latent 的通道数(对 Latte-1 通常为 4,对应 VAE 的 latent 通道)
out_channelsint,可选输出通道数;为None时默认取in_channels
num_layersint,默认1空间/时间 Transformer 块的层数(两层各堆叠num_layers个块)
dropoutfloat,默认0.0Transformer 块内的 dropout 概率
cross_attention_dimint,可选文本encoder_hidden_states的维度,用于空间块的交叉注意力
attention_biasbool,默认FalseTransformer 块注意力是否包含 bias 参数
sample_sizeint,默认64latent 图的宽/高(正方形),用于学习位置嵌入,训练时固定
patch_sizeint,可选patch 嵌入层中的 patch 尺寸
activation_fnstr,默认"geglu"前馈网络使用的激活函数(测试中亦使用"gelu-approximate"
num_embeds_ada_normint,可选训练时使用的扩散步数,用于学习 AdaLN 的时间嵌入数量;推理时最多去噪不超过该步数
norm_typestr,默认"layer_norm"归一化类型,可选"layer_norm""ada_layer_norm"(Latte-1 实际使用"ada_norm_single"
norm_elementwise_affinebool,默认True归一化层是否使用逐元素仿射参数
norm_epsfloat,默认1e-5归一化层的 epsilon
caption_channelsint,可选文本(caption)嵌入的通道数,对应 T5 的输出维度
video_lengthint,默认16视频帧数,用于生成时间位置编码

其中几个参数需要特别说明其设计意图:

  • sample_sizepatch_sizePatchEmbed需要根据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 张量;
  • timesteptorch.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_attentionsbool,默认True):是否启用时间注意力,关闭后模型仅执行空间建模;
  • return_dictbool,默认True):为True时返回Transformer2DModelOutput,否则返回裸 tuple。

逐步数据流

  1. 张量重排(B, C, F, H, W)→ 置换为(B, F, C, H, W)并展平为(B*F, C, H, W),把每一帧当作独立的二维图;
  2. Patch 嵌入pos_embed将每帧切成 patch 并加上二维位置编码,得到(B*F, num_patches, inner_dim)的 token 序列;
  3. 时间条件注入adaln_single(timestep, ...)生成timestepembedded_timestep(后者用于输出端的调制);这里added_cond_kwargs被置为{"resolution": None, "aspect_ratio": None}
  4. 文本投影与广播caption_projection把 T5 文本嵌入投影到inner_dim,并通过repeat_interleave(num_frame, ...)把文本 token 复制到每一帧上(encoder_hidden_states_spatial);同理,timestep也被按帧数、按 patch 数复制为timestep_spatialtimestep_temp,分别喂给空间块和时间块;
  5. 空间块:对每一帧 token 执行自注意力 + 文本交叉注意力 + 前馈(transformer_blocks),支持梯度检查点(gradient_checkpointing)路径;
  6. 时间块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)供下一层空间块使用;
  7. 输出调制与 unpatchifyembedded_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),并且tokenizertext_encoder被声明为_optional_components——这意味着你可以提前用encode_prompt把 prompt 编码为prompt_embeds,之后移除这两个组件、仅凭嵌入直接驱动管线(测试test_save_load_optional_components验证了这一行为,见 test_latte.py)。

5.2 使用 torch.compile 加速

在管线内部,transformerhidden_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=2temporal_attention_block_skip_range=2cross_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)的重要佐证。

七、从源码验证到实践要点小结

回顾整个分析过程,可以沉淀出以下可直接用于工程的要点:

  1. 输入形状约定:模型接受 5 维视频 latent(B, C, F, H, W),输出同形状预测;sample_size决定 patch 位置编码的固定分辨率,推理分辨率变化依赖PatchEmbed的插值机制;
  2. 时空解耦建模:空间块承担文本交叉注意力,时间块为纯自注意力并依赖一维 sincos 时间位置编码;enable_temporal_attentions=False可关闭时间建模(对应纯图像/逐帧生成场景);
  3. 条件注入:timestep 经AdaLayerNormSingle编码后既进入 AdaLN 调制(输出端 scale/shift),也作为各 Transformer 块的 AdaLN 输入;文本嵌入经PixArtAlphaTextProjection投影并逐帧广播;
  4. 与管线协作:作为LattePipeline的可量化、可编译、可 offload 的transformer组件,支持 8-bit 量化、torch.compilechannels_last、CPU offload 等一揽子加速/省显存手段;
  5. 可验证性:所有上述结论均可通过 源码实现、管线调用点 与 测试用例 逐一核对,是学习"如何在 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 11:35:46

【单片机课设毕设项目】基于 STM32 或 51 单片机的 LCD1602 显示植物环境无线管控系统设计 基于 STM32 或 51 单片机的声光报警温室温湿度光照自动控制系统(020607)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于嵌入式单片机,Java、小程序技术领域和毕业项目实战 ✌️…

作者头像 李华
网站建设 2026/9/10 11:35:02

2026年靠谱的旧衣服回收平台怎么找?主流平台实测对比来了

换季时节,天气转冷,当冬装夏装堆在一起,衣柜早已爆满,收纳空间告急;搬家的时候,清理出几大袋旧衣服,带走超重、扔掉又实在可惜;想开展一场断舍离,清理闲置衣物&#xff0…

作者头像 李华
网站建设 2026/9/10 11:34:41

Cesium中Entity与Primitive核心概念与性能优化指南

1. Cesium中的Entity与Primitive核心概念解析在三维地理可视化领域,Cesium作为当前最强大的WebGL地球引擎之一,其图形渲染体系主要围绕Entity和Primitive两大核心概念构建。我刚接触Cesium时曾被这两个概念困扰许久——它们看似都能实现相似的可视化效果…

作者头像 李华
网站建设 2026/9/10 11:31:59

WebView封装APK技术解析与实践指南

1. 项目概述:网址与本地文件封装为APK的核心价值在移动应用开发领域,将现有网站或本地HTML文件快速转换为安卓安装包(APK)的需求日益增长。Website 2 APK Builder Pro这类工具的出现,为不具备原生开发能力的内容提供者…

作者头像 李华