PrefetchDataset无class_names属性:如何获取TensorFlow数据集类别名
解决
PrefetchDataset对象无class_names属性的问题 问题场景
使用tf.keras.utils.image_dataset_from_directory创建训练/验证数据集后,调用train_ds.class_names获取类别名称时触发如下错误:
--------------------------------------------------------------------------- AttributeError Traceback (most recent call last) Cell In [29], line 1 ----> 1 train_ds.class_names AttributeError: 'PrefetchDataset' object has no attribute 'class_names'
数据集创建代码:
train_ds = tf.keras.utils.image_dataset_from_directory( data_dir, validation_split = 0.2, subset = "training", seed = 123, image_size = (img_height, img_width), batch_size = batch_size) val_ds = tf.keras.utils.image_dataset_from_directory( data_dir, validation_split = 0.2, subset = "validation", seed = 123, image_size = (img_height, img_width), batch_size = batch_size)
核心原因
image_dataset_from_directory返回的原始对象是BatchDataset,自带class_names属性;但如果后续对数据集执行了prefetch()(比如常见的train_ds = train_ds.prefetch(tf.data.AUTOTUNE)),会将其转换为PrefetchDataset类,该类没有class_names属性,因此报错。
解决方法
方法1:提前保存类别名称
在对数据集做任何预处理(包括prefetch)之前,先将类别名称存入变量:
# 创建数据集后立即提取并保存类别名 class_names = train_ds.class_names # 再执行后续预处理操作 train_ds = train_ds.prefetch(tf.data.AUTOTUNE) val_ds = val_ds.prefetch(tf.data.AUTOTUNE) # 后续直接使用class_names变量即可 print(class_names)
方法2:从数据集标签推导类别(已预处理后补救)
如果已经完成预处理无法回溯,可以通过遍历数据集标签提取唯一类别,或直接读取数据集目录下的文件夹名:
# 方式A:从数据集标签提取 import numpy as np all_labels = [] for _, labels in train_ds: all_labels.extend(labels.numpy()) # 去重并排序,保证和原始class_names顺序一致 unique_labels = np.unique(all_labels) # 方式B:直接读取数据集目录下的文件夹名 import os class_names = sorted([name for name in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, name))])
方法3:避免覆盖原始数据集对象
如果不需要复用原始数据集,可以用新变量名存储预处理后的数据集,保留原始对象的class_names属性:
# 用新变量存储预处理后的数据集,原始train_ds仍保留class_names processed_train_ds = train_ds.prefetch(tf.data.AUTOTUNE) processed_val_ds = val_ds.prefetch(tf.data.AUTOTUNE) # 仍可通过原始对象获取类别名 print(train_ds.class_names)
内容的提问来源于stack exchange,提问作者totames
相关产品推荐
相关产品推荐

