跳到正文

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. An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设你要做一个电商图片搜索系统,用户上传一张商品图,系统得返回相似款。如果我用ViT做图像编码,那输入图片怎么拆成一个个小图块再变成向量?你具体说说这个过程。

  2. 问法 2 · 层层追问

    ViT把图像当句子处理,那图像是怎么变成token序列的?……分块之后每个patch怎么变成向量?……用卷积怎么实现?……最后维度怎么变化的?

  3. 问法 3 · 直球架构

    请讲一下ViT中图像分块和patch embedding的具体实现:输入图像怎么切块,每个块怎么展平并投影到embedding维度,以及最终如何构造token序列(包括CLS token和位置编码)。

同模块相关题目