VPT源码架构 从ViT到VPT
一 本文目的
在VPT之前,这里先讲解ViT架构(基础),但实际上,很多小白了解ViT的架构但并不了解架构的具体实现,也不清楚QKV等具体的维度变化,本文的目标是对ViT的解析,并且学会ViT的对hidden_size的各种数据操作,譬如
x.shape[0]、expand(B, -1, -1)、flatten、transpose(-1, -2)、cat等,完成本人对以前阶段学习的汇总和反思。完成ViT的解析后,了解怎么从ViT到VPT,如何继承并修改ViT的源码,也是对后面自己改别人代码的启示。预计下一坑是从CLIP到CoOp。

二 从ViT的数据流到ViT的详细实现
ViT的数据流是
-
embedding:从图片切割成patch,把patch通过resnet等网络转换成patch_embedding,拼接上CLS token后加上position_embeddings。 维度示例:这里输入是224*224的图片,经过16*16的patch切割成14*14=196个patch,通过ResNet 转换成768维(hidden_size),接下来拼接CLS token,现在的token数量就是196+1,再进行position_embeddings的初始化并加到197维上去,所以输出是B*197*768(B指图片数量)。
-
进入transformer encoder,这里transformer encoder由L块儿Block构成,每个Block包含LayerNorm、多头注意力以及MLP层。 这里每个Block的输入输出维度是固定的,也就是说,从12块之前到之后的维度不变,一直是B*197*768.
-
接下来是MLP Head,输入是把CLS token的对应维度的数据拿来,然后映射到分类目标数量即可,即768 - > num _targets.

