基于InceptionNet与ResNet50的自定义SSD目标检测模型构建咨询
自定义SSD模型的Keras损失函数与评估指标方案
一、损失函数选择(基于Keras原生工具)
SSD目标检测任务包含类别分类和边框回归两个子任务,需将两类损失加权组合,完全可以用Keras现有函数封装实现,无需从零开发:
- 分类损失:选用
keras.losses.SparseCategoricalCrossentropy(from_logits=True)(标签为整数格式时)或CategoricalCrossentropy(标签为one-hot格式时),适配多类别分类需求,同时仅对正样本(非背景类)计算损失。 - 回归损失:SSD标准用Smooth L1损失,Keras中的
keras.losses.Huber(delta=1.0)与Smooth L1等价,完美适配边框回归的误差处理。 - 组合损失:自定义Keras Loss类封装上述两种损失的加权求和(通常分类损失权重设为1,回归损失权重设为10,可根据任务调整),示例代码如下:
import tensorflow as tf from tensorflow import keras class SSDLoss(keras.losses.Loss): def __init__(self, cls_weight=1.0, reg_weight=10.0, name="ssd_loss"): super().__init__(name=name) self.cls_weight = cls_weight self.reg_weight = reg_weight self.cls_loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction='none') self.reg_loss_fn = keras.losses.Huber(delta=1.0, reduction='none') def call(self, y_true, y_pred): # 拆分真实标签与预测结果的回归、分类部分 y_true_reg = y_true[..., :4] y_true_cls = tf.argmax(y_true[..., 4:], axis=-1) y_pred_reg = y_pred[..., :4] y_pred_cls = y_pred[..., 4:] # 生成正样本掩码(仅计算非背景类的损失) pos_mask = tf.cast(y_true_cls != 0, tf.float32) # 计算分类损失 cls_loss = self.cls_loss_fn(y_true_cls, y_pred_cls) cls_loss = tf.reduce_sum(cls_loss * pos_mask) / tf.maximum(tf.reduce_sum(pos_mask), 1.0) # 计算回归损失 reg_loss = self.reg_loss_fn(y_true_reg, y_pred_reg) reg_loss = tf.reduce_sum(reg_loss * pos_mask[:, :, tf.newaxis]) / tf.maximum(tf.reduce_sum(pos_mask), 1.0) # 加权组合总损失 return self.cls_weight * cls_loss + self.reg_weight * reg_loss
二、评估指标选择
目标检测核心指标可基于Keras的Metric类封装实现,无需依赖外部工具:
- 平均精度均值(mAP):这是目标检测的核心评估指标,需对每个类别计算精确率-召回率(PR)曲线下面积,再求均值。可通过Keras
Metric类结合tf.image.non_max_suppression实现,示例简化版代码如下:
class MeanAveragePrecision(keras.metrics.Metric): def __init__(self, num_classes, iou_threshold=0.5, name="mAP"): super().__init__(name=name) self.num_classes = num_classes self.iou_threshold = iou_threshold self.class_ap = self.add_weight(name="class_ap", shape=(num_classes,), initializer="zeros") self.batch_count = self.add_weight(name="batch_count", initializer="zeros") def update_state(self, y_true, y_pred, sample_weight=None): for batch_idx in range(tf.shape(y_true)[0]): # 筛选当前样本的真实正框(非背景) true_boxes = y_true[batch_idx] pos_true_mask = tf.cast(tf.argmax(true_boxes[...,4:], axis=-1) != 0, tf.bool) true_boxes = tf.boolean_mask(true_boxes, pos_true_mask) # 对预测框做非极大值抑制(NMS) pred_scores = tf.reduce_max(y_pred[batch_idx][...,4:], axis=-1) pred_classes = tf.argmax(y_pred[batch_idx][...,4:], axis=-1) nms_indices = tf.image.non_max_suppression( y_pred[batch_idx][...,:4], pred_scores, max_output_size=100, iou_threshold=0.5 ) pred_boxes = tf.gather(y_pred[batch_idx], nms_indices) pred_classes = tf.gather(pred_classes, nms_indices) # 计算每个类别的AP(简化版,实际可按VOC标准实现PR曲线积分) for cls in range(1, self.num_classes): cls_true = true_boxes[tf.argmax(true_boxes[...,4:], axis=-1) == cls] cls_pred = pred_boxes[pred_classes == cls] if tf.shape(cls_true)[0] == 0 and tf.shape(cls_pred)[0] == 0: ap = 1.0 elif tf.shape(cls_true)[0] == 0 or tf.shape(cls_pred)[0] == 0: ap = 0.0 else: # 计算IoU匹配 ious = tf.map_fn(lambda pred: tf.reduce_max(tf.image.iou(pred[:4], cls_true[...,:4])), cls_pred) matches = ious >= self.iou_threshold tp = tf.reduce_sum(tf.cast(matches, tf.float32)) fp = tf.shape(cls_pred)[0] - tp fn = tf.shape(cls_true)[0] - tp precision = tp / (tp + fp + 1e-6) recall = tp / (tp + fn + 1e-6) ap = (precision + recall) / 2 # 简化计算,建议替换为PR曲线积分 self.class_ap.assign_add(tf.one_hot(cls, self.num_classes) * ap) self.batch_count.assign_add(1.0) def result(self): return tf.reduce_mean(self.class_ap / tf.maximum(self.batch_count, 1.0)) def reset_state(self): self.class_ap.assign(tf.zeros_like(self.class_ap)) self.batch_count.assign(0.0)
- 辅助指标:还可实现IoU均值(mIoU)、单类别精确率/召回率,均基于Keras
Mean或Precision/Recall类封装即可,核心逻辑是先通过IoU匹配正样本,再统计对应指标。
三、实用建议
- 骨干拼接前需确保特征图尺寸/通道数匹配,可使用
tf.keras.layers.Resizing或1×1卷积层调整特征维度后再拼接。 - 训练时搭配Keras原生回调函数,如
ModelCheckpoint保存最优模型、EarlyStopping防止过拟合、TensorBoard可视化训练过程。 - 数据增强可采用
tf.keras.layers.experimental.preprocessing中的随机裁剪、翻转、亮度调整等工具,适配目标检测任务需求。
内容的提问来源于stack exchange,提问作者Parkirat Sandhu
相关产品推荐
相关产品推荐

