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
相关产品推荐
相关产品推荐

