微调640x640图像的自定义ViT,双三次插值处理位置嵌入可行吗?
ViT位置嵌入尺寸不匹配的双三次插值方案验证
问题背景
预训练的ViT-base-patch16-384模型针对384×384图像训练,对应的位置嵌入包含1个CLS token嵌入 + 24×24个补丁嵌入,总长度为577;而自定义模型使用640×640图像,对应1个CLS token嵌入 + 40×40个补丁嵌入,总长度为1601,加载预训练权重时出现如下尺寸不匹配错误:
size mismatch for pos_embed: copying a param with shape torch.Size([1,577,768]) from checkpoint, the shape in current model is torch.Size([1,1601,768])
你尝试用双三次插值解决该问题,但给出的代码存在逻辑错误,无法正确完成位置嵌入的适配。
原代码的问题
原代码直接对包含CLS token的整个位置嵌入序列做插值:
import torch a = torch.rand(1, 577, 768) # 模拟预训练pos_embed a_temp = a.unsqueeze(0) # 转为4D b = torch.nn.functional.interpolate(a_temp, [1601,768], mode='bicubic') b = torch.squeeze(b,0)
这种做法会把CLS token的嵌入和图像补丁的嵌入混在一起插值,破坏CLS token的预训练语义——CLS token是ViT中负责全局特征聚合的特殊token,其位置嵌入不需要适配图像尺寸,应该直接保留预训练的原始权重,仅对图像补丁部分的位置嵌入做插值适配新的网格尺寸。
正确的双三次插值实现
正确的做法是拆分CLS token和补丁位置嵌入,仅对补丁部分做2D网格插值,再重新拼接:
import torch # 模拟预训练的位置嵌入:shape [1, 577, 768] pretrained_pos_embed = torch.rand(1, 577, 768) # 1. 拆分CLS token嵌入和图像补丁的位置嵌入 cls_embed = pretrained_pos_embed[:, 0:1, :] # shape [1, 1, 768] patch_pos_embed = pretrained_pos_embed[:, 1:, :] # shape [1, 576, 768] # 2. 将补丁位置嵌入重塑为2D网格,并调整维度顺序以适配插值要求 patch_pos_embed = patch_pos_embed.reshape(1, 24, 24, 768) # 对应24×24的补丁网格 patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2) # 转为 [1, 768, 24, 24],通道在前 # 3. 双三次插值到目标网格尺寸40×40 new_patch_pos_embed = torch.nn.functional.interpolate( patch_pos_embed, size=(40, 40), mode='bicubic', align_corners=False ) # 4. 恢复维度顺序并展平,得到新的补丁位置嵌入 new_patch_pos_embed = new_patch_pos_embed.permute(0, 2, 3, 1).flatten(1, 2) # shape [1, 1600, 768] # 5. 拼接CLS token嵌入和新的补丁位置嵌入,得到最终适配的位置嵌入 new_pos_embed = torch.cat([cls_embed, new_patch_pos_embed], dim=1) # shape [1, 1601, 768]
关键说明
- 仅对图像补丁的位置嵌入做插值:因为补丁的位置嵌入对应图像的空间网格,插值可以平滑适配新的网格密度;
- 保留CLS token的原始嵌入:CLS token不对应图像的空间位置,其预训练权重已经学习到全局特征聚合的能力,不需要修改;
- 插值时设置
align_corners=False:这是ViT适配不同尺寸时的常用做法,能避免边界处的特征畸变。
内容的提问来源于stack exchange,提问作者Preetom Saha Arko
相关产品推荐
相关产品推荐

