拆解Stable Diffusion:交叉注意力如何听懂提示词
拆解Stable Diffusion交叉注意力如何听懂提示词【免费下载链接】stable-diffusionA latent text-to-image diffusion model项目地址: https://gitcode.com/GitHub_Trending/st/stable-diffusion同一句提示词换种写法出图风格天差地别想让画面里恰好出现提示词强调的细节却总被忽略这些对不上的感觉根源是文本条件进入扩散模型的方式太弱。Stable Diffusion 的答案是交叉注意力Cross-Attention在 U-Net 的去噪网络里让图像特征主动去查文本特征把语义对齐做成每一步去噪都能感知的软检索。交叉注意力在生成链路里的位置整条链路是这样分段的冻结的 CLIP 文本编码器把提示词变成 77 个词元向量每个 768 维由 ldm/modules/encoders/modules.py 的FrozenCLIPEmbedder完成图像被 AutoencoderKL 压成 4 通道的低分辨率潜空间张量扩散过程在这里进行去噪用的 U-Net 在 configs/stable-diffusion/v1-inference.yaml 里声明了关键两行conditioning_key: crossattn context_dim: 768crossattn告诉 U-Net条件注入方式用交叉注意力768就是 CLIP 输出的词元维度也是 K、V 投影的输入维度。交叉注意力不是一次性的条件相加而是嵌在 U-Net 各层级的 Transformer 块里每个去噪时间步都参与一次。白话版原理一次带权重的查字典类比查字典每个图像位置拿着一句问题Query查询先翻目录页Key键看哪个词条最相关再按相关度把词条内容Value值混合进自己。写成公式每一步配一句人话Q X·W_q图像特征 X 经投影变成问题向量K C·W_k文本特征 C 经另一组投影变成目录页V C·W_v同一批文本特征再投影一次变成词条正文A softmax(Q·Kᵀ / √d)每个问题与每条目录算相似度点积除以 √d 压住数值softmax 归一成总和为 1 的权重out A·V每个位置拿到的是所有词条的加权平均关键差别自注意力里 Q、K、V 来自同一个输入交叉注意力里 Q 来自图像、K/V 来自文本——两套向量走的是两套线性层谁提问和被提问彻底分开。源码走读CrossAttention 的三个线性层核心类在 ldm/modules/attention.py构造部分class CrossAttention(nn.Module): def __init__(self, query_dim, context_dimNone, heads8, dim_head64, dropout0.): inner_dim dim_head * heads context_dim default(context_dim, query_dim) self.scale dim_head ** -0.5 # 缩放因子 1/√d self.to_q nn.Linear(query_dim, inner_dim, biasFalse) # 图像侧投影 self.to_k nn.Linear(context_dim, inner_dim, biasFalse) # 文本侧键 self.to_v nn.Linear(context_dim, inner_dim, biasFalse) # 文本侧值 ...注意to_k、to_v吃的是context_dim768与to_q的query_dim无关——两个空间各投影各的。前向传播def forward(self, x, contextNone, maskNone): q self.to_q(x) context default(context, x) # 没传 context 时退化为自注意力 k self.to_k(context) v self.to_v(context) q, k, v map(lambda t: rearrange(t, b n (h d) - (b h) n d, hh), (q, k, v)) # 拆8个头 sim einsum(b i d, b j d - b i j, q, k) * self.scale # 相似度矩阵 × 缩放 ... attn sim.softmax(dim-1) # 归一化成权重 out einsum(b i j, b j d - b i d, attn, v) # 按权重混合文本内容 out rearrange(out, (b h) n d - b n (h d), hh) return self.to_out(out)文本怎么和特征图对上SpatialTransformer负责把二维图变一维序列x self.proj_in(x) # 1x1卷积升到内维度 x rearrange(x, b c h w - b (h w) c) # 32x32图 → 1024个位置 for block in self.transformer_blocks: x block(x, contextcontext) x self.proj_out(x) return x x_in # 残差输出每个 Transformer 块内部则是自注意力 → 交叉注意力 → 前馈三连def _forward(self, x, contextNone): x self.attn1(self.norm1(x)) x # 自注意力像素间互通 x self.attn2(self.norm2(x), contextcontext) x # 交叉注意力查文本 x self.ff(self.norm3(x)) x return x设计取舍为什么这么写缩放因子为什么是 1/√d。点积结果是 d 个随机项之和方差随 d 线性增长。d64 时不缩放相似度动辄 ±8softmax 会把权重几乎全压给单一词元——所有位置都去抄同一个词注意力变成独裁。self.scale dim_head ** -0.5把方差拉回 1 附近让 softmax 保持众数而非独裁。多头拆分买的是什么。rearrange把 512 维劈成 8 组 64 维并行计算相当于 8 个读者同时查字典有的头学会盯颜色词有的盯空间关系有的盯材质。单头只能表达一种相关多头能表达多种互不干扰的语义对齐。零初始化投影层的冷启动技巧。proj_out被zero_module清零意味着训练刚启动时SpatialTransformer输出恒为 0整个模块退化成恒等映射——U-Net 先按普通卷积网络跑交叉注意力再慢慢接管避免随机初始化在初期炸掉整个去噪网络。动手验证一条命令出图几行代码看权重先用仓库自带脚本跑一张权重文件需自备python scripts/txt2img.py --prompt a painting of a fire, oil on canvas --ckpt 你的ckpt路径想看到交叉注意力的中间量不必跑完整扩散构造一个块直接打印注意力矩阵即可import torch from ldm.modules.attention import BasicTransformerBlock block BasicTransformerBlock(320, n_heads8, d_head40, context_dim768, checkpointFalse) x torch.randn(1, 1024, 320) # 一张 32x32 潜图的展平序列 c torch.randn(1, 77, 768) # 77 个 CLIP 词元 ca block.attn2 # 取出交叉注意力层 q ca.to_q(x); k ca.to_k(c); v ca.to_v(c) sim q[0] k[0].T * (40 ** -0.5) print(sim.shape) # torch.Size([1024, 77]) print(sim.argmax(dim1)[:10]) # 前10个像素位置各自选中了哪个词元打印出的第 2 维索引就是每个图像位置最关注的词元下标。换不同提示词重跑分布会明显变化——这就是听没听懂最直接的证据。要点回顾交叉注意力解决的是文本条件如何进入去噪网络Q 来自图像、K/V 来自文本三个独立线性层完成两空间的桥接本质是软检索每个像素位置对 77 个词元算相似度softmax 后按权重混合词元内容1/√d缩放防 softmax 饱和8 头并行表达多语义对齐zero_module让新模块冷启动时等价于恒等配置层面只需conditioning_key: crossattn与context_dim: 768两行U-Net 即在各层级挂上SpatialTransformer验证路径跑 scripts/txt2img.py 出图或直接打印sim矩阵观察词元选择延伸阅读仓库内文件注意力全部实现ldm/modules/attention.pyU-Net 挂接 SpatialTransformer 的位置ldm/modules/diffusionmodules/openaimodel.py条件注入总控LatentDiffusionldm/models/diffusion/ddpm.pyCLIP 文本编码器ldm/modules/encoders/modules.py模型超参配置configs/stable-diffusion/v1-inference.yaml【免费下载链接】stable-diffusionA latent text-to-image diffusion model项目地址: https://gitcode.com/GitHub_Trending/st/stable-diffusion创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考