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

同步Django集成神经网络:解决TensorFlow加载耗时问题

解决同步Django架构下神经网络预测的重复加载问题

方案一:Django Worker启动时预加载模型

这是最低开销的实现方式,利用WSGI服务器(如Gunicorn、uWSGI)的多worker机制,让每个worker在启动时一次性加载TensorFlow模型,后续请求直接复用已加载的模型,无需重复初始化。

实现步骤

  1. 在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()
  1. 在视图中直接调用全局模型:
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异步架构。

实现步骤

  1. 创建模型服务脚本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
  1. 在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的同步架构。

实现步骤

  1. 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)
  1. 独立的模型服务进程监听队列:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 05:54:51