资讯动态

告别CNN依赖:用PyTorch从零实现ViT模型,手把手教你理解‘图像即词’的Transformer思想

发布时间:2026/8/14 18:44:42 来源:尧图企业网站定制
从零构建ViT模型用PyTorch拆解Transformer的视觉革命当一张224×224的RGB图像被切割成196个16×16的小方块时每个方块突然变成了类似NLP中的单词。这种将图像视为视觉句子的大胆想法正是Vision TransformerViT颠覆计算机视觉领域的起点。本文将带你用PyTorch从第一行代码开始亲手搭建这个抛弃CNN依赖的革新架构。1. 视觉Transformer的设计哲学传统卷积神经网络CNN通过局部感受野、权重共享等归纳偏置inductive bias处理图像而ViT的核心突破在于完全摒弃这些视觉专属假设。就像人类阅读文章时不会预先假设某个单词必须与相邻单词相关一样ViT让模型自己学习图像块patch间的所有关系。关键设计对比特性CNNViT处理单元像素局部邻域图像块全局关系空间信息处理卷积核自动捕获位置编码显式注入计算复杂度O(HWk²C)O(N²D)数据依赖性小数据表现良好需要大规模预训练其中N是图像块数量D是嵌入维度k是卷积核尺寸。ViT的全局注意力机制虽然计算量随序列长度平方增长但在现代GPU上利用矩阵运算优势实际效率可能优于CNN的层次化计算。# 图像分块示例假设输入为3x224x224图像 def split_to_patches(image, patch_size16): # image shape: (C, H, W) patches image.unfold(1, patch_size, patch_size)\ .unfold(2, patch_size, patch_size) return patches.contiguous().view(3, -1, patch_size, patch_size)这段代码展示了如何将图像分解为16×16的块。实际ViT实现中我们会用更高效的线性投影代替这种显式分块# 实际使用的线性投影层 patch_embed nn.Conv2d(3, embed_dim, kernel_sizepatch_size, stridepatch_size)2. 核心组件实现详解2.1 图像到序列的魔法Patch EmbeddingViT的第一项创新是将2D图像转换为1D序列。不同于NLP中现成的单词图像需要经过特殊处理分块展平将H×W×C的图像划分为N个P×P×C的块线性投影用全连接层将每个P²C维的块映射到D维空间位置编码添加可学习的位置嵌入保留空间信息class PatchEmbedding(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size (img_size, img_size) self.patch_size (patch_size, patch_size) self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): B, C, H, W x.shape assert H self.img_size[0], f输入图像高度应为{self.img_size[0]} assert W self.img_size[1], f输入图像宽度应为{self.img_size[1]} x self.proj(x) # (B, D, H/P, W/P) x x.flatten(2).transpose(1, 2) # (B, N, D) return x提示使用卷积实现分块投影比手动切片更高效这也是许多实现中的优化技巧2.2 Transformer编码器实现ViT的编码器与原始Transformer几乎相同但有两个关键调整前置层归一化将LayerNorm放在注意力层和MLP之前Pre-LN残差连接每个子层后都保留残差连接class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., dropout0.): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn MultiHeadAttention(dim, num_heads, dropout) self.norm2 nn.LayerNorm(dim) self.mlp MLP(dim, int(dim * mlp_ratio), dropout) def forward(self, x): # 注意力子层 x x self.attn(self.norm1(x)) # MLP子层 x x self.mlp(self.norm2(x)) return x多头注意力实现细节class MultiHeadAttention(nn.Module): def __init__(self, dim, num_heads8, dropout0.): super().__init__() assert dim % num_heads 0, dim必须能被num_heads整除 self.num_heads num_heads self.scale (dim // num_heads) ** -0.5 self.qkv nn.Linear(dim, dim * 3) self.proj nn.Linear(dim, dim) self.dropout nn.Dropout(dropout) def forward(self, x): B, N, D x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, D // self.num_heads) q, k, v qkv.permute(2, 0, 3, 1, 4) attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.dropout(attn) x (attn v).transpose(1, 2).reshape(B, N, D) x self.proj(x) return x3. 完整ViT模型组装将各个组件组合成完整模型时需要注意几个关键设计选择分类令牌CLS Token借鉴BERT的[CLS]标记聚合全局信息位置编码使用可学习的1D位置嵌入而非正弦编码MLP头预训练和微调阶段使用不同的分类头class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4., dropout0.1): super().__init__() self.patch_embed PatchEmbedding(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(dropout) self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth)]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) def forward(self, x): B x.shape[0] x self.patch_embed(x) # (B, N, D) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat((cls_tokens, x), dim1) # (B, N1, D) x x self.pos_embed x self.pos_drop(x) for blk in self.blocks: x blk(x) x self.norm(x) return self.head(x[:, 0]) # 仅使用CLS token进行分类4. 训练技巧与实战建议4.1 数据增强策略ViT对数据增强极为敏感推荐组合使用以下方法RandAugment自动选择最佳增强组合MixUp图像混合增强CutMix区域替换增强随机擦除模拟遮挡场景from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.RandAugment(), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]), ])4.2 学习率调度ViT训练通常采用带热身的余弦退火调度def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps): def lr_lambda(current_step): if current_step num_warmup_steps: return float(current_step) / float(max(1, num_warmup_steps)) progress float(current_step - num_warmup_steps) / \ float(max(1, num_training_steps - num_warmup_steps)) return 0.5 * (1.0 math.cos(math.pi * progress)) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)4.3 梯度裁剪与精度训练由于ViT的深度结构训练时需要特别注意梯度问题# 混合精度训练示例 scaler torch.cuda.amp.GradScaler() for inputs, targets in dataloader: optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(inputs) loss criterion(outputs, targets) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update()在ImageNet上训练基础版ViTViT-B/16的典型配置超参数值批次大小512基础学习率3e-3热身步数10,000权重衰减0.3Dropout0.1训练周期3005. 模型变体与性能优化5.1 高效ViT架构原始ViT计算复杂度随图像分辨率平方增长以下是几种改进方案金字塔结构PVT引入下采样构建多尺度特征滑动窗口Swin局部注意力降低计算量混合架构HybridCNNTransformer组合# 混合架构示例 class HybridViT(nn.Module): def __init__(self): super().__init__() self.cnn_backbone torchvision.models.resnet50(pretrainedTrue) self.patch_embed nn.Conv2d(2048, embed_dim, kernel_size1) self.transformer TransformerBlocks(...) def forward(self, x): x self.cnn_backbone.conv1(x) x self.cnn_backbone.layer1(x) # ... 更多CNN层 x self.patch_embed(x) # 后续与ViT相同5.2 知识蒸馏技巧使用CNN教师模型指导ViT训练能显著提升小数据集表现class DistillationWrapper(nn.Module): def __init__(self, student, teacher): super().__init__() self.student student self.teacher teacher self.teacher.eval() # 固定教师模型 def forward(self, x): student_logits self.student(x) with torch.no_grad(): teacher_logits self.teacher(x) # 计算KL散度损失 loss F.kl_div( F.log_softmax(student_logits / T, dim1), F.softmax(teacher_logits / T, dim1), reductionbatchmean) * T * T return loss5.3 部署优化建议实际部署时可考虑以下优化TensorRT加速转换模型为FP16或INT8格式剪枝量化移除不重要的注意力头ONNX导出实现跨平台部署# ONNX导出示例 dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export(model, dummy_input, vit.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}})在构建ViT模型时最令人惊讶的发现是当移除所有视觉特有的归纳偏置后模型反而能从数据中学习到更通用的特征表示。这种少即是多的哲学正是Transformer在多个领域展现惊人能力的深层原因。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价