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

如何提取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 22:27:03