ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch MNIST实战:FC与CNN参数传递与数据管道深度解析

PyTorch MNIST实战:FC与CNN参数传递与数据管道深度解析 简介本资源是一份面向深度学习初学者与PyTorch入门者的实战教学包聚焦手写数字识别任务完整呈现全连接网络FC与卷积神经网络CNN两种主流模型在MNIST数据集上的从零实现、训练与对比分析过程。资源共104个文件涵盖11个核心Python脚本含数据预处理、模型定义、训练循环、33张JPG/PNG格式的训练可视化图表如损失/准确率曲线、8个.gz压缩的原始MNIST数据文件、4个.pth/.pt模型权重文件以及requirements.txt依赖清单和README说明文档整体压缩包大小为67.2MB。已有323人下载学习适合希望深入理解PyTorch数据加载、模型构建、训练监控与结果可视化的开发者。读者可直接复现双模型训练流程对比性能差异并基于清晰的目录结构FC/CNN/dataset/my_mnist_dataset/images等快速定位模块代码掌握工业级项目组织规范。1. 为什么用 PyTorch 实现 FC 和 CNN 训练 MNIST不是“跑通就行”而是要拆清每一层参数传递逻辑很多初学者把 MNIST 当成“Hello World”式任务下载数据、搭个模型、train() 一下acc 到 98% 就收工。但真实项目里一旦换数据、调 batch_size、加 dropout 或迁移到新硬件模型突然不收敛、loss 飙升、GPU 显存爆满——问题往往出在 FC 和 CNN 的张量形状衔接处、DataLoader 的 collate_fn 行为、甚至torch.nn.CrossEntropyLoss对 label 的隐式类型转换上。这个 PyTorch 项目不是玩具代码它用两个并行目录FC/和CNN/强制你对比全连接网络如何把 28×28 图像展平成 784 维向量再逐层线性变换而卷积网络如何用nn.Conv2d(1, 32, kernel_size3)在通道维度上保留空间局部性它用重复出现的lossacc.csv文件告诉你训练过程必须可复现、可回溯——不是只看最终 acc而是要能定位第 47 轮 epoch 时 val_loss 突然上升是因学习率衰减过早还是验证集采样偏差。适合正在从 Keras 迁移到 PyTorch、或需要独立调试模型结构的工程师你得亲手改forward()里的x F.relu(self.conv1(x))观察x.shape变化而不是依赖model.summary()。2. 数据加载与预处理为什么make_ours_dataset.py不只是解压.gz而是重构数据管道2.1 MNIST 原始二进制格式解析绕过 torchvision 下载失败的硬核方案PyTorch 官方torchvision.datasets.MNIST在国内常因 CDN 限流返回 404尤其当torchvision版本与 PyTorch 不匹配时如torch2.1.0torchvision0.16.0。本项目直接提供原始.gz文件train-images-idx3-ubyte.gz,train-labels-idx1-ubyte.gz等并通过make_ours_dataset.py手动解析。关键不是“怎么读”而是“读完怎么对齐”。MNIST 图像文件头结构固定前 16 字节为 magic number4B、num_images4B、rows4B、cols4B标签文件头为 magic number4B、num_labels4B。解析代码需严格校验import gzip import numpy as np def read_idx3_ubyte(filename): with gzip.open(filename, rb) as f: magic int.from_bytes(f.read(4), big) assert magic 2051, fInvalid magic number: {magic} # 图像文件 magic num_images int.from_bytes(f.read(4), big) rows int.from_bytes(f.read(4), big) cols int.from_bytes(f.read(4), big) # 读取所有像素数据num_images * rows * cols转为 uint8 buf f.read() data np.frombuffer(buf, dtypenp.uint8).reshape(num_images, rows, cols) return data def read_idx1_ubyte(filename): with gzip.open(filename, rb) as f: magic int.from_bytes(f.read(4), big) assert magic 2049, fInvalid magic number: {magic} # 标签文件 magic num_labels int.from_bytes(f.read(4), big) buf f.read() labels np.frombuffer(buf, dtypenp.uint8) assert len(labels) num_labels, Label count mismatch return labels提示assert magic 2051是防错关键。若跳过此校验解压后得到全零矩阵或 shape 错误后续训练必然崩溃。np.frombuffer比np.loadtxt快 10 倍以上且避免字符串解析开销。2.2 自定义 Dataset 类my_mnist_dateset目录下的__getitem__如何控制数据增强粒度my_mnist_dateset/目录下应存在继承torch.utils.data.Dataset的类其__getitem__方法决定单样本输出格式。常见错误是直接return image, label忽略 PyTorch 对 tensor 类型和维度的要求。正确实现需图像归一化MNIST 像素范围 0–255必须缩放到[0,1]或[-1,1]后者需transforms.Normalize((0.5,), (0.5,))增加通道维度原始image是(28,28)CNN 输入需(1,28,28)FC 输入需(784,)Label 类型torch.LongTensor非int或np.int64否则CrossEntropyLoss报错from torch.utils.data import Dataset import torch import torchvision.transforms as T class CustomMNIST(Dataset): def __init__(self, images, labels, trainTrue, transformNone): self.images images.astype(np.float32) / 255.0 # 归一化 self.labels labels.astype(np.int64) self.train train # 定义 transform训练时加随机旋转测试时不加 self.transform transform or T.Compose([ T.ToTensor(), # 自动增加 channel dim 并转 float32 T.RandomRotation(degrees10) if train else T.Lambda(lambda x: x) ]) def __len__(self): return len(self.images) def __getitem__(self, idx): img self.images[idx] # shape: (28, 28) label self.labels[idx] # ToTensor 将 (H,W) - (C,H,W)自动归一化到 [0,1] img_tensor self.transform(torch.from_numpy(img).unsqueeze(0)) # (1,28,28) return img_tensor, torch.tensor(label, dtypetorch.long)2.2.1transforms.Compose中ToTensor()的隐式行为解析T.ToTensor()不是简单torch.tensor()它执行三步np.array→torch.Tensordtype 自动转float32HWC或HW→CHW对灰度图unsqueeze(0)加通道uint8→float32并/255.0注意这是唯一归一化步骤不可省略若手动写torch.tensor(img, dtypetorch.float32)会丢失通道维度且未归一化导致 CNN 输入为[0,255]权重爆炸。2.3 DataLoader 的关键参数batch_size64为何不是越大越好DataLoader的batch_size直接影响显存占用和梯度更新稳定性。本项目requirements.txt应指定torch2.0利用其persistent_workersTrue和pin_memoryTrue优化 IOfrom torch.utils.data import DataLoader train_dataset CustomMNIST(train_images, train_labels, trainTrue) val_dataset CustomMNIST(val_images, val_labels, trainFalse) train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers4, # 子进程数需 CPU 核心数 pin_memoryTrue, # 将 tensor 锁页内存加速 GPU 传输 persistent_workersTrue # 避免每个 epoch 重建 worker 进程 )参数推荐值作用错误示例后果num_workersmin(4, os.cpu_count())并行加载减少 GPU 等待设为 0单线程加载GPU 利用率 30%pin_memoryTrue将 host memory 锁页DMA 直传 GPUFalse数据拷贝慢 2–3 倍persistent_workersTrueworker 进程复用避免 fork 开销False每个 epoch 重建进程启动延迟高注意batch_size64是平衡点。若设128FC 模型显存占用从 1.2GB 升至 2.1GB但梯度噪声增大收敛变慢CNN因卷积运算特性batch_size32时 loss 曲线更平滑。3. 模型构建与训练循环FC 与 CNN 的 forward 函数差异如何决定反向传播路径3.1 全连接网络FC的forward()展平操作是性能瓶颈也是调试入口FC/目录下的模型继承nn.Module核心在于flatten层的位置和Linear的输入维度计算import torch.nn as nn import torch.nn.functional as F class FCNet(nn.Module): def __init__(self, input_dim28*28, hidden_dim128, num_classes10): super().__init__() self.fc1 nn.Linear(input_dim, hidden_dim) # 784 - 128 self.fc2 nn.Linear(hidden_dim, hidden_dim) # 128 - 128 self.fc3 nn.Linear(hidden_dim, num_classes) # 128 - 10 self.dropout nn.Dropout(0.2) def forward(self, x): # x shape: (B, 1, 28, 28) → 需展平为 (B, 784) x x.view(x.size(0), -1) # 关键-1 自动计算剩余维度 x F.relu(self.fc1(x)) x self.dropout(x) x F.relu(self.fc2(x)) x self.fc3(x) # 最后一层不激活交由 CrossEntropyLoss 处理 return x3.1.1x.view(x.size(0), -1)的底层机制与替代方案view()要求张量内存连续。若之前有transpose或narrow操作需先contiguous()。更安全的写法是x.flatten(1)# 等价但更鲁棒 x x.flatten(1) # 从 dim1 开始展平保留 batch dimflatten(1)与view(B, -1)的区别view仅重排 shape不改变内存布局失败时抛RuntimeErrorflatten保证返回连续张量兼容性更强3.2 卷积神经网络CNN的forward()卷积层输出尺寸必须手算验证CNN/目录下的模型必须显式计算每层输出 H/W否则nn.Linear输入维度错误。以LeNet-5变体为例class CNNNet(nn.Module): def __init__(self, num_classes10): super().__init__() # Conv1: in1, out32, k3, s1, p0 → H_out floor((28-3)/1)1 26 self.conv1 nn.Conv2d(1, 32, kernel_size3, stride1, padding0) self.pool1 nn.MaxPool2d(kernel_size2, stride2) # 26→13 # Conv2: in32, out64, k3, s1, p0 → 13→11, pool→5 self.conv2 nn.Conv2d(32, 64, kernel_size3, stride1, padding0) self.pool2 nn.MaxPool2d(kernel_size2, stride2) # 11→5 # 全连接层输入64 * 5 * 5 1600 self.fc1 nn.Linear(64 * 5 * 5, 128) self.fc2 nn.Linear(128, num_classes) def forward(self, x): x F.relu(self.conv1(x)) # (B,1,28,28) → (B,32,26,26) x self.pool1(x) # → (B,32,13,13) x F.relu(self.conv2(x)) # → (B,64,11,11) x self.pool2(x) # → (B,64,5,5) x x.view(x.size(0), -1) # → (B,1600) x F.relu(self.fc1(x)) x self.fc2(x) return x3.2.1 卷积输出尺寸公式与调试技巧通用公式H_out floor((H_in 2×p - k) / s) 1W_out floor((W_in 2×p - k) / s) 1调试时在forward中插入print(x.shape)def forward(self, x): print(Input:, x.shape) # torch.Size([64, 1, 28, 28]) x F.relu(self.conv1(x)) print(After conv1:, x.shape) # torch.Size([64, 32, 26, 26]) x self.pool1(x) print(After pool1:, x.shape) # torch.Size([64, 32, 13, 13]) # ... 后续同理提示若pool2后 shape 为(B,64,6,6)说明conv2输出 H/W 计算错误需检查padding是否漏设。3.3 训练循环中的损失与优化CrossEntropyLoss为何要求 label 为 LongTensor训练主循环中损失计算是易错点criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) for epoch in range(10): for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) # GPU 加速 optimizer.zero_grad() output model(data) # output shape: (B, 10) loss criterion(output, target) # target 必须是 LongTensor loss.backward() optimizer.step()nn.CrossEntropyLoss内部执行log_softmax nll_loss要求output:(N, C)float32target:(N,)long即torch.int64每个元素为 class index0–9若target是float或int报错Expected object of scalar type Long but got scalar type Int。CustomMNIST.__getitem__中torch.tensor(label, dtypetorch.long)即为此而设。4. 训练监控与结果分析从lossacc.csv提取可复现的收敛证据4.1 CSV 日志的结构化写入为什么不用print()而用csv.writerlossacc.csv文件包含epoch,train_loss,val_loss,train_acc,val_acc五列由训练循环实时追加。手动print()无法结构化分析而csv支持 Pandas 直接绘图import csv # 初始化 CSV 文件 with open(lossacc.csv, w, newline) as f: writer csv.writer(f) writer.writerow([epoch, train_loss, val_loss, train_acc, val_acc]) # 训练中每 epoch 写入 with open(lossacc.csv, a, newline) as f: writer csv.writer(f) writer.writerow([epoch, train_loss, val_loss, train_acc, val_acc])4.1.1 验证集准确率计算的正确姿势准确率不能仅用output.argmax(1) target的均值需考虑 batch size 不整除总样本数def calculate_accuracy(output, target): pred output.argmax(dim1, keepdimTrue) # (B,1) correct pred.eq(target.view_as(pred)).sum().item() return 100. * correct / len(target) # 用 len(target) 而非 batch_size # 在验证循环中 val_acc 0 with torch.no_grad(): for data, target in val_loader: data, target data.to(device), target.to(device) output model(data) val_acc calculate_accuracy(output, target) val_acc / len(val_loader) # 平均每个 batch 的 acc4.2 用 Pandas 分析lossacc.csv识别过拟合与欠拟合信号加载 CSV 后关键指标需交叉验证import pandas as pd import matplotlib.pyplot as plt df pd.read_csv(lossacc.csv) # 绘制 loss 曲线 plt.figure(figsize(12,4)) plt.subplot(1,2,1) plt.plot(df[epoch], df[train_loss], labelTrain Loss) plt.plot(df[epoch], df[val_loss], labelVal Loss, linestyle--) plt.xlabel(Epoch) plt.ylabel(Loss) plt.legend() plt.title(Loss Curve) # 绘制 acc 曲线 plt.subplot(1,2,2) plt.plot(df[epoch], df[train_acc], labelTrain Acc) plt.plot(df[epoch], df[val_acc], labelVal Acc, linestyle--) plt.xlabel(Epoch) plt.ylabel(Accuracy (%)) plt.legend() plt.title(Accuracy Curve) plt.tight_layout() plt.show()4.2.1 从曲线诊断模型状态现象FC 模型典型表现CNN 模型典型表现应对措施Train loss ↓, Val loss ↑第 5 轮后明显发散第 8 轮后缓慢上升加Dropout(0.5)或Weight DecayTrain/Val loss 平行下降两者 gap 0.01gap 0.005模型容量不足增宽网络如 FC hidden_dim→256Train loss 不降初始 learning_rate0.01 时 loss≈2.3log(10)conv1权重全零检查nn.init.kaiming_normal_()是否调用注意lossacc.csv中val_loss若持续高于train_loss且 gap 0.1说明验证集分布与训练集不一致——可能make_ours_dataset.py中 train/val 划分未打乱或DataLoader的shuffleFalse。5. 模型部署前的关键验证用单张图像测试 FC 与 CNN 的推理一致性5.1 构建最小推理脚本脱离训练循环验证model.eval()行为训练后保存的模型.pt需在独立脚本中加载并测试排除DataLoader干扰# test_inference.py import torch from FC import FCNet # 或 from CNN import CNNNet from my_mnist_dateset import CustomMNIST # 加载模型 model FCNet() model.load_state_dict(torch.load(fc_best.pth)) model.eval() # 关闭 dropout/batchnorm # 加载单张测试图像 test_dataset CustomMNIST(test_images[:1], test_labels[:1], trainFalse) img, label test_dataset[0] # img: (1,28,28), label: tensor(7) img img.unsqueeze(0) # 增加 batch dim → (1,1,28,28) # 推理 with torch.no_grad(): output model(img) pred output.argmax(dim1).item() print(fTrue label: {label.item()}, Predicted: {pred}) print(fOutput logits: {output.squeeze().tolist()})5.1.1model.eval()与model.train()的实际影响model.train()启用Dropout随机置零、BatchNorm用 batch 统计model.eval()Dropout失效输出原值、BatchNorm用 running_mean/var若推理时忘记eval()FC 模型因Dropout导致输出波动CNN 因BatchNorm用 mini-batch 统计而非全局统计而精度下降 5–10%。5.2 对比 FC 与 CNN 的中间特征用hook提取 conv1 输出热力图CNN 的优势在于局部特征提取可通过register_forward_hook可视化# 提取 conv1 输出 activation {} def get_activation(name): def hook(model, input, output): activation[name] output.detach() return hook model.conv1.register_forward_hook(get_activation(conv1)) # 前向传播 output model(img) act activation[conv1] # shape: (1,32,26,26) # 可视化前 4 个通道 plt.figure(figsize(12,3)) for i in range(4): plt.subplot(1,4,i1) plt.imshow(act[0,i].cpu(), cmapviridis) plt.title(fChannel {i}) plt.show()此时可观察FC 模型无中间特征图所有信息压缩在fc1.weight中而 CNN 的conv1输出已呈现边缘响应如数字“1”的竖直线条被高亮证明其学习到了空间不变特征——这正是它比 FC 在 MNIST 上高 2–3% 准确率的根本原因。本文还有配套的精品资源点击获取
返回列表