You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 23:45:55