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

TensorFlow TFRecord格式下k-hot编码多分类实现方法问询

在TFRecord中实现k-hot多分类标签的最优方案

嘿,针对你想在目标检测TFRecord数据集里用k-hot编码替代one-hot做多分类的需求,完全不需要重复存储同一个Bounding Box,这里有两个更高效的方案:

方案1:直接存储k-hot向量字段

你可以在每个tf.train.Example的目标对象字段里,新增一个专门存储k-hot标签的特征,直接把[0,1,0,1,0]这样的向量存进去。这种方式最直观,后续模型解析也很直接。

生成TFRecord时的代码示例:

def create_tf_example(annotation):
    # 先处理图像路径、尺寸、bbox坐标等常规字段
    img_path = annotation['image_path']
    xmins = [annotation['bbox']['xmin']]
    xmaxs = [annotation['bbox']['xmax']]
    ymins = [annotation['bbox']['ymin']]
    ymaxs = [annotation['bbox']['ymax']]
    
    # 定义当前bbox的k-hot标签
    k_hot_label = [0, 1, 0, 1, 0]  # 对应你的多分类需求
    
    # 构建TF Example
    tf_example = tf.train.Example(features=tf.train.Features(feature={
        'image/height': dataset_util.int64_feature(annotation['height']),
        'image/width': dataset_util.int64_feature(annotation['width']),
        'image/filename': dataset_util.bytes_feature(img_path.encode('utf8')),
        'image/source_id': dataset_util.bytes_feature(img_path.encode('utf8')),
        'image/encoded': dataset_util.bytes_feature(open(img_path, 'rb').read()),
        'image/format': dataset_util.bytes_feature('jpeg'.encode('utf8')),
        'image/object/bbox/xmin': dataset_util.float_list_feature(xmins),
        'image/object/bbox/xmax': dataset_util.float_list_feature(xmaxs),
        'image/object/bbox/ymin': dataset_util.float_list_feature(ymins),
        'image/object/bbox/ymax': dataset_util.float_list_feature(ymaxs),
        # 新增k-hot标签字段
        'image/object/class/k_hot_label': dataset_util.int64_list_feature(k_hot_label),
    }))
    return tf_example

之后在模型的输入解析环节,你需要把这个字段解析成张量,并且把损失函数替换为多标签分类损失(比如BinaryCrossentropy),因为每个类别是独立的二分类任务,不再是互斥的one-hot分类。

方案2:存储多类别索引列表(更节省空间)

如果你的类别数量很多,直接存k-hot向量会浪费存储空间,这时候可以只存储当前bbox所属的类别索引(比如[2,4]表示属于第2和第4类),后续在模型端再转换成k-hot向量。

生成TFRecord时的代码示例:

def create_tf_example(annotation):
    # 常规字段处理...
    # 存储类别索引列表
    multi_class_indices = [2, 4]  # 注意类别索引的起始值(是从0还是1开始)
    
    tf_example = tf.train.Example(features=tf.train.Features(feature={
        # 其他常规字段...
        # 新增多类别索引字段
        'image/object/class/multi_indices': dataset_util.int64_list_feature(multi_class_indices),
    }))
    return tf_example

在模型解析时,你可以用tf.one_hot结合tf.reduce_max来生成k-hot向量:

def parse_tfrecord_fn(example_proto):
    feature_description = {
        # 其他特征描述...
        'image/object/class/multi_indices': tf.io.VarLenFeature(tf.int64),
    }
    parsed_features = tf.io.parse_single_example(example_proto, feature_description)
    
    # 将稀疏张量转为密集张量,再生成k-hot向量
    multi_indices = tf.sparse.to_dense(parsed_features['image/object/class/multi_indices'])
    k_hot_label = tf.reduce_max(tf.one_hot(multi_indices, depth=5), axis=0)  # depth是总类别数
    
    # 其他解析逻辑...
    return parsed_features, k_hot_label

为什么不推荐重复Bounding Box?

重复存储同一个bbox会带来两个明显的问题:

  • 数据冗余:同一个bbox的坐标、图像信息等重复存储,会大幅增大TFRecord的体积,降低数据加载效率。
  • 训练逻辑混乱:后续训练时,模型会把重复的bbox当成不同的目标,计算IOU、损失时都会出现重复计算,容易引入不必要的误差,还需要额外处理去重逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 07:11:44