一 本文目的

在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的数据流是

  1. 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指图片数量)。

  2. 进入transformer encoder,这里transformer encoder由L块儿Block构成,每个Block包含LayerNorm、多头注意力以及MLP层。 这里每个Block的输入输出维度是固定的,也就是说,从12块之前到之后的维度不变,一直是B*197*768.

  3. 接下来是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

为什么要这样?

  1. 效率更高(一次 GEMM)

  2. 和原始 Transformer / BERT 实现一致

  3. 后续用 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)

发生了什么?

  1. 保留 CLS

  2. 插入当前层的 prompt

  3. 删掉上一层的 prompt

  4. 保留 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

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