第2章:认识数据 —— Tensor 与 Dataset

🎉摘要:手写数字识别从零开始教程,带你开启深度学习之旅!零基础讲解MNIST数据集、神经网络原理、图像分类核心知识,手把手搭建手写数字识别模型,新手也能轻松入门深度学习实战项目。

这一章的目标

  • 搞清楚 PyTorch 里最核心的数据类型 Tensor(张量) 到底是什么

  • 亲手把 MNIST 数据集下载下来,看看里面到底长什么样

  • 理解 Dataset 的概念

再聊聊张量(Tensor)

第 1 章我随口说了句"张量就是数组",这句话不严谨,但对你理解 80% 的代码已经够了。这里稍微展开说清楚:

张量 = 包含数据的多维数组 + 能在 GPU 上运算的能力

普通 Python 列表 [1, 2, 3] 也能存数据,那为什么要用 Tensor?

# Python 列表做不了的事
import torch

a = torch.tensor([1.0, 2.0, 3.0])   # 张量
b = torch.tensor([4.0, 5.0, 6.0])
print(a + b)        # → tensor([5., 7., 9.])  对位相加
print(a * b)        # → tensor([4., 10., 18.]) 对位相乘
print(a @ b)        # → tensor(32.)  点积(内积)

运行代码,输出如下:

tensor([5., 7., 9.])
tensor([ 4., 10., 18.])
tensor(32.)

同样的操作用 Python 列表写会很啰嗦。更重要的是,Tensor 的运算在底层是 C/C++ 实现的,比 Python 循环快几十上百倍。当你的数据量是几万张图片时,这个速度差异就是 "1 秒" 和 "10 分钟" 的区别。

疑问:pytorch 中的 tensor 张量存在的目的是为了在 GPU 上面更快速的运算?

不只是为了 GPU 加速。GPU 高速运算只是张量的一大优势,但不是它存在的全部目的。

PyTorch 的 tensor = 多维数组容器

  • 0 维:标量 torch.tensor(3.0)

  • 1 维:向量

  • 2 维:矩阵

  • 3/4 / 更高维:图像、批量图像、时序数据

Python 原生也有列表 list、numpy 数组,那为什么还要搞 tensor?

对比 list /numpy/tensor

  1. Python list:普通动态列表,没有专门矩阵运算,很慢,不能上 GPU

  2. numpy array:CPU 上高效多维数组,数学运算很快,但 numpy 不能放到 GPU 显存

  3. torch.Tensor:

    • ✅ CPU 上可以做高速数学运算(类似 numpy)

    • ✅ 支持迁移到 GPU (CUDA) 做并行加速(numpy 做不到)

    • ✅ 自带自动求导 (autograd)!这是深度学习最核心功能

张量的形状(shape)

shape 告诉你这个张量"几行几列"(几维,每维多大):

import torch

t1 = torch.tensor([1, 2, 3])           # shape: (3,)       — 1维,有3个元素
t2 = torch.tensor([[1, 2], [3, 4]])    # shape: (2, 2)     — 2维,2行2列
t3 = torch.randn(2, 3, 4)              # shape: (2, 3, 4)  — 3维

print(t1.shape, t2.shape, t3.shape)

运行代码,输出:

torch.Size([3]) torch.Size([2, 2]) torch.Size([2, 3, 4])

📌 术语:维度(dimension)

在深度学习中,"维度"这个词有歧义:

  • 有时指张量有几维(比如 3 维张量)

  • 有时指某一维的大小(比如 "这个向量是 256 维的")

结合上下文理解就行。比如 MNIST 的一张图是 (1, 28, 28)的 3 维张量,意思是:1 个通道(灰度图)× 28 像素高 × 28 像素宽。

下载并查看 MNIST 数据集

MNIST 是一个经典的手写数字数据集,里面有 6 万张训练图片1 万张测试图片。每张图是 28×28 像素的灰度手写数字(0~9)。

📌 术语:训练集 vs 测试集

  • 训练集(training set):用来"教"模型的数据。相当于学生的练习题。

  • 测试集(test set):用来"考"模型的数据。这些图模型从没见过,如果考得好说明真的学会了,不是死记硬背。

自动下载

PyTorch 的 torchvision 已经帮我们封装好了,一行代码就能下载:

import torch
from torchvision import datasets, transforms

# 下载 MNIST 训练集
train_dataset = datasets.MNIST(
    root="./data",        # 数据存哪(会自动创建目录)
    train=True,           # True=训练集,False=测试集
    download=True,        # 本地没有就自动下载
    transform=transforms.ToTensor()  # 把PIL图片转成Tensor
)

print(f"训练集大小:{len(train_dataset)} 张图片")

运行代码,输出如下:

  0%|          | 32.8k/9.91M [00:00<01:15, 131kB/s]
...
  0%|          | 0.00/28.9k [00:00<?, ?B/s]
