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

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,本质也是从文件读取数据,数据不进入图定义),所以图的体积只包含计算逻辑,自然很小。

从非TFRecord格式的图像文件创建数据集的正确方法?

分两种常见场景来给你说:

场景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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 04:15:00