news 2026/8/5 8:59:39

从NLP到CV:手把手教你用PyTorch实现ViT图像分类(附完整代码)

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从NLP到CV:手把手教你用PyTorch实现ViT图像分类(附完整代码)

从零构建Vision Transformer:PyTorch实战图像分类与性能深度剖析

最近在复现一些视觉Transformer的经典论文时,我重新把ViT的代码从头敲了一遍。这个过程让我意识到,很多教程虽然提供了代码,但往往忽略了模块设计背后的“为什么”,以及在实际训练中可能遇到的“坑”。今天,我想抛开那些复杂的公式,从一个实践者的角度,和你聊聊如何用PyTorch一步步搭建一个可用的ViT,并分享一些在CIFAR-10和ImageNet子集上对比CNN的心得。如果你已经熟悉PyTorch的基本操作,但对Transformer如何“看懂”图像感到好奇,那么这篇文章正是为你准备的。

我们将不止步于模型的搭建,更会深入到数据准备、训练技巧、可视化分析以及性能对比的完整闭环。你会发现,理解一个模型最好的方式,就是亲手把它“造”出来,然后“用”起来,最后再和别的模型“比一比”

1. 环境准备与核心思想解构

在开始写代码之前,我们需要明确ViT最核心的直觉:将图像视为一个由小块(Patch)组成的序列。这与卷积神经网络(CNN)逐层提取局部特征的思路截然不同。CNN依靠卷积核的滑动来捕捉空间相关性,其归纳偏置(Inductive Bias)假设了图像的局部性和平移不变性。而Transformer最初为序列数据设计,其自注意力机制擅长捕捉长距离依赖关系,但对数据本身的结构没有先验假设。

提示:ViT的这种“无先验”特性是一把双刃剑。在数据量不足时,它可能不如CNN学得快、学得好;但在海量数据上,它摆脱了卷积核大小的限制,能更自由地建立图像任意两个区域之间的联系,潜力巨大。

我们的实验环境基于Python 3.8+和PyTorch 1.12+。建议使用虚拟环境管理依赖。

# 创建并激活虚拟环境(以conda为例) conda create -n vit_tutorial python=3.8 conda activate vit_tutorial # 安装核心依赖 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # 根据CUDA版本调整 pip install matplotlib seaborn pandas scikit-learn tqdm tensorboard

准备好环境后,我们先从最关键的图像分块嵌入(Patch Embedding)模块开始。这是将2D图像转换为1D序列的桥梁。

2. 核心模块的PyTorch实现与细节剖析

2.1 Patch Embedding:从图像到序列的转换

传统Transformer处理的是单词序列,每个单词被映射为一个向量(词嵌入)。对于图像,ViT的做法是将一张图切割成多个固定大小的非重叠小块(例如16x16像素),然后将每个小块展平,并通过一个线性层(或一个卷积层)投影到一个固定的维度D。这个D就是Transformer隐藏层的大小。

为什么用卷积层实现?虽然论文中描述为“线性投影”,但在代码实现上,使用一个卷积核大小和步长都等于patch size的卷积层更为高效。它可以一次性对所有patch进行相同的线性变换。

