如何将PyTorch神经网络部署为服务?张量形状报错解决
解决PyTorch CNN单张测试形状不匹配问题及Flask API部署方案
一、解决单张图片测试的形状不匹配错误
报错mat1和mat2形状无法相乘(10x36和780x70)的核心原因是单张图片的预处理流程与训练/数据加载器测试时不一致,导致卷积输出的特征图展平后维度与全连接层输入不匹配。解决步骤如下:
1. 严格对齐训练时的预处理流程
训练阶段使用的transforms(如尺寸缩放、通道转换、归一化等)必须完全复用到单张测试中:
- 训练时用了
Resize((28,28)),单张测试必须将图片缩放到相同尺寸; - 训练时输入是单通道灰度图,测试时也要把RGB图转成灰度;
- 训练时用了
Normalize(mean=[0.5], std=[0.5]),测试时必须用相同参数做归一化。
2. 添加Batch维度
数据加载器输出的张量形状是(batch_size, channels, H, W),但单张图片经ToTensor()处理后是(channels, H, W),需要用unsqueeze(0)手动添加batch维度,否则网络会误将通道维度当作batch维度,导致后续计算混乱。
3. 正确的单张测试代码示例
import torch from PIL import Image from torchvision import transforms # 完全复用训练时的transforms train_transform = transforms.Compose([ transforms.Resize((28, 28)), transforms.Grayscale(num_output_channels=1), transforms.ToTensor(), transforms.Normalize(mean=[0.5], std=[0.5]) ]) # 加载并预处理单张图片 img = Image.open("test_image.jpg") img_tensor = train_transform(img) img_tensor = img_tensor.unsqueeze(0) # 形状变为(1, 1, 28, 28) # 加载训练好的模型 model = torch.load("trained_cnn.pth", map_location=torch.device('cpu')) model.eval() # 推理(关闭梯度计算) with torch.no_grad(): output = model(img_tensor) prediction = torch.argmax(output, dim=1).item() print(f"预测结果:{prediction}")
二、Flask API部署步骤
将模型封装为Flask API,实现接收图片请求并返回预测结果,步骤如下:
1. 安装依赖
pip install flask torch torchvision pillow
2. 编写Flask服务代码
from flask import Flask, request, jsonify import torch from PIL import Image from torchvision import transforms import io app = Flask(__name__) # 全局加载模型(避免每次请求重新加载,提升效率) model = torch.load("trained_cnn.pth", map_location=torch.device('cpu')) model.eval() # 复用训练时的transforms train_transform = transforms.Compose([ transforms.Resize((28, 28)), transforms.Grayscale(num_output_channels=1), transforms.ToTensor(), transforms.Normalize(mean=[0.5], std=[0.5]) ]) @app.route('/predict', methods=['POST']) def predict(): # 检查请求中是否包含图片文件 if 'file' not in request.files: return jsonify({'error': '未提供图片文件'}), 400 file = request.files['file'] if file.filename == '': return jsonify({'error': '未选择图片文件'}), 400 try: # 读取并预处理图片 img = Image.open(io.BytesIO(file.read())) img_tensor = train_transform(img) img_tensor = img_tensor.unsqueeze(0) # 模型推理 with torch.no_grad(): output = model(img_tensor) prediction = torch.argmax(output, dim=1).item() return jsonify({'prediction': prediction}) except Exception as e: return jsonify({'error': str(e)}), 500 if __name__ == '__main__': app.run(host='0.0.0.0', port=5000, debug=False)
3. 测试API服务
启动服务后,可通过curl或Postman发送POST请求测试:
curl -X POST -F "file=@test_image.jpg" http://localhost:5000/predict
内容的提问来源于stack exchange,提问作者SillySlimeSimon
相关产品推荐
相关产品推荐

