SSD的简单实现 SSD的实现定义一个分类预测器对特征图上的每个空间位置每个像素做预测num_inputs输入特征图的通道数num_anchors每个像素位置生成的锚框数量num_classes目标类别数返回一个 2D 卷积层输出通道数num_anchors * (num_classes 1)1 是为了包含背景类每个锚框输出 (num_classes 1) 个值偏移预测器输出通道数num_anchors * 4每个锚框需要预测 4 个偏移量每个锚框独立拥有一组 4 个偏移量互不共享。一个包装函数forward返回block的值问题当输入尺寸不一样时我们的输出尺寸也不一样不好统一处理统一格式我们首先将通道维移到最后一维。因为不同尺度下批量大小仍保持不变我们可以将预测结果转成二维的批量大小高*宽*通道数的格式以方便之后在维度1上的连结。对列表中的每个尺度预测调用 flatten_pred然后在 dim1样本内的预测维度上拼接为了在多个尺度下检测目标我们在下面定义了高和宽减半块down_sample_blk该模块将输入特征图的高度和宽度减半。下采样块两遍Conv2d卷积层-BN-ReLU最后的MaxPool2d(2)使高宽减半基础特征提取网络从原始图像中抽取特征的骨干网。*真实的 SSD 通常用 VGG-16 或 ResNet 作为基础网络num_filters 定义了每一层的通道数3输入是 RGB 三通道图像16 → 32 → 64每经过一个下采样块通道数翻倍串联所有的下采样块打包返回一个顺序采样的块通道数变多尺寸变小我们整个模型定义为五个块组成定义每个块的前向传播函数输入 X 经过当前模块 blk得到输出特征图 Y。基于 Y 的空间尺寸生成锚框。分类头对 Y 做卷积输出每个锚框的类别分数。只根据特征图预测回归头对 Y 做卷积输出每个锚框的 4 个偏移量。返回一个元组包含4个东西Y当前模块输出的特征图传给下一个模块anchors当前尺度生成的所有锚框cls_preds这些锚框的类别预测bbox_preds这些锚框的位置偏移预测完整的模型继承 nn.Module接收类别数。Idx_to_in_channels5 个模块输出特征图的通道数用setattr动态创建属性循环 5 次每次创建 3 个东西最终有3*5个模块*fblk_{i}格式化字符串生成 blk_0, blk_1 ... blk_4创建特征提取块等价于self.blk_0 base_net()创建分类预测头创建边界框回归头创建 3 个长度为 5 的列表准备存放 5 个模块的输出anchors[i]第 i 个模块生成的锚框cls_preds[i]第 i 个模块的类别预测bbox_preds[i]第 i 个模块的偏移预测getattr(self, blk_0)取下来得到 base_net() 实例cat拼接所有锚框拼接并reshape分类预测拼接回归预测训练模型读取参数初始化参数定义优化算法定义损失函数分类交叉熵回归L1绝对值定义两个基础损失cls_preds模型输出的类别预测 (B, N, num_classes)cls_labels每个锚框的真实类别标签 (B, N)bbox_preds模型输出的偏移预测 (B, N*4)bbox_labels每个正样本锚框的真实偏移量 (B, N*4)bbox_masks掩码标记哪些锚框参与回归损失计算cls_loss(...)计算每个锚框的交叉熵损失返回形状 (B×N,)*这里不断地reshape是为了符合函数的参数要求bbox_preds * bbox_masks用掩码屏蔽掉不参与回归的锚框。bbox_labels * bbox_masks同理真实偏移也做掩码bbox_loss(...)计算 L1 损失绝对值差返回形状 (B, N*4)。两个损失直接相加也可以加权但这里简化为等权。最终返回 (B,)即每个样本的总损失。定义评价函数cls_preds.argmax(dim-1)在最后一维类别维度取最大值索引得到预测的类别 (B, N)。.type(cls_labels.dtype)把预测结果的数据类型转成和标签一致比如从 long 转 int 等。 cls_labels逐元素比较预测对的位置为 True即 1错的为 False即 0。.sum()统计预测正确的总个数。边界框平均绝对误差bbox_labels - bbox_preds真实偏移与预测偏移的差。整体的训练模型代码训练前准备外层循环遍历 epoch内层循环遍历每个batch调用模型前向传播multibox_target 把图像级别的真实标注转换成锚框级别的训练标签计算损失反向传播参数更新累加器把 4 个数累加起来用于最后算 epoch 级别的平均指标评估函数对每个锚框的类别分数做 softmax变成概率调整维度顺序从 (B, N, C) 变成 (B, C, N)output d2l.multibox_detection(cls_probs, bbox_preds, anchors)真实框锚框偏移量修正