如何在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(...)
几个关键注意点:
- 绝对不要用
global变量存储模型、图或会话,Django的多线程环境下,全局变量是所有请求共享的,必然会导致冲突。 - 用pickle加载模型时,确保保存模型的方式和加载方式一致(比如你用
pickle.dump(clf, f)保存,就用pickle.load(f)加载)。 - 如果用线程本地存储,不要在请求结束后调用
K.clear_session(),否则线程复用的时候会丢失会话状态,导致下一次请求出错。
内容的提问来源于stack exchange,提问作者Mike
相关产品推荐
相关产品推荐