import torch import torch.nn as nn import torch.nn.functional as F class PatchEmbed(nn.Module): """ 将2D图像转换为Patch Embedding序列。 使用卷积层实现,效率更高。 """ def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.img_size = (img_size, img_size) if isinstance(img_size, int) else img_size self.patch_size = (patch_size, patch_size) if isinstance(patch_size, int) else patch_size self.grid_size = (self.img_size[0] // self.patch_size[0], self.img_size[1] // self.patch_size[1]) self.num_patches = self.grid_size[0] * self.grid_size[1] # 核心:一个卷积层同时完成分块和投影 # kernel_size=patch_size, stride=patch_size 确保了不重叠的切割 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=self.patch_size, stride=self.patch_size) def forward(self, x): B, C, H, W = x.shape # 确保输入图像尺寸符合预期 assert H == self.img_size[0] and W == self.img_size[1], \ f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." # (B, C, H, W) -> (B, embed_dim, H/patch, W/patch) x = self.proj(x) # 展平空间维度 -> (B, embed_dim, num_patches) x = x.flatten(2) # 调整维度顺序,符合Transformer输入: (B, num_patches, embed_dim) x = x.transpose(1, 2) return x

维度变换解析: 假设输入图像为(B, 3, 224, 224)patch_size=16embed_dim=768

  1. 经过self.proj卷积后,输出形状为(B, 768, 14, 14)。因为224/16 = 14
  2. flatten(2)将最后两个维度展平:(B, 768, 14*14)->(B, 768, 196)
  3. transpose(1, 2)交换维度:(B, 196, 768)。现在,我们有196个token,每个token是768维的向量。

2.2 可学习的分类Token与位置编码

Transformer Encoder处理的是这个长度为196的序列。但我们需要一个代表整张图片的向量用于最终的分类。ViT借鉴了BERT的[CLS]token,引入了一个可学习的分类token。这个token会与patch tokens拼接在一起,送入Transformer。经过多层自注意力计算后,这个分类token的输出就聚合了全局信息,用于分类。

同时,由于自注意力机制本身是置换不变(Permutation-invariant)的——打乱输入序列的顺序,输出序列的顺序也会被打乱,但内容不变——模型无法感知patch之间的空间位置关系。因此,我们必须显式地加入位置编码(Positional Encoding)

class VisionTransformer(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4., qkv_bias=True, drop_rate=0., attn_drop_rate=0.): super().__init__() self.num_classes = num_classes self.embed_dim = embed_dim self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches = self.patch_embed.num_patches # 1. 可学习的分类token self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) # 2. 可学习的位置编码 self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) # 3. Transformer Encoder之前的Dropout self.pos_drop = nn.Dropout(p=drop_rate) # 堆叠Transformer Encoder Blocks self.blocks = nn.ModuleList([ Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, drop=drop_rate, attn_drop=attn_drop_rate) for i in range(depth)]) # LayerNorm self.norm = nn.LayerNorm(embed_dim) # 分类头 self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity() # 初始化权重 self._init_weights() def _init_weights(self): # 分类token和位置编码用截断正态分布初始化 nn.init.trunc_normal_(self.cls_token, std=.02) nn.init.trunc_normal_(self.pos_embed, std=.02) # 线性层和卷积层用Xavier初始化,LayerNorm用默认初始化 self.apply(self._init_vit_weights) def _init_vit_weights(self, m): if isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) if m.bias is not None: nn.init.constant_(m.bias, 0) elif isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') elif isinstance(m, nn.LayerNorm): nn.init.constant_(m.bias, 0) nn.init.constant_(m.weight, 1.0) def forward(self, x): B = x.shape[0] # batch size # 获取patch embeddings x = self.patch_embed(x) # (B, num_patches, embed_dim) # 扩展并添加分类token cls_tokens = self.cls_token.expand(B, -1, -1) # (1, 1, D) -> (B, 1, D) x = torch.cat((cls_tokens, x), dim=1) # (B, num_patches+1, embed_dim) # 添加位置编码 x = x + self.pos_embed x = self.pos_drop(x) # 通过Transformer Encoder for blk in self.blocks: x = blk(x) # 取分类token对应的输出 x = self.norm(x) cls_output = x[:, 0] # (B, embed_dim) # 分类 logits = self.head(cls_output) # (B, num_classes) return logits

关键点讨论

  • 分类Token的扩展self.cls_token初始形状是(1, 1, D)。在forward时,我们用.expand(B, -1, -1)将其扩展到batch中每个样本都有一个,这是PyTorch的高效广播操作。
  • 位置编码的可学习性:ViT使用的是可学习的位置编码,即一个与(num_patches+1, embed_dim)形状相同的参数矩阵。这与原始Transformer使用固定的正弦余弦编码不同。在实践中,可学习的位置编码通常表现良好且更简单。
  • 为什么是加法?位置编码通过加法与token嵌入结合。这是一种简单而有效的融合方式,让模型在后续的计算中能同时利用语义信息和位置信息。

2.3 Transformer Encoder Block的实现

这是模型的核心,包含多头自注意力(MSA)和前馈网络(MLP),并伴有残差连接和层归一化。

class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0.): super().__init__() self.norm1 = nn.LayerNorm(dim) self.attn = Attention(dim, num_heads=num_heads, qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop) self.norm2 = nn.LayerNorm(dim) mlp_hidden_dim = int(dim * mlp_ratio) self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, drop=drop) def forward(self, x): # 残差连接1: 注意力模块 x = x + self.attn(self.norm1(x)) # 残差连接2: 前馈网络模块 x = x + self.mlp(self.norm2(x)) return x class Attention(nn.Module): def __init__(self, dim, num_heads=8, qkv_bias=False, attn_drop=0., proj_drop=0.): super().__init__() self.num_heads = num_heads head_dim = dim // num_heads self.scale = head_dim ** -0.5 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_drop) def forward(self, x): B, N, C = x.shape # 生成Q, K, V qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] # 每个形状: (B, num_heads, N, head_dim) # 计算注意力分数 attn = (q @ k.transpose(-2, -1)) * self.scale # (B, num_heads, N, N) attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) # 加权求和 x = (attn @ v).transpose(1, 2).reshape(B, N, C) # (B, N, C) x = self.proj(x) x = self.proj_drop(x) return x class Mlp(nn.Module): def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.): super().__init__() out_features = out_features or in_features hidden_features = hidden_features or in_features self.fc1 = nn.Linear(in_features, hidden_features) self.act = act_layer() self.fc2 = nn.Linear(hidden_features, out_features) self.drop = nn.Dropout(drop) def forward(self, x): x = self.fc1(x) x = self.act(x) x = self.drop(x) x = self.fc2(x) x = self.drop(x) return x

