如何使用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图形化工具
- 新建请求,请求类型选择
POST,输入接口地址http://192.168.178.13:5000/predict - 切换到
Headers标签,添加键Content-Type,值为application/json - 切换到
Body标签,选择raw格式,右侧格式选择JSON,粘贴你的JSON内容:
{ "url" : "https://socialistmodernism.com/wp-content/uploads/2017/07/placeholder-image.png" }
- 点击
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
相关产品推荐
相关产品推荐

