使用Flask部署Teachable Machine的Keras模型遇维度不匹配错误
问题分析
错误提示明确:模型的Conv1_pad层期望4维输入(格式为(batch_size, height, width, channels)),但实际传入的是1维张量(None,)。根源是你在upload函数里没有调用model_predict处理上传的图片,直接使用了全局变量data的初始值[0],这个初始值是1维数组,完全不符合模型的输入要求。
修复步骤
- 移除不必要的全局变量
data,让model_predict直接返回处理好的图片张量,避免依赖全局状态。 - 在
upload函数中调用model_predict,确保上传的图片被正确预处理后再传入模型。 - 修正
model_predict里的变量名冲突:原代码中image被重复赋值,会导致后续操作出错。 - 清理标签格式:
labels.txt中的标签带格式符号,需要提取纯文本内容。
修正后的完整代码片段
from tensorflow.keras.models import load_model from PIL import Image, ImageOps import numpy as np import os from flask import Flask, request, secure_filename app = Flask(__name__) MODEL_PATH = 'your_model_path.h5' # 替换为你的模型实际路径 UPLOAD_FOLDER = 'uploads' os.makedirs(UPLOAD_FOLDER, exist_ok=True) # 加载模型和标签 model = load_model(MODEL_PATH, compile=False) # 读取并清理标签,去掉换行和多余符号 class_names = [line.strip().split(':')[-1].strip() for line in open('labels.txt', 'r').readlines()] def model_predict(img_path): # 修正变量名冲突:将重复的image改为img_resized img = Image.open(img_path).convert('RGB') size = (224, 224) img_resized = ImageOps.fit(img, size, Image.Resampling.LANCZOS) image_array = np.asarray(img_resized) normalized_image_array = (image_array.astype(np.float32) / 127.0) - 1 # 构造模型需要的4维输入张量:(1, 224, 224, 3) processed_data = np.expand_dims(normalized_image_array, axis=0) return processed_data @app.route('/', methods=['GET']) def index(): return ''' <!doctype html> <title>Image Prediction</title> <h1>Upload an Image</h1> <form method=post enctype=multipart/form-data action=/predict> <input type=file name=file accept="image/*"> <input type=submit value=Predict> </form> ''' @app.route('/predict', methods=['POST']) def upload(): if request.method == 'POST': f = request.files['file'] file_path = os.path.join(UPLOAD_FOLDER, secure_filename(f.filename)) f.save(file_path) # 调用预处理函数得到符合要求的输入 processed_data = model_predict(file_path) preds = model.predict(processed_data) index = np.argmax(preds) class_name = class_names[index] return f"Prediction Result: {class_name}" return "Invalid Request" if __name__ == '__main__': app.run(debug=True)
关键修复点说明
- 移除全局变量:避免全局状态导致的不可预期问题,让函数逻辑更独立可靠。
- 强制执行图片预处理:通过调用
model_predict,确保上传的图片被调整为224x224尺寸、归一化,并转换为模型需要的4维张量。 - 变量名冲突修复:原代码中
image被重复赋值,导致ImageOps.fit调用出错,修改后彻底解决该问题。 - 标签格式清理:自动提取
labels.txt中的纯标签文本,避免返回带格式符号的结果。
内容的提问来源于stack exchange,提问作者jay
相关产品推荐
相关产品推荐