代码细节

  • Pre-Norm vs Post-Norm:ViT采用的是Pre-Norm结构(x + submodule(norm(x))),这与原始Transformer的Post-Norm不同。Pre-Norm通常训练更稳定。
  • 注意力计算优化:通过一个线性层self.qkv同时计算Q、K、V,然后通过reshapepermute拆分成多头,这是常见的优化写法,比分别用三个线性层更高效。
  • 缩放因子self.scale = head_dim ** -0.5用于在点积后缩放注意力分数,防止梯度消失。

3. 训练流程、可视化与实战调优

模型搭建好了,但让它真正工作起来还需要完整的训练循环、数据加载和评估。我们以CIFAR-10数据集为例,因为它规模适中,适合快速实验。

3.1 数据准备与增强策略

对于ViT,尤其是patch size较小(如16)时,图像分辨率需要是patch size的整数倍。CIFAR-10原始为32x32,如果patch size=4,则得到8x8=64个patch,序列长度尚可。但更常见的做法是将图像上采样到224x224,以匹配ImageNet预训练模型的输入尺寸。这里我们展示一个灵活的预处理流程。

import torchvision.transforms as transforms from torchvision.datasets import CIFAR10 from torch.utils.data import DataLoader def get_dataloaders(data_dir='./data', batch_size=128, img_size=224): """ 获取CIFAR-10的训练和测试数据加载器。 包含针对ViT的数据增强。 """ # 训练数据增强:更强,防止过拟合 train_transform = transforms.Compose([ transforms.RandomResizedCrop(img_size, scale=(0.8, 1.0)), # 随机裁剪并缩放 transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet统计量 ]) # 测试/验证集转换:只有中心裁剪和归一化 test_transform = transforms.Compose([ transforms.Resize(int(img_size * 1.05)), # 稍大一点再中心裁剪 transforms.CenterCrop(img_size), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) train_dataset = CIFAR10(root=data_dir, train=True, download=True, transform=train_transform) test_dataset = CIFAR10(root=data_dir, train=False, download=True, transform=test_transform) train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True) return train_loader, test_loader

注意:对于ViT,数据增强尤为重要。因为模型本身缺乏CNN的平移不变性等归纳偏置,更需要通过丰富的数据来学习这些不变性。RandAugment、MixUp、CutMix等高级增强技术对ViT的性能提升非常关键。

3.2 训练循环与优化器配置

ViT的训练通常需要更长的预热(Warmup)和更精细的学习率调度。

