model.predict在Celery/uWSGI中挂起问题求助
问题复现
以下是用于NSFW图片检测的Python代码,在Python Shell和Django Shell中调用model.predict仅需约100ms,但在uWSGI和Celery环境中,预测环节会挂起10分钟以上,加锁无法解决该问题:
import numpy as np import tensorflow as tf import tensorflow_hub as hub from apps.common.utils.error_handling import suppress_callable_to_sentry from django.conf import settings from threading import Lock MODEL_PATH = settings.BASE_DIR / "apps/core/utils/nsfw_detector/nsfw.299x299.h5" model = tf.keras.models.load_model(MODEL_PATH, custom_objects={"KerasLayer": hub.KerasLayer}, compile=False) IMAGE_DIM = 299 TOTAL_THRESHOLD = 0.9 INDIVIDUAL_THRESHOLD = 0.7 predict_lock = Lock() @suppress_callable_to_sentry(Exception, return_value=False) def is_nsfw(image): if image.mode == "RGBA": image = image.convert("RGB") image = image.resize((IMAGE_DIM, IMAGE_DIM)) image = np.array(image) / 255.0 image = np.expand_dims(image, axis=0) with predict_lock: preds = model.predict(image)[0] categories = ["drawings", "hentai", "neutral", "porn", "sexy"] probabilities = {cat: float(pred) for cat, pred in zip(categories, preds)} individual_nsfw_prob = max(probabilities["porn"], probabilities["hentai"], probabilities["sexy"]) total_nsfw_prob = probabilities["porn"] + probabilities["hentai"] + probabilities["sexy"] return (individual_nsfw_prob > INDIVIDUAL_THRESHOLD) or (total_nsfw_prob > TOTAL_THRESHOLD)
原因分析
- 多进程模型继承冲突:uWSGI和Celery默认通过
fork创建子进程,主进程加载的模型实例会被所有子进程继承,但TensorFlow的计算图、会话以及底层CUDA资源无法在fork后的进程中正常共享,导致predict调用时触发资源死锁或无限等待。 - 全局实例的跨进程线程竞争:模块级别的全局模型在多进程环境中,TensorFlow默认的多线程推理会与uWSGI/Celery的进程管理机制产生资源冲突,加锁仅能解决单进程内的线程冲突,无法处理跨进程的资源争抢。
- 未初始化的子进程TensorFlow上下文:主进程加载模型后,fork出的子进程没有重新初始化TensorFlow上下文,导致推理时无法正确获取计算资源,陷入阻塞。
解决方案
方案1:在每个Worker进程中单独加载模型
避免在主进程全局加载模型,改为在每个Worker进程启动时或第一次调用函数时加载模型,确保每个进程拥有独立的模型实例和TensorFlow上下文。
Celery环境适配
利用Celery的worker_init信号,在Worker启动时加载模型:
import numpy as np import tensorflow as tf import tensorflow_hub as hub from apps.common.utils.error_handling import suppress_callable_to_sentry from django.conf import settings from threading import Lock from celery.signals import worker_init MODEL_PATH = settings.BASE_DIR / "apps/core/utils/nsfw_detector/nsfw.299x299.h5" model = None predict_lock = Lock() IMAGE_DIM = 299 TOTAL_THRESHOLD = 0.9 INDIVIDUAL_THRESHOLD = 0.7 @worker_init.connect def load_model_on_worker_init(**kwargs): global model # 清理旧的TensorFlow上下文 tf.keras.backend.clear_session() model = tf.keras.models.load_model(MODEL_PATH, custom_objects={"KerasLayer": hub.KerasLayer}, compile=False) @suppress_callable_to_sentry(Exception, return_value=False) def is_nsfw(image): global model # 非Celery环境下的懒加载逻辑 if model is None: tf.keras.backend.clear_session() model = tf.keras.models.load_model(MODEL_PATH, custom_objects={"KerasLayer": hub.KerasLayer}, compile=False) if image.mode == "RGBA": image = image.convert("RGB") image = image.resize((IMAGE_DIM, IMAGE_DIM)) image = np.array(image) / 255.0 image = np.expand_dims(image, axis=0) with predict_lock: preds = model.predict(image)[0] categories = ["drawings", "hentai", "neutral", "porn", "sexy"] probabilities = {cat: float(pred) for cat, pred in zip(categories, preds)} individual_nsfw_prob = max(probabilities["porn"], probabilities["hentai"], probabilities["sexy"]) total_nsfw_prob = probabilities["porn"] + probabilities["hentai"] + probabilities["sexy"] return (individual_nsfw_prob > INDIVIDUAL_THRESHOLD) or (total_nsfw_prob > TOTAL_THRESHOLD)
uWSGI环境适配
使用uWSGI的postfork钩子,在每个Worker进程fork完成后加载模型:
# 在Django的wsgi.py文件中添加 import uwsgidecorators import tensorflow as tf import tensorflow_hub as hub from django.conf import settings MODEL_PATH = settings.BASE_DIR / "apps/core/utils/nsfw_detector/nsfw.299x299.h5" model = None @uwsgidecorators.postfork def load_model_postfork(): global model tf.keras.backend.clear_session() model = tf.keras.models.load_model(MODEL_PATH, custom_objects={"KerasLayer": hub.KerasLayer}, compile=False)
方案2:限制TensorFlow的线程数
禁用TensorFlow的多线程并行推理,减少与uWSGI/Celery进程管理的资源竞争:
在模型加载前添加如下配置:
# 限制TensorFlow内部的线程数,避免多进程下的资源冲突 tf.config.threading.set_intra_op_parallelism_threads(1) tf.config.threading.set_inter_op_parallelism_threads(1)
方案3:切换为TensorFlow Lite模型
将Keras H5模型转换为TensorFlow Lite格式,Lite引擎更轻量,在多进程环境下兼容性更好:
模型转换代码
import tensorflow as tf import tensorflow_hub as hub MODEL_PATH = "apps/core/utils/nsfw_detector/nsfw.299x299.h5" model = tf.keras.models.load_model(MODEL_PATH, custom_objects={"KerasLayer": hub.KerasLayer}, compile=False) # 转换为TFLite模型 converter = tf.lite.TFLiteConverter.from_keras_model(model) tflite_model = converter.convert() with open("apps/core/utils/nsfw_detector/nsfw.tflite", "wb") as f: f.write(tflite_model)
修改推理代码
import numpy as np import tensorflow as tf from apps.common.utils.error_handling import suppress_callable_to_sentry from django.conf import settings from threading import Lock MODEL_PATH = settings.BASE_DIR / "apps/core/utils/nsfw_detector/nsfw.tflite" interpreter = None input_details = None output_details = None predict_lock = Lock() IMAGE_DIM = 299 TOTAL_THRESHOLD = 0.9 INDIVIDUAL_THRESHOLD = 0.7 def init_tflite_model(): global interpreter, input_details, output_details interpreter = tf.lite.Interpreter(model_path=str(MODEL_PATH)) interpreter.allocate_tensors() input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() # 根据环境在Worker启动时调用init_tflite_model # 比如Celery的worker_init或uWSGI的postfork @suppress_callable_to_sentry(Exception, return_value=False) def is_nsfw(image): global interpreter, input_details, output_details if interpreter is None: init_tflite_model() if image.mode == "RGBA": image = image.convert("RGB") image = image.resize((IMAGE_DIM, IMAGE_DIM)) image = np.array(image) / 255.0 image = np.expand_dims(image, axis=0).astype(np.float32) # TFLite需要匹配输入类型 with predict_lock: interpreter.set_tensor(input_details[0]['index'], image) interpreter.invoke() preds = interpreter.get_tensor(output_details[0]['index'])[0] categories = ["drawings", "hentai", "neutral", "porn", "sexy"] probabilities = {cat: float(pred) for cat, pred in zip(categories, preds)} individual_nsfw_prob = max(probabilities["porn"], probabilities["hentai"], probabilities["sexy"]) total_nsfw_prob = probabilities["porn"] + probabilities["hentai"] + probabilities["sexy"] return (individual_nsfw_prob > INDIVIDUAL_THRESHOLD) or (total_nsfw_prob > TOTAL_THRESHOLD)
内容的提问来源于stack exchange,提问作者Işık Kaplan
相关产品推荐
相关产品推荐

