使用TensorFlow Dataset API训练CNN后,保存图时.meta文件过大问题求助
我之前在用TensorFlow的Dataset API搭配可馈送迭代器训练图像模型时,也碰到过一模一样的问题——训练完保存的.meta文件大得离谱,后来排查下来发现几个关键原因,对应的解决方法分享给你:
这是最常见的原因!如果你是用tf.data.Dataset.from_tensor_slices()直接把CIFAR-10的numpy数组(比如加载到内存的所有训练图和标签)转成Dataset,那这些数据会被序列化到计算图里——CIFAR-10光是训练集就有5万张32×32×3的图片,算下来近150MB,加上其他数据,.meta文件自然暴大。
解决办法:改用从磁盘文件读取数据的方式构建Dataset,比如CIFAR-10的二进制格式文件,用tf.data.FixedLengthRecordDataset来读取:
# 定义CIFAR-10二进制文件的结构 record_bytes = 1 + 32*32*3 # 1字节标签 + 32×32×3字节图像数据 def parse_cifar10_record(record): # 解析二进制记录 features = tf.io.decode_raw(record, tf.uint8) label = tf.cast(features[0], tf.int32) image = tf.reshape(features[1:], [32, 32, 3]) image = tf.cast(image, tf.float32) / 255.0 # 归一化 return image, label # 从磁盘文件构建Dataset train_dataset = tf.data.FixedLengthRecordDataset( filenames=['cifar-10-batches-bin/data_batch_1.bin', 'cifar-10-batches-bin/data_batch_2.bin', ...], record_bytes=record_bytes ).map(parse_cifar10_record).shuffle(10000).batch(32)
这样计算图里只有读取文件和解析的操作,不会包含整个数据集的原始数据,.meta文件体积会立刻降下来。
如果你坚持要用内存中的数据,或者数据处理链比较复杂,可以把模型的输入做成独立的占位符,训练时把迭代器输出的batch数据feed给占位符,这样模型结构和数据处理逻辑在计算图里是分开的,保存模型时只保存模型相关的部分:
# 第一步:构建独立的模型图,输入用占位符 input_images = tf.placeholder(tf.float32, shape=[None, 32, 32, 3], name='input_images') input_labels = tf.placeholder(tf.int32, shape=[None], name='input_labels') # 构建你的CNN模型 def cnn_model(inputs): x = tf.layers.conv2d(inputs, 32, (3,3), activation='relu') x = tf.layers.max_pooling2d(x, (2,2), 2) x = tf.layers.conv2d(x, 64, (3,3), activation='relu') x = tf.layers.max_pooling2d(x, (2,2), 2) x = tf.layers.flatten(x) x = tf.layers.dense(x, 128, activation='relu') return tf.layers.dense(x, 10) logits = cnn_model(input_images) loss = tf.losses.sparse_softmax_cross_entropy(labels=input_labels, logits=logits) train_op = tf.train.AdamOptimizer().minimize(loss) # 第二步:数据处理部分(训练时才会执行,不影响模型图) train_dataset = tf.data.Dataset.from_tensor_slices((train_x, train_y)).shuffle(1000).batch(32) train_iter = train_dataset.make_one_shot_iterator() next_batch = train_iter.get_next() # 训练与保存 with tf.Session() as sess: sess.run(tf.global_variables_initializer()) for step in range(10000): batch_x, batch_y = sess.run(next_batch) sess.run(train_op, feed_dict={input_images: batch_x, input_labels: batch_y}) # 只保存模型的可训练变量(权重、偏置等),不保存数据处理相关节点 saver = tf.train.Saver(var_list=tf.trainable_variables()) saver.save(sess, './cifar10_cnn_model')
这样保存的.meta文件只会包含模型的结构和变量,数据处理的迭代器、Dataset节点都不会被保存。
TensorFlow的SavedModel是更现代、更高效的模型保存格式,它可以明确指定要导出的模型输入输出签名,自动过滤掉训练相关和数据处理的冗余节点,而且体积更小,部署也更方便:
# 训练完成后,用SavedModelBuilder导出模型 builder = tf.saved_model.builder.SavedModelBuilder('./saved_model') # 定义模型的输入输出签名 input_signature = tf.saved_model.utils.build_tensor_info(input_images) label_signature = tf.saved_model.utils.build_tensor_info(input_labels) output_signature = tf.saved_model.utils.build_tensor_info(logits) prediction_signature = tf.saved_model.signature_def_utils.build_signature_def( inputs={'images': input_signature, 'labels': label_signature}, outputs={'logits': output_signature}, method_name=tf.saved_model.signature_constants.PREDICT_METHOD_NAME ) with tf.Session() as sess: sess.run(tf.global_variables_initializer()) # ... 训练过程 ... builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.TRAINING, tf.saved_model.tag_constants.SERVING], signature_def_map={'predict': prediction_signature} ) builder.save()
导出的SavedModel文件夹里只会保留模型推理和训练必要的部分,数据处理的冗余节点会被自动排除。
如果已经训练完,想拯救现有的大.meta文件,可以用graph_util工具清理掉训练相关的冗余节点,生成冻结的.pb文件(体积远小于.meta):
from tensorflow.python.framework import graph_util with tf.Session() as sess: # 加载已保存的模型 saver = tf.train.import_meta_graph('./cifar10_cnn_model.meta') saver.restore(sess, './cifar10_cnn_model') # 指定模型的输出节点名称(比如你的logits节点名,可通过tensorboard查看) output_node_names = ['dense_1/BiasAdd'] # 替换成你实际的输出节点名 graph_def = tf.get_default_graph().as_graph_def() # 移除训练节点,只保留推理必要的节点 trimmed_graph_def = graph_util.remove_training_nodes(graph_def) # 把变量转成常量,生成冻结图 frozen_graph_def = graph_util.convert_variables_to_constants( sess, trimmed_graph_def, output_node_names ) # 保存冻结后的图 with tf.gfile.GFile('./frozen_cifar10_model.pb', 'wb') as f: f.write(frozen_graph_def.SerializeToString())
这个.pb文件可以直接用于推理,体积会小很多,而且没有冗余的训练和数据处理节点。
内容的提问来源于stack exchange,提问作者pjao4512

