尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

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 与 LatteTransformer3DModelLatteLatent 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_blocksnn.ModuleList包裹的BasicTransformerBlock列表负责对每一帧内部的空间 token 做自注意力与可选的文本交叉注意力时间 Transformer 块self.temporal_transformer_blocks同样由BasicTransformerBlock组成但cross_attention_dimNone纯自注意力负责建模帧与帧之间的时序关系输出层norm_outLayerNormscale_shift_tableAdaLN 调制proj_out把 patch 特征线性映射回像素通道最后通过unpatchify恢复出视频 latent 张量Latte 专属辅助模块adaln_singleAdaLayerNormSingle用于注入 timestep 条件与caption_projectionPixArtAlphaTextProjection用于把 T5 文本嵌入投影到 Transformer 隐层维度以及注册为 buffer 的时间位置嵌入temp_pos_embed由get_1d_sincos_pos_embed_from_grid生成非持久化persistentFalse。这种空间块 时间块交替堆叠的结构正是 Latte 在潜空间中对视频时空维度解耦建模的核心思想。三、完整配置参数与默认值LatteTransformer3DModel的所有构造参数都通过register_to_config注册到config中因此既可以在实例化时传入也可以通过from_pretrained从 checkpoint 的config.json加载。下表汇总了 源码 docstring 与签名 中的全部参数参数类型 / 默认值含义与影响num_attention_headsint默认16多头注意力的头数inner_dim num_attention_heads * attention_head_dimattention_head_dimint默认88每个注意力头的通道数in_channelsint可选输入 latent 的通道数对 Latte-1 通常为 4对应 VAE 的 latent 通道out_channelsint可选输出通道数为None时默认取in_channelsnum_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-approximatenum_embeds_ada_normint可选训练时使用的扩散步数用于学习 AdaLN 的时间嵌入数量推理时最多去噪不超过该步数norm_typestr默认layer_norm归一化类型可选layer_norm或ada_layer_normLatte-1 实际使用ada_norm_singlenorm_elementwise_affinebool默认True归一化层是否使用逐元素仿射参数norm_epsfloat默认1e-5归一化层的 epsiloncaption_channelsint可选文本caption嵌入的通道数对应 T5 的输出维度video_lengthint默认16视频帧数用于生成时间位置编码其中几个参数需要特别说明其设计意图sample_size与patch_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_size8, patch_size2, attention_head_dim8, num_attention_heads3, caption_channels32, in_channels4, cross_attention_dim24, out_channels8, attention_biasTrue, activation_fngelu-approximate, num_embeds_ada_norm1000, norm_typeada_norm_single, norm_elementwise_affineFalse, norm_eps1e-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。逐步数据流张量重排(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_attentionsTrue时先把(B*F, num_patches, D)重排为(B*num_patches, F, D)——即每个空间位置的所有帧聚在一起第一层时间块之前还会加上temp_pos_embed时间位置编码然后由temporal_transformer_blocks做纯自注意力无交叉注意力最后再重排回(B*F, num_patches, D)供下一层空间块使用输出调制与 unpatchifyembedded_timestep复制到每帧后与scale_shift_table相加并chunk(2)拆出 scale/shift对norm_out后的隐状态做hidden * (1 scale) shift调制经proj_out映射为(patch_size * patch_size * out_channels)的 patch 预测最后reshapeeinsum(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组件。管线还包含AutoencoderKLVAE、T5EncoderModelLatte 使用 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_dtypetorch.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_stateslatent_model_input, encoder_hidden_statesprompt_embeds, timestepcurrent_timestep, enable_temporal_attentions..., return_dictFalse的方式被调用见 pipeline_latte.py因此可以安全地对 transformer 与 VAE 的解码器做编译import torch from diffusers import LattePipeline pipeline LattePipeline.from_pretrained(maxin-cn/Latte-1, torch_dtypetorch.float16).to(cuda) # 将 transformer 与 vae 切换为 channels_last 内存布局 pipeline.transformer.to(memory_formattorch.channels_last) pipeline.vae.to(memory_formattorch.channels_last) # 编译组件 pipeline.transformer torch.compile(pipeline.transformer) pipeline.vae.decode torch.compile(pipeline.vae.decode) video pipeline(promptA 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_8bitTrue) text_encoder_8bit T5EncoderModel.from_pretrained( maxin-cn/Latte-1, subfoldertext_encoder, quantization_configquant_config, torch_dtypetorch.float16, ) quant_config DiffusersBitsAndBytesConfig(load_in_8bitTrue) transformer_8bit LatteTransformer3DModel.from_pretrained( maxin-cn/Latte-1, subfoldertransformer, quantization_configquant_config, torch_dtypetorch.float16, ) pipeline LattePipeline.from_pretrained( maxin-cn/Latte-1, text_encodertext_encoder_8bit, transformertransformer_8bit, torch_dtypetorch.float16, device_mapbalanced, ) 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, subfoldertransformer, ...)加载官方权重印证了该模型类与ModelMixin/ConfigMixin的完整集成保存、加载、量化均开箱可用。5.4 直接实例化与独立前向若需在自定义脚本中独立使用该模型例如做模型研究或接入自定义管线可参考测试中的构造方式import torch from diffusers import LatteTransformer3DModel transformer LatteTransformer3DModel( sample_size8, num_layers1, patch_size2, attention_head_dim8, num_attention_heads3, caption_channels32, in_channels4, cross_attention_dim24, out_channels8, attention_biasTrue, activation_fngelu-approximate, num_embeds_ada_norm1000, norm_typeada_norm_single, norm_elementwise_affineFalse, norm_eps1e-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_stateslatents, timesteptimestep, encoder_hidden_statesprompt_embeds, enable_temporal_attentionsTrue, ) 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.pyPyramid Attention BroadcastPAB通过spatial_attention_block_skip_range2、temporal_attention_block_skip_range2、cross_attention_block_skip_range2等配置在指定 timestep 区间如空间注意力(100, 700)、时间注意力(100, 800)内跳过部分注意力计算并复用相邻步的注意力图以提速其块标识符明确指向transformer_blocks空间/交叉与temporal_transformer_blocks时间印证了本文对模块命名的解读FasterCache基于块级与时间步级跳过skip_range2以及无条件批次跳过配合注意力权重回调实现缓存式加速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_attentionsFalse可关闭时间建模对应纯图像/逐帧生成场景条件注入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),仅供参考
返回列表