import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR import time def train_one_epoch(model, dataloader, criterion, optimizer, scheduler, device, epoch): model.train() running_loss = 0.0 correct = 0 total = 0 for batch_idx, (inputs, targets) in enumerate(dataloader): inputs, targets = inputs.to(device), targets.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, targets) loss.backward() # 可选:梯度裁剪,防止梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step() # 每个step后更新学习率 running_loss += loss.item() _, predicted = outputs.max(1) total += targets.size(0) correct += predicted.eq(targets).sum().item() if batch_idx % 50 == 0: print(f'Epoch: {epoch} | Batch: {batch_idx}/{len(dataloader)} | Loss: {loss.item():.4f}') epoch_loss = running_loss / len(dataloader) epoch_acc = 100. * correct / total return epoch_loss, epoch_acc def validate(model, dataloader, criterion, device): model.eval() running_loss = 0.0 correct = 0 total = 0 with torch.no_grad(): for inputs, targets in dataloader: inputs, targets = inputs.to(device), targets.to(device) outputs = model(inputs) loss = criterion(outputs, targets) running_loss += loss.item() _, predicted = outputs.max(1) total += targets.size(0) correct += predicted.eq(targets).sum().item() val_loss = running_loss / len(dataloader) val_acc = 100. * correct / total return val_loss, val_acc # 配置优化器和调度器 def get_optimizer_scheduler(model, train_loader, epochs, lr=3e-4, warmup_epochs=5): optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=0.05) # AdamW是ViT训练标配 # 1. 线性预热 warmup_scheduler = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=warmup_epochs*len(train_loader)) # 2. 余弦退火 cosine_scheduler = CosineAnnealingLR(optimizer, T_max=(epochs - warmup_epochs)*len(train_loader), eta_min=1e-6) # 组合调度器 scheduler = optim.lr_scheduler.SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[warmup_epochs*len(train_loader)]) return optimizer, scheduler

关键训练技巧

  • 优化器AdamW优于Adam,因为它正确地实现了权重衰减(Weight Decay),这对ViT的稳定训练至关重要。
  • 学习率调度线性预热(Warmup)可以避免训练初期的不稳定。随后接余弦退火(Cosine Annealing)缓慢降低学习率。
  • 梯度裁剪:虽然不总是必要,但在训练深度Transformer时加上梯度裁剪(如max_norm=1.0)是个好习惯。

3.3 注意力可视化:模型在看哪里?

理解ViT如何工作的一大乐趣是可视化其注意力图。我们可以提取中间层的注意力权重,看看模型在分类时关注图像的哪些部分。

import numpy as np import matplotlib.pyplot as plt def visualize_attention(model, img_tensor, patch_size=16, layer_index=-1, head_index=0): """ 可视化指定层和头部的注意力图(针对分类token)。 img_tensor: 形状为 (1, C, H, W) 的单个图像张量 """ model.eval() with torch.no_grad(): # 前向传播,并注册钩子获取注意力权重 attn_weights = [] def hook_fn(module, input, output): # output 是 (attn_matrix, weighted_values) attn_weights.append(output[0].detach().cpu()) # 注册钩子到指定层的注意力模块 target_layer = model.blocks[layer_index].attn handle = target_layer.register_forward_hook(hook_fn) # 前向传播 logits = model(img_tensor) handle.remove() # 移除钩子 # attn_weights[0] 形状: (1, num_heads, num_patches+1, num_patches+1) attn_map = attn_weights[0][0, head_index, 0, 1:] # 取分类token对所有patch tokens的注意力 (num_patches,) # 将注意力权重重塑为2D网格 grid_size = int(np.sqrt(attn_map.shape[0])) attn_map_2d = attn_map.reshape(grid_size, grid_size).numpy() # 可视化 fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 5)) # 原始图像 img = img_tensor[0].cpu().permute(1, 2, 0).numpy() img = img * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406]) # 反归一化 img = np.clip(img, 0, 1) ax1.imshow(img) ax1.set_title('Original Image') ax1.axis('off') # 注意力热力图叠加 im = ax2.imshow(attn_map_2d, cmap='hot', interpolation='nearest') ax2.set_title(f'Attention Map (Layer {layer_index}, Head {head_index})') ax2.axis('off') plt.colorbar(im, ax=ax2, fraction=0.046, pad=0.04) plt.tight_layout() plt.show() return attn_map_2d, logits.argmax().item()

运行这个函数,你可能会发现,浅层的注意力图往往比较分散,关注局部边缘和纹理;而深层的注意力图则更加集中,聚焦于与类别语义相关的关键物体区域。这直观地展示了ViT如何从局部信息整合到全局语义理解。

4. ViT与CNN的实战性能对比与分析

理论说再多,不如跑个实验看看。我们在CIFAR-10数据集上,对比一个轻量级ViT(如ViT-Tiny)和一个经典的CNN(如ResNet-18)的表现。为了公平,我们使用相同的训练策略(数据增强、优化器、调度器、训练轮数)。

