ViT 图像块嵌入原理与实现
图像分割为 patch 并映射为 token 序列的完整流程
原题:在视觉Transformer(ViT)中,图像是如何被分割并转换为输入token序列的?请说明图像块嵌入(patch embedding)的具体实现过程。
多模态 · 百度真题
回答与解析
从图像到token
输入按PyTorch布局记为 [B,C,H,W]。若patch大小为 P_h×P_w,且先不考虑重叠,patch网格为 H/P_h × W/P_w,token数:
N = (H/P_h) × (W/P_w)
正方形patch时才可简写为 N=HW/P²。若H或W不能整除patch大小,需要明确选择裁剪、padding或可变分辨率策略。
每个patch展平后维度是 P_h×P_w×C,再乘可学习矩阵 E ∈ R^((P_hP_wC)×D),得到D维patch token。把所有patch堆叠后形状为 [B,N,D]。
卷积等价实现
使用 Conv2d(C,D,kernel_size=P,stride=P) 时,每个卷积窗口正好覆盖一个不重叠patch,每个输出通道的卷积核可展平成一行投影权重,因此与“切patch、展平、线性投影”在权重映射上等价:
proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch, stride=patch)
tokens = proj(x).flatten(2).transpose(1, 2) # [B,N,D]
卷积是常见且便于调用优化kernel的实现,但不能脱离硬件和框架宣称它一定“远高于”unfold+GEMM。
原始ViT在序列前加入可学习CLS token,并加可学习位置embedding;分类读取CLS输出。其他模型也可用平均池化、二维/相对位置编码。Self-attention提供单层全局token交互,而CNN通过堆叠、下采样或大卷积核同样可以获得全局感受野,区别是归纳偏置和交互路径,不是“CNN做不到全局”。
口语版讲法(约4分钟)
- 从输入形状推导patch数量
- 解释展平和线性投影
- 证明卷积实现的等价关系
- 说明CLS与位置编码
- 讨论尺寸边界和CNN比较
我从维度开始讲。假设输入图像在PyTorch里是B、C、H、W,patch大小是Ph乘Pw,并且patch之间不重叠。那么高方向有H除以Ph个位置,宽方向有W除以Pw个位置,总token数N等于这两项相乘。只有在正方形patch时,才可以简写成H乘W除以P平方,不能把宽度那一项误写成W除以W。
得到每个patch后,把Ph乘Pw乘C个像素值展平为一个向量,再乘一个可学习投影矩阵,从Ph乘Pw乘C维映射到模型维度D。所有patch堆起来,形状就是B、N、D。这就是patch embedding的数学本质:先把局部图像块变成向量,再投影成Transformer能处理的token。
工程里通常用一个Conv2d完成这两步。卷积核大小和步幅都设成patch大小,每个窗口正好覆盖一个不重叠patch;每个输出通道的卷积核展平后,就对应线性投影矩阵的一列或一行。卷积输出是B、D、网格高、网格宽,再flatten空间维并transpose成B、N、D。它是常见的等价实现,也容易利用框架优化,但不能不看框架和硬件就宣称一定比unfold加矩阵乘快很多。
原始ViT还会在patch序列前加一个可学习CLS token,再加可学习的位置embedding。经过Transformer编码器后,用CLS位置的表示做分类。这里也要留边界:CLS不是唯一方案,有的视觉Transformer会用全局平均池化;位置也可以用二维、相对或其他形式,取决于架构。
图像尺寸不是patch大小整数倍时,必须显式决定怎么做。可以resize或中心裁剪,也可以padding到可整除,或者使用支持可变分辨率和位置插值的实现。每种方案都可能影响边缘信息和位置分布,训练与推理要保持一致。patch越小,token越多,局部细节更细,但标准attention计算和显存会随token数平方增长;patch越大则更省算力,但可能丢失小目标细节。
最后我不会说self-attention能做全局关系而CNN做不到。更准确的是,标准ViT的一层self-attention就能让任意patch直接交互,而传统CNN通常从局部卷积开始,通过多层堆叠、下采样或更大卷积核逐渐扩大感受野。两者差别在交互路径和归纳偏置。面试里把公式、实现等价、尺寸边界和计算取舍讲完整,比只背“图像像句子”更有说服力。
还要注意位置embedding的长度要和最终token数匹配。固定分辨率训练后换分辨率,常见做法是对二维位置网格插值,但插值并不保证零损失,需要在目标分辨率回归验证。若使用重叠patch或卷积stem,token数公式和简单的不重叠分块也会变化,回答时应先说明当前假设。
关键一句:patch大小同时影响token数量、细节保留和attention二次成本。
核验来源
面试官还可能这样问
- 问法 1 · 场景切入
假设你要做一个电商图片搜索系统,用户上传一张商品图,系统得返回相似款。如果我用ViT做图像编码,那输入图片怎么拆成一个个小图块再变成向量?你具体说说这个过程。
- 问法 2 · 层层追问
ViT把图像当句子处理,那图像是怎么变成token序列的?……分块之后每个patch怎么变成向量?……用卷积怎么实现?……最后维度怎么变化的?
- 问法 3 · 直球架构
请讲一下ViT中图像分块和patch embedding的具体实现:输入图像怎么切块,每个块怎么展平并投影到embedding维度,以及最终如何构造token序列(包括CLS token和位置编码)。