接下来详细讲述数据流各部分的详细流程,从代码层面理解:
Embedding层
class Embeddings(nn.Module): """Construct the embeddings from patch, position embeddings. """ # 将图像切成 patch,并加上 CLS token 和位置编码 def __init__(self, config, img_size, in_channels=3): super(Embeddings, self).__init__() self.hybrid = None img_size = _pair(img_size) if config.patches.get("grid") is not None: grid_size = config.patches["grid"] patch_size = (img_size[0] // 16 // grid_size[0], img_size[1] // 16 // grid_size[1]) n_patches = (img_size[0] // 16) * (img_size[1] // 16) self.hybrid = True else: patch_size = _pair(config.patches["size"]) n_patches = (img_size[0] // patch_size[0]) * (img_size[1] // patch_size[1]) self.hybrid = False if self.hybrid: self.hybrid_model = ResNetV2(block_units=config.resnet.num_layers, width_factor=config.resnet.width_factor) in_channels = self.hybrid_model.width * 16 self.patch_embeddings = Conv2d(in_channels=in_channels, out_channels=config.hidden_size, kernel_size=patch_size, stride=patch_size) self.position_embeddings = nn.Parameter(torch.zeros(1, n_patches+1, config.hidden_size)) self.cls_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size)) self.dropout = Dropout(config.transformer["dropout_rate"]) def forward(self, x): # 生成 ViT 的 token 序列:[CLS] + patch tokens + position embedding # 将卷积输出的 patch 特征 (B, C, H, W) 拉平成 patch 序列 (B, N, C) B = x.shape[0] cls_tokens = self.cls_token.expand(B, -1, -1) if self.hybrid: x = self.hybrid_model(x) x = self.patch_embeddings(x) x = x.flatten(2) x = x.transpose(-1, -2) x = torch.cat((cls_tokens, x), dim=1) embeddings = x + self.position_embeddings embeddings = self.dropout(embeddings) return embeddings
这里有config.patches.get("grid") self.hybrid 控制图片是不是需要ResNet来提供 inductive bias,也就是在切patch之前先引入局部性\平移不变性\边缘纹理先验.
通过conv2d这个卷积层来完成patch embedding的转换,输出通道数即为对应的hidden_size. 输入[B, 3, 224, 224],输出[B, 768, 14, 14].初始化cls_token position_embeddings
接下来是embedding的forward,
下面介绍一些常用的维度操作( 后面关于维度的操作也会形成文档 ).
首先对于shape[0](shape函数的结果是维度列表),得到一共有多少张图片,即 B = batch size.
expand 的规则是 -1:这一维不变 正数:扩展到指定大小 不拷贝内存,只改 view 注意区分repeat,repeat会复制内存的!
flatten(2) 从第 2 维开始,全部压平[B, 768, 14, 14]->[B, 768, 196]
transpose(-1, -2)交换倒数第一维和倒数第二维得到 [B,196,768]
torch.cat((cls_tokens, x), dim=1) 意思是从dim=1的维度上拼接,也就是得到从cls_tokens: [B, 1, 768] x: [B, 196, 768]得到[B, 197, 768]
最后利用广播机制,position_embeddings 形状[1, 197, 768]相加的时候广播[B, 197, 768]给每个token加上一个“我在第几个位置”的信息. 常规的使用dropout正则化,防止过拟合.
Encoder层的基础构件
encoder包含若干个Block块 每个Block块分别包含 MLP\Layernorm\MHA 接下来分别解释
MLP
class Mlp(nn.Module): def __init__(self, config): # Transformer Block 中的前馈网络(Feed-Forward Network) super(Mlp, self).__init__() self.fc1 = Linear(config.hidden_size, config.transformer["mlp_dim"]) self.fc2 = Linear(config.transformer["mlp_dim"], config.hidden_size) self.act_fn = ACT2FN["gelu"] self.dropout = Dropout(config.transformer["dropout_rate"]) self._init_weights() def _init_weights(self): nn.init.xavier_uniform_(self.fc1.weight) nn.init.xavier_uniform_(self.fc2.weight) nn.init.normal_(self.fc1.bias, std=1e-6) nn.init.normal_(self.fc2.bias, std=1e-6) def forward(self, x): x = self.fc1(x) x = self.act_fn(x) x = self.dropout(x) x = self.fc2(x) x = self.dropout(x) return x
MLP(x)=W2⋅GELU(W1x+b1)+b2
这里不过多解释,代码比较简单,这里添加一点初始化的知识
全连接 + GELU这样的配置 常用xavier_uniform_初始化
针对于Relu的激活函数,基本使用He initialization,pytorch也是使用kaiming 初始化卷积层参数的.
(B, N, hidden_size) @ (hidden_size, mlp_dim)
→ (B, N, mlp_dim) 在第一个线性层一般先升维,第二个降维.
Multi-Head Self-Attention
一、Attention 的整体作用
输入:hidden_states: (B, N, hidden_size) 输出:attention_output: (B, N, hidden_size)
核心流程:
hidden_states → Q, K, V 线性映射 768->768 → reshape 成多头 768/12 -> 64 → QKᵀ / √d 12 12 → softmax → 加权求和 V 12 64 → 拼回 hidden_size 12*68 -> 768 → Linear + Dropout
二、init:注意力模块的参数定义
class Attention(nn.Module): def __init__(self, config, vis): super(Attention, self).__init__()
1️⃣ vis:是否保存注意力权重(可视化用)self.vis = vis
-
vis=True:返回 attention map(画热力图) -
vis=False:只训练,不保存权重
2️⃣ 多头注意力的核心参数
self.num_attention_heads = config.transformer["num_heads"] self.attention_head_size = int(config.hidden_size / self.num_attention_heads) self.all_head_size = self.num_attention_heads * self.attention_head_size
假设:
-
hidden_size = 768 -
num_heads = 12
那么:
表格 还在加载中,请等待加载完成后再尝试复制
关键思想:多头不是“多个小 Linear”,而是 一个大 Linear + reshape
三、Q / K / V 的 Linear(重点)
self.query = Linear(config.hidden_size, self.all_head_size) self.key = Linear(config.hidden_size, self.all_head_size) self.value = Linear(config.hidden_size, self.all_head_size)
❗ 非常重要的一点
❌ 不是
Linear(768, 64)此处64指的是每个头的维度 ✅ 而是Linear(768, 768),一次性算出所有 head 的 Q / K / V
为什么要这样?
-
效率更高(一次 GEMM)
-
和原始 Transformer / BERT 实现一致
-
后续用
view + permute切分 head
四、输出映射 & Dropout
self.out = Linear(config.hidden_size, config.hidden_size) self.attn_dropout = Dropout(config.transformer["attention_dropout_rate"]) self.proj_dropout = Dropout(config.transformer["attention_dropout_rate"]) self.softmax = Softmax(dim=-1)
-
out:把拼回的多头结果再线性映射一次 -
attn_dropout:作用在注意力权重上 -
proj_dropout:作用在最终输出上
五、!!!transpose_for_scores:最关键的 shape 操作
def transpose_for_scores(self, x):
输入 x 的 shape = (B, N, hidden) = (B, N, 768)我们现在要做的事情只有一个:把 768 拆成 (num_heads, head_dim)
1️⃣ 计算新 shape(拆 head)
new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)
等价于:
(B, N) + (num_heads, head_dim) → (B, N, num_heads, head_dim)
2️⃣ view:重解释内存(不拷贝)
x = x.view(*new_x_shape)
此时:
(B, N, 12, 64)
3️⃣ permute:换维度顺序
return x.permute(0, 2, 1, 3)
维度变化:
(B, N, num_heads, head_dim)→(B, num_heads, N, head_dim)
因为后面要做:
Q @ Kᵀ 而 PyTorch 的 batch matmul 要求:(B, H, N, D) @ (B, H, D, N) padmask
六、!!! forward:前向传播逐行拆解
输入
def forward(self, hidden_states):
hidden_states.shape = (B, N, hidden_size)
1️⃣ 计算 Q / K / V
mixed_query_layer = self.query(hidden_states) mixed_key_layer = self.key(hidden_states) mixed_value_layer = self.value(hidden_states)
shape:
(B, N, 768)
2️⃣ 拆成多头
query_layer = self.transpose_for_scores(mixed_query_layer) key_layer = self.transpose_for_scores(mixed_key_layer) value_layer = self.transpose_for_scores(mixed_value_layer)
shape 统一为:
(B, num_heads, N, head_dim)
3️⃣ 计算注意力分数 QKᵀ
attention_scores = torch.matmul( query_layer, key_layer.transpose(-1, -2) )
关键 transpose(-1, -2)key_layer: (B, H, N, D)→ transpose(B, H, D, N)
matmul 后:(B, H, N, N)
含义:
每个 head、每个 query token 对所有 key token 的相似度
4️⃣ 缩放(防止数值爆炸)
attention_scores = attention_scores / math.sqrt(self.attention_head_size)
这是 Scaled Dot-Product Attention 的核心公式。/d 此处d指hidden_size
5️⃣ Softmax → 概率
attention_probs = self.softmax(attention_scores)
shape 不变:
(B, H, N_query, N_key)
每一行和为 1,表示“关注分布”
6️⃣ 是否保存权重(可视化)
weights = attention_probs if self.vis else None
7️⃣ Dropout(注意力层)
attention_probs = self.attn_dropout(attention_probs)
8️⃣ 加权求和 Value
context_layer = torch.matmul(attention_probs, value_layer)
shape:
(B, H, N, D)
每个 token 得到一个新的表示
9️⃣ permute:把 head 维度换回
context_layer = context_layer.permute(0, 2, 1, 3).contiguous()
(B, H, N, D)→(B, N, H, D)
为什么要 .contiguous()?
-
permute只改 stride,不改内存 -
view()要求内存连续 -
所以必须
.contiguous()
⚠️ 这是很多人 debug 的坑点
🔟 拼回 hidden_size
new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,) context_layer = context_layer.view(*new_context_layer_shape)
(B, N, H, D)→(B, N, H×D)→(B, N, 768)
1️⃣1️⃣ 输出映射
attention_output = self.out(context_layer) #这里又加了一个线性层 增加非线性能力 attention_output = self.proj_dropout(attention_output)
最终输出:
(B, N, hidden_size)
1️⃣2️⃣ 返回结果
return attention_output, weights
七、一句话总结 Attention 的“本质”
Attention 做的只有一件事: 用当前 token(Q)去“问”所有 token(K)然后“按重要性加权汇总信息”(V)
八、现在研究 VPT 的直接关系
-
VPT 的 Prompt token
-
会参与 Q / K / V
-
影响 attention_scores
-
-
Attention 是 Prompt 最有效的注入点
-
MLP 只是“自己加工自己”
这就是为什么:Prompt 调 Attention,比调 MLP 有用得多
九、你现在已经到了什么层次?
你已经可以:
✅ 看懂 Attention 的 数学公式 ✅ 对照代码解释 每一次 view / permute 的目的 ✅ 理解 为什么一定是 (B, H, N, D)
LayerNorm
Block块
class Block(nn.Module): def __init__(self, config, vis): super(Block, self).__init__() self.hidden_size = config.hidden_size self.attention_norm = LayerNorm(config.hidden_size, eps=1e-6) self.ffn_norm = LayerNorm(config.hidden_size, eps=1e-6) self.ffn = Mlp(config) self.attn = Attention(config, vis) def forward(self, x): h = x x = self.attention_norm(x) # Attention 之前的层归一化 x, weights = self.attn(x) x = x + h h = x x = self.ffn_norm(x) # MLP 之前的层归一化 x = self.ffn(x) x = x + h return x, weights def load_from(self, weights, n_block): ROOT = f"Transformer/encoderblock_{n_block}" with torch.no_grad(): query_weight = np2th(weights[pjoin(ROOT, ATTENTION_Q, "kernel")]).view(self.hidden_size, self.hidden_size).t() key_weight = np2th(weights[pjoin(ROOT, ATTENTION_K, "kernel")]).view(self.hidden_size, self.hidden_size).t() value_weight = np2th(weights[pjoin(ROOT, ATTENTION_V, "kernel")]).view(self.hidden_size, self.hidden_size).t() out_weight = np2th(weights[pjoin(ROOT, ATTENTION_OUT, "kernel")]).view(self.hidden_size, self.hidden_size).t() query_bias = np2th(weights[pjoin(ROOT, ATTENTION_Q, "bias")]).view(-1) key_bias = np2th(weights[pjoin(ROOT, ATTENTION_K, "bias")]).view(-1) value_bias = np2th(weights[pjoin(ROOT, ATTENTION_V, "bias")]).view(-1) out_bias = np2th(weights[pjoin(ROOT, ATTENTION_OUT, "bias")]).view(-1) self.attn.query.weight.copy_(query_weight) self.attn.key.weight.copy_(key_weight) self.attn.value.weight.copy_(value_weight) self.attn.out.weight.copy_(out_weight) self.attn.query.bias.copy_(query_bias) self.attn.key.bias.copy_(key_bias) self.attn.value.bias.copy_(value_bias) self.attn.out.bias.copy_(out_bias) mlp_weight_0 = np2th(weights[pjoin(ROOT, FC_0, "kernel")]).t() mlp_weight_1 = np2th(weights[pjoin(ROOT, FC_1, "kernel")]).t() mlp_bias_0 = np2th(weights[pjoin(ROOT, FC_0, "bias")]).t() mlp_bias_1 = np2th(weights[pjoin(ROOT, FC_1, "bias")]).t() self.ffn.fc1.weight.copy_(mlp_weight_0) self.ffn.fc2.weight.copy_(mlp_weight_1) self.ffn.fc1.bias.copy_(mlp_bias_0) self.ffn.fc2.bias.copy_(mlp_bias_1) self.attention_norm.weight.copy_(np2th(weights[pjoin(ROOT, ATTENTION_NORM, "scale")])) self.attention_norm.bias.copy_(np2th(weights[pjoin(ROOT, ATTENTION_NORM, "bias")])) self.ffn_norm.weight.copy_(np2th(weights[pjoin(ROOT, MLP_NORM, "scale")])) self.ffn_norm.bias.copy_(np2th(weights[pjoin(ROOT, MLP_NORM, "bias")]))
$$\begin{aligned} x' &= x + \text{MHSA}(\text{LN}(x)) \\ x'' &= x' + \text{MLP}(\text{LN}(x')) \end{aligned}$$
有了前面的MLP/LN/MHA的基础,这些就不难了,重点需要看的是 load_from(),目前我还不知道其他源码的撰写方法. 不过可以知道的是load_from和resume_from都是mmdetection库中的函数,用于加载模型文件。分别用于 加载完整的模型文件 和 继续加载上一次未完成的训练。
Encoder
class Encoder(nn.Module): def __init__(self, config, vis): super(Encoder, self).__init__() self.vis = vis self.layer = nn.ModuleList() self.encoder_norm = LayerNorm(config.hidden_size, eps=1e-6) for _ in range(config.transformer["num_layers"]): layer = Block(config, vis) self.layer.append(copy.deepcopy(layer)) def forward(self, hidden_states): attn_weights = [] for layer_block in self.layer: hidden_states, weights = layer_block(hidden_states) if self.vis: attn_weights.append(weights) encoded = self.encoder_norm(hidden_states) return encoded, attn_weights def forward_cls_layerwise(self, hidden_states): # hidden_states: B, 1+n_patches, dim # 提取每一层的 [CLS] token if hidden_states.size(0) != 1: raise ValueError('not support batch-wise cls forward yet') cls_embeds = [] cls_embeds.append(hidden_states[0][0]) for i,layer_block in enumerate(self.layer): hidden_states, _ = layer_block(hidden_states) if i < len(self.layer)-1: cls_embeds.append(hidden_states[0][0]) encoded = self.encoder_norm(hidden_states) cls_embeds.append(hidden_states[0][0]) cls_embeds = torch.stack(cls_embeds) # 12, dim return cls_embeds
一、Encoder 在 ViT 里的角色
一句话:Encoder = N 个 Transformer Block 顺序堆叠 + 最后一层 LayerNorm
二、init:Encoder 的结构定义
class Encoder(nn.Module): def __init__(self, config, vis): super(Encoder, self).__init__() self.vis = vis self.layer = nn.ModuleList() self.encoder_norm = LayerNorm(config.hidden_size, eps=1e-6)
1️⃣ self.vis
-
控制 是否保存 attention map
-
用于:可视化 分析 attention
-
训练时通常
False
2️⃣ self.layer = nn.ModuleList()
为什么是 ModuleList,而不是普通 list?
因为:
-
PyTorch 只有 注册到 Module 的子模块 才能:
-
被
.to(device) -
被
.parameters() -
被保存 / 加载权重
-
📌 ModuleList 是:
“我有一堆结构一样、参数不同的子模块”
3️⃣ self.encoder_norm
self.encoder_norm = LayerNorm(config.hidden_size)
这是 整个 Encoder 最后的 LayerNorm。
📌 对应 ViT 论文中的:
“A final LayerNorm is applied after the Transformer encoder”
4️⃣ 构建 N 个 Block
for _ in range(config.transformer["num_layers"]): layer = Block(config, vis) self.layer.append(copy.deepcopy(layer))
为什么要 copy.deepcopy?
因为:
-
每个 Block 结构一样
-
但 参数必须完全独立
如果不用 deepcopy:
-
所有层会共享一套参数 ❌(灾难)
维度小结(初始化阶段)
假设:
-
num_layers = 12 -
hidden_size = 768
那么:
Encoder ├── Block 0 ├── Block 1 ├── ... ├── Block 11 └── LayerNorm
三、forward:标准 Encoder 前向传播
def forward(self, hidden_states):
输入
hidden_states: (B, N, hidden)
-
B:batch size -
N:token 数(1 + patch 数) -
hidden:768(ViT-B)
1️⃣ attention 权重容器
attn_weights = []
只在 vis=True 时才真正用到。
2️⃣ 逐层通过 Block
for layer_block in self.layer: hidden_states, weights = layer_block(hidden_states)
关键点:Block 不改变维度每一层都是:(B, N, hidden) → (B, N, hidden)
3️⃣ 是否保存 attention map
if self.vis: attn_weights.append(weights)
此时:
attn_weights: List[ (B, num_heads, N, N) ] # 长度 = num_layers
4️⃣ Encoder 最后的 LayerNorm
encoded = self.encoder_norm(hidden_states)
📌 这是 Pre-LN Transformer 的最后一步稳定化
5️⃣ 返回
return encoded, attn_weights
返回内容
-
encoded:(B, N, hidden) -
attn_weights:(num_layers, B, num_heads, N, N) 或 []
四、forward_cls_layerwise:标注重点函数
这个函数 不是标准 ViT forward,而是一个:
分析 / 研究用接口:提取每一层的 CLS token 表示
1️⃣ 输入检查(为什么只支持 batch=1)
if hidden_states.size(0) != 1: raise ValueError('not support batch-wise cls forward yet')
原因很简单:
后面代码直接写死了:
hidden_states[0][0]
-
第一个
[0]:batch index -
第二个
[0]:CLS token
如果 batch > 1,就要写成循环,作者偷懒了 😄
2️⃣ cls_embeds 是什么?
cls_embeds = [] cls_embeds.append(hidden_states[0][0])
此时:
-
hidden_states[0]:
(N, hidden)
-
hidden_states[0][0]:
(hidden,) # CLS token
📌 这是 Embedding + Position 后,还没进任何 Block 的 CLS
3️⃣ 逐层跑 Block,并收集 CLS
for i, layer_block in enumerate(self.layer): hidden_states, _ = layer_block(hidden_states) if i < len(self.layer)-1: cls_embeds.append(hidden_states[0][0])
逻辑拆解:
-
每经过一层 Block:
-
CLS token 会被 Attention + MLP 更新
-
-
把 中间层 CLS 表示存下来
注意这个条件:
if i < len(self.layer)-1:
👉 最后一层 CLS 留给 encoder_norm 之后再存(但实际代码没存)
4️⃣ 最后的 encoder_norm + CLS
encoded = self.encoder_norm(hidden_states) cls_embeds.append(hidden_states[0][0])
现在你得到了:
-
Embedding 后 CLS
-
Block1 后 CLS
-
Block2 后 CLS
-
...
-
BlockN + LN 后 CLS
5️⃣ torch.stack
cls_embeds = torch.stack(cls_embeds) # 12, dim
假设 ViT-B/16:
cls_embeds.shape = (12, 768)
📌 这是一个 “CLS 表示随层数演化的轨迹”
五、这个函数通常用来干嘛?
非常重要,这里是研究级用法:
1️⃣ 观察 CLS 表示是否逐层“语义化”
-
浅层:纹理、边缘
-
深层:语义、类别
2️⃣ VPT / Prompt / Adapter 论文常用
比如:
-
Prompt 是否只影响前几层?
-
CLS 在哪一层开始分离?
3️⃣ 可用于:
-
线性探针(layer-wise probing)
-
表征分析
-
可视化(t-SNE / PCA)
六、Encoder 的一句话总结
Encoder 是由多个 Transformer Block 顺序堆叠而成的模块,每个 Block 通过多头自注意力和前馈网络逐步更新 token 表示,最后通过 LayerNorm 输出稳定的特征表示。
Transformer
class Transformer(nn.Module): def __init__(self, config, img_size, vis): super(Transformer, self).__init__() self.embeddings = Embeddings(config, img_size=img_size) self.encoder = Encoder(config, vis) def forward(self, input_ids): embedding_output = self.embeddings(input_ids) encoded, attn_weights = self.encoder(embedding_output) return encoded, attn_weights def forward_cls_layerwise(self, input_ids): embedding_output = self.embeddings(input_ids) cls_embeds = self.encoder.forward_cls_layerwise(embedding_output) return cls_embeds
简单 无需多言,forward的输入是需要批处理的图片,输出是[B, 197(Patch + cls), 768(hidden_size)]
ViT
好,这一段已经是 完整的 ViT 分类模型封装 了,我们继续保持前面的节奏: 逐层结构 → forward 数据流 → load_from 预训练权重(重点)
一、整体定位:VisionTransformer 是什么?
class VisionTransformer(nn.Module): # Vision Transformer 分类模型: # Patch Embedding + Transformer Encoder + 分类头
一句话总结:
VisionTransformer = ViT backbone + 分类头(Linear)
结构层级关系是:
VisionTransformer ├── Transformer │ ├── Embeddings │ └── Encoder (N 个 Block) └── Head (Linear 分类层)
二、init:模型初始化
def __init__(self, model_type, img_size=224, num_classes=21843, vis=False):
1️⃣ model_type
config = CONFIGS[model_type]
-
决定:
-
hidden_size(768 / 1024)
-
num_layers(12 / 24)
-
num_heads
-
patch_size
-
-
本质是 ViT-B / ViT-L / Hybrid ViT 的配置表
2️⃣ num_classes
self.num_classes = num_classes self.classifier = config.classifier
-
num_classes > 0→ 分类任务 -
num_classes == 0→ 只做特征提取(常见于迁移学习 / VPT)
3️⃣ Transformer 主体
self.transformer = Transformer(config, img_size, vis)
你前面已经完整分析过:
Transformer ├── Embeddings └── Encoder
4️⃣ 分类头(Head)
self.head = Linear(config.hidden_size, num_classes) if num_classes > 0 else nn.Identity()
👉 非常关键的一行
-
分类任务:
CLS token → Linear → logits
-
非分类任务(VPT / Adapter):
CLS token → Identity(不动)
三、forward:标准分类前向传播
def forward(self, x, vis=False): x, attn_weights = self.transformer(x) logits = self.head(x[:, 0])
我们逐步拆开。
Step 1:Transformer Backbone
x, attn_weights = self.transformer(x)
输出:
x: [B, 1 + N_patches, hidden_dim]
Step 2:取 CLS token
x[:, 0]
📌 语义:
ViT 里 CLS token 被当作全局图像表示
shape:
[B, hidden_dim]
Step 3:分类头
logits = self.head(x[:, 0])
-
如果
num_classes = C:
logits: [B, C]
Step 4:是否返回 attention map
if not vis: return logits return logits, attn_weights
attn_weights 结构
你注释写得非常准确:
attn_weights: [num_layers, B, num_heads, num_patches, num_patches]
常用于:Attention 可视化 Prompt / Patch 分析 可解释性论文
四、forward_cls_layerwise:逐层 CLS(重点)
def forward_cls_layerwise(self, x): cls_embeds = self.transformer.forward_cls_layerwise(x) return cls_embeds
你已经看过底层实现,这里只强调用途:
输出:
cls_embeds: [num_layers + 1, hidden_dim]
语义:
表格 还在加载中,请等待加载完成后再尝试复制
📌 VPT / Adapter / Prompt learning 必备接口
五、load_from:加载预训练权重(最容易踩坑的地方)
这部分非常工程化,我们拆成 5 个子阶段。
① Patch Embedding 权重
self.transformer.embeddings.patch_embeddings.weight.copy_( np2th(weights["embedding/kernel"], conv=True) ) self.transformer.embeddings.patch_embeddings.bias.copy_( np2th(weights["embedding/bias"]) )
把 JAX / numpy 的 Conv patch embedding 转成 PyTorch 的 Conv2d
② CLS token & Encoder Norm
self.transformer.embeddings.cls_token.copy_(np2th(weights["cls"])) self.transformer.encoder.encoder_norm.weight.copy_( np2th(weights["Transformer/encoder_norm/scale"]) ) self.transformer.encoder.encoder_norm.bias.copy_( np2th(weights["Transformer/encoder_norm/bias"]) )
📌 注意:
-
CLS 是 可训练参数
-
Encoder 最后有一个 LayerNorm
③ 位置编码(🔥重点🔥)
posemb = np2th(weights["Transformer/posembed_input/pos_embedding"]) posemb_new = self.transformer.embeddings.position_embeddings
情况 1:尺寸一致(最理想)
if posemb.size() == posemb_new.size(): self.transformer.embeddings.position_embeddings.copy_(posemb)
情况 2:尺寸不一致(最常见)
比如:
-
预训练:224×224 → 14×14 patch
-
微调:384×384 → 24×24 patch
处理流程:
1. 拆 CLS token 2. 还原为 2D grid 3. 双线性插值 resize 4. 再拼回 CLS
代码核心:
posemb_grid = posemb_grid.reshape(gs_old, gs_old, -1) posemb_grid = ndimage.zoom(posemb_grid, zoom, order=1) posemb_grid = posemb_grid.reshape(1, gs_new * gs_new, -1) posemb = np.concatenate([posemb_tok, posemb_grid], axis=1)
📌 这是 ViT 迁移学习最经典的一段代码
④ Transformer Blocks 权重
for bname, block in self.transformer.encoder.named_children(): for uname, unit in block.named_children(): unit.load_from(weights, n_block=uname)
这里会调用你之前分析过的:
Block.load_from(...)
加载内容包括:
-
Q / K / V / Out
-
MLP FC1 / FC2
-
LayerNorm 权重
⑤ Hybrid ViT(ResNet + ViT)
if self.transformer.embeddings.hybrid:
👉 Hybrid ViT = CNN stem + Transformer
会额外加载:
-
ResNet root conv
-
GroupNorm
-
每个 ResNet block 的权重
📌 如果你用的是 ViT-B/16(纯 ViT),这段不会执行。
六、最终一句话总结
VisionTransformer是完整的 ViT 分类模型封装,包含 Transformer 主干和分类头。forward用于标准分类任务,基于 CLS token 预测类别;forward_cls_layerwise提供逐层 CLS 表示,支持 Prompt Learning 和特征分析;load_from实现从 JAX 预训练权重到 PyTorch 的完整迁移,并支持位置编码自适应 resize。
好了!我们已经学会ViT了!(这里没有花过多笔墨描述将ResNet应用到patch embedding之前的过程,实际上如果用ViT-B/16(纯 ViT)不会用到ResNet,所以感兴趣的同学可以自行了解)
让我们接着来看如何把ViT改成VPT吧~
三 从VPT的构建到完整的训练流程
class PromptedTransformer(Transformer): # 在标准 ViT Transformer 的基础上,引入可学习的 Prompt Token,并将其插入到输入 token 序列中实现参数高效微调。 def __init__(self, prompt_config, config, img_size, vis): # 初始化 Prompt Token 的数量、维度、投影方式以及其参数初始化策略。 assert prompt_config.LOCATION == "prepend" assert prompt_config.INITIATION == "random" assert prompt_config.NUM_DEEP_LAYERS is None assert not prompt_config.DEEP_SHARED super(PromptedTransformer, self).__init__( config, img_size, vis) self.prompt_config = prompt_config self.vit_config = config img_size = _pair(img_size) patch_size = _pair(config.patches["size"]) num_tokens = self.prompt_config.NUM_TOKENS self.num_tokens = num_tokens # number of prompted tokens self.prompt_dropout = Dropout(self.prompt_config.DROPOUT) # if project the prompt embeddings if self.prompt_config.PROJECT > -1: # only for prepend / add prompt_dim = self.prompt_config.PROJECT self.prompt_proj = nn.Linear( prompt_dim, config.hidden_size) nn.init.kaiming_normal_( self.prompt_proj.weight, a=0, mode='fan_out') # ??? else: prompt_dim = config.hidden_size self.prompt_proj = nn.Identity() # initiate prompt(把prompt加进去): if self.prompt_config.INITIATION == "random": val = math.sqrt(6. / float(3 * reduce(mul, patch_size, 1) + prompt_dim)) # noqa self.prompt_embeddings = nn.Parameter(torch.zeros( 1, num_tokens, prompt_dim)) # xavier_uniform initialization 是 xavier uniform初始化,防止 prompt 一开始“压死” patch 特征 nn.init.uniform_(self.prompt_embeddings.data, -val, val) if self.prompt_config.DEEP: # noqa total_d_layer = config.transformer["num_layers"]-1 self.deep_prompt_embeddings = nn.Parameter(torch.zeros( total_d_layer, num_tokens, prompt_dim)) # xavier_uniform initialization nn.init.uniform_(self.deep_prompt_embeddings.data, -val, val) else: raise ValueError("Other initiation scheme is not supported") def incorporate_prompt(self, x): # 将可学习的 Prompt Token 插入到 CLS token 与 Patch token 之间,形成新的 Transformer 输入序列。 # combine prompt embeddings with image-patch embeddings B = x.shape[0] # after CLS token, all before image patches # 用的是父类ViT 的 embedding函数: x → patch embedding → 加 CLS → 加 position embedding ) x = self.embeddings(x) # (batch_size, 1 + n_patches, hidden_dim) # x[:,:,:]的这三个维度分别是什么?一个token512维不应该很长吗?怎么3维就写完了呢 x = torch.cat(( x[:, :1, :], self.prompt_dropout(self.prompt_proj(self.prompt_embeddings).expand(B, -1, -1)), x[:, 1:, :] ), dim=1) # (batch_size, cls_token + n_prompt + n_patches, hidden_dim) # 这地方为什么加进去的是prompt_dropout return x def train(self, mode=True): # 重写训练模式,使 backbone 编码器冻结,仅允许 prompt 相关模块参与训练。 # set train status for this class: disable all but the prompt-related modules if mode: # training: self.encoder.eval() self.embeddings.eval() self.prompt_proj.train() self.prompt_dropout.train() else: # eval: for module in self.children(): module.train(mode) def forward_deep_prompt(self, embedding_output): # 在每一层 Transformer 中动态插入 Deep Prompt Token,实现层级级别的提示微调。 attn_weights = [] hidden_states = None weights = None B = embedding_output.shape[0] num_layers = self.vit_config.transformer["num_layers"] for i in range(num_layers): if i == 0: hidden_states, weights = self.encoder.layer[i](embedding_output) else: if i <= self.deep_prompt_embeddings.shape[0]: deep_prompt_emb = self.prompt_dropout(self.prompt_proj( self.deep_prompt_embeddings[i-1]).expand(B, -1, -1)) hidden_states = torch.cat(( hidden_states[:, :1, :], deep_prompt_emb, hidden_states[:, (1+self.num_tokens):, :] # 去掉旧 prompt 后的 patch tokens ), dim=1) hidden_states, weights = self.encoder.layer[i](hidden_states) if self.encoder.vis: attn_weights.append(weights) encoded = self.encoder.encoder_norm(hidden_states) return encoded, attn_weights def forward(self, x): # 完成 Prompt 插入后的 Transformer 前向传播,支持 shallow prompt 与 deep prompt 两种模式。 # this is the default version: embedding_output = self.incorporate_prompt(x) if self.prompt_config.DEEP: encoded, attn_weights = self.forward_deep_prompt( embedding_output) else: encoded, attn_weights = self.encoder(embedding_output) return encoded, attn_weights
VPT的transformer构建(基于ViT transformer)
一、PromptedTransformer 在整个 ViT / VPT 里的位置
先给你一个全局定位,不然后面容易迷路。
原始 ViT
Image → Embeddings(patch + cls + pos) → Encoder (L 个 Transformer Block) → CLS → 分类
VPT(这里的 PromptedTransformer)
Image → Embeddings → 【插入 Prompt Tokens】 → Encoder(冻结) → 只训练 Prompt
📌 核心思想:
不改 ViT 主干参数,只通过“可学习 token”影响注意力分布
二、类定义与约束(为什么一上来 assert)
class PromptedTransformer(Transformer):
它是 Transformer 的子类 不是重写 ViT,而是“在 ViT 上加东西”
1️⃣ 这些 assert 在干嘛?
assert prompt_config.LOCATION == "prepend" assert prompt_config.INITIATION == "random" assert prompt_config.NUM_DEEP_LAYERS is None assert not prompt_config.DEEP_SHARED
这是在限制实验设置,确保:
表格 还在加载中,请等待加载完成后再尝试复制
📌 这说明作者只实现/验证了这一种 VPT setting
2️⃣ 调用父类 Transformer
super(PromptedTransformer, self).__init__(config, img_size, vis)
这一步做了什么?
创建了:
-
self.embeddings -
self.encoder
PromptedTransformer ≠ 重写 ViT,而是“包在外面”
三、Prompt Token 的核心参数
num_tokens = self.prompt_config.NUM_TOKENS self.num_tokens = num_tokens
这就是 prompt token 数量(如 5 / 10 / 20)
每个 prompt token:
shape = [hidden_dim]
Prompt Dropout
self.prompt_dropout = Dropout(self.prompt_config.DROPOUT)
非常重要的一点:
Prompt 是“外挂信号”,如果不正则化,极容易过拟合
所以:
-
patch token:正常 dropout
-
prompt token:单独 dropout
四、Prompt Projection
if self.prompt_config.PROJECT > -1:
为什么要 projection?
两种设计:
情况 A:prompt_dim == hidden_size
prompt: [*, *, 768] 直接用
情况 B:prompt_dim < hidden_size
prompt: [*, *, 128] → Linear → [*, *, 768]
减少 prompt 参数量(VPT-small)
这个初始化
nn.init.kaiming_normal_(self.prompt_proj.weight, a=0, mode='fan_out')
为什么不是 Xavier?
-
这是一个 Linear + 后接 attention
-
kaiming 对深层网络更稳
-
只用于 projection,不是 prompt 本体
Prompt 本体还是 Xavier-like(下面你会看到)
五、!Prompt Embedding 的初始化
val = math.sqrt(6. / float(3 * reduce(mul, patch_size, 1) + prompt_dim))
这行在干嘛?
它是 Xavier Uniform 的变体
原始 Xavier:
sqrt(6 / (fan_in + fan_out))
这里的 fan_in:
3 * patch_size_h * patch_size_w
👉 等价于: 一个 patch 的输入像素数
📌 含义非常深刻:
Prompt 的尺度 ≈ Patch embedding 的尺度 👉 防止 prompt 一开始“压死”图像特征
Prompt 参数本体
self.prompt_embeddings = nn.Parameter( torch.zeros(1, num_tokens, prompt_dim) )
为什么是 3 维?
[1, num_prompt, dim]
-
1:方便 batch expand -
num_prompt:prompt token 数量 -
dim:token embedding 维度
初始化方式
nn.init.uniform_(self.prompt_embeddings.data, -val, val)
这是 VPT 成功的关键之一
六、incorporate_prompt:Prompt 真正插进去的地方(🔥🔥🔥)
def incorporate_prompt(self, x):
1️⃣ x = self.embeddings(x) 这一步非常重要
x = self.embeddings(x)
此时:
x.shape = [B, 1 + N, hidden_dim]
表格 还在加载中,请等待加载完成后再尝试复制
2️⃣ 非常关键的问题
x[:,:,:] 的三个维度是什么? 一个 token 768 维不应该很长吗?
回答这个“本质问题”:
768 是 embedding 维度,不是 token 数量
所以:
x[B, token_index, embedding_dim]
这就是 Transformer 的基本数据结构。
3️⃣ expand(B, -1, -1) 再解释一次(结合上下文)
self.prompt_embeddings.expand(B, -1, -1)
原始:
[1, num_prompt, dim]
expand 后:
[B, num_prompt, dim]
📌 不是复制参数,只是广播视图
4️⃣ 拼接顺序(非常重要)
torch.cat(( x[:, :1, :], # CLS prompt_tokens, # Prompt x[:, 1:, :] # Patch ), dim=1)
结果:
[B, 1 + num_prompt + num_patches, dim]
📌 Prompt 插在 CLS 和 patch 之间
5️⃣ 为什么 prompt 要 dropout?
self.prompt_dropout(...)
原因一句话总结:
Prompt 是“软指令”,不是图像本体,必须正则化
否则会:快速 overfit 骗 attention 只看 prompt
七、train():冻结 ViT,只训练 Prompt(🔥VPT 精髓🔥)
def train(self, mode=True):
训练时:
self.encoder.eval() self.embeddings.eval()
❌ ViT backbone 不训练
self.prompt_proj.train() self.prompt_dropout.train()
✅ 只训练 prompt
📌 这就是 参数高效微调(PEFT)
八、! Deep Prompt(层级 Prompt)的 forward
def forward_deep_prompt(self, embedding_output): # 在每一层 Transformer 中动态插入 Deep Prompt Token,实现层级级别的提示微调。 attn_weights = [] hidden_states = None weights = None B = embedding_output.shape[0] num_layers = self.vit_config.transformer["num_layers"] for i in range(num_layers): if i == 0: hidden_states, weights = self.encoder.layer[i](embedding_output) else: if i <= self.deep_prompt_embeddings.shape[0]: deep_prompt_emb = self.prompt_dropout(self.prompt_proj( self.deep_prompt_embeddings[i-1]).expand(B, -1, -1)) hidden_states = torch.cat(( hidden_states[:, :1, :], deep_prompt_emb, hidden_states[:, (1+self.num_tokens):, :] # 去掉旧 prompt 后的 patch tokens ), dim=1) hidden_states, weights = self.encoder.layer[i](hidden_states) if self.encoder.vis: attn_weights.append(weights) encoded = self.encoder.encoder_norm(hidden_states) return encoded, attn_weights
Deep Prompt 的思想:不只在输入层加 prompt而是在 每一层Transformer之前都加
为什么要把第一层单独拿出来呢?这里是因为前面会调用incorporate_prompt 相当于把第一个block块之前的 embedding之后的 hidden_states加完了。
之后则是deep_prompt_emb =self.prompt_dropout(self.prompt_proj(self.deep_prompt_embeddings[i-1]).expand(B, -1, -1)) 这段代码应用了expand ,即 prompt 对 batch 内所有样本共享
核心逻辑(重点看这一段)
hidden_states = torch.cat(( hidden_states[:, :1, :], # CLS deep_prompt_emb, # 新 prompt hidden_states[:, (1+self.num_tokens):, :] ), dim=1)
发生了什么?
-
保留 CLS
-
插入当前层的 prompt
-
删掉上一层的 prompt
-
保留 patch tokens
📌 prompt 是“层内使用,不跨层累积”
为什么不累积?如果每层都累积:token 数量会爆炸 attention失控
九、整体一句话总结
PromptedTransformer 在标准 ViT 的 embedding 与 encoder 之间引入可学习的 Prompt Token,通过 prepend 的方式影响 Transformer 的注意力计算。在训练过程中冻结 ViT 主干,仅更新 prompt 参数,实现参数高效微调(VPT)。该实现同时支持 shallow prompt(输入层)与 deep prompt(逐层插入)两种形式。
VPT封装
一、PromptedVisionTransformer
class PromptedVisionTransformer(VisionTransformer): # 将带 Prompt 的 Transformer 封装为完整 ViT 分类模型,并使用 CLS token 进行最终分类。 def __init__( self, prompt_cfg, model_type, img_size=224, num_classes=21843, vis=False ): assert prompt_cfg.VIT_POOL_TYPE == "original" super(PromptedVisionTransformer, self).__init__( model_type, img_size, num_classes, vis) if prompt_cfg is None: raise ValueError("prompt_cfg cannot be None if using PromptedVisionTransformer") self.prompt_cfg = prompt_cfg vit_cfg = CONFIGS[model_type] self.transformer = PromptedTransformer( prompt_cfg, vit_cfg, img_size, vis) def forward(self, x, vis=False): # 使用 CLS token 的输出作为图像全局表示并完成分类预测。 x, attn_weights = self.transformer(x) x = x[:, 0] logits = self.head(x) if not vis: return logits return logits, attn_weights
这是“完整模型封装层” 前面的 PromptedTransformer 是骨干网络改造 这里是:真正拿来做分类的 VPT
二、类的角色(一句话定位)
PromptedVisionTransformer =Prompted Transformer + CLS Pooling + Classification Head
三、init:模型结构构建
def __init__( self, prompt_cfg, model_type, img_size=224, num_classes=21843, vis=False ):
参数逐个解释:
表格 还在加载中,请等待加载完成后再尝试复制
1️⃣ 断言 pool 方式
assert prompt_cfg.VIT_POOL_TYPE == "original"
什么意思?
ViT 有几种 pooling:
-
"original"→ 用 CLS token -
"avg"→ 平均 patch -
"prompt"→ prompt pooling(有些论文)
📌 VPT 原论文只用 CLS token 所以这里强制要求。
2️⃣ 初始化原始 VisionTransformer
super(PromptedVisionTransformer, self).__init__( model_type, img_size, num_classes, vis)
这一步做了什么?
在 原始 ViT 中,这一步会初始化:
-
self.transformer(⚠️马上会被覆盖) -
self.head(分类头) -
embedding / encoder / CLS token / pos embedding
3️⃣ prompt_cfg 合法性检查
if prompt_cfg is None: raise ValueError(...)
防止:用户误用普通 ViT forward 却没传 prompt 配置
4️⃣ 保存配置
self.prompt_cfg = prompt_cfg vit_cfg = CONFIGS[model_type]
-
vit_cfg:标准 ViT 的结构配置-
hidden_size
-
num_layers
-
num_heads
-
patch_size
-
5️⃣ 核心:替换 transformer
self.transformer = PromptedTransformer( prompt_cfg, vit_cfg, img_size, vis)
🔥 这是最关键的一行
你要注意:
❗ 原始
VisionTransformer里已经有一个self.transformer❗ 这里 直接覆盖成 PromptedTransformer
结果是:
VisionTransformer ├── embeddings ├── encoder ├── head └── transformer ← PromptedTransformer(带 prompt)
📌 所以:
forward 时自动走 prompt 逻辑,而不是原始 ViT
四、forward:真正的推理路径
def forward(self, x, vis=False):
1️⃣ Transformer forward(带 prompt)
x, attn_weights = self.transformer(x)
输出的 x 是:
(B, 1 + prompt + patch, hidden_dim)
2️⃣ 取 CLS token
x = x[:, 0]
👉 标准 ViT 做法:
-
CLS 被训练成 全局语义表示
-
prompt / patch 不直接用于分类
3️⃣ 分类头
logits = self.head(x)
-
self.head是一个 Linear:
(hidden_dim → num_classes)
4️⃣ 是否返回 attention
if not vis: return logits return logits, attn_weights
用于:
-
训练 / 测试:只要 logits
-
可视化:看 attention + prompt 影响
五、总结整个 PromptedVisionTransformer
这是一个“对外接口完全等价于 ViT”,但在内部用 prompt 进行参数高效调制的分类模型。
-
不改 head
-
不改 loss
-
不改训练 pipeline
-
只动 prompt
六、你现在已经走到哪一步了?
你已经能:
-
看懂 VPT / Deep Prompt 的完整 forward
-
理解 prompt 如何进入 attention
-
知道 prompt 在“结构上”如何替代 finetuning
更多推荐
所有评论(0)