问题与目标
语言 Transformer 接收 Token 序列,而图像是 C × H × W 的像素张量。Vision Transformer 把图像切成固定 Patch,将每个 Patch 投影成向量,再加入位置与类别信息。本篇用张量形状跑通这条转换链路。
核心概念
对 H × W 图像和 P × P Patch,Patch 数量为:
text
N = (H / P) × (W / P)
每个 Patch 原始维度是 C × P × P,经过线性投影得到 D 维视觉 Token。加入一个可学习的 [CLS] 后,编码器输入形状为 B × (N+1) × D。
Patch 越小,空间细节越多,但注意力计算随 Token 数平方增长。位置编码告诉模型 Patch 的空间顺序,否则打乱 Patch 后序列表示缺少位置关系。
可运行实现
bash
python -m pip install "torch>=2.3,<3"
python
import torch
from torch import nn
class PatchEmbedding(nn.Module):
def __init__(self, image_size=32, patch_size=8, channels=3, dim=64):
super().__init__()
self.grid = image_size // patch_size
self.projection = nn.Conv2d(channels, dim, kernel_size=patch_size, stride=patch_size)
self.cls = nn.Parameter(torch.zeros(1, 1, dim))
self.position = nn.Parameter(torch.zeros(1, self.grid**2 + 1, dim))
def forward(self, images):
patches = self.projection(images) # B, D, grid, grid
tokens = patches.flatten(2).transpose(1, 2) # B, N, D
cls = self.cls.expand(images.size(0), -1, -1)
return torch.cat([cls, tokens], dim=1) + self.position
images = torch.randn(2, 3, 32, 32)
tokens = PatchEmbedding()(images)
print(images.shape) # [2, 3, 32, 32]
print(tokens.shape) # [2, 17, 64]
Conv2d 的卷积核与步幅都等于 Patch 大小,因此等价于同时切块和线性投影。这个例子验证输入契约,不宣称已经训练出图像分类器。
常见问题与排查
图像尺寸不能整除 Patch
明确 Resize、Padding 或动态分辨率策略。静默裁掉边缘会丢失证据。
把通道顺序写成 NHWC
PyTorch 默认卷积输入是 NCHW。读取图片后打印 shape,并确认归一化统计对应 RGB 通道。
Patch 越小越准确
更细 Token 会显著增加计算和显存。应在细节召回、输入分辨率与延迟之间实验。
ViT 等于视觉语言模型
ViT 只负责视觉编码。视觉语言模型还需要把视觉特征连接到语言模型,并进行图文训练或对齐。
小结
ViT 的关键转换是像素 → Patch → 视觉 Token。跟踪每一步张量形状和位置编码后,多模态模型中的图像 Token 就不再是抽象概念。
License: CC BY-NC 4.0
Updated 2 hours ago
Was this article helpful? Give it a like.
0 comments


