使用TensorFlow+Flask开发图像预测网站遇两类报错求助
问题:TensorFlow+Flask图像预测网站的远程图片读取错误
我正在用TensorFlow、Flask和Python开发一个支持图像预测的网站,相关代码及遇到的问题如下:
主应用代码
from flask import Flask, render_template import os import numpy as np import pandas as pd app = Flask(__name__) @app.route('/') def index(): return render_template('index.html') import tensorflow as tf import tensorflow_hub as hub model = tf.keras.models.load_model(MODEL_PATH) IMG_SIZE = 224 BATCH_SIZE = 32 custom_path = "http://t1.gstatic.com/licensed-image?q=tbn:ANd9GcQd6lM4HtInRF3cxw6h3MgUZIIiJCdMgFvXKrhaJrbw61tN3aYpMIVBi0dx0KPv1sdCrLk0sBhPeNVt8m0" custom_data = create_data_batches(custom_path, test_data=True) custom_preds = model.predict(custom_data) # Get custom image prediction labels custom_pred_labels = [get_pred_label(custom_preds[i]) for i in range(len(custom_preds))] print(custom_pred_labels) @app.route('/my-link/') def my_link(): return f"The predictions are: {custom_pred_labels}" if __name__ == '__main__': app.run(host="localhost", port=3000, debug=True)
辅助函数代码
process_image函数
def process_image(image_path, img_size=IMG_SIZE): """ Takes an image file path and turns the image into a Tensor. """ image = tf.io.read_file(image_path) image = tf.image.decode_jpeg(image, channels=3) image = tf.image.convert_image_dtype(image, tf.float32) image = tf.image.resize(image, size=[img_size, img_size]) return image
create_data_batches函数(测试数据部分)
def create_data_batches(X, y=None, batch_size=BATCH_SIZE, valid_data=False, test_data=False): """ Creates batches out of data out of image (X) and label (y) pairs. Shuffles the data if it's training data but doesn't shuffle if it's validation data. Also accepts test data as input (no labels) """ if test_data: print("Creating test data batches...") data = tf.data.Dataset.from_tensor_slices((tf.constant(X))) # only filepaths (no labels) data_batch = data.map(process_image).batch(BATCH_SIZE) return data_batch
get_pred_label函数
def get_pred_label(prediction_probabilites): """ Turns an array of prediction probabilities into a label. """ return unique_breeds[np.argmax(prediction_probabilites)]
遇到的错误
- 初始错误:
ValueError: Unbatching a tensor is only supported for rank >= 1
将custom_path改为列表后,又出现新错误:
UNIMPLEMENTED: File system scheme 'http' not implemented (file: 'http://t1.gstatic.com/licensed-image?q=tbn:ANd9GcQd6lM4HtInRF3cxw6h3MgUZIIiJCdMgFvXKrhaJrbw61tN3aYpMIVBi0dx0KPv1sdCrLk0sBhPeNVt8m0')
解决方案
问题根源
tf.io.read_file仅支持读取本地文件系统的路径,无法直接处理HTTP/HTTPS远程链接;第一个错误是因为custom_path不是列表,导致张量维度不符合要求,已解决。
修改步骤
- 添加远程图片下载函数,先将图片从URL下载到内存
- 修改
process_image函数,支持同时处理本地路径和远程URL - 调整Flask逻辑,将预测代码移到路由函数内(避免应用启动时就执行预测)
修改后完整代码
from flask import Flask, render_template import os import numpy as np import pandas as pd import tensorflow as tf import tensorflow_hub as hub import requests from io import BytesIO app = Flask(__name__) # 全局配置(替换为你的实际参数) MODEL_PATH = "你的模型文件路径" IMG_SIZE = 224 BATCH_SIZE = 32 unique_breeds = [] # 替换为训练时使用的类别列表 # 加载模型 model = tf.keras.models.load_model(MODEL_PATH) def download_image(url): """从远程URL下载图片,返回字节流""" response = requests.get(url) response.raise_for_status() # 捕获HTTP请求错误 return BytesIO(response.content) def process_image(image_source, img_size=IMG_SIZE): """ 支持本地文件路径或远程URL,将图片转换为Tensor """ if isinstance(image_source, str) and image_source.startswith(('http://', 'https://')): # 处理远程URL img_bytes = download_image(image_source) image = tf.image.decode_jpeg(img_bytes.getvalue(), channels=3) else: # 处理本地文件 image = tf.io.read_file(image_source) image = tf.image.decode_jpeg(image, channels=3) image = tf.image.convert_image_dtype(image, tf.float32) image = tf.image.resize(image, size=[img_size, img_size]) return image def create_data_batches(X, y=None, batch_size=BATCH_SIZE, valid_data=False, test_data=False): """创建数据批次,支持测试数据的URL列表""" if test_data: print("Creating test data batches...") data = tf.data.Dataset.from_tensor_slices((tf.constant(X))) data_batch = data.map(process_image).batch(BATCH_SIZE) return data_batch def get_pred_label(prediction_probabilities): """将预测概率转换为类别标签""" return unique_breeds[np.argmax(prediction_probabilities)] @app.route('/') def index(): return render_template('index.html') @app.route('/my-link/') def my_link(): # 可改为从请求参数动态获取图片URL custom_path = ["http://t1.gstatic.com/licensed-image?q=tbn:ANd9GcQd6lM4HtInRF3cxw6h3MgUZIIiJCdMgFvXKrhaJrbw61tN3aYpMIVBi0dx0KPv1sdCrLk0sBhPeNVt8m0"] custom_data = create_data_batches(custom_path, test_data=True) custom_preds = model.predict(custom_data) custom_pred_labels = [get_pred_label(custom_preds[i]) for i in range(len(custom_preds))] return f"The predictions are: {custom_pred_labels}" if __name__ == '__main__': app.run(host="localhost", port=3000, debug=True)
注意事项
- 确保安装
requests库:执行pip install requests unique_breeds必须替换为你训练模型时使用的类别列表,否则预测标签会出错- 可以添加额外错误处理(如URL无效、图片下载失败等场景)
内容的提问来源于stack exchange,提问作者Yusuf Saad
相关产品推荐
相关产品推荐

