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

如何在TF Lite Model Maker中为EfficientDet-Lite添加类别权重?

问题描述

使用TFLite Model Maker的EfficientDet-Lite系列模型进行目标检测时,希望通过类别权重解决类别不平衡问题,但直接在model.fit()中添加class_weight参数时触发报错:

ValueError: `class_weight` is only supported for Models with a single output.

原因是EfficientDet-Lite属于多输出模型(包含分类、回归等分支),原生不支持class_weight参数,需通过修改损失函数的方式注入类别权重。

解决方案

可以通过给EfficientDet的分类损失添加类别权重实现需求,仅需少量修改核心代码,具体步骤如下:

1. 计算类别权重

先根据训练数据统计每个类别的样本占比,计算平衡权重(示例采用balanced策略,即总样本数/(类别数×该类样本数)):

import numpy as np
from sklearn.utils.class_weight import compute_class_weight

# 从训练数据中提取所有有效类别标签(过滤背景类)
labels = []
for sample in train_data.gen_dataset():
    # groundtruth_classes形状为[batch_size, max_detections]
    batch_labels = sample[1]['groundtruth_classes'].numpy().flatten().tolist()
    labels.extend([label for label in batch_labels if label != 0])

# 计算类别权重
class_weights = compute_class_weight(
    class_weight='balanced',
    classes=np.unique(labels),
    y=labels
)

# 生成包含背景类的完整权重张量(背景类权重设为1)
num_classes = len(train_data.label_map)
full_class_weights = np.ones(num_classes)
for idx, weight in enumerate(class_weights, start=1):
    full_class_weights[idx] = weight
full_class_weights = tf.convert_to_tensor(full_class_weights, dtype=tf.float32)

2. 修改EfficientDet损失函数

找到TFLite Model Maker依赖的EfficientDet损失文件(路径通常为tensorflow_examples/lite/model_maker/third_party/efficientdet/loss.py),修改DetectionLoss类以支持类别权重:

class DetectionLoss(tf.keras.losses.Loss):
    def __init__(self, config, class_weights=None, **kwargs):
        super().__init__(**kwargs)
        self.config = config
        # 默认权重全为1,即不修改原有损失逻辑
        self.class_weights = class_weights if class_weights is not None else tf.ones(config.num_classes)

    def call(self, y_true, y_pred):
        """Computes detection loss."""
        cls_outputs = y_pred[0]
        box_outputs = y_pred[1]

        cls_targets = y_true[0]
        box_targets = y_true[1]
        num_positives = y_true[2]

        # 计算原始分类损失
        cls_loss = sigmoid_focal_loss(
            cls_outputs, cls_targets, alpha=self.config.alpha, gamma=self.config.gamma)
        
        # 应用类别权重:根据真实类别索引匹配对应权重,乘到分类损失上
        if self.class_weights is not None:
            class_indices = tf.argmax(cls_targets, axis=-1)
            weights = tf.gather(self.class_weights, class_indices)
            weights = tf.expand_dims(weights, axis=-1)  # 扩展维度匹配损失形状
            cls_loss = cls_loss * weights

        # 回归损失保持原有逻辑不变
        box_loss = huber_loss(box_outputs, box_targets, delta=self.config.delta)

        # 按正样本数归一化总损失
        cls_loss = tf.math.divide_no_nan(tf.reduce_sum(cls_loss), num_positives)
        box_loss = tf.math.divide_no_nan(tf.reduce_sum(box_loss), num_positives)
        total_loss = cls_loss + box_loss
        return total_loss

3. 在Model Maker中传入权重参数

修改tensorflow_examples/lite/model_maker/core/task/model_spec/object_detector_spec.py:

  1. 在EfficientDetLiteSpec基类中添加class_weights属性:
class EfficientDetLiteSpec(BaseSpec):
    def __init__(self, model_name, model_dir=None, hparams='', **kwargs):
        # 原有初始化代码保持不变
        self.class_weights = None  # 添加类别权重属性
        # 处理自定义参数
        for key, value in kwargs.items():
            setattr(self, key, value)
  1. 在build_model方法中,初始化损失函数时传入权重:
def build_model(self, num_classes):
    """Builds the EfficientDet-Lite model."""
    # 原有代码保持不变
    
    # 初始化带类别权重的损失函数
    loss_fn = loss_lib.DetectionLoss(
        self.config, class_weights=self.class_weights)
    
    model.compile(
        optimizer=self.optimizer,
        loss=loss_fn,
        steps_per_execution=self.steps_per_execution)
    return model

4. 在训练代码中配置权重

在你的训练脚本中,给创建的spec对象设置类别权重:

# (接原有训练代码)
spec = object_detector.EfficientDetLite0Spec(
    model_name = model_name,
    model_dir='/home/alex/checkpoints/',
    hparams='grad_checkpoint=true,strategy=gpus',
    epochs=epochs, batch_size=batch_size,
    steps_per_execution=1, moving_average_decay=0,
    var_freeze_expr='(efficientnet|fpn_cells|resample_p6)',
    tflite_max_detections=25
)

# 传入预计算的类别权重
spec.class_weights = full_class_weights

train_data, validation_data, test_data = object_detector.DataLoader.from_csv(file_path)

model = object_detector.create(train_data, model_spec=spec, train_whole_model=True, validation_data=validation_data)
说明
  • 该方案仅修改3处核心文件,无需重构整个训练流程,符合"少量改动"的需求。
  • 类别权重仅作用于分类损失,回归损失保持原有逻辑,符合目标检测的训练规律。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 12:57:14