部署含自定义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
相关产品推荐
相关产品推荐

