同步Django集成神经网络:解决TensorFlow加载耗时问题
解决同步Django架构下神经网络预测的重复加载问题
方案一:Django Worker启动时预加载模型
这是最低开销的实现方式,利用WSGI服务器(如Gunicorn、uWSGI)的多worker机制,让每个worker在启动时一次性加载TensorFlow模型,后续请求直接复用已加载的模型,无需重复初始化。
实现步骤
- 在Django项目的
wsgi.py中预加载模型:
import os from django.core.wsgi import get_wsgi_application import tensorflow as tf os.environ.setdefault('DJANGO_SETTINGS_MODULE', 'your_project.settings') # 仅在worker启动时执行一次模型加载 MODEL = tf.keras.models.load_model('path/to/your/trained_model.h5') application = get_wsgi_application()
- 在视图中直接调用全局模型:
from django.http import JsonResponse from your_project.wsgi import MODEL def predict(request): # 解析请求中的输入数据(示例) input_data = request.POST.getlist('data', type=float) processed_input = tf.convert_to_tensor([input_data]) # 直接使用已加载的模型预测 prediction = MODEL.predict(processed_input).tolist() return JsonResponse({'prediction': prediction[0]})
注意事项
- 生产环境需使用Gunicorn/uWSGI等成熟WSGI服务器,避免用Django自带的
runserver(它会启动双进程,可能重复加载模型)。 - 根据服务器内存调整worker数量,每个worker会独立加载一份模型,内存占用为
worker数 × 模型大小。
方案二:本地进程间通信(IPC)单模型服务
如果模型体积过大,多worker重复加载会占用过多内存,可启动一个独立进程加载模型,Django视图通过进程队列同步调用预测服务,无需修改Django异步架构。
实现步骤
- 创建模型服务脚本
model_service.py:
import tensorflow as tf from multiprocessing import Queue, Process import sys import time def model_worker(task_queue): # 仅加载一次模型 model = tf.keras.models.load_model('path/to/your/model.h5') while True: task_id, input_data = task_queue.get() if task_id == 'EXIT': break # 处理预测并返回结果 prediction = model.predict(input_data).tolist() task_queue.put((task_id, prediction)) # 在Django启动时初始化进程和队列 task_queue = Queue() worker_process = Process(target=model_worker, args=(task_queue,)) worker_process.daemon = True worker_process.start() # 将队列暴露给Django视图 sys.modules[__name__].PREDICTION_QUEUE = task_queue
- 在Django视图中调用:
import uuid import time from django.http import JsonResponse from .model_service import PREDICTION_QUEUE def predict(request): input_data = request.POST.getlist('data', type=float) processed_input = tf.convert_to_tensor([input_data]) # 生成唯一任务ID区分请求 task_id = str(uuid.uuid4()) PREDICTION_QUEUE.put((task_id, processed_input)) # 同步等待结果(设置超时避免无限阻塞) timeout = 15 start_time = time.time() while time.time() - start_time < timeout: response = PREDICTION_QUEUE.get() if response[0] == task_id: return JsonResponse({'prediction': response[1][0]}) return JsonResponse({'error': 'Prediction timeout'}, status=500)
方案三:Redis消息队列同步调用
如果需要分布式部署(如多台Django服务器共享模型服务),可使用Redis作为中间消息队列,实现Django与模型服务的解耦,仍保持Django的同步架构。
实现步骤
- Django视图发送任务并轮询结果:
import redis import uuid import time from django.http import JsonResponse # 初始化Redis连接 redis_client = redis.Redis(host='localhost', port=6379, db=0) def predict(request): input_data = request.POST.getlist('data', type=float) task_id = str(uuid.uuid4()) # 将任务存入Redis队列 redis_client.rpush('prediction_tasks', f"{task_id}:{','.join(map(str, input_data))}") # 轮询Redis获取结果 timeout = 10 start_time = time.time() while time.time() - start_time < timeout: result = redis_client.get(f"pred_result:{task_id}") if result: redis_client.delete(f"pred_result:{task_id}") return JsonResponse({'prediction': eval(result.decode())}) time.sleep(0.1) return JsonResponse({'error': 'Timeout'}, status=500)
- 独立的模型服务进程监听队列:
import redis import tensorflow as tf redis_client = redis.Redis(host='localhost', port=6379, db=0) model = tf.keras.models.load_model('path/to/your/model.h5') while True: # 阻塞等待队列任务 _, task_str = redis_client.blpop('prediction_tasks', timeout=0) task_id, input_str = task_str.decode().split(':', 1) input_data = list(map(float, input_str.split(','))) processed_input = tf.convert_to_tensor([input_data]) # 预测并将结果存入Redis prediction = model.predict(processed_input).tolist() redis_client.set(f"pred_result:{task_id}", str(prediction[0]))
内容的提问来源于stack exchange,提问作者Lex Podgorny
相关产品推荐
相关产品推荐

