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

如何使用Flask接收JSON字符串并调用PyTorch模型返回图像预测结果

解决Flask传递JSON获取图像预测结果的方案

一、先修复现有代码的隐藏问题

原有代码存在几处会导致运行报错、结果异常的问题,需要先修正:

  • Img类get_prediction方法的Softmax参数错误:nn.Softmax初始化需要传入dim维度参数,不能直接传入模型输出tensor,修正后逻辑见后续完整代码。
  • 批量URL预测逻辑无结果返回:原有批量处理只返回类型标识,没有返回每个URL的预测结果,需要把结果存入集合后返回。
  • 缺少请求参数校验:未传JSON、缺少url字段会直接触发接口报错,需要增加异常分支处理。

二、3种常用的JSON请求发送方式

方法1:使用curl命令行工具

Linux/macOS直接执行以下命令:

curl -X POST -H "Content-Type: application/json" -d '{"url": "https://socialistmodernism.com/wp-content/uploads/2017/07/placeholder-image.png"}' http://192.168.178.13:5000/predict

Windows cmd环境执行需要转义引号:

curl -X POST -H "Content-Type: application/json" -d "{\"url\": \"https://socialistmodernism.com/wp-content/uploads/2017/07/placeholder-image.png\"}" http://192.168.178.13:5000/predict

方法2:使用Python requests库

编写Python脚本发送请求:

import requests

url = "http://192.168.178.13:5000/predict"
payload = {"url": "https://socialistmodernism.com/wp-content/uploads/2017/07/placeholder-image.png"}
headers = {"Content-Type": "application/json"}

response = requests.post(url, json=payload, headers=headers)
print(response.json())

方法3:使用Postman图形化工具

  1. 新建请求,请求类型选择POST,输入接口地址http://192.168.178.13:5000/predict
  2. 切换到Headers标签,添加键Content-Type,值为application/json
  3. 切换到Body标签,选择raw格式,右侧格式选择JSON,粘贴你的JSON内容:
{
    "url" : "https://socialistmodernism.com/wp-content/uploads/2017/07/placeholder-image.png"
}
  1. 点击Send按钮即可看到返回的预测结果。

三、修复后的完整代码参考

modelLoader.py

from torch import device, load, nn

class Model:
    def __init__(self, path, class_list=None, dvc=None):
        if class_list is None:
            class_list = [0, 1]
        if dvc is None:
            dvc = 'cpu'
        self.class_list = class_list 
        self.model = load(path, map_location=device(dvc))
        num_ftrs = self.model.fc.in_features                  
        self.model.fc = nn.Linear(num_ftrs, len(class_list))  
        self.model.eval()

imageLoader.py

import torchvision.transforms as transforms
import io
from PIL import Image
from torch import nn
import requests

class Img:
    def __init__(self, url=None, image=None, image_bytes=None):
        if url:
            img = Image.open(requests.get(url, stream=True).raw)
            img_byte_arr = io.BytesIO()
            img.save(img_byte_arr, format=img.format)
            self.image_bytes = img_byte_arr.getvalue()
        elif image:
            f = image.read()
            self.image_bytes = bytearray(f)
        elif image_bytes:
            self.image_bytes = image_bytes

    def transform_image(self):
        data_transforms = transforms.Compose([transforms.Resize((224, 224)),
                                              transforms.CenterCrop(224), 
                                              transforms.ToTensor(),
                                              transforms.Normalize([0.485, 0.456, 0.406], 
                                                                   [0.229, 0.224, 0.225])])
        image = Image.open(io.BytesIO(self.image_bytes)).convert('RGB')
        return data_transforms(image).unsqueeze(0)

    def get_prediction(self, model):
        tensor = self.transform_image()
        output = model.model(tensor)
        # 修正Softmax调用方式
        sm = nn.Softmax(dim=1)(output)
        best = output.argmax().item()
        return model.class_list[best], sm.tolist()

app.py

from modelLoader import Model
from imageLoader import Img
import os
from flask import Flask, jsonify, request

app = Flask(__name__)
classes = ["dummy-image", "product-image"]
model_path = os.path.join("data", "models", "model1709", "model1709")
m1 = Model(model_path, classes, "cpu")

@app.route('/predict', methods=['POST'])
def predict():
    # 增加参数校验
    if not request.is_json:
        return jsonify({"error": "请求格式必须为JSON"}), 400
    input_json = request.get_json()
    if "url" not in input_json:
        return jsonify({"error": "缺少必填参数url"}), 400
    input_data = input_json['url']
    # 批量处理逻辑修改为返回结果
    if isinstance(input_data, list):
        result = {}
        for elem in input_data:
            img_elem = Img(url=elem)
            res, info = img_elem.get_prediction(m1)
            result[elem] = {"res": res, "info": info}
        return jsonify({"type": "bulk_upload", "result": result})
    # 单张处理
    img_inpdata = Img(url=input_data)
    res, info = img_inpdata.get_prediction(m1)
    return jsonify({input_data: res, "info": info})


if __name__ == '__main__':
    # 增加0.0.0.0绑定,支持局域网访问
    app.run(debug=True, host='0.0.0.0')

内容的提问来源于stack exchange,提问作者lxg95

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 16:54:03