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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 21:50:51