• Python PyTorch 获取 MNIST 数据


    1 PyTorch 获取 MNIST 数据

    import torch
    import numpy as np
    import matplotlib.pyplot as plt # type: ignore
    from torchvision import datasets, transforms
    
    def mnist_get():
        print(torch.__version__)
        # 定义数据转换
        transform = transforms.Compose([
            transforms.ToTensor(),  # 将图像转换为张量
            transforms.Normalize((0.5,), (0.5,))  # 归一化图像数据
        ])
        # 获取数据
        train_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
        test_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
        # 训练数据
        train_image = train_data.data.numpy()
        train_label = train_data.targets.numpy()
        # 测试数据
        test_image = test_data.data.numpy()
        test_label = test_data.targets.numpy()
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21

    2 PyTorch 保存 MNIST 数据

    import torch
    import numpy as np
    import matplotlib.pyplot as plt # type: ignore
    from torchvision import datasets, transforms
    
    def mnist_save(mnist_path):
        print(torch.__version__)
        # 定义数据转换
        transform = transforms.Compose([
            transforms.ToTensor(),  # 将图像转换为张量
            transforms.Normalize((0.5,), (0.5,))  # 归一化图像数据
        ])
        # 获取数据
        train_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
        test_data = datasets.MNIST(root='./data', train=False, download=True, transform=transform)
        # 训练数据
        train_image = train_data.data.numpy()
        train_label = train_data.targets.numpy()
        # 测试数据
        test_image = test_data.data.numpy()
        test_label = test_data.targets.numpy()
        np.savez(mnist_path, train_data=train_image, train_label=train_label, test_data=test_image, test_label=test_label)
    
    mnist_path = 'C:\\Users\\Hyacinth\\Desktop\\mnist.npz'
    mnist_save(mnist_path)
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21
    • 22
    • 23
    • 24
    • 25

    3 PyTorch 显示 MNIST 数据

    import torch
    import numpy as np
    import matplotlib.pyplot as plt # type: ignore
    from torchvision import datasets, transforms
    
    def mnist_show(mnist_path):
        data = np.load(mnist_path)
        image = data['train_data'][0:100]
        label = data['train_label'].reshape(-1, )
        plt.figure(figsize = (10, 10))
        for i in range(100):
            print('%f, %f' % (i, label[i]))
            plt.subplot(10, 10, i + 1)
            plt.imshow(image[i])
        plt.show()
    
    mnist_path = 'C:\\Users\\Hyacinth\\Desktop\\mnist.npz'
    mnist_show(mnist_path)
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18

    在这里插入图片描述

  • 相关阅读:
    OpenCV自学笔记二十:图像分割和提取
    移动端的布局
    数据结构(栈和队列)
    圆满收官!华秋电子亮相2022慕尼黑华南电子展,数字化平台赋能智能制造
    微服务框架 SpringCloud微服务架构 4 Ribbon 4.1 负载均衡原理
    大数据之DStream 转换 完整使用 (第十四章)
    【论文笔记】UNet
    ubuntu 16.04成功安装meteor
    vue 动态数字效果 vue-animate-number
    数据预处理(预备知识)
  • 原文地址:https://blog.csdn.net/qq_34814092/article/details/138168623