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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 04:10:15