如何在TensorFlow的TakeDataset对象上使用file_paths属性
报错原因
tf.keras.utils.image_dataset_from_directory 返回的原生 BatchDataset 对象自带 file_paths 属性,但调用 take()/skip() 方法后生成的 TakeDataset/SkipDataset 属于衍生数据集,不会继承该属性,因此直接调用会抛出属性不存在的错误。
可行解决方案
方案1:从原始验证集路径中截取(适配现有代码,实现最简单)
你创建验证集时固定了seed=123,数据集的shuffle顺序完全固定,拆分后的测试集、验证集顺序和原始验证集的file_paths顺序完全对应,可以直接按拆分的样本数截取:
# 先获取原始验证集的全量文件路径 val_all_paths = val_ds.file_paths # 计算测试集的样本总数:测试集批次数量 * 单批次大小 test_sample_count = (val_batches // 5) * batch_size # 按顺序截取测试集、剩余验证集的路径 test_file_paths = val_all_paths[:test_sample_count] new_val_file_paths = val_all_paths[test_sample_count:]
得到的test_file_paths和你拆分后的test_ds样本顺序完全一一对应,预测后可以直接和结果拼接生成DataFrame。
方案2:构造带路径的数据集(逻辑更稳妥,避免顺序错乱)
如果后续要修改shuffle逻辑,也可以直接把路径嵌入数据集的返回值中,拆分后也能直接获取:
# 构造和验证集顺序完全匹配的路径数据集 path_ds = tf.data.Dataset.from_tensor_slices(val_ds.file_paths) # 将路径和原始验证集的(图像、标签)数据绑定 val_ds_with_path = tf.data.Dataset.zip((val_ds, path_ds)) # 再执行拆分操作,拆分后的数据集自带路径字段 val_batches = tf.data.experimental.cardinality(val_ds_with_path) test_ds_with_path = val_ds_with_path.take(val_batches // 5) val_ds_with_path = val_ds_with_path.skip(val_batches // 5)
预测时遍历test_ds_with_path即可同时拿到单批次的图像、标签、对应文件路径,不需要额外维护顺序映射。
内容的提问来源于stack exchange,提问作者Xanthan Gum
相关产品推荐
相关产品推荐

