mmsegmentation 项目实践:AdaBins 自适应分箱单目深度估计(EfficientNet-B5 + mViT 实现解析)
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
导读
本文围绕 mmsegmentation 仓库中projects/Adabins这一官方示例项目展开,系统讲解 AdaBins(Adaptive Bins)这一基于 Transformer 的单目深度估计方法:从论文核心思想(将深度范围自适应划分为可学习分箱、以箱中心线性组合得到最终深度)出发,结合仓库内真实可运行的源码与配置,逐一拆解其 EfficientNet-B5 编码器解码器骨干网络、mViT 解码头、NYU / KITTI 训练配置与评测指标,并说明如何在 mmseg 生态中通过DepthEstimator、DepthMetric等基础组件完成训练与验证。读完本文,你将掌握 AdaBins 的完整实现链路、关键超参数的含义,以及如何在 mmsegmentation 中复现或扩展该深度估计任务。
AdaBins 方法概述
AdaBins 由 S. A. Bhat、I. Alhashim 和 P. Wonka 提出,发表于 CVPR 2021,论文为《AdaBins: Depth Estimation Using Adaptive Bins》(arxiv: 2011.14141)。其核心思路是:在编码器-解码器卷积网络基础上,引入一个基于 Transformer 的构建模块,把预测的深度范围划分为若干分箱(bins),且箱宽随输入图像自适应变化;最终深度值由这些箱中心的线性组合得到。该模块被命名为AdaBins。
其论文摘要(见 projects/Adabins/README.md)指出:作者以 baseline 编码器-解码器 CNN 为起点,研究全局信息处理如何帮助提升整体深度估计精度,并提出一个 Transformer 构建块,将深度范围划分为箱,其中心值按每张图像自适应估计,最终深度为箱中心的线性组合。实验表明该方法在多个主流深度数据集上的所有指标上相较当时 SOTA 均有显著提升,并配套提供了代码与预训练权重。
论文总结的三大贡献:
- 提出一种执行场景信息全局处理的架构构建块,将预测深度范围划分为随图像变化的宽度可学分箱,最终深度估计是箱中心值的线性组合;
- 在 NYU 与 KITTI 两个最流行数据集的监督式单目深度估计上,所有指标均取得显著提升;
- 系统分析了所提 AdaBins 块的不同改进方式及其对深度估计精度的影响。
仓库中的工程化落地
projects/Adabins是 mmsegmentation 的独立可运行子项目,其目录结构与 mmseg 的"backbone + decode_head + 配置文件"约定完全一致:
projects/Adabins/ ├── README.md ├── backbones/ │ ├── __init__.py │ └── adabins_backbone.py # EfficientNet-B5 编码器 + 多级跳跃解码 ├── decode_head/ │ ├── __init__.py │ └── adabins_head.py # mViT 解码头(自适应分箱核心) └── configs/ ├── _base_/ │ ├── datasets/nyu.py # NYU 数据管线与 DepthMetric 评测配置 │ ├── models/Adabins.py # 模型骨架配置(backbone + decode_head) │ └── default_runtime.py # 运行时、日志、可视化等通用配置 └── adabins/ ├── adabins_efficient_b5_4x16_25e_NYU_416x544.py └── adabins_efficient_b5_4x16_25e_kitti_352x704.py其中backbones/adabins_backbone.py与decode_head/adabins_head.py均通过@MODELS.register_module()注册到 mmseg 的MODELS注册表,因此可以在配置文件中以字符串type='AdabinsBackbone'/type='AdabinsHead'直接引用。
骨干网络:EfficientNet-B5 编码器与跳跃式上采样解码器
骨干网络实现位于 backbones/adabins_backbone.py,整体是"编码器-解码器"结构:
编码器(Encoder)使用timm.create_model('tf_efficientnet_b5_ap', pretrained=True)加载 EfficientNet-B5 预训练权重,随后将global_pool与classifier替换为nn.Identity(),去掉分类头,只保留特征提取部分;forward会逐个模块前向并保存每一阶段的中间特征(features列表),供解码端跳跃连接使用。
解码器由 1×1 卷积conv2与四级UpSampleBN模块组成。每个UpSampleBN先将低分辨率特征双线性插值到与跳跃特征相同尺寸,再沿通道拼接后经过两个ConvModule(默认norm_cfg=dict(type='BN')、act_cfg=dict(type='LeakyReLU'))。各层跳跃输入通道数固定拼接了 EfficientNet-B5 各阶段的特征(112+64、40+24、24+16、16+8),最终由 3×3 卷积conv3输出num_classes=128通道的稠密特征图,供解码头使用。
关键参数(来自 Adabins.py 模型配置):
| 参数 | 取值 | 说明 |
|---|---|---|
basemodel_name | tf_efficientnet_b5_ap | timm 中带抗锯齿池化的 EfficientNet-B5 |
num_features | 2048 | 中间特征通道数(逐级减半:2048→1024→512→256→128) |
num_classes | 128 | 最终输出特征通道数,必须与解码头in_channels一致 |
bottleneck_features | 2048 | EfficientNet-B5 末层特征通道数 |
解码头:mViT 自适应分箱模块
解码头是 AdaBins 的核心,实现在 decode_head/adabins_head.py,由三个子模块协同完成"分箱预测 + 逐像素深度回归":
PatchTransformerEncoder(全局上下文编码)
将输入特征图用embedding_convPxP(核大小与步长均为patch_size的卷积)切成 patch 序列,得到形状为n, embedding_dim, s的嵌入;加上可学习的positional_encodings(形状(500, embedding_dim)的参数)后,转置为 Transformer 要求的S, N, E格式,送入 4 层nn.TransformerEncoderLayer(embedding_dim=128, num_heads=4, dim_feedforward=1024)。输出序列中:
- 第 0 个 token(
tgt[0])作为regression_head,用于全局回归 bin 宽度; - 第 1 到
n_query_channels+1个 token(tgt[1:n_query_channels+1])作为queries,参与逐像素注意力图计算。
PixelWiseDotProduct(逐像素查询点积)
将骨干输出的特征图x与queries做矩阵点积,得到形状为n, n_query_channels, h, w的range_attention_maps,即每个像素在各查询(深度区间)上的响应图。代码中通过assert c == ck保证特征通道数与查询嵌入维度一致(均为 128)。
分箱回归与深度合成
regressor(Linear(128,256) → LeakyReLU → Linear(256,256) → LeakyReLU → Linear(256, n_bins))从全局 token 回归出n_bins个归一化箱宽:
norm='linear'时先relu再加eps=0.1防止除零与负箱宽;norm='softmax'时直接返回torch.softmax(y, dim=1)与注意力图;- 其余情况走
sigmoid。
随后bin_widths_normed = y / y.sum(dim=1, keepdim=True)做归一化,乘以深度范围(max_val - min_val)得到绝对箱宽,左侧 pad 一个min_val后cumsum得到箱边界bin_edges,箱中心centers = 0.5 * (bin_edges[:, :-1] + bin_edges[:, 1:])。另一路conv_out(1×1 卷积 + Softmax)把注意力图映射为每个像素在各箱上的权重out。最终深度图:
pred = Σ_c out[c] * centers[c]即逐像素地对箱中心加权求和,这正是论文"最终深度值是箱中心线性组合"的工程化表达。
解码头关键参数(来自模型配置):
| 参数 | 默认值 | NYU 配置 | KITTI 配置 | 说明 |
|---|---|---|---|---|
in_channels | - | 128 | 128 | 骨干输出通道数 |
n_query_channels | 128 | 128 | 128 | 查询 token 数量 |
patch_size | 16 | 16 | 16 | 全局上下文 patch 大小 |
embedding_dim | 128 | 128 | 128 | Transformer 嵌入维度 |
num_heads | 4 | 4 | 4 | 注意力头数 |
n_bins | 100 | 256 | 256 | 分箱数量 |
min_val/max_val | 0.1 / 10 | 0.001 / 10 | 0.001 / 80 | 深度范围,KITTI 深达 80m |
norm | linear | linear | linear | 箱宽归一化方式 |
推理阶段predict()会执行forward并取最后一个输出(深度图),随后torch.clamp(pred, min_val, max_val)把预测裁剪到合法深度范围,并将inf替换为max_val、nan替换为min_val。
模型封装:基于 DepthEstimator 的深度估计任务
AdaBins 在 mmseg 中并不走语义分割的EncoderDecoder,而是复用仓库新增的深度估计器 mmseg/models/segmentors/depth_estimator.py 中的DepthEstimator(同样注册于MODELS,继承自EncoderDecoder)。其类注释给出了完整的调用链:
- 训练:
loss()→extract_feat()→_decode_head_forward_train()→decode_head.loss(); - 推理:
predict()→inference()→whole_inference()/slide_inference()/slide_flip_inference()→encode_decode()→decode_head.predict(); - 后处理:
postprocess_result()会依据img_meta去除 padding 区域、处理翻转,并将深度图双线性 resize 回ori_shape,最终以SegDataSample.pred_depth_map(PixelData)形式输出。
配置文件中model = dict(type='DepthEstimator', ...)即指定该封装,test_cfg=dict(mode='whole')表示整图推理。
数据管线与评测指标
NYU 数据配置
configs/base/datasets/nyu.py 定义了 NYU 深度估计数据管线:
- 数据集类型
NYUDataset,data_root='data/nyu',测试图像位于images/test、深度标注位于annotations/test; - 管线包含
LoadImageFromFile(to_float32=True)、LoadDepthAnnotation(depth_rescale_factor=1e-3,即把毫米级深度缩放到米)与PackSegInputs(meta_keys 中携带depth_map_path等); - 验证/测试评估器为
DepthMetric,配置max_depth_eval=10.0, crop_type='nyu_crop',其中nyu_crop是 NYU 官方评测采用的中心裁剪方式,剔除图像边界无效区域。
评测指标定义
mmseg/evaluation/metrics/depth_metric.py 中的DepthMetric支持 9 个标准深度估计指标:d1/d2/d3(阈值准确率 δ1/δ2/δ3)、abs_rel(相对绝对误差)、sq_rel(相对平方误差)、rmse、rmse_log、log10与silog,可通过depth_metrics参数按需选择;还可配置min_depth_eval/max_depth_eval过滤评估深度范围、depth_scale_factor缩放深度、crop_type选择裁剪策略。
训练配置与复现实验
两个训练配置均通过_base_继承Adabins.py模型配置、default_runtime.py运行时配置(default_scope='mmseg'、SyncBN 归一化等),并通过custom_imports显式导入projects.Adabins.backbones与projects.Adabins.decode_head,使 mmseg 能够解析到注册的组件:
NYU 配置 adabins_efficient_b5_4x16_25e_NYU_416x544.py:crop_size=(416, 544),data_preprocessor指定输入尺寸,模型沿用min_val=0.001, max_val=10。
KITTI 配置 adabins_efficient_b5_4x16_25e_kitti_352x704.py:crop_size=(352, 704),仅继承模型配置,并将解码头改为decode_head=dict(min_val=0.001, max_val=80),以匹配 KITTI 更深的测量范围(可达 80m)。
默认运行时配置(default_runtime.py)中log_processor = dict(by_epoch=False)表示按迭代数记录日志,tta_model = dict(type='SegTTAModel')为测试时增强预留接口。
训练与测试可复用仓库根目录的标准入口脚本:
python tools/train.py projects/Adabins/configs/adabins/adabins_efficient_b5_4x16_25e_NYU_416x544.py python tools/test.py projects/Adabins/configs/adabins/adabins_efficient_b5_4x16_25e_NYU_416x544.py <权重路径> --out <结果目录>多卡场景可改用tools/dist_train.sh/tools/dist_test.sh。需要注意:EfficientNet-B5 预训练权重通过 timm 自动下载,训练前请保证网络可达或已配置好本地权重缓存。
复现性能与参考数值
仓库 README 中给出的 NYU 与 KITTI 复现结果如下(训练 25 epoch、batchsize 16,参数规模约 78M):
| 模型 | 编码器 | 训练轮数 | Batchsize | 训练分辨率 | δ1 | δ2 | δ3 | REL | RMS | RMS log |
|---|---|---|---|---|---|---|---|---|---|---|
| AdaBins_nyu | EfficientNet-B5 | 25 | 16 | 416x544 | 0.903 | 0.984 | 0.997 | 0.103 | 0.364 | 0.044 |
| AdaBins_kitti | EfficientNet-B5 | 25 | 16 | 352x764 | 0.964 | 0.995 | 0.999 | 0.058 | 2.360 | 0.088 |
其中 δ1/δ2/δ3 为阈值准确率(越大越好),REL、RMS、RMS log 为误差指标(越小越好)。这两组权重以 third-party 形式提供,可直接用于tools/test.py的推理验证,具体下载地址见 projects/Adabins/README.md 中的 Links 列。
扩展阅读与引用
AdaBins 的官方开源实现与论文细节可分别通过 projects/Adabins/README.md 中提供的官方仓库链接与 arXiv 论文访问。若在论文或项目中引用该方法,可使用 README 中给出的 BibTeX 条目:
@article{10.1109/cvpr46437.2021.00400, author = {Bhat, S. A. and Alhashim, I. and Wonka, P.}, title = {Adabins: depth estimation using adaptive bins}, journal = {2021 IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)}, year = {2021}, doi = {10.1109/cvpr46437.2021.00400} }对于希望在 mmsegmentation 中扩展深度估计能力的开发者,可以从projects/Adabins出发,参照 mmseg/datasets/nyu.py、mmseg/models/segmentors/depth_estimator.py 与 mmseg/evaluation/metrics/depth_metric.py 了解数据集、模型封装与指标注册方式,从而将该框架迁移到自定义的深度估计数据与任务上。
【免费下载链接】mmsegmentationOpenMMLab Semantic Segmentation Toolbox and Benchmark.项目地址: https://gitcode.com/GitHub_Trending/mm/mmsegmentation
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考