实验设置

  • 模型
    • ViT-Tiny:patch_size=4,embed_dim=192,depth=12,num_heads=3(约6M参数)
    • ResNet-18: 标准结构 (约11M参数)
  • 数据集: CIFAR-10 (50k训练,10k测试)
  • 图像尺寸: 上采样至 32x32 (ViT patch=4, 得到64个token) 或 224x224 (使用预训练权重微调)
  • 训练: 100个epoch,AdamW优化器,线性预热+余弦退火,Batch Size=128。

以下是一个简化的对比实验框架和可能的结果分析:

def compare_vit_cifar10(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f'Using device: {device}') # 1. 加载数据 train_loader, test_loader = get_dataloaders(img_size=32, batch_size=128) # 小尺寸快速实验 # 2. 初始化模型 vit_model = VisionTransformer(img_size=32, patch_size=4, in_chans=3, num_classes=10, embed_dim=192, depth=12, num_heads=3, mlp_ratio=4).to(device) cnn_model = torchvision.models.resnet18(num_classes=10).to(device) # 需要调整第一层卷积适应32x32输入 cnn_model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False) # 修改第一层 cnn_model.maxpool = nn.Identity() # 移除初始的maxpool # 3. 训练与评估函数 (略,同上文) # ... 分别训练两个模型 ... # 4. 结果对比表格 results = { 'Model': ['ViT-Tiny', 'ResNet-18'], 'Params (M)': [sum(p.numel() for p in vit_model.parameters())/1e6, sum(p.numel() for p in cnn_model.parameters())/1e6], 'Train Acc (%)': [vit_train_acc, cnn_train_acc], # 训练后获得 'Test Acc (%)': [vit_test_acc, cnn_test_acc], # 训练后获得 'Training Time (s)': [vit_time, cnn_time] } import pandas as pd df = pd.DataFrame(results) print(df.to_markdown(index=False))

假设我们得到如下(模拟)结果:

ModelParams (M)Train Acc (%)Test Acc (%)Training Time (s)
ViT-Tiny5.799.592.1850
ResNet-1811.299.894.3620

结果分析

  1. 小数据集的困境:在仅有5万张图像的CIFAR-10上,ResNet-18的测试精度显著高于ViT-Tiny。这验证了ViT的核心论点:缺乏归纳偏置使得它在中小型数据集上需要更多数据才能达到与CNN相当的性能。CNN的卷积结构天然适合图像,学习效率更高。
  2. 过拟合现象:ViT的训练精度接近100%,但测试精度有约7%的差距,过拟合比ResNet更严重。这说明需要更强的正则化(如DropPath、Label Smoothing)和数据增强。
  3. 训练时间:ViT的训练时间更长。这是因为自注意力机制的计算复杂度与序列长度的平方成正比(O(N²))。尽管我们的序列长度只有64,但Transformer block中的矩阵运算仍然比ResNet的卷积更耗时。
  4. 参数效率:ViT-Tiny参数更少,但性能却更低,这进一步说明在小数据上,参数效率不等于学习效率。

那么,ViT的优势在哪里?

当我们切换到更大的数据集(如ImageNet-1K,120万张图像),并使用大规模预训练时,故事就完全不同了。下表展示了在ImageNet-21K(1400万张图像)上预训练后,再在ImageNet-1K上微调的结果(引用自论文及后续研究):

ModelPretrain DataImageNet Top-1 Acc (%)Throughput (img/s)
ResNet-50ImageNet-1K76.51220
ViT-Base/16ImageNet-21K84.0290
ViT-Large/16ImageNet-21K85.385

可以看到,在大规模预训练数据的加持下,ViT的性能实现了对CNN的显著超越。虽然其推理速度(吞吐量)仍低于高度优化的CNN,但差距在缩小,并且研究界在不断地改进ViT的计算效率(如Swin Transformer的窗口注意力、DeiT的蒸馏技术)。

给实践者的建议

  • 如果你的数据量有限(<100万),优先考虑使用CNN(如ResNet, EfficientNet)或使用在大型数据集上预训练好的ViT进行微调。从头训练一个小型ViT很可能不如CNN。
  • 如果你有海量数据,或者可以获取到大规模的预训练模型(如来自Hugging Facetimm库或官方发布的权重),那么ViT及其变体(Swin, DeiT, BEiT)是值得探索的强大工具,尤其在需要建模长距离依赖的任务上(如高分辨率图像分类、图像分割)。
  • 不要忽视混合模型:将CNN的局部特征提取能力与Transformer的全局建模能力结合(如CoAtNet, ConViT),往往能在效率和性能之间取得更好的平衡。

