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

如何加载已训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 06:18:03