Sora/DiT架构:扩散Transformer视频生成技术
一、引言
2024年OpenAI Sora改变了视频生成:1分钟连贯视频、3D一致性、模拟物理世界。Sora背后是 DiT(Diffusion Transformer) 架构——用Transformer替代UNet做扩散模型的骨干网络。本文将深入 DiT/Sora 的技术原理。
二、DiT架构设计
2.1 Patch化 + Transformer
# DiT将视频帧划分为Patches(类似ViT)
# 输入:原始视频 B×T×C×H×W → B×(T·H·W/P³)×D
class PatchEmbed3D(nn.Module):
"""3D Patch嵌入:时间+空间分块"""
def __init__(self, patch_size=(2, 16, 16), in_channels=3, embed_dim=1152):
super().__init__()
self.patch_size = patch_size
# 3D卷积分块
self.proj = nn.Conv3d(
in_channels, embed_dim,
kernel_size=patch_size,
stride=patch_size
)
def forward(self, x):
# x: [B, C, T, H, W] 例如 [1, 3, 64, 256, 256]
x = self.proj(x) # → [1, 1152, 32, 16, 16]
B, C, T, H, W = x.shape
x = x.flatten(2).transpose(1, 2) # → [1, 8192, 1152]
return x
2.2 自适应Layer Norm (adaLN-Zero)
class DiTBlock(nn.Module):
"""DiT核心块:adaLN + 多头注意力 + 点式前馈"""
def __init__(self, hidden_size=1152, num_heads=16):
super().__init__()
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False)
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False)
self.attn = nn.MultiheadAttention(hidden_size, num_heads, batch_first=True)
# MLP (SwiGLU)
self.mlp = nn.Sequential(
nn.Linear(hidden_size, 4 * hidden_size),
nn.GELU(),
nn.Linear(4 * hidden_size, hidden_size),
)
# ★ adaLN: 从时间步嵌入c生成调制参数
# 输出: 6×scale + 6×shift = 12×hidden_size
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 12 * hidden_size)
)
def forward(self, x, c):
# x: [B, N, D], c: [B, D] 时间步条件
# 生成6组调制参数
params = self.adaLN_modulation(c).chunk(12, dim=–1)
shift_msa, scale_msa, gate_msa = params[:3]
shift_mlp, scale_mlp, gate_mlp = params[3:6]
# 多头注意力(带调制)
x_norm = self.norm1(x) * (1 + scale_msa.unsqueeze(1)) + shift_msa.unsqueeze(1)
attn_out, _ = self.attn(x_norm, x_norm, x_norm)
x = x + gate_msa.unsqueeze(1) * attn_out
# MLP(带调制)
x_norm = self.norm2(x) * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1)
mlp_out = self.mlp(x_norm)
x = x + gate_mlp.unsqueeze(1) * mlp_out
return x
2.3 Full DiT Model
class DiT(nn.Module):
"""完整DiT模型"""
def __init__(self, input_size=32, patch_size=2, in_channels=4,
hidden_size=1152, depth=28, num_heads=16):
super().__init__()
# Patch嵌入
self.x_embedder = PatchEmbed3D(
patch_size=(1, patch_size, patch_size),
embed_dim=hidden_size
)
# 时间步嵌入(傅里叶特征)
self.t_embedder = TimestepEmbedder(hidden_size)
# 条件嵌入(文本)
self.y_embedder = LabelEmbedder(768, hidden_size) # T5-XXL
# Transformer块
self.blocks = nn.ModuleList([
DiTBlock(hidden_size, num_heads) for _ in range(depth)
])
# 输出投影
self.final_layer = FinalLayer(
hidden_size, patch_size, in_channels
)
# 权重初始化(DiT专用)
self.initialize_weights()
def forward(self, x, t, y):
"""x: [B,C,T,H,W], t: [B], y: [B,L,D_text]"""
# 1. Patch嵌入 + 位置编码
x = self.x_embedder(x) # [B, N, D]
# 2. 条件嵌入
c = self.t_embedder(t) # 时间步 → [B, D]
if y is not None:
c = c + self.y_embedder(y) # 文本 → [B, D]
# 3. Transformer处理
for block in self.blocks:
x = block(x, c)
# 4. 输出:预测噪声
x = self.final_layer(x, c) # [B, N, patch_size²*C]
return x
# DiT-S/2: depth=12, hidden=384, heads=6, 33M params
# DiT-B/2: depth=12, hidden=768, heads=12, 130M
# DiT-L/2: depth=24, hidden=1024, heads=16, 458M
# DiT-XL/2:depth=28, hidden=1152, heads=16, 675M
四、Sora专有特性
# Sora的额外设计:
# 1. 原生可变分辨率
# 不同于固定256×256,Sora处理各种分辨率1280×720、1920×1080
# 使用3D Rotary Position Embedding (3D-RoPE)
# 2. 时空Compression
# 先通过VAE压缩到潜空间(T×H×W → T/4×H/8×W/8)
# 降维8×8×4=256x后再做扩散
# 3. LLM驱动的视频理解
# 用GPT-4/DALL-E 3 re-caption所有训练视频
# 高质量文本条件(200+ tokens)替代简单标签
# 4. 联合训练
# 图像+视频联合训练 → 图像质量提升视频细节
# 训练时动态padding到统一patch数
五、推理优化
# DDIM加速:50步 → 20步
from diffusers import DPMSolverMultistepScheduler
scheduler = DPMSolverMultistepScheduler.from_config(
model.scheduler.config,
algorithm_type="dpmsolver++",
solver_order=3
)
# 模型并行(DiT-XL单卡放不下)
# Tensor Parallelism (TP) + Sequence Parallelism (SP)
# TP: 每个head在不同GPU
# SP: 长序列分散到多GPU的LayerNorm/Dropout
# 内存优化
pipe.enable_model_cpu_offload() # 不活跃部分放CPU
pipe.enable_vae_slicing() # VAE分块解码
pipe.enable_vae_tiling() # 像素空间分块
# 编译加速
import torch._dynamo
pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead")
六、DiT vs UNet
| 架构 | CNN+CrossAttn | Pure Transformer |
| 分辨率 | 固定 | 原生可变 |
| 规模扩展 | 次线性 | 线性(DiT scaling law) |
| 训练效率 | 高(小规模) | 高(大规模) |
| FID (256×256) | 2.2 | 2.27 |
七、总结
DiT/Sora的核心设计:
网硕互联帮助中心




评论前必须登录!
注册