TensorFlow:如何替换导入GraphDef中的Placeholder以连接数据集提供器
替换导入GraphDef中的Placeholder并连接到数据集提供器的方法
结合我用slim API做模型评估的经验(尤其是参考eval_image_classifier.py的思路),给你一步步拆解具体操作:
1. 先锁定默认计算图
首先要确保所有操作都在同一个计算图中进行,避免图混乱:
import tensorflow as tf from tensorflow.contrib import slim # 设置当前图为默认图 with tf.Graph().as_default() as graph: # 后续所有操作都在这个图里执行
2. 配置数据集提供器与预处理逻辑
这部分可以直接参考eval_image_classifier.py里的实现,核心是加载数据集、定义预处理函数,得到可以输出批量数据的张量:
# 1. 选择目标数据集(比如ImageNet、自定义数据集) dataset = slim.dataset.Dataset( data_sources=your_data_sources, reader=tf.TFRecordReader, decoder=your_decoder, num_samples=your_num_samples, items_to_descriptions=your_item_descriptions, num_classes=your_num_classes ) # 2. 创建数据集提供器,负责加载批量数据 provider = slim.dataset_data_provider.DatasetDataProvider( dataset, shuffle=False, # 评估阶段一般不打乱数据 common_queue_capacity=32, common_queue_min=8 ) # 3. 从提供器获取原始数据,比如图像和标签 [image, label] = provider.get(['image', 'label']) # 4. 应用预处理函数(和训练/导出模型时的预处理逻辑保持一致!) processed_image = image_preprocessing_fn( image, image_height=224, # 要和模型输入尺寸匹配 image_width=224, is_training=False # 评估阶段关闭训练相关操作(比如BN的更新、dropout) ) # 批量处理,得到批量输入张量 batch_images, batch_labels = tf.train.batch( [processed_image, label], batch_size=32, num_threads=4, capacity=32 * 2 )
3. 核心:导入GraphDef并替换Placeholder
这是关键步骤,要把导入的计算图里的Placeholder替换成我们上面生成的batch_images张量:
# 1. 读取预导出的GraphDef文件(比如.pb格式) with tf.gfile.GFile('your_model.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) # 2. 找到原模型中的Placeholder名称(比如导出时定义的"input_images:0") # 注意:要和导出模型时的Placeholder名称完全一致,包括后缀":0" original_placeholder_name = "input_images:0" # 3. 导入GraphDef时,通过input_map参数替换Placeholder # 把原Placeholder映射到我们的批量预处理图像张量 tf.import_graph_def( graph_def, input_map={original_placeholder_name: batch_images}, name='' # 保持原模型的节点名称空间,避免前缀干扰 ) # 4. 现在可以获取导入模型的输出节点,进行后续评估 logits = graph.get_tensor_by_name("logits:0") # 替换成你模型的输出节点名称 predictions = tf.argmax(logits, 1)
注意事项
- 张量匹配:预处理后的
batch_images的形状、数据类型(比如float32)必须和原Placeholder完全一致,否则会报错。 - 节点名称:导入GraphDef时如果设置了
name参数,后续获取节点时要加上对应的前缀,比如name='import'的话,节点名称会变成import/logits:0。 - 预处理一致性:评估时的预处理逻辑必须和训练、导出模型时完全相同,否则会影响结果准确性。
内容的提问来源于stack exchange,提问作者mateja
相关产品推荐
相关产品推荐

