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

如何在Django中实现多Keras分类器会话同时预测?

兄弟,我太懂你这种Django里用Keras模型遇到会话冲突的痛苦了!之前在做一个多视图调用同个预训练模型的项目时,也踩过一模一样的坑,给你几个亲测有效的解决方案,你可以根据自己的场景选:

方案一:每个请求加载独立模型实例(简单直接,适合低并发)

既然共享全局模型会导致会话冲突,那干脆每个请求单独加载模型、创建独立的TensorFlow图和会话,彻底隔离状态。虽然每次请求加载模型会有一点性能开销,但胜在逻辑简单,不容易出问题。

代码示例:

from keras import backend as K
import tensorflow as tf
import pickle

def predict_view1(request):
    # 1. 每个请求单独加载模型
    with open(modelfilenameandpath, "rb") as f:
        clf = pickle.load(f)
    
    # 2. 创建独立的图和会话
    with tf.Graph().as_default(), tf.Session() as sess:
        K.set_session(sess)
        # 直接执行预测,预训练模型不需要重新初始化变量
        results = clf.predict_proba(the_new_vecs)
        # 预测完成后清理当前会话
        K.clear_session()
    
    # 处理结果并返回
    return HttpResponse(...)

def predict_view2(request):
    # 和view1完全一致的逻辑
    with open(modelfilenameandpath, "rb") as f:
        clf = pickle.load(f)
    
    with tf.Graph().as_default(), tf.Session() as sess:
        K.set_session(sess)
        results = clf.predict_proba(the_new_vecs)
        K.clear_session()
    
    return HttpResponse(...)

方案二:用线程本地存储复用模型(性能友好,适合高并发)

Django默认是每个请求分配一个线程,我们可以用threading.local()来为每个线程维护独有的模型实例、图和会话,这样每个线程只加载一次模型,后续请求复用,既隔离了会话,又避免了重复加载的性能损耗。

代码示例:

import threading
from keras import backend as K
import tensorflow as tf
import pickle

# 创建线程本地存储对象,每个线程的变量都是独立的
thread_local = threading.local()

def get_thread_model():
    # 检查当前线程是否已经加载过模型
    if not hasattr(thread_local, 'clf'):
        # 加载模型到当前线程的本地存储
        with open(modelfilenameandpath, "rb") as f:
            thread_local.clf = pickle.load(f)
        
        # 为当前线程创建独立的图和会话
        thread_local.graph = tf.Graph()
        with thread_local.graph.as_default():
            thread_local.sess = tf.Session()
            K.set_session(thread_local.sess)
    
    return thread_local.clf, thread_local.graph, thread_local.sess

def predict_view1(request):
    clf, graph, sess = get_thread_model()
    
    # 用当前线程的图作为上下文执行预测
    with graph.as_default():
        K.set_session(sess)
        results = clf.predict_proba(the_new_vecs)
    
    # 这里不要调用K.clear_session(),因为线程可能会被复用处理下一个请求
    return HttpResponse(...)

def predict_view2(request):
    clf, graph, sess = get_thread_model()
    
    with graph.as_default():
        K.set_session(sess)
        results = clf.predict_proba(the_new_vecs)
    
    return HttpResponse(...)

方案三:把模型做成独立服务(生产环境首选,彻底解耦)

如果你的项目是生产环境,或者并发量比较大,最稳妥的方式是把模型部署成独立的API服务,Django只需要通过HTTP请求调用这个服务获取预测结果,彻底把TensorFlow的会话和Django的Web服务隔离开。

第一步:写一个简单的模型服务(用Flask举例)

# model_service.py
from flask import Flask, request, jsonify
import pickle
import tensorflow as tf
from keras import backend as K

app = Flask(__name__)

# 服务启动时只加载一次模型
with open(modelfilenameandpath, "rb") as f:
    clf = pickle.load(f)
graph = tf.get_default_graph()

@app.route('/predict', methods=['POST'])
def predict():
    data = request.get_json()
    the_new_vecs = data['vecs']
    # 用预加载的图执行预测
    with graph.as_default():
        results = clf.predict_proba(the_new_vecs).tolist()
    return jsonify({'results': results})

if __name__ == '__main__':
    app.run(port=5000, debug=False)

第二步:Django视图调用这个服务

# views.py
import requests
import numpy as np

def predict_view1(request):
    # 准备你的向量数据,转成列表方便JSON传输
    the_new_vecs = ... # 你的向量处理逻辑
    vecs_list = the_new_vecs.tolist()
    
    # 调用模型服务
    response = requests.post(
        'http://localhost:5000/predict',
        json={'vecs': vecs_list}
    )
    results = response.json()['results']
    
    # 把结果转回numpy数组(如果需要)
    results_np = np.array(results)
    # 处理结果并返回
    return HttpResponse(...)

def predict_view2(request):
    # 和view1一致的调用逻辑
    the_new_vecs = ...
    vecs_list = the_new_vecs.tolist()
    
    response = requests.post(
        'http://localhost:5000/predict',
        json={'vecs': vecs_list}
    )
    results = response.json()['results']
    
    return HttpResponse(...)

几个关键注意点:

  1. 绝对不要用global变量存储模型、图或会话,Django的多线程环境下,全局变量是所有请求共享的,必然会导致冲突。
  2. 用pickle加载模型时,确保保存模型的方式和加载方式一致(比如你用pickle.dump(clf, f)保存,就用pickle.load(f)加载)。
  3. 如果用线程本地存储,不要在请求结束后调用K.clear_session(),否则线程复用的时候会丢失会话状态,导致下一次请求出错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 07:41:19