如何加载已训练TensorFlow模型并使用不同batch size进行预测?
解决TensorFlow模型推理时的Batch Size适配问题
我太懂这种卡壳的感觉了!明明只是想调整个batch size,结果卡在frozen model的固定输入上,着实头疼。我来给你拆解两种模型的解决方案,再教你怎么根据设备自动选最优的batch size。
一、针对Frozen Model的Batch Size修改方法
Frozen model的输入shape被硬编码在GraphDef里,直接加载没法改,但我们可以手动修改GraphDef的节点属性,把batch size改成动态维度(也就是-1或者None):
步骤1:加载并修改GraphDef
import tensorflow as tf import numpy as np # 加载你的frozen model with tf.io.gfile.GFile('your_frozen_model.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 找到你的输入节点(替换成你实际的输入节点名称,比如'input_1') input_node_name = 'input_node_name' for node in graph_def.node: if node.name == input_node_name: # 把第一个维度(batch size)改成-1(动态) node.attr['shape'].shape.dim[0].size = -1 break
步骤2:导入修改后的Graph并测试推理
with tf.compat.v1.Session() as sess: # 导入修改后的GraphDef tf.import_graph_def(graph_def, name='') # 获取输入输出张量 input_tensor = sess.graph.get_tensor_by_name(f'{input_node_name}:0') output_tensor = sess.graph.get_tensor_by_name('your_output_node_name:0') # 替换成你的输出节点名 # 测试任意batch size的输入,比如batch=4 test_input = np.random.randn(4, 299, 299, 3).astype(np.float32) predictions = sess.run(output_tensor, feed_dict={input_tensor: test_input}) print(f"推理结果shape: {predictions.shape}") # 应该是(4, ...),符合你的输出维度
注意事项
- 如果你的frozen模型里有依赖固定batch size的节点(比如硬编码了shape的全连接层、池化层),修改输入shape后可能会报错。这种情况最好回到训练阶段,把输入shape设为
[None, 299,299,3]再重新冻结模型。 - 如果你用了TensorRT优化,一定要先修改GraphDef的输入shape,再进行TensorRT转换,并且要启用动态shape模式:
from tensorflow.python.compiler.tensorrt import trt_convert as trt # 设置TensorRT转换参数,启用动态shape conversion_params = trt.DEFAULT_TRT_CONVERSION_PARAMS._replace( max_workspace_size_bytes=1 << 30, # 分配1GB工作空间 precision_mode=trt.TrtPrecisionMode.FP16, dynamic_shape=True, allow_build_at_runtime=True ) # 转换修改后的GraphDef converter = trt.TrtGraphConverterV2( input_graph_def=graph_def, conversion_params=conversion_params ) converter.convert(input_shapes={input_node_name: [None, 299,299,3]}) trt_graph_def = converter.get_graph_def() # 之后就可以用这个TRT优化后的GraphDef进行动态batch推理了
二、针对Saved Model的Batch Size适配
Saved Model本身就支持动态shape,操作起来简单很多,直接加载后传入任意batch size的输入即可:
import tensorflow as tf # 加载Saved Model loaded_model = tf.saved_model.load('your_saved_model_dir') # 获取推理签名(通常是'serving_default') infer_fn = loaded_model.signatures['serving_default'] # 测试batch size=8的输入 test_input = tf.random.normal([8, 299, 299, 3]) predictions = infer_fn(test_input) # 查看输出结果的shape print(f"推理结果shape: {list(predictions.values())[0].shape}")
三、根据设备环境动态选择最优Batch Size
想要根据GPU/CPU的内存和性能选最合适的batch size,可以写个小工具函数,测试不同batch size的平均推理时间,选最快的那个:
import time def find_optimal_batch_size(input_shape, infer_func, device='GPU'): # 根据设备内存设置候选batch size(可以自行调整范围) batch_candidates = [1, 2, 4, 8, 16, 32, 64] best_batch = 1 fastest_avg_time = float('inf') height, width, channels = input_shape for batch_size in batch_candidates: try: # 生成测试输入 test_input = np.random.randn(batch_size, height, width, channels).astype(np.float32) # 先预热模型,避免初始化时间影响结果 infer_func(test_input) # 多次推理取平均时间 start_time = time.time() for _ in range(100): infer_func(test_input) avg_time = (time.time() - start_time) / 100 print(f"Batch size {batch_size}: 平均推理时间 {avg_time:.4f} 秒") if avg_time < fastest_avg_time: fastest_avg_time = avg_time best_batch = batch_size except tf.errors.ResourceExhaustedError: print(f"Batch size {batch_size} 超出设备内存,跳过") continue print(f"\n{device} 最优batch size为: {best_batch}") return best_batch # 用法示例(针对Saved Model) loaded_model = tf.saved_model.load('your_saved_model_dir') infer_func = lambda x: loaded_model.signatures['serving_default'](tf.convert_to_tensor(x)) optimal_batch = find_optimal_batch_size((299,299,3), infer_func)
内容的提问来源于stack exchange,提问作者oscarriddle
相关产品推荐
相关产品推荐

