📑 本页目录(点开跳转)
07 · 去噪主干:U-Net 与 DiT
⏱ 34 分钟 | ⭐ 前面六章一直在说「让模型预测噪声」,这一章打开那个模型
🎯 一句话
去噪网络只是一个接口固定的函数:进 (带噪的图, 第几步, 条件),出「这里面加了什么噪声」—— 里面用什么结构实现,扩散本身完全不关心。
所以主干从 U-Net 换成 Transformer(DiT),上一章的 CFG、之前的加噪公式和 DDIM 跳步,一行都不用改。
🧱 一、先把接口钉死:三进一出
上一章的 CFG 要跑两次去噪网络,一次带条件一次不带。那两次调用长这样:
| 输入 | 是什么 | 形状 |
|---|---|---|
x_t |
第 t 步的带噪图(潜空间里就是潜变量) | 和输出同形,比如 4×64×64 |
t |
走到第几步了 | 一个整数,先变成向量再送进去 |
c |
条件:文本嵌入、类别、别的图…… | 一串向量 |
输出只有一个:一张和 x_t 同形状的预测噪声。
⭐ 这个接口就是「扩散」和「网络」之间的那道缝。 加噪是固定公式、采样是固定循环,只有这个函数是学出来的 —— 换主干时调用方一个字都不用动。
⚠️ 有件事接口里没写出来:t 不能只在输入端塞一次。
它要影响的是每一层的行为(t 大时该输出接近纯噪声,t 小时只该修细节)。
所以标准做法是把 t 变成向量,在每个块内部注入。下面两种主干注入方式不同,目的一样。
🪜 二、U-Net:先压下去,再抬回来
左臂(编码器)几轮「卷积块 → 下采样」,分辨率一路减半、通道一路翻倍;底部(瓶颈)分辨率最低,⭐ 自注意力通常只插在这附近 —— 注意力的代价是 token 数的平方,压小之后才付得起;右臂(解码器)几轮「上采样 → 卷积块」抬回原分辨率。条件走交叉注意力进来:Query 来自图像特征,Key/Value 来自文本嵌入。
跳跃连接把左臂每一级送到右臂同分辨率那一级。⭐ 它用的是拼接,不是残差那种相加:
# 跳跃连接把通道数变成了多少
import torch
import torch.nn as nn
torch.manual_seed(0)
x = torch.randn(1, 4, 32, 32)
h1 = nn.Conv2d(4, 64, 3, padding=1)(x) # 编码器一级:存下来
h2 = nn.Conv2d(64, 128, 3, stride=2, padding=1)(h1) # 下采样:32×32 → 16×16
cat = torch.cat([nn.Upsample(scale_factor=2)(h2), h1], 1) # ⭐ 拼接,不是相加
print("编码器 h1", tuple(h1.shape), "\n瓶颈 h2", tuple(h2.shape),
"\n拼接后 ", tuple(cat.shape), "← 128 + 64 = 192")
编码器 h1 (1, 64, 32, 32)
瓶颈 h2 (1, 128, 16, 16)
拼接后 (1, 192, 32, 32) ← 128 + 64 = 192
所以解码器每一级的输入通道是「上采样上来的 + 编码器同级存下的」。写代码时最容易在这里对不上。
✂️ 三、跳跃连接到底救了什么
标准答案是「保留细节」。这句话可以量出来 —— 下面这段几秒钟跑完:
# 下采样到底压掉了什么:把细纹和粗块一起送进 8 倍下采样再还原
import torch
import torch.nn.functional as F
torch.manual_seed(0)
H = 64
yy, xx = torch.meshgrid(torch.arange(H), torch.arange(H), indexing="ij")
coarse = ((xx - 32) ** 2 + (yy - 32) ** 2 < 20 ** 2).float() # 低频:一个大圆
fine = ((xx % 2 == 0).float() - 0.5) * 0.5 # 高频:2 像素周期条纹
img = (coarse + fine).view(1, 1, H, H)
small = F.avg_pool2d(img, 8) # ⭐ 8 倍下采样
back = F.interpolate(small, size=(H, H), mode="bilinear", align_corners=False)
def keep(sig): # 还原后这部分信号还剩多少(1.0 = 完好,0 = 全没了)
return float((back * sig).sum() / (sig * sig).sum())
print(f"低频(大圆)保留 {keep(coarse.view(1,1,H,H)):.3f}")
print(f"高频(细条纹)保留 {keep(fine.view(1,1,H,H)):.4f}")
低频(大圆)保留 0.830
高频(细条纹)保留 0.0000
⭐ 高频那一行是
0.0000,不是「变小了」,是【一点不剩】。 走瓶颈那条路只带得回「这是一只猫、它在画面中间」这类低频信息; 「猫毛往哪个方向长」在下采样那一步就没了,再精巧的解码器也变不回来。
⚠️ 而扩散模型在 t 很小的那几步,干的正好就是补高频(轮廓早在大 t 时就定了)。
所以跳跃连接不是锦上添花,是去噪这个任务本身要求的。
🧩 四、DiT:把潜空间切成 token
DiT 把整个 U 形扔掉,换成一摞完全一样的 Transformer 块。关键是开头和结尾那两步转换:
# DiT 的第一步和最后一步:把潜空间切成 token,过完 Transformer 再拼回去
import torch
import torch.nn as nn
torch.manual_seed(0)
B, C, H, W, P, D = 1, 4, 32, 32, 2, 384 # 潜空间 4×32×32,patch 边长 2,token 384 维
x = torch.randn(B, C, H, W)
tok = nn.Conv2d(C, D, P, stride=P)(x).flatten(2).transpose(1, 2) # ⭐ 切块+投影一步完成
print("潜空间", tuple(x.shape), "→ token", tuple(tok.shape),
f" N = ({H}//{P})×({W}//{P}) = {(H//P)*(W//P)}")
tok = nn.TransformerEncoderLayer(D, 6, 4 * D, batch_first=True, norm_first=True)(tok)
print("过完一层", tuple(tok.shape), " 形状不变,所以能一直堆下去")
out = nn.Linear(D, P * P * C)(tok).view(B, H // P, W // P, P, P, C)
out = out.permute(0, 5, 1, 3, 2, 4).reshape(B, C, H, W) # 拼回整图
print("拼回去 ", tuple(out.shape), " 和输入同形 =", out.shape == x.shape)
潜空间 (1, 4, 32, 32) → token (1, 256, 384) N = (32//2)×(32//2) = 256
过完一层 (1, 256, 384) 形状不变,所以能一直堆下去
拼回去 (1, 4, 32, 32) 和输入同形 = True
⭐ patch 边长是这里最贵的旋钮:边长从 2 减到 1,token 数 256 → 1024(×4),
而注意力是 token 数的平方 —— 代价 ×16。DiT-XL/2 那个 /2 说的就是这个数。
t 和 c 怎么进:不用交叉注意力,而是把两者合成一个向量,去生成每一层 LayerNorm 的缩放平移、以及每个残差分支上的门控系数(adaLN-Zero)。
⭐ 那个门控是零初始化的,训练刚开始时每个块都是恒等映射 —— 下一章 ControlNet 的零初始化卷积是同一个想法。
🏁 五、为什么后来是 DiT 赢了
不是因为「Transformer 比卷积聪明」。三条都很具体:
| 理由 | 具体是什么 |
|---|---|
| scaling 更可预测 ⭐ | 加宽、加深、加数据能换回多少收益,Transformer 这条线上有大量现成曲线;U-Net「通道加到多少算够」从来没有类似共识 |
| 工程栈白捡 ⭐ | FlashAttention、序列并行、混合精度配方、各种融合 kernel —— 都是为语言模型造的,DiT 拿来即用 |
| 归纳偏置更少 | 数据足够多时「少假设 + 更多数据」通常赢;换分辨率对 DiT 也只是 token 数变了 |
⚠️ 代价很实在:注意力是平方复杂度,token 一多就付不起。所以 DiT 几乎总是长在潜空间里(第 5 章那一套)—— 直接在 512×512 像素上切 patch,token 数是六位数,没人跑得起。是 Latent Diffusion 先把图压小了,DiT 才成立。 💡 所以也不是「U-Net 过时了」:数据少算力少时卷积那两条先验是白送的,老一代开源模型周边的控制生态(下一章的主题)也大多是给 U-Net 写的。想往大了堆才是 DiT 的主场。
🚧 六、一个特别容易搞混的点
DiT 不是「用 Transformer 生成图像」。 自回归图像生成也用 Transformer —— 把图切成 token 一个接一个地吐 —— 但那是另一种范式:
| 自回归图像生成 | DiT | |
|---|---|---|
| 训练目标 | 预测下一个 token | 预测噪声(和前面几章一模一样) |
| 怎么出图 | 逐 token 解码 | DDPM / DDIM 那个去噪循环,一行不改 |
⭐ 换主干就像把一段程序里的排序从快排换成归并:复杂度和常数变了,调用方一个字不用动。 DiT 换掉的只是那个预测噪声的函数的实现,扩散那套数学一条都没动。
🔗 这一章连到哪里
| 去哪 | 为什么 |
|---|---|
| 机器学习与深度学习基础 12 · CNN 处理图像 | U-Net 的左臂就是那一章的零件堆起来的。⭐ 尤其去看它「感受野」那一节:加一次 stride=2 的下采样,之后每层的感受野翻倍增长 —— 这既是 U-Net 敢把分辨率一路压小的理由,也是本章第三节那个 0.0000 的代价来源 |
| 大模型全景导论 主线 2 · 模型怎样看懂一句话 | DiT 的每一块就是那一章那张图:注意力 → 残差 → FFN,结构一层没改,差别只在序列里装的不是词而是图像 patch。它讲的 Q/K/V 各自的角色,本章的交叉注意力(Query 来自图像、Key/Value 来自文本)直接用得上 |
| 大模型全景导论 10b · 图像生成 | 14 分钟的索引版。它的组件表里「U-Net / DiT」只占一行、写着「DiT = 用 Transformer 做去噪,现在的主流」—— 本章是那一行的展开,也是对它最容易被误读成「Transformer 生成图像」的纠偏 |
✅ 检查点
- 去噪网络的输入有哪三样?输出的形状和谁一样?
- 为什么时间步
t不能只在输入端塞进去一次? - U-Net 的跳跃连接用的是拼接还是相加?这对解码器的通道数有什么影响?
- 第三节那段代码里,8 倍下采样之后高频信号保留了多少?这个数说明跳跃连接在防什么?
- DiT 里 patch 边长从 2 减到 1,token 数和注意力代价各变成几倍?
- 为什么说「DiT 不是用 Transformer 生成图像」?它又为什么几乎总要配合潜空间?
👀 答案
x_t(带噪图)、t(第几步)、c(条件)。输出和x_t同形,比如 4×64×64 进、4×64×64 出。- 因为
t要影响每一层的行为 ——t大时该输出接近纯噪声、t小时只修细节。所以它变成向量后在每个块内部注入(U-Net 加进残差块,DiT 用 adaLN 生成每层的缩放平移)。 - 拼接(
torch.cat,通道维)。所以解码器每一级的输入通道 = 上采样上来的 + 编码器同级存下的;代码里是128 + 64 = 192。 0.0000,一点不剩(同一次实验里低频保留 0.830)。说明瓶颈那条路只带得动低频,高频细节必须绕过瓶颈直接送到解码器 —— 而扩散在t小的那几步干的正好是补高频。- token 数 256 → 1024(×4),注意力是 token 数的平方,所以代价 ×16。
DiT-XL/2的/2说的就是 patch 边长。 - 因为它仍然是扩散:训练目标是预测噪声,出图还是 DDPM/DDIM 那个去噪循环;自回归图像生成的目标是预测下一个 token。至于潜空间:注意力是 token 数的平方复杂度,直接在 512×512 像素上切 patch,token 数是六位数,跑不起 —— 先由 Latent Diffusion 把图压小,DiT 才成立。
🛑 可以停在这里
读到这里,你已经能看懂一份扩散模型配置在说什么:unet 那串通道数是左臂的每一级,attention_resolutions 是注意力插在哪几级,DiT 那边的 patch_size 和 hidden_size 各自在换什么代价。
什么时候回来:生成图大结构对、细节糊时回看第三节(那通常是跳跃连接这条路上的问题);要换分辨率或估显存时回看第四节的 token 数和第五节那条平方复杂度。
⚡ 走神救援
先记住这几件事
- U-Net 和 DiT 都实现去噪网络;改变主干不等于改变整个扩散训练目标。
- U-Net 用多尺度与跳跃连接保留信息,DiT 把空间块作为 token 处理。
- 核对时间和条件怎样进入各层,再比较形状、计算量与细节保留。
下一节 👉 08-ControlNet与风格LoRA.md