news 2026/9/7 8:21:47

transformer用于图像分类

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
transformer用于图像分类

计算机视觉与transformer(Vision transformer

1、VIT:

(1)原理:将图像分割成小块,通过线性变换得到patch embedding,加上位置编码,输入到encoder中,最后用分类头进行分类。

(2)代码部分:patch embedding

position encoding

transformer encoder block

Muti-head attention

MLP

最终分类头

  1. 代码步骤:

patch embedding

position encoding

transformer encoder block

Muti-head attention

MLP

最终分类头

  1. 代码步骤:

import torch

from torch import nn

from einops import rearrange#用于张量操作

第一部分(转变为编码)

class PatchEmbedding(nn.Module):

def __init__(self, img_size=32, patch_size=4, in_channels=3, embed_dim=128):

"""

将图像分割为小块并线性投影到嵌入空间

参数:

img_size: 输入图像尺寸(假设为正方形)

patch_size: 每个小块的尺寸

in_channels: 输入通道数(RGB为3)

embed_dim: 嵌入维度

embed_dim 全称 embedding dimension(嵌入维度),表示:

每个图像块(patch)被映射到的向量维度

模型内部特征表示的统一维度(所有 Transformer 层的输入/输出维度)

super().__init__()

self.img_size = img_size

self.patch_size = patch_size#将参数赋予给变量

# 计算分块数量:32/4=8 → 8x8=64个patch

self.num_patches = (img_size // patch_size) ** 2平方就得到了数量

# 定义卷积层代替线性投影(更高效)

self.projection = nn.Conv2d(

in_channels=in_channels,

out_channels=embed_dim,

kernel_size=patch_size,

stride=patch_size

)这里不是很理解为什么要这么做呢,为什么卷积核和步长要这样设置

def forward(self, x):#定义前向传播,将上边的函数用起来

# x形状: (batch_size, 3, 32, 32)

x = self.projection(x)# 投影后形状: (batch_size, 128, 8, 8),这里的128就是转换后的特征维度

x = rearrange(x, 'b c h w -> b (h w) c')

# 展平为序列:形状变为 (batch_size, 64, 128),一直到这里,一张图由原来的3*32*32,变为了64*128,也就是真正的变成了64个向量。

return x

batch_size, 64, 128

第二部分(多头注意力机制)

class MultiHeadAttention(nn.Module):

def __init__(self, embed_dim=128, num_heads=4):

super().__init__()

self.num_heads = num_heads

self.head_dim = embed_dim // num_heads

#计算每个注意力头的维度,128/4=32,每个注意力头处理32维信息

# 合并计算QKV的线性投影

self.qkv = nn.Linear(embed_dim, embed_dim * 3)

#使用一个线性层,输入维度是128,输出是128*3,也就是一个线性层同时生成Q/K/V,然后将输出分割为三部分,这样可以减少计算

self.attention_dropout = nn.Dropout(0.1)

#注意力分数计算完成后使用,随机关闭一些权重。

self.projection = nn.Linear(embed_dim, embed_dim)

#将多头注意力后的结果进行投影,保持维度一致性

前期工作完成,开始使用函数

def forward(self, x):

batch_size, seq_len, embed_dim = x.shape

#前面x的形状是bactch*64*128

# 生成QKV:形状 (batch_size, seq_len, 3*embed_dim)

qkv = self.qkv(x)在这里变成了batch*64*3*128

# 拆分为多头:形状 (3, batch_size, num_heads, seq_len, head_dim)

将原先的batch*64*3*128变成batch*64*3*4*32的拆成3*batch*64*4*32

qkv = rearrange(qkv, 'b s (n h d) -> n b h s d', n=3, h=self.num_heads)#输出为3*batch*4*64*32

q, k, v = qkv[0], qkv[1], qkv[2]

#现在对三个矩阵取位置,得到Q/K/V对应的矩阵形式

# 计算注意力分数

scores = torch.matmul(q, k.transpose(-1, -2)) / (self.head_dim ** 0.5)

#这里是在计算Q与K的点积相似度。Q的维度batch*4*64*32,K的维度也是batch*4*64*32,transpose(-1, -2)将K的最后两个维度,可以计算点积,Scores的形状是batch*4*64*64

attention = torch.softmax(scores, dim=-1)

#将scores转变为概率分布,每个问询变量与key的注意力权重和为1,dim=-1表示在最后一个维度上进行softmax,也就是在64上进行,主要是确保每个问询变量与key的注意力权重和为1。得到的应该是batch*4*64*64个。

attention = self.attention_dropout(attention)

#这里是为了防止过拟合

# 加权求和

x = torch.matmul(attention, v)

#attention的形状是batch*4*64*64,v的形状是batch*4*64*32,进行点积计算,也就是用注意力权重对v进行加权平均,得到新的V值

x = rearrange(x, 'b h s d -> b s (h d)')

# 合并多头,这里X的形状是batch*4*64*32,转变成为batch*64*128,也就是将多头的信息合并了

x = self.projection(x)

#这里将上边合并后的信息进行线性投影,将多头的信息融合

return x

batch*64*128

第三部分(定义前向传播)

class MLP(nn.Module):#开始定义transformer中的前向传播

def __init__(self, embed_dim=128, hidden_dim=512):

super().__init__()

self.net = nn.Sequential(

nn.Linear(embed_dim, hidden_dim),

nn.GELU(),# ViT中常用GELU激活函数

nn.Dropout(0.1),

nn.Linear(hidden_dim, embed_dim),

nn.Dropout(0.1)

)这里定义了一个线性层函数

def forward(self, x):

return self.net(x)#运用前边的函数进行前向传播

batch*64*128

第四部分(定义BLOCK)

class TransformerBlock(nn.Module):#开始定义transformer的编码器块

def __init__(self, embed_dim=128, num_heads=4):

super().__init__()

self.norm1 = nn.LayerNorm(embed_dim)#这里应该是进行了归一化吧

self.attn = MultiHeadAttention(embed_dim, num_heads)#这里用的是前边的类

self.norm2 = nn.LayerNorm(embed_dim)

self.mlp = MLP(embed_dim)#这里用到的是前面的前向传播网络

def forward(self, x):

# 残差连接 + 层归一化(Pre-Norm结构)

x = x + self.attn(self.norm1(x))

x = x + self.mlp(self.norm2(x))

return x

batch*64*128

第五部分(开始使用前面的东西)

class VisionTransformer(nn.Module):

def __init__(self, num_classes=10, depth=4):

super().__init__()

self.patch_embed = PatchEmbedding()#将图片进行编码batch*64*128

embed_dim = 128

# 可学习的分类token([CLS] token)它的作用是汇总整个序列的信息,用于最终的分类任务

self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))

生成一个形状为1*1*128的张量,1表示初始批次大小为1,1表示序列中只有一个可学习的token,embed_dim表示与图像的编码维度一致,也就是使用正态分布进行标准化。nn.Parameter表示该张量是模型的可训练参数。

# 位置编码(可学习参数)

self.pos_embed = nn.Parameter(torch.randn(1, self.patch_embed.num_patches + 1, embed_dim))

#这里为什么是65,是因为还要为前边的cls_token进行编码,这里的1表示所有样本共享同一组位置编码,不随批次变化而变化

# 创建Transformer编码器堆叠

self.blocks = nn.Sequential(*[

TransformerBlock() for _ in range(depth)

])

#depth=4,nn.Sequential(*[block1, block2, block3])等价于nn.Sequential(block1, block2, block3),for _ in range(depth):重复创建 depth 个相同的 TransformerBlock。batch*64*128

# 分类头

self.head = nn.Sequential(

nn.LayerNorm(embed_dim),

nn.Linear(embed_dim, num_classes)

)输出为batch*10,这里的输入的形状是batch*128

定义前向传播,开始使用前边的类

def forward(self, x):

batch_size = x.shape[0]

# 生成patch嵌入

x = self.patch_embed(x) #这个类下边定义的,但用了前边的类,形状 (batch_size, 64, 128)

# 添加分类token

cls_token = self.cls_token.expand(batch_size, -1, -1)#将cls_token的形状变为实际批次的大小

x = torch.cat([cls_token, x], dim=1) # 形状 (batch_size, 65, 128),将cls_token拼接到x中去。

# 添加位置编码

x += self.pos_embed#每个x都在自身编码的基础上加上位置编码

# 通过Transformer编码器

x = self.blocks(x)

# 取出分类token的特征,只有这个是可以用来回归或者使用的特征

cls_token_final = x[:, 0]这个意思就是说在(batch_size, 65, 128)中取出,也就是每65个取出第一个张量,得到(batch_size, 128)

# 分类头

return self.head(cls_token_final)得到最终的分类结果

# 使用示例

if __name__ == "__main__":

# 创建虚拟输入(batch_size=4)

dummy_img = torch.randn(4, 3, 32, 32) # 4张32x32的RGB图片

# 初始化ViT模型

vit = VisionTransformer(num_classes=10)

# 前向传播

output = vit(dummy_img)

print("输出形状:", output.shape) # 应该得到 torch.Size([4, 10]

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

IOPaint:AI图片修复的免费完整指南

IOPaint:AI图片修复的免费完整指南 【免费下载链接】IOPaint Image inpainting tool powered by SOTA AI Model. Remove any unwanted object, defect, people from your pictures or erase and replace(powered by stable diffusion) any thing on your pictures. …

作者头像 李华
网站建设 2026/9/7 8:19:29

双闭环直流调速系统从原理到SIMULINK仿真完整设计指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 8:17:40

T+转U8+工具全解析:历史数据迁移的实操指南与避坑经验

简介:面向用友T与U8数据迁移场景,提供从T账套转换至U8所需的核心程序与配套脚本。包内共13个文件,含8个SQL脚本、2个配置文件、2个OCX控件及1个主程序exe,压缩包仅426KB,轻量易部署。SQL脚本覆盖总账、应收、应付、存货…

作者头像 李华
网站建设 2026/9/7 8:16:31

企业级AI Agent平台搭建实战:从架构设计到落地避坑

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/7 8:14:20

PDFViewer OCX控件从注册到集成:WinForms/VB6 PDF预览实战指南

简介:PDFViewer OCX 控件开发包面向需要在 C#、C、HTML 等环境中集成 PDF 显示与交互功能的 Windows 开发者,解决应用内嵌 PDF 阅读能力的问题。包内共 82 个文件,压缩后约 2.8MB,涵盖 h/cpp/cs 等多语言源码、ocx/dll/exe 控件及…

作者头像 李华