PyTorch深度学习笔记(二十)(模型验证测试)

前言

到这一章节为止,依据小土堆课程的PyTorch深度学习笔记基础部分已经完结了,接下来将依据李沐动手学深度学习课程进行PyTorch深度学习笔记的进阶部分

预测图片

完整的模型验证(测试,demo)套路,利用已经训练好的模型,然后给它提供输入。

输入狗的图片,并打开

python 复制代码
image_path = "imgs/dog.png"
image = Image.open(image_path)

4通道的RGBA转为3通道的RGB图片

python 复制代码
image = image.convert("RGB")

转换图像格式并设定网络

python 复制代码
transform = torchvision.transforms.Compose([torchvision.transforms.Resize((32,32)),   
                                            torchvision.transforms.ToTensor()])

image = transform(image)
print(image.shape)

class Tudui(nn.Module):
    def __init__(self):
        super(Tudui, self).__init__()        
        self.model1 = nn.Sequential(
            nn.Conv2d(3,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,64,5,1,2),
            nn.MaxPool2d(2),
            nn.Flatten(),
            nn.Linear(64*4*4,64),
            nn.Linear(64,10)
        )
        
    def forward(self, x):
        x = self.model1(x)
        return x

GPU上训练的东西映射到CPU上

python 复制代码
model = torch.load("model/tudui_29.pth",map_location=torch.device('cpu'))

转为四维,符合网络输入需求

python 复制代码
image = torch.reshape(image,(1,3,32,32))

将模型转为测试类型

python 复制代码
model.eval()

不进行梯度计算,减少内存计算

python 复制代码
with torch.no_grad():
    output = model(image)

概率最大类别的输出

python 复制代码
print(output.argmax(1))

完整代码

python 复制代码
import torchvision
from PIL import Image
from torch import nn
import torch

image_path = "imgs/dog.png"
image = Image.open(image_path)
image = image.convert("RGB") 
print(image)

transform = torchvision.transforms.Compose([torchvision.transforms.Resize((32,32)),   
                                            torchvision.transforms.ToTensor()])

image = transform(image)
print(image.shape)

class Tudui(nn.Module):
    def __init__(self):
        super(Tudui, self).__init__()        
        self.model1 = nn.Sequential(
            nn.Conv2d(3,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,64,5,1,2),
            nn.MaxPool2d(2),
            nn.Flatten(),
            nn.Linear(64*4*4,64),
            nn.Linear(64,10)
        )
        
    def forward(self, x):
        x = self.model1(x)
        return x

model = torch.load("model/tudui_29.pth",map_location=torch.device('cpu'))
print(model)
image = torch.reshape(image,(1,3,32,32)) 
model.eval()
with torch.no_grad():
    output = model(image)
output = model(image)
print(output)
print(output.argmax(1)) 
相关推荐
Joshua-a33 分钟前
工程数学-理解卷积的含义_线性时不变系统的冲激响应与卷积
人工智能·深度学习
song50140 分钟前
React Native for OpenHarmony 实战:三方库 react-native-css-transformer 的鸿蒙化适配指南
css·人工智能·深度学习·react native·transformer
老金带你玩AI5 小时前
ChatGPT dots,到底能替你操多少心?
人工智能
AliCloudROS6 小时前
计算巢 X DeepSeek Harness — 云端智能体工作台
人工智能·阿里云
深蓝AI6 小时前
GitHub的Push一年涨4.9倍:AI Agent为什么逼它重做Git存储?
人工智能·github
stereohomology6 小时前
大模型的观点谨:StoryTold 「Crafting Apps」四件套 · 深度总览
人工智能·llm
运维行者_7 小时前
网络性能监控怎么做?从自动发现到根因分析的4个环节
运维·服务器·网络·人工智能·支持向量机
Su米苏7 小时前
Spring AI 中MCP 与普通 @Tool的区别
java·人工智能·spring
美狐美颜sdk7 小时前
直播APP源码可以直接接入视频美颜sdk吗?技术方案详解
android·人工智能·音视频·美颜sdk·直播美颜sdk
问商十三载7 小时前
AI引擎生成式优化实践:刑事辩护律师团队IP构建
大数据·人工智能