如何为PyTorch的Conv2d层格式化MNIST手绘输入数据?
PyTorch Conv2d输入格式与推理函数实现
问题背景
我之前手动实现过卷积神经网络,现在改用PyTorch提升速度,但对Conv2d层的输入格式不明确。已基于MNIST训练并加载好模型,代码如下:
import torch import torch.nn as nn import torch.nn.functional as F class NeuralNetwork(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 8, 3, stride = 1, padding = 1) self.pool = nn.MaxPool2d(2, stride = 2) self.conv2 = nn.Conv2d(8, 8, 3, stride = 1, padding = 1) self.linear1 = nn.Linear(7 * 7 * 8, 128) self.linear2 = nn.Linear(128, 128) self.linear3 = nn.Linear(128, 10) def forward(self, x): x = self.pool(F.relu(self.conv1(x))) x = self.pool(F.relu(self.conv2(x))) x = torch.flatten(x, 1) x = F.relu(self.linear1(x)) x = F.relu(self.linear2(x)) x = self.linear3(x) return x my_model = NeuralNetwork() my_model.load_state_dict(torch.load("model_weights.pth", weights_only=True)) my_model.eval()
当前Web应用流程:
- 用户在28×28黑白画布绘图;
- 绘图内容转为784长度的扁平数组,0代表白色区域,1代表黑色区域。
需要实现some_formatting_function和some_reading_function完成模型推理。
解决方案
1. 实现some_formatting_function
PyTorch的Conv2d层要求输入格式为**(batch_size, channels, height, width)**,针对用户的扁平数组,需要做以下转换:
- 将784维数组重塑为28×28的二维结构;
- 添加单通道维度(MNIST为灰度图,通道数=1);
- 添加batch维度(模型默认接受批量输入,这里单样本batch_size=1);
- 转换为float类型(模型训练时通常使用float32张量)。
注意:MNIST训练集原始像素是黑色为0,白色为255,归一化后是0(黑)1(白)。而用户的输入是0(白)1(黑),需要反转像素值才能匹配训练时的输入分布,否则预测会出错。
纯PyTorch实现版本:
def some_formatting_function(flattened_array): # 转换为float32张量,反转像素值(匹配MNIST训练分布) tensor = 1 - torch.tensor(flattened_array, dtype=torch.float32) # 调整形状为(batch=1, channel=1, height=28, width=28) return tensor.view(1, 1, 28, 28)
依赖NumPy的实现版本:
import numpy as np def some_formatting_function(flattened_array): # 转换为NumPy数组并反转像素值 arr = 1 - np.array(flattened_array, dtype=np.float32) # 重塑为28×28,再添加batch和通道维度 return arr.reshape(1, 1, 28, 28)
2. 实现some_reading_function
模型输出是形状为**(batch_size, 10)**的张量,每个元素对应0-9的预测得分,取得分最高的类别索引就是预测的数字。
实现代码:
def some_reading_function(pred): # 在类别维度(dim=1)取最大值对应的索引,转换为Python整数 return pred.argmax(dim=1).item()
完整推理代码
# 假设flattened_array_of_0_and_1是用户输入的784维数组 formatted_tensor = some_formatting_function(flattened_array_of_0_and_1) # 推理时关闭梯度计算,节省内存 with torch.no_grad(): pred = my_model(formatted_tensor) guessed_digit = some_reading_function(pred) print(guessed_digit)
内容的提问来源于stack exchange,提问作者Ming Lin
相关产品推荐
相关产品推荐

