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

如何修改单类目标检测训练代码以支持两类训练?

两类目标检测模型训练的代码修改方案

要把原单类训练代码改成两类,核心是调整类别定义、标签存储和独热编码逻辑,具体步骤如下:

1. 更新类别参数与索引

首先明确两类的ID(遵循非背景类从1开始的惯例),并更新category_index:

# 定义两类的class ID
duck_class_id = 1
second_class_id = 2  # 替换成你的第二类ID,比如对应自定义类别名称
num_classes = 2

# 构建包含两类的category_index
category_index = {
    duck_class_id: {'id': duck_class_id, 'name': 'rubber_ducky'},
    second_class_id: {'id': second_class_id, 'name': 'toy_car'}  # 替换成你的第二类名称
}

2. 核心:为每个标注框关联类别标签

原代码默认所有框都是同一类,现在你必须为每张图片的每个标注框记录对应的类别ID:

  • 新增一个gt_classes列表,其中gt_classes[i]是第i张图片中所有标注框的类别ID数组,形状和gt_boxes[i]的行数一致(每个框对应一个类别ID)。

然后修改数据预处理循环,处理类别标签:

label_id_offset = 1
train_image_tensors = []
gt_classes_one_hot_tensors = []
gt_box_tensors = []

# 同时遍历图片、框坐标和对应类别标签
for train_image_np, gt_box_np, gt_class_ids in zip(train_images_np, gt_boxes, gt_classes):
  train_image_tensors.append(tf.expand_dims(tf.convert_to_tensor(train_image_np, dtype=tf.float32), axis=0))
  gt_box_tensors.append(tf.convert_to_tensor(gt_box_np, dtype=tf.float32))
  
  # 将原始类别ID(1/2)转换为模型要求的0索引(0/1)
  zero_indexed_classes = tf.convert_to_tensor(gt_class_ids - label_id_offset, dtype=tf.int32)
  
  # 生成对应2类的独热编码
  gt_classes_one_hot_tensors.append(tf.one_hot(zero_indexed_classes, num_classes))
print('Done prepping data.')

3. 其他必要调整

  • 标注环节:标注图片时,要给每个框选择对应的类别(比如鸭子选1,第二类选2),并把这些ID整理成gt_classes列表,这是区分两类的前提。
  • 模型配置:后续加载预训练模型时,要将模型配置中的num_classes参数从1改为2,保证模型输出适配两类检测。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 19:40:29