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

部署含自定义CTCLayer的TensorFlow模型的Flask应用响应超时排查

问题分析:Flask部署后验证码预测超时的原因及解决方案

问题背景

开发了一个简易Flask应用,支持上传PNG格式验证码图片并预测文本内容。应用从h5文件加载自定义CTCLayer,本地运行正常,但部署到服务器后预测响应极慢,最终触发Gunicorn Worker超时错误。疑问是模型加载方式存在问题,还是服务器CPU、GPU等资源不足导致的?

Flask代码

import os
from flask import Flask, render_template, request
import tensorflow as tf
import numpy as np
from tensorflow import keras
from keras import layers
from ctc_layer import CTCLayer


app = Flask(__name__)

UPLOAD_FOLDER = 'uploads'
ALLOWED_EXTENSIONS = {'png'}
app.config['UPLOAD_FOLDER'] = UPLOAD_FOLDER


img_width = 177
img_height = 40
max_length = 5
characters = ['0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'e', 'p']


char_to_num = layers.StringLookup(
    vocabulary=list(characters), mask_token=None
)


num_to_char = layers.StringLookup(
    vocabulary=char_to_num.get_vocabulary(), mask_token=None, invert=True
)


def encode_single_prediction(img_path):
    img = tf.io.read_file(img_path)
    img = tf.io.decode_png(img, channels=1)
    img = tf.image.convert_image_dtype(img, tf.float32)
    img = tf.image.resize(img, [img_height, img_width])
    img = tf.transpose(img, perm=[1, 0, 2])
    img = tf.expand_dims(img, axis=0)
    return img


def decode_single_prediction(pred):
    input_len = np.array([pred.shape[1]])
    results = tf.keras.backend.ctc_decode(pred, input_length=input_len, greedy=True)[0][0][:, :max_length]
    output_text = tf.strings.reduce_join(num_to_char(results[0])).numpy().decode("utf-8")
    return output_text


@app.route('/')
def home():
    return render_template("index.html")


@app.route('/predict', methods=['POST'])
def predict():
    uploaded_file = request.files['file']
    temp_file_path = os.path.join(app.config['UPLOAD_FOLDER'], uploaded_file.filename)
    uploaded_file.save(temp_file_path)
    image = encode_single_prediction(temp_file_path)
    model = keras.models.load_model('trained_model.h5', custom_objects={'CTCLayer': CTCLayer})
    prediction_model = keras.models.Model(
        model.get_layer(name="image").input, model.get_layer(name="dense2").output
    )
    prediction = prediction_model.predict(image)
    result = decode_single_prediction(prediction)
    return result


if __name__ == '__main__':
    app.run()

CTCLayer类代码

from tensorflow import keras
import tensorflow as tf


class CTCLayer(keras.layers.Layer):
    def __init__(self, name=None):
        super().__init__(name=name)
        self.loss_fn = keras.backend.ctc_batch_cost

    def call(self, y_true, y_pred):
        batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64")
        input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64")
        label_length = tf.cast(tf.shape(y_true)[1], dtype="int64")

        input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64")
        label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64")

        loss = self.loss_fn(y_true, y_pred, input_length, label_length)
        self.add_loss(loss)

        return y_pred

核心原因:模型重复加载导致的性能瓶颈

你的代码里最大的问题是把模型加载和预测模型构建放在了/predict请求处理函数内部——每收到一次预测请求,就会重新读取磁盘上的h5文件、初始化整个模型、构建预测模型分支,这会带来巨大的性能开销:

  • 磁盘读取模型文件本身耗时极长
  • 模型初始化会占用大量CPU/内存资源,服务器并发处理请求时会迅速耗尽资源,直接导致响应超时
  • 本地运行时因为只有单次请求,这个问题被掩盖,但服务器多请求场景下会被无限放大

服务器资源不足可能是次要因素,但模型加载方式错误是导致超时的主要原因。

修复步骤

1. 将模型加载移到全局作用域

在Flask应用初始化时就完成模型加载,而不是每次请求都重复加载:

# 放在app初始化之后,路由定义之前
model = keras.models.load_model('trained_model.h5', custom_objects={'CTCLayer': CTCLayer})
prediction_model = keras.models.Model(
    model.get_layer(name="image").input, model.get_layer(name="dense2").output
)

2. 修改predict函数,复用已加载的模型

@app.route('/predict', methods=['POST'])
def predict():
    uploaded_file = request.files['file']
    temp_file_path = os.path.join(app.config['UPLOAD_FOLDER'], uploaded_file.filename)
    uploaded_file.save(temp_file_path)
    image = encode_single_prediction(temp_file_path)
    # 直接复用全局初始化好的prediction_model
    prediction = prediction_model.predict(image)
    result = decode_single_prediction(prediction)
    # 可选:删除临时文件,避免磁盘空间被占用
    os.remove(temp_file_path)
    return result

3. 额外优化建议

  • 临时文件清理:每次请求生成的PNG临时文件记得删除,避免服务器磁盘被占满
  • GPU加速检查:如果服务器配备GPU,启动时打印tf.config.list_physical_devices('GPU')确认TensorFlow是否正确识别并使用GPU(CPU推理速度远低于GPU)
  • Gunicorn配置调整:可适当增加worker数量(建议为CPU核心数的2-4倍),延长超时时间(例如--timeout 30),但这只是辅助手段,核心优化还是模型加载逻辑
  • 模型格式转换:将h5模型转换为TensorFlow SavedModel格式,加载速度更快;或对模型进行量化压缩,减少内存占用和推理时间

资源问题排查(若修改后仍超时)

  • 用top/htop查看CPU、内存使用率,确认是否存在资源耗尽情况
  • 检查服务器GPU状态,确认TensorFlow是否启用GPU加速
  • 查看磁盘IO指标,确认是否为磁盘读取速度瓶颈(但模型移到全局加载后这部分仅发生一次)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 16:33:16