
DiT图像分块Patchify全解3步把像素变Token【免费下载链接】DiTOfficial PyTorch Implementation of Scalable Diffusion Models with Transformers项目地址: https://gitcode.com/GitHub_Trending/di/DiTDiT把扩散模型骨干换成Transformer绕不开的动作就是图像分块Patchify图先切块喂给网络算完再拼回完整图像。读完这篇你能在白板上独立画出这条像素到Token再回到像素的数据流。为什么Transformer吃不下一张完整图像DiT图像分块的动机Transformer的注意力是两两配对的序列长度是N注意力矩阵就要算N²个组合。你可以把它想象成N个人两两握手人一多握手次数是平方级爆炸的。把64×64的图直接展平喂进去序列长度4096注意力矩阵要算约1600万个组合显存和算力瞬间爆表。而且逐像素展平还有一个更隐蔽的问题每个位置只代表一个像素信息太碎Token之间也做不出有意义的交互。DiT的解法是先分块——把图像按patch_size切成网格每个小块当成一个Token序列长度直接除以patch_size的平方。这就是DiT图像分块要解决的事既保住空间结构又把序列压到Transformer吃得下的规模。 数据流拆解DiT图像分块步骤——从像素到Token再回像素整条前向链路就是四段接力x_embedder切块嵌入 → 28个Transformer Block →final_layer投影 →unpatchify还原。下面按数据流动的方向逐段拆。x_embedder做了什么图像分块嵌入的三步DiT Patchify的入口是x_embedder它由三个动作串成先把输入图按patch_size×patch_size切成不重叠的小块再把每个小块内部的所有像素展平成一维向量最后用一个线性层把向量投影到隐藏维度D。拿默认配置算一笔账32×32的输入、patch_size2时整张图被切成256个Patch输入张量(N, C, 32, 32)变成(N, T, D)其中T256、D1152。也就是说原来1024个位置的二维图像被压缩成了256个语义更稠密的序列元素——这才是Transformer能高效处理图像的关键一步。位置编码为什么必须加Patchify丢了空间坐标切块、展平之后每个Token住在画面哪个位置彻底丢了(N, T, D)里的T只是个计数张量本身不携带任何坐标相当于把一桌按棋盘摆好的棋子全倒进了袋子里。所以DiT在分块嵌入后立刻补上一份位置嵌入x self.x_embedder(x) self.pos_embed # (N, T, D)这份pos_embed是用固定的2D sin-cos位置编码初始化的给第(i, j)个Patch一个独一无二的门牌号。换句话说图像分块嵌入负责内容位置编码负责地址两者相加Transformer才知道每个Token在画面里的真实位置。unpatchify还原Patchify还原算法里einsum那步在挪什么经过28个DiTBlock和final_layer之后每个Token被投影回patch_size²×C个数值形状是(N, T, p²·C)——相当于每个Token手里攥着自己那一小块的原材料但排布是交错的块的行列位置藏在T里块内像素和通道混在一起。unpatchify还原要做的就是逆着来models.py里只有三行核心x x.reshape(shape(x.shape[0], h, w, p, p, c)) x torch.einsum(nhwpqc-nchpwq, x) imgs x.reshape(shape(x.shape[0], c, h * p, h * p))第一步reshape把T拆回网格(N, T, p²·C)变成(N, h, w, p, p, c)每个Patch重新知道自己在哪一行哪一列。中间那行einsum是整段最妙的一笔它把通道维c挪到最前面把块内像素(p, p)和网格位置(h, w)重新排好队——你可以把它想象成把一筐按格子码好的货先按品类归堆再按货架顺序上架。最后reshape把相邻小块在两个方向上拼接(N, c, h·p, w·p)就还原成了完整图像。整条链路严格对称x_embedder怎么切unpatchify还原就怎么拼回去。patch_size选大还是选小对生成质量的影响patch_size是DiT图像分块里最直接的超参数models.py里提供了/2、/4、/8三档变体。选小的patch2×232×32的输入切出256个Token空间细节保得最完整但注意力要算256²个组合算力吃紧选大的patch8×8只剩16个Token注意力组合降到256个训练飞快可每个Token要代表64个像素高频细节容易糊。你可以把它想象成报纸的印刷分辨率小patch是细版印刷清楚但贵大patch是粗版速写便宜但丢细节。取舍其实很简单——分辨率高、显存富余时偏小追求吞吐时偏大。 一张图看懂DiT完整前向链路这批样例图正是切块→嵌入→逐层Transformer→unpatchify还原整条链路跑完的产物你能在图里感受到的细节差异本质就是patch_size在细节保留与算力开销之间的取舍。带走这三件事图像分块让注意力开销降一个平方位置编码给Token补回空间坐标unpatchify对称还原完整图像【免费下载链接】DiTOfficial PyTorch Implementation of Scalable Diffusion Models with Transformers项目地址: https://gitcode.com/GitHub_Trending/di/DiT创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考