云计算百科
云计算领域专业知识百科平台

Sora/DiT架构:扩散Transformer视频生成技术

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

特性UNet (SD3)DiT (Sora)
架构 CNN+CrossAttn Pure Transformer
分辨率 固定 原生可变
规模扩展 次线性 线性(DiT scaling law)
训练效率 高(小规模) 高(大规模)
FID (256×256) 2.2 2.27

七、总结

DiT/Sora的核心设计:

  • Patch化 → 统一的token表示
  • adaLN-Zero → 条件注入的优雅方式
  • 3D压缩 → 时空潜空间降维
  • 原生可变分辨率 → 告别固定尺寸
  • Scaling Law → Transformer的规模红利
  • 赞(0)
    未经允许不得转载:网硕互联帮助中心 » Sora/DiT架构:扩散Transformer视频生成技术
    分享到: 更多 (0)

    评论 抢沙发

    评论前必须登录!