如何提取TFRecords所有标签并使用Python TensorFlow绘制标签直方图
解决方案
现有代码的问题
- 仅调用了
ds_train.take(10),只能获取前10个批次的标签,没有遍历整个TFRecords对应的数据集 - 直接追加TensorFlow张量对象到列表,绘图时会出现类型不兼容问题
sns.distplot已在seaborn v0.11版本后被弃用,更建议用sns.histplot替代
完整实现代码
首先确保你已经正确解析TFRecords为(image, label)格式的tf.data.Dataset对象,再执行以下代码:
import tensorflow as tf import seaborn as sns import matplotlib.pyplot as plt import numpy as np all_label = [] # 直接遍历整个数据集,不要加take()限制即可获取全部标签 for image, label in ds_train: # 将张量转为numpy数组后再添加到列表 all_label.append(label.numpy()) # 把嵌套的数组拉平为一维数组,适配绘图要求 all_label = np.concatenate(all_label).ravel() # 绘制直方图 sns.histplot(all_label, kde=True) plt.xlabel("标签值") plt.ylabel("样本数量") plt.title("TFRecords全量标签分布直方图") plt.show()
性能优化方案
如果数据集非常大,不想占用过多内存,可以跳过图像加载环节,大幅提升提取速度:
# 定义仅提取标签的TFRecords解析函数,替换原有全量解析函数 def parse_tfrecord_only_label(example_proto): feature_description = { # 仅保留标签的特征定义,删除图像的特征定义 'label': tf.io.FixedLenFeature([], tf.int64), } example = tf.io.parse_single_example(example_proto, feature_description) return example['label'] # 重新加载TFRecords文件,使用上述解析函数得到仅含标签的数据集 ds_only_label = tf.data.TFRecordDataset(你的TFRecords文件路径列表).map(parse_tfrecord_only_label) # 提取标签速度会提升数倍 all_label = [label.numpy() for label in ds_only_label]
内容的提问来源于stack exchange,提问作者user3452134
相关产品推荐
相关产品推荐

