TensorFlow自定义验证数据集访问首个元素时无限加载问题求助
问题原因
- 核心问题是数据集流水线的操作顺序错误:你先执行了
map预处理操作(包含图片IO读取、解码),再使用take/skip拆分数据集。 - TensorFlow的
tf.data是懒执行机制,调用skip(n_train_items)获取验证集数据时,每一次迭代都会从头遍历整个数据集,先执行前n_train_items个样本的完整预处理流程(包含40多万次图片IO读取操作),跳过这些样本后才会返回验证集数据,这就是验证集首元素加载耗时极长的根因。 - 当你简化
prep_image仅返回文件名时,预处理没有IO开销,skip操作的耗时可以忽略,所以验证集可以正常访问;一旦加入tf.io.read_file这类IO操作,数十万次IO的开销直接导致进程看起来像无限加载。
解决方案
调整数据集流水线的操作顺序,先拆分路径数据集,再分别做预处理,修改后的代码如下:
buffer_size = 1024 batch_size = 32 def prep_image(filename, img_shape=(144,144), channels=3, dtype=tf.float32): image_string = tf.io.read_file(filename) image = tf.image.decode_jpeg(image_string, channels=channels) image = tf.image.convert_image_dtype(image, dtype) image = tf.image.resize(image, img_shape) return image def preprocess_triplet(a, p, n): return prep_image(a), prep_image(p), prep_image(n) # 先对原始路径数据集做shuffle和拆分,此时仅操作字符串无IO开销 # 如果需要固定拆分结果,可给shuffle设置固定seed,比如shuffle(buffer_size, seed=123) dataset = dataset.shuffle(buffer_size) split = .8 n_train_items = round(len(dataset) * split) # 先拆分纯路径数据集 train_path_dataset = dataset.take(n_train_items) val_path_dataset = dataset.skip(n_train_items) # 拆分后再分别对训练集、验证集做预处理 # 训练集流水线 train_dataset = train_path_dataset.map(preprocess_triplet, num_parallel_calls=tf.data.AUTOTUNE) train_dataset = train_dataset.shuffle(buffer_size) train_dataset = train_dataset.batch(batch_size, drop_remainder=False) train_dataset = train_dataset.prefetch(tf.data.AUTOTUNE) # 验证集流水线 val_dataset = val_path_dataset.map(preprocess_triplet, num_parallel_calls=tf.data.AUTOTUNE) val_dataset = val_dataset.batch(batch_size, drop_remainder=False) val_dataset = val_dataset.prefetch(tf.data.AUTOTUNE)
- 额外优化点:给
map操作加上num_parallel_calls=tf.data.AUTOTUNE参数,可以并行执行预处理操作,进一步提升数据加载速度。 - 若需要保证拆分结果可复现,只需给拆分前的
shuffle操作设置固定随机种子即可。
内容的提问来源于stack exchange,提问作者user7354365
相关产品推荐
相关产品推荐