最后,我想分享一个在调试ViT时遇到的真实问题:位置编码的插值。当你微调一个在224x224图像上预训练的ViT,而你的任务输入是384x384时,直接加载预训练的位置编码会因序列长度不匹配而报错。这时需要对位置编码进行2D插值。很多开源库(如timm)已经内置了这个功能,但自己实现时需要注意。

# 位置编码插值示例(简化) def resize_pos_embed(pos_embed, new_shape, num_extra_tokens=1): """ 调整位置编码的尺寸。 pos_embed: 原始位置编码,形状 (1, N_old+1, D) new_shape: 新图像的网格大小 (H_new//P, W_new//P) num_extra_tokens: 额外的token数量(如分类token),默认为1 """ pos_embed_tokens, pos_embed_grid = pos_embed[:, :num_extra_tokens], pos_embed[0, num_extra_tokens:] old_shape = int(math.sqrt(pos_embed_grid.shape[0])) # 假设原来是正方形网格 pos_embed_grid = pos_embed_grid.reshape(1, old_shape, old_shape, -1).permute(0, 3, 1, 2) pos_embed_grid = F.interpolate(pos_embed_grid, size=new_shape, mode='bicubic', align_corners=False) pos_embed_grid = pos_embed_grid.permute(0, 2, 3, 1).reshape(1, -1, pos_embed.shape[-1]) new_pos_embed = torch.cat([pos_embed_tokens, pos_embed_grid], dim=1) return new_pos_embed

这个细节看似微小,却是在实际部署和微调中经常遇到的“拦路虎”。理解每个模块的输入输出形状,并准备好应对各种尺寸变化,是工程化实现中不可或缺的一环。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/5 8:55:31

3步释放90%内存:让旧电脑秒变新设备的秘密武器

3步释放90%内存&#xff1a;让旧电脑秒变新设备的秘密武器 【免费下载链接】memreduct Lightweight real-time memory management application to monitor and clean system memory on your computer. 项目地址: https://gitcode.com/gh_mirrors/me/memreduct 你是否曾遇…

作者头像 李华
网站建设 2026/8/5 10:54:38

Cobalt Strike后渗透技巧:从WiFi密码获取到屏幕截图的实战演示

Cobalt Strike后渗透实战&#xff1a;从凭证窃取到屏幕监控的深度操作指南 如果你已经跨过了Cobalt Strike的入门门槛&#xff0c;搭建好了环境&#xff0c;也成功让几个目标上线&#xff0c;那么接下来真正考验技术深度的时刻才刚刚开始。后渗透阶段&#xff0c;远不止是看着一…

作者头像 李华
网站建设 2026/8/5 9:30:36

从零到一:使用深度学习项目训练环境,轻松复现CycleGAN图像转换

从零到一&#xff1a;使用深度学习项目训练环境&#xff0c;轻松复现CycleGAN图像转换 你是不是也遇到过这样的情况&#xff1f;在网上看到一个很酷的AI项目&#xff0c;比如能把马变成斑马的CycleGAN&#xff0c;兴致勃勃地下载了代码&#xff0c;结果发现环境配置就卡住了半…

作者头像 李华
网站建设 2026/8/5 16:23:25

PyTorch 2.6效果展示:基于CUDA镜像的GPU加速训练实测

PyTorch 2.6效果展示&#xff1a;基于CUDA镜像的GPU加速训练实测 1. 引言&#xff1a;当深度学习遇见GPU加速 如果你正在学习或使用深度学习&#xff0c;一定对“训练太慢”这个问题深有体会。一个简单的图像分类模型&#xff0c;用CPU跑上几个小时甚至几天都是家常便饭。这种…

作者头像 李华
网站建设 2026/8/5 14:57:08

YOLOv8与PDF-Extract-Kit-1.0联合应用:精准定位PDF文档元素

YOLOv8与PDF-Extract-Kit-1.0联合应用&#xff1a;精准定位PDF文档元素 1. 效果惊艳的开场 如果你曾经尝试过从PDF文档中提取内容&#xff0c;特别是那些包含复杂排版、表格和公式的学术论文或技术文档&#xff0c;你就会知道这有多让人头疼。传统的OCR工具往往把整个页面当作…

作者头像 李华