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

如何打印TensorFlow数据集标签及获取其全部可用标签?

解答:打印TensorFlow数据集标签及获取全部可用标签

嘿,这两个需求其实都能通过你已经加载的dataset_info对象轻松实现,我来一步步给你说明:

1. 获取并打印数据集中的全部可用标签

你已经通过tfds.load的with_info=True参数拿到了dataset_info,这里面就包含了标签的所有元信息。要获取所有标签名称,直接访问dataset_info.features['label'].names即可——它会返回一个包含所有标签名称的列表,完美解决你不知道所有标签的问题。

在你的代码里添加这段就能打印所有标签:

# 获取所有标签名称列表
all_labels = dataset_info.features['label'].names
print("数据集中的全部可用标签:")
for idx, label_name in enumerate(all_labels):
    print(f"标签{idx}: {label_name}")

2. 打印数据集中单个样本的标签

因为你设置了as_supervised=True,所以数据集里的每个example都是(图像张量, 标签张量)的元组。你可以把标签张量转成整数,再对应到上面的标签名称列表,就能得到易读的标签信息。

比如修改你遍历训练集的代码,同时打印前几个样本的标签:

# 打印训练集前5个样本的标签(含名称)
print("\n训练集部分样本的标签信息:")
for idx, (image, label) in enumerate(training_set.take(5)):
    label_id = label.numpy()  # 将张量转为整数
    label_name = all_labels[label_id]
    print(f"第{idx+1}个样本:标签ID={label_id},标签名称={label_name}")

整合后的完整代码

把上面的逻辑加到你原有的代码里,最终代码如下:

import tensorflow_datasets as tfds

splits = tfds.Split.ALL.subsplit(weighted=(70,30))
(training_set, validation_set), dataset_info = tfds.load(
    'tf_flowers', 
    with_info=True, 
    as_supervised=True,
    split=splits
)

num_classes = dataset_info.features['label'].num_classes
num_training_examples = 0
num_validation_examples = 0

# 计数训练集样本
for example in training_set:
    num_training_examples += 1
# 计数验证集样本
for example in validation_set:
    num_validation_examples += 1

# 打印基础统计信息
print('Total Number of Classes: {}'.format(num_classes))
print('Total Number of Training Images: {}'.format(num_training_examples))
print('Total Number of Validation Images: {} \n'.format(num_validation_examples))

# 1. 获取并打印全部可用标签
all_labels = dataset_info.features['label'].names
print("数据集中的全部可用标签:")
for idx, label_name in enumerate(all_labels):
    print(f"标签{idx}: {label_name}")

# 2. 打印训练集前5个样本的标签
print("\n训练集部分样本的标签信息:")
for idx, (image, label) in enumerate(training_set.take(5)):
    label_id = label.numpy()
    label_name = all_labels[label_id]
    print(f"第{idx+1}个样本:标签ID={label_id},标签名称={label_name}")

简单解释下:dataset_info.features['label']是一个*CategoryFeature*对象,它的names属性存储了所有标签的人类可读名称,num_classes则是标签的总数。通过label.numpy()可以把TensorFlow的标签张量转换成Python整数,再和all_labels对应就能得到标签名称啦。

内容的提问来源于stack exchange,提问作者Kulothungan Senthil

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:16:51