news 2026/7/22 22:51:59

[深度学习]Vision Transformer

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
[深度学习]Vision Transformer

Pytorch实现Vision Transformer

importtorchimporttorch.nnasnnclassPatchEmbedding(nn.Module):def__init__(self,img_size=224,patch_size=16,in_channels=3,embed_dim=768):super().__init__()self.img_size=img_size self.patch_size=patch_size self.n_patches=(img_size//patch_size)**2# 使用卷积层实现patch分割和嵌入self.proj=nn.Conv2d(in_channels=in_channels,out_channels=embed_dim,kernel_size=patch_size,stride=patch_size)defforward(self,x):# 输入x形状: [batch_size, in_channels, img_size, img_size]# 输出形状: [batch_size, n_patches, embed_dim]x=self.proj(x)# [batch_size, embed_dim, n_patches^0.5, n_patches^0.5]x=x.flatten(2)# [batch_size, embed_dim, n_patches]x=x.transpose(1,2)# [batch_size, n_patches, embed_dim]returnxclassPositionEmbedding(nn.Module):def__init__(self,n_patches,embed_dim,dropout=0.1):super().__init__()self.pos_embed=nn.Parameter(torch.zeros(1,n_patches+1,embed_dim))# +1 for class tokenself.dropout=nn.Dropout(dropout)defforward(self,x):# x形状: [batch_size, n_patches+1, embed_dim]x=x+self.pos_embed# 添加位置编码x=self.dropout(x)returnxclassMultiHeadAttention(nn.Module):def__init__(self,embed_dim,num_heads,dropout=0.1):super().__init__()self.embed_dim=embed_dim self.num_heads=num_heads self.head_dim=embed_dim//num_headsassertself.head_dim*num_heads==embed_dim,"Embedding dimension must be divisible by number of heads"self.qkv=nn.Linear(embed_dim,embed_dim*3)# 同时计算Q,K,Vself.attn_dropout=nn.Dropout(dropout)self.proj=nn.Linear(embed_dim,embed_dim)self.proj_dropout=nn.Dropout(dropout)self.scale=self.head_dim**-0.5defforward(self,x):batch_size,n_patches,embed_dim=x.shape# 计算Q,K,V [batch_size, n_patches, num_heads, head_dim]qkv=self.qkv(x).reshape(batch_size,n_patches,3,self.num_heads,self.head_dim).permute(2,0,3,1,4)q,k,v=qkv[0],qkv[1],qkv[2]# 计算注意力分数 [batch_size, num_heads, n_patches, n_patches]attn=(q @ k.transpose(-2,-1))*self.scale attn=attn.softmax(dim=-1)attn=self.attn_dropout(attn)# 应用注意力权重到V上 [batch_size, num_heads, n_patches, head_dim]out=attn @ v out=out.transpose(1,2).reshape(batch_size,n_patches,embed_dim)# 线性投影和dropoutout=self.proj(out)out=self.proj_dropout(out)returnoutclassMLP(nn.Module):def__init__(self,in_features,hidden_features,out_features,dropout=0.1):super().__init__()self.fc1=nn.Linear(in_features,hidden_features)self.act=nn.GELU()self.fc2=nn.Linear(hidden_features,out_features)self.dropout=nn.Dropout(dropout)defforward(self,x):x=self.fc1(x)x=self.act(x)x=self.dropout(x)x=self.fc2(x)x=self.dropout(x)returnxclassTransformerBlock(nn.Module):def__init__(self,embed_dim,num_heads,mlp_ratio=4,dropout=0.1):super().__init__()self.norm1=nn.LayerNorm(embed_dim)self.attn=MultiHeadAttention(embed_dim,num_heads,dropout)self.norm2=nn.LayerNorm(embed_dim)self.mlp=MLP(in_features=embed_dim,hidden_features=embed_dim*mlp_ratio,out_features=embed_dim,dropout=dropout)defforward(self,x):# 残差连接和层归一化x=x+self.attn(self.norm1(x))x=x+self.mlp(self.norm2(x))returnxclassVisionTransformer(nn.Module):def__init__(self,img_size=224,patch_size=16,in_channels=3,n_classes=1000,embed_dim=768,depth=12,num_heads=12,mlp_ratio=4,dropout=0.1):super().__init__()self.patch_embed=PatchEmbedding(img_size,patch_size,in_channels,embed_dim)n_patches=self.patch_embed.n_patches# 分类token和位置编码self.cls_token=nn.Parameter(torch.zeros(1,1,embed_dim))self.pos_embed=PositionEmbedding(n_patches,embed_dim,dropout)# Transformer编码器self.blocks=nn.Sequential(*[TransformerBlock(embed_dim,num_heads,mlp_ratio,dropout)for_inrange(depth)])# 分类头self.norm=nn.LayerNorm(embed_dim)self.head=nn.Linear(embed_dim,n_classes)# 初始化权重nn.init.trunc_normal_(self.cls_token,std=0.02)defforward(self,x):batch_size=x.shape[0]# 生成patch嵌入x=self.patch_embed(x)# [batch_size, n_patches, embed_dim]# 添加class tokencls_token=self.cls_token.expand(batch_size,-1,-1)x=torch.cat([cls_token,x],dim=1)# [batch_size, n_patches+1, embed_dim]# 添加位置编码x=self.pos_embed(x)# 通过Transformer编码器x=self.blocks(x)# 分类x=self.norm(x)cls_token_final=x[:,0]# 只取class token对应的输出x=self.head(cls_token_final)returnxif__name__=='__main__':x=torch.rand(1,3,224,224)model=VisionTransformer(img_size=224,patch_size=16,)y=model(x)print('y.shape = ',y.shape)print(y)

参考资料

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

C++之【深入理解Vector】三部曲之二

前言:我们已经理解了vector的初始化和迭代器初始化,那么接下来要继续深入理解vector,它是如何扩容的,空间及数据个数是如何存储的。 vector空间增长问题 容量空间接口说明size获取数据个数capacity获取容量大小empty判断是否为空…

作者头像 李华
网站建设 2026/7/22 0:38:45

港科校友|李铭鸿,李泓曦:一脉相承

以信任和爱作为家庭的基石,校友李铭鸿Thomas和儿子李泓曦Conan先后踏上科大的教育之路,体现了大学一直培养的探索精神与独特个性。Conan全心投入本科学习,而父母灌输给他的自由、幸福和相互尊重的价值观继续引导着他,展示了科大一…

作者头像 李华
网站建设 2026/7/20 20:52:56

ava面试速成版,背这份八股文(含答案)就对了!

别再拿旧资料瞎准备了!看看我们这份联合2025-2026届成功入职头部企业的12位准大厂人,深挖近3个月一线互联网、科技公司的真实面经反馈、核心考察重点,把大厂面试官的提问逻辑、评分标准、高频考点全拆解,耗时打磨出这份「最新大厂…

作者头像 李华
网站建设 2026/7/19 12:12:18

NAS的大内存有必要吗?到底需不需要 SSD 缓存?核心逻辑一次讲清

NAS的大内存有必要吗?到底需不需要 SSD 缓存?核心逻辑一次讲清 哈喽小伙伴们好,我是Stark-C~ 前段时间有个粉丝在我的推荐下入手了极空间Z4Pro ,当时的好价仅需两千出头,确实挺划算的,只不过到手的是8GB内…

作者头像 李华
网站建设 2026/7/21 19:16:15

橙色工作汇报PPT模板

扫描下载文档详情页: https://www.didaidea.com/wenku/16415.html

作者头像 李华
网站建设 2026/7/19 17:00:43

本地搭建 Clawdbot + ZeroNews 访问

最近,一个名为 ClawdBot(现已更名 OpenClaw) 的项目在技术圈引起了广泛讨论。许多人称其为“真正能做实事的 AI”、“个人 AI 助理的未来形态”。它不仅仅是一个聊天机器人,更是一个能够接入日常工作、生活,直接在用户…

作者头像 李华