MRL推理与评估实战:如何用pytorch_inference.py测出每个嵌套维度的准确率(附TTA技巧)
MRL推理与评估实战:如何用pytorch_inference.py测出每个嵌套维度的准确率(附TTA技巧)
MRL(Matryoshka Representation Learning,套娃表示学习)是同名论文的官方开源实现,其评估脚本 pytorch_inference.py 可以一条命令测出同一个 ResNet50 模型在 8~2048 各嵌套维度下的 Top-1/Top-5 准确率,还支持 TTA 测试时增强与 ImageNet 鲁棒性基准评测。本文带你走通从环境配置到读懂输出的 MRL 推理评估完整流程。🪆
为什么能测出"每个嵌套维度"的准确率
Matryoshka 的核心思想:一个 2048 维特征向量的前缀本身就是一个可用的短表示。官方实现用 MRL_Linear_Layer 在 8、16、32、64、128、256、512、1024、2048 这 9 个维度上同时训练分类头,因此评估时天然可以逐维度拆解:
环境准备:三步装好依赖与数据
git clone https://gitcode.com/gh_mirrors/mrl/MRL
cd MRL
pip3 install -r requirements.txt
评测数据需要序列化为 FFCV 格式,write_imagenet.sh 可一键完成:
cd train
export IMAGENET_DIR=/path/to/imagenet
export WRITE_DIR=/your/write/dir
./write_imagenet.sh 500 0.50 90
一键评估命令:快速测出各嵌套维度准确率
cd inference
python pytorch_inference.py --path <final_weight.pt> --dataset V1 --mrl
三种典型模型形态,命令只需微调参数:
# MRL-E 高效版(单头按维度切分)
python pytorch_inference.py --path ckpt.pt --dataset V1 --mrl --efficient
# 固定维度 FF 基线:用 --rep_size 指定维度,如 512
python pytorch_inference.py --path ckpt.pt --dataset V1 --rep_size 512
# 官方上传的旧版 checkpoint(命名如 r50_mrl1_e0_ff2048.pt)
python pytorch_inference.py --path r50_mrl1_e0_ff2048.pt --dataset V1 --mrl --old_ckpt
常用参数速查表:
| 参数 | 作用 |
|---|---|
--path | 模型 checkpoint(.pt)路径 |
--dataset | V1(1K 验证集)/ V2 / A / R / sketch |
--mrl | 以 MRL 模型加载,一次评估全部 9 个维度 |
--efficient | MRL-E 变体 |
--rep_size | 固定维度基线的维度(MRL 模型不需要) |
--tta | 开启 TTA 测试时增强 |
--old_ckpt | 兼容官方旧版 checkpoint |
--workers | dataloader 进程数,默认 12 |
读懂输出:逐维度 Top-1 / Top-5 与单图耗时
脚本会为每个嵌套维度依次打印三段指标:
Rep. Size 8
Top-1 accuracy for 8 : xx.xx
Top-5 accuracy for 8 : xx.xx
Total time: xx.x (average time per image: xx.xx ms)
逐维度统计全部由 utils.py 中的 evaluate_model_nesting() 完成:对每个维度做 top-5 排序、累计 Top-1/Top-5 命中数,最后按图片总数归一化。想评估其他维度,只需修改脚本顶部的 NESTING_LIST 常量(pytorch_inference.py)。
TTA技巧:水平翻转 + logits融合,准确率再提一截
加上 --tta 参数即可开启 TTA(Test-Time Augmentation)。实现非常简洁:原图与水平翻转图各推理一次,两组 logits 相加融合:
# utils.py 中 TTA 的核心逻辑
logits = model(img_input)
logits += model(torch.flip(img_input, dims=[3]))
torch.flip(..., dims=[3]) 生成镜像图,两组 logits 求和后取 softmax(2 倍缩放不影响 argmax 排序),从而稳定带来小幅准确率提升。⚠️ 注意:官方论文报告的分类结果不含 TTA,TTA 主要服务于自适应分类的模型级联场景,可参考 model_analysis/ 目录。
进阶一:鲁棒性基准一键切换
--dataset 支持四个鲁棒性测试集,命令结构完全一致:
V2:ImageNetV2(经 imagenetv2_pytorch 自动下载,无需本地路径)A:ImageNet-AR:ImageNet-Rsketch:ImageNet-Sketch
后三者需放在 ROOT 目录(默认 ../../IMAGENET/)。ImageNet-A/R 只有 200 个类,脚本通过 imagenet_id.py 中的索引映射自动选取对应 logits 列,无需任何手工处理。
进阶二:保存推理结果,衔接模型分析
--save_logits、--save_softmax、--save_gt、--save_predictions 四个开关会把每张图的 logits、概率、真值标签与预测结果存成 .pth 文件,命名自动组合为 mrl=1_efficient=0_dataset=V1_tta=False_logits.pth 形式,供 model_analysis/ 下的分析 notebook 直接使用:
- GradCAM:可视化各嵌套维度的注意力,可见小维度更易混淆同一超类内的类别
- Custom SuperClass:基于 WordNet 层级的 30 个超类性能分析
- Oracle Upper Bound:为每张图片寻找最优维度,计算理论上限
- Model Cascades:模型级联策略
扩展:同一个脚本做图像检索评估
加 --retrieval 参数,同一脚本即可转做特征导出:把数据库与查询集的特征向量 dump 成 .npy 数组,再配合 retrieval/ 目录下的 notebook(faiss_nn.ipynb、reranking.ipynb)完成近邻检索与重排序。各嵌套维度下的检索质量对比:
常见问题 FAQ
- checkpoint 里带
module.前缀:训练端用 DDP 保存,utils.py 的get_ckpt()会自动剥掉前 7 个字符,通常无需处理。 --rep_size没生效:MRL 模型会忽略该参数,--mrl下默认评估全部 9 个维度。- 加载官方 checkpoint 报错:旧版头结构(MultiHead/SingleHead)必须加
--old_ckpt才能正确构建。 - 只能在 GPU 上跑:脚本内部固定调用
model.cuda(),需 CUDA 环境。 - 想改嵌套维度:修改 pytorch_inference.py 顶部的
NESTING_LIST即可。
小结
- 一条命令
pytorch_inference.py --mrl即可输出 8→2048 共 9 个嵌套维度的 Top-1/Top-5 准确率与单图耗时 - 加
--tta用水平翻转测试时增强,轻松再提一点准确率 - 加
--save_*系列开关保存推理结果,衔接模型分析与 GradCAM 可视化 - 鲁棒性基准(V2/A/R/sketch)与图像检索评估复用同一套脚本,参数切换即可
更多推荐






所有评论(0)