TensorFlow中使用Dataset.from_tensor_slices()导致日志过大的问题咨询
刚好踩过这个坑,来给你详细解答这两个问题:
tf.data.Dataset.from_tensor_slices()会导致日志和图文件暴涨? 你观察得很准,from_tensor_slices()确实会把传入的numpy数组(或其他张量)转换成tf.constant()嵌入到计算图中。虽然MNIST原始数据只有13M,但tf.constant()在图中存储的是张量的完整结构和每个元素的细节,而且TensorFlow序列化图的时候,会把这些常量的所有数据都包含进去——这就相当于把整个数据集“硬编码”进了图里。
更关键的是,事件文件(events.out.tfevents)和graph.pbtxt会完整记录图的所有节点信息,包括这些嵌入的常量数据,所以哪怕是13M的原始数据,转成图节点的序列化格式后体积会膨胀好几倍,最终达到几百兆并不奇怪。
而官方示例代码之所以没有这个问题,是因为它没有把数据集嵌到图里,而是用了从文件流读取的方式(哪怕是旧的队列API,本质也是从文件读取数据,数据不进入图定义),所以图的体积只包含计算逻辑,自然很小。
分两种常见场景来给你说:
场景1:单张独立的图像文件(如PNG/JPG)
如果你的图像是按类别存放在不同文件夹的单文件(比如class_0/img1.png、class_1/img2.jpg这种结构),最常用的方法是用tf.data.Dataset.list_files()获取所有文件路径,再通过map()解析每个文件:
import tensorflow as tf import os def parse_single_image(file_path): # 从文件名提取标签(假设路径结构是 ./data/class_name/image.png) label = tf.strings.split(file_path, os.sep)[-2] label = tf.strings.to_number(label, out_type=tf.int32) # 读取并解码图像 image_raw = tf.io.read_file(file_path) # 根据图像格式选择decode_png或decode_jpeg image = tf.image.decode_png(image_raw, channels=1) # MNIST是单通道灰度图 # 预处理:归一化、调整尺寸等 image = tf.cast(image, tf.float32) / 255.0 image = tf.reshape(image, (28, 28, 1)) return image, label # 构建数据集 file_dataset = tf.data.Dataset.list_files("./mnist_images/*/*.png") dataset = file_dataset.map(parse_single_image, num_parallel_calls=tf.data.AUTOTUNE) # 打乱、分批、预取优化,提升训练效率 dataset = dataset.shuffle(buffer_size=10000).batch(32).prefetch(tf.data.AUTOTUNE)
场景2:二进制打包的批量图像(如MNIST的idx文件)
像MNIST这种把所有图像打包成单个二进制文件的情况,用tf.data.FixedLengthRecordDataset是非常合适的——因为每个样本的字节长度是固定的,我们可以跳过文件头部,直接按固定长度读取每个样本:
import tensorflow as tf # MNIST二进制文件的头部信息:图像文件头部16字节,标签文件头部8字节 IMAGE_HEADER_BYTES = 16 LABEL_HEADER_BYTES = 8 IMAGE_BYTES = 28 * 28 # 每个MNIST图像的字节数 LABEL_BYTES = 1 # 每个标签的字节数 def parse_mnist_batch(image_bytes, label_bytes): # 解码图像字节为张量 image = tf.io.decode_raw(image_bytes, tf.uint8) image = tf.reshape(image, (28, 28, 1)) image = tf.cast(image, tf.float32) / 255.0 # 解码标签字节为张量 label = tf.io.decode_raw(label_bytes, tf.uint8) label = tf.reshape(label, ()) label = tf.cast(label, tf.int32) return image, label # 分别加载图像和标签数据集 image_dataset = tf.data.FixedLengthRecordDataset( "./train-images-idx3-ubyte", record_bytes=IMAGE_BYTES, header_bytes=IMAGE_HEADER_BYTES ) label_dataset = tf.data.FixedLengthRecordDataset( "./train-labels-idx1-ubyte", record_bytes=LABEL_BYTES, header_bytes=LABEL_HEADER_BYTES ) # 合并图像和标签 dataset = tf.data.Dataset.zip((image_dataset, label_dataset)) dataset = dataset.map(parse_mnist_batch, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.shuffle(60000).batch(32).prefetch(tf.data.AUTOTUNE)
核心总结
不管哪种场景,核心原则都是让数据从外部文件读取,而不是嵌入到计算图中,这样既能避免图文件和事件文件的异常膨胀,也能支持更大的数据集(不会因为把数据放进图而导致内存溢出)。
内容的提问来源于stack exchange,提问作者3voC

