TensorFlow并行双模型:调用类方法时设置图/会话上下文的实现
并行运行独立Graph/Session的TensorFlow模型包装方案
我之前也碰到过类似的问题,用独立Graph和Session做隔离确实能完美解决模型并行时的图/操作冲突问题。你已经搭好了ModelWrapper的架子,我来把它补全成可直接复用的实现,顺便讲讲使用中的注意事项:
完整的ModelWrapper实现
import tensorflow as tf import numpy as np class ModelWrapper: def __init__(self): self.graph = None self.sess = None self.model = None def load_model(self, pth_model=None): # 创建完全独立的Graph实例 self.graph = tf.Graph() # 严格在当前Graph的上下文内完成模型加载和Session初始化 with self.graph.as_default(): # 配置GPU显存按需分配(多模型共享GPU时必备) config = tf.ConfigProto(gpu_options=tf.GPUOptions(allow_growth=True)) self.sess = tf.Session(graph=self.graph, config=config) with self.sess.as_default(): # 替换为你实际的模型加载逻辑 if pth_model: # 示例1:加载SavedModel格式的模型 self.model = tf.saved_model.loader.load(self.sess, ["serve"], pth_model) # 示例2:加载自定义Keras模型权重 # self.model = MyCustomKerasModel() # self.model.load_weights(pth_model) def predict(self, np_x): # 必须同时绑定当前模型的Graph和Session上下文 with self.graph.as_default(): with self.sess.as_default(): # 替换为你实际的预测逻辑 predictions = self.model.predict(np_x) return predictions def close(self): # 用完记得释放Session资源,避免内存泄漏 if self.sess: self.sess.close()
使用示例(并行调用两个模型)
你可以初始化两个独立的Wrapper实例,分别加载不同模型,然后通过多线程/多进程实现并行预测:
# 初始化两个模型包装器 model_a = ModelWrapper() model_b = ModelWrapper() # 分别加载不同的模型文件 model_a.load_model(pth_model="./path/to/model_a") model_b.load_model(pth_model="./path/to/model_b") # 准备测试输入 test_input = np.random.rand(8, 224, 224, 3) # 示例:8张224x224的RGB图片 predict_results = {} # 用多线程实现并行预测 import threading def run_prediction(model, input_data, result_dict, key): result_dict[key] = model.predict(input_data) # 启动两个并行线程 thread_a = threading.Thread(target=run_prediction, args=(model_a, test_input, predict_results, "model_a")) thread_b = threading.Thread(target=run_prediction, args=(model_b, test_input, predict_results, "model_b")) thread_a.start() thread_b.start() thread_a.join() thread_b.join() # 查看预测结果 print("模型A预测结果:", predict_results["model_a"].shape) print("模型B预测结果:", predict_results["model_b"].shape) # 释放资源 model_a.close() model_b.close()
关键注意事项
- 上下文必须严格绑定:每次调用
predict时,必须同时进入self.graph.as_default()和self.sess.as_default()上下文,确保所有TensorFlow操作都归属到当前模型的独立Graph中,绝对不能跨上下文执行操作。 - GPU显存配置:如果是GPU环境,一定要在创建Session时设置
allow_growth=True,否则单个模型会占用全部显存,导致另一个模型无法启动。 - 资源释放:长期运行的服务中,记得在模型不再使用时调用
close()方法关闭Session,避免显存/内存泄漏。
内容的提问来源于stack exchange,提问作者sccrthlt
相关产品推荐
相关产品推荐

