重生学AI第十六集:线性层nn.Linear

nn.Linear

1.背景知识

今天学习nn.Linear(线性层),也叫全连接层,我们在卷积层通过卷积操作获取图像特征,然后在线性层,将图像特征转化为最终的输出结果,在图像分类任务中,就是每个类别的打分,后续在输出层,根据打分来输出每个类别的概率。

比如说我们的目标任务是将一些图片分为十个类别,那么out_features 就是10

参数

  • in_features :输入数据的特征数
  • out_features :输出维度 / 输出类别数
  • bias :偏置

先来看一组代码

python 复制代码
import torchvision
from torch.utils.data import DataLoader
from torchvision import transforms

datasets = torchvision.datasets.CIFAR10("../dataset",train=False,download=True,transform=transforms.ToTensor())
dataloader = DataLoader(datasets,batch_size=64)



for data in dataloader:
    img, target = data
    print(img.shape)

控制台输出是这样的

到这里我们要先理解一下这个张量数据:(64,4,32,32), 通过前面的学习我们已经知道这个张量数据的形状是(batch_size,channels,height,width),宽32,代表一行有32个像素,每个像素代表不同的颜色,高也是一样,那么这张图就有32*32个像素,也可以把他们看成是数据,如下图所示,下面这张图是一个5*5的格子,你也可以看成是宽=5,高=5的一个图像,按照张量的写法就是,(1,1,5,5)

好,我们已经学会理解一张图像了,再说回前面那个数据,(64,3,32,32)代表这张图一个通道有32*32个数据,也就是1024个数据,一共有3个通道,1024*3=3072,就是说一张图有3072个数据

2.展平的多种方式

那么,我们就要通过展平的方式,将这个图像的张量数据转化成形状为(batch_size,features)的数据,以便我们将数据送入线性层,in_features和out_features的形状都是如此,那么展平的方式有哪些呢?

1.reshape

python 复制代码
output1 = img.reshape(img.size(0),-1)

2.view

python 复制代码
output1 = img.view(img.size(0),-1)

3.flatten

python 复制代码
output1 = torch.flatten(img,start_dim=1)

参数:

  • start_dim: 展平开始的维度(默认从零开始,但要保留批次,所以需要设置为1)
  • end_dim:展平到哪一维(默认到最后)

这三种用哪种都可以,区别不大,选择喜欢的就好


img.size()是一个对象,返回的数据和shape属性值一样,img.size(0)等价于img.shape[0],取得是第一个数据64,也就是图像的批次(batch_size)


完整代码:

python 复制代码
import torchvision

from torch.utils.data import DataLoader
from torchvision import transforms

datasets = torchvision.datasets.CIFAR10("../dataset", train=False, download=True, transform=transforms.ToTensor())
dataloader = DataLoader(datasets, batch_size=64)

for data in dataloader:
    img, target = data
    input1 = img.reshape(img.size(0), -1)
    print(input1.size())

我们通过展平得到了能够送入线性层的数据,如下图所示:

3.线性模型

通过上面的打印,我们已经知道了特征数是3072个,所以第一个参数就写3072,然后需要10个结果类别,偏置就不写了,让他默认即可

python 复制代码
import torchvision
from torch import nn
from torch.utils.data import DataLoader
from torchvision import transforms

datasets = torchvision.datasets.CIFAR10("../dataset", train=False, download=True, transform=transforms.ToTensor())
dataloader = DataLoader(datasets, batch_size=64)


class LinearModel(nn.Module):
    def __init__(self):
        super(LinearModel, self).__init__()
        self.linear1 = nn.Linear(3072, 10)

    def forward(self, input1):
        output1 = self.linear1(input1)
        return output1


model = LinearModel()

for data in dataloader:
    img, target = data
    input1 = img.reshape(img.size(0), -1)
    output1 = model(input1)
    print(output1.size())

运行结果:

成功的把原来的每张图片的3072个特征转换成了10个结果类别

相关推荐
devpotato15 小时前
人工智能(十八)- 大模型幻觉生产风险治理
人工智能
TangGeeA15 小时前
Hermes Agent 安全约束实现分析:模型层、提示词层、Agent 层与 Tool 层
人工智能·ai
狐狐生风15 小时前
LangGraph 生产级部署全解:FastAPI + Docker
python·docker·langchain·prompt·fastapi·langgraph·agentai
happyprince15 小时前
04-FlagEmbedding 微调模块详细分析
人工智能
cd_9492172115 小时前
2026做标书用哪个AI工具好?深挖标书AI核心竞争力与实测对比
人工智能
派拉软件15 小时前
AI 网关:重塑企业级大模型服务治理架构
大数据·人工智能·架构
CLX050515 小时前
C#怎么实现全局异常过滤器_C#如何捕获控制器报错【核心】
jvm·数据库·python
江汉似年15 小时前
强化学习中的 On-policy 与 Off-policy 全面解析
人工智能·深度学习·算法·rl
性野喜悲15 小时前
python将excel中的链接转成图片并替换链接展示在excel中【将pdf的第一页插入excel并将对应信息获取到插入签名等位置】
开发语言·python·excel
sunneo15 小时前
03-从Chat到Act-Agent行动闭环的产品心理学拆解
人工智能·产品运营·aigc·产品经理·ai-native