100%|██████████| 28.9k/28.9k [00:01<00:00, 23.6kB/s]
100%|██████████| 28.9k/28.9k [00:01<00:00, 23.6kB/s]

  0%|          | 0.00/1.65M [00:00<?, ?B/s]
  2%|▏         | 32.8k/1.65M [00:01<01:15, 21.5kB/s]
...
 99%|█████████▉| 1.64M/1.65M [02:09<00:00, 13.7kB/s]
100%|██████████| 1.65M/1.65M [02:09<00:00, 12.7kB/s]

  0%|          | 0.00/4.54k [00:00<?, ?B/s]
100%|██████████| 4.54k/4.54k [00:00<00:00, 8.25MB/s]
训练集大小:60000 张图片

运行这段代码,PyTorch 会自动从网上下载 MNIST 数据并存到 ./data/MNIST/ 目录。下载只需一次,以后运行就不会再下载了。如下图:

💡 下载可能需要 1~5 分钟,取决于网速。如果下载失败,关掉 VPN 再试试,或者参考第 1 章的镜像源方案。

看看数据长什么样

import torch
from torchvision import datasets, transforms

# 下载 MNIST 训练集(如果已经下载,则不重复下载)
train_dataset = datasets.MNIST(
    root="./data",        # 数据存哪(会自动创建目录)
    train=True,           # True=训练集,False=测试集
    download=True,        # 本地没有就自动下载
    transform=transforms.ToTensor()  # 把PIL图片转成Tensor
)
print(f"训练集大小:{len(train_dataset)} 张图片")

# 取第一张图片
image, label = train_dataset[0]
print(f"图片的 Tensor 形状: {image.shape}")  # → torch.Size([1, 28, 28])
print(f"标签(哪个数字): {label}")           # → 5

运行代码,输出如下:

训练集大小:60000 张图片
图片的 Tensor 形状: torch.Size([1, 28, 28])
标签(哪个数字): 5

解释 shape: (1, 28, 28),其格式为 (第 0 维, 第 1 维, 第 2 维),每一维的含义如下:

  • 第 0 维 = 1    →  通道数(灰度图只有一个通道,RGB彩图是3)

  • 第 1 维 = 28   →  图片高度(像素)

  • 第 2 维 = 28   →  图片宽度(像素)

📌 术语:标签(Label)

标签就是"正确答案"。图片是手写的 "5",它的标签就是 5。模型训练的目标就是:给它一张图片,它能输出正确的标签。

可视化一张图片

知道数据是什么样子很重要。用 matplotlib 画出来:

import torch
from torchvision import datasets, transforms
import matplotlib.pyplot as plt

# 解决中文乱码 这两行加在plt导入之后
plt.rcParams["font.family"] = ["SimHei", "WenQuanYi Micro Hei", "Heiti TC"]
plt.rcParams["axes.unicode_minus"] = False  # 解决负号显示异常

# 下载 MNIST 训练集
train_dataset = datasets.MNIST(
    root="./data",        # 数据存哪(会自动创建目录)
    train=True,           # True=训练集,False=测试集
    download=True,        # 本地没有就自动下载
    transform=transforms.ToTensor()  # 把PIL图片转成Tensor
)

# 查看第一张图片内容
image, label = train_dataset[0]

# image 的形状是 (1, 28, 28),但 imshow 需要 (28, 28) 或 (28, 28, 1)
# squeez() 会去掉所有大小为 1 的维度
plt.imshow(image.squeeze(), cmap='gray')
plt.title(f"这是数字:{label}")
plt.show()

运行代码,如果你看到一张白底黑字的手写数字图片(或者黑底白字,取决于显示设置),并且标题显示的是正确的数字,那就对了。如下图:

💡 matplotlib 是什么?

是一款应用十分广泛、基于 Python 语言开发的专业数据可视化开源工具库,能够灵活绘制各类静态图表、交互式图表与动画图形,也是数据分析、科研绘图领域最常用的基础绘图工具之一。

理解 Dataset

Dataset 是 PyTorch 中表示数据集的一个抽象。你可以把它想象成一个 "智能列表":

  • 它知道总共有多少条数据(len(dataset))

  • 它知道怎么取第 i 条数据(dataset[i])

  • 它帮你处理了文件读取、格式转换这些脏活

PyTorch 内置了很多常用数据集(MNIST、CIFAR-10、ImageNet 等),你不需要手动去网上下载、解压、读取文件。

本章小结

  • Tensor 是 PyTorch 的核心数据类型,本质是能高效运算的多维数组

  • shape 表示张量的维度和大小

  • MNIST 数据集有 6 万训练图 + 1 万测试图,每张 28×28 像素

  • Dataset 是 PyTorch 的数据集抽象,支持自动下载和索引访问

点击查看完整的下载 MNIST、查看张量形状、可视化样本示例代码。

 

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号