MRL推理与评估实战:如何用pytorch_inference.py测出每个嵌套维度的准确率(附TTA技巧)

【免费下载链接】MRL Code repository for the paper - "Matryoshka Representation Learning" 【免费下载链接】MRL 项目地址: https://gitcode.com/gh_mirrors/mrl/MRL

MRL(Matryoshka Representation Learning,套娃表示学习)是同名论文的官方开源实现,其评估脚本 pytorch_inference.py 可以一条命令测出同一个 ResNet50 模型在 8~2048 各嵌套维度下的 Top-1/Top-5 准确率,还支持 TTA 测试时增强与 ImageNet 鲁棒性基准评测。本文带你走通从环境配置到读懂输出的 MRL 推理评估完整流程。🪆

MRL套娃表示学习原理图:一个特征向量同时支持8到2048各嵌套维度

为什么能测出"每个嵌套维度"的准确率

Matryoshka 的核心思想:一个 2048 维特征向量的前缀本身就是一个可用的短表示。官方实现用 MRL_Linear_Layer 在 8、16、32、64、128、256、512、1024、2048 这 9 个维度上同时训练分类头,因此评估时天然可以逐维度拆解:

MRL嵌套维度Top-1准确率曲线:ResNet50在8到2048每个维度上的分类准确率

环境准备:三步装好依赖与数据

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)路径
--datasetV1(1K 验证集)/ V2 / A / R / sketch
--mrl以 MRL 模型加载,一次评估全部 9 个维度
--efficientMRL-E 变体
--rep_size固定维度基线的维度(MRL 模型不需要)
--tta开启 TTA 测试时增强
--old_ckpt兼容官方旧版 checkpoint
--workersdataloader 进程数,默认 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/ 目录。

MRL自适应分类级联:TTA测试时增强与嵌套维度级联结合提升推理精度

进阶一:鲁棒性基准一键切换

--dataset 支持四个鲁棒性测试集,命令结构完全一致:

  • V2:ImageNetV2(经 imagenetv2_pytorch 自动下载,无需本地路径)
  • A:ImageNet-A
  • R:ImageNet-R
  • sketch: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:模型级联策略

MRL模型分析:各嵌套维度GradCAM注意力可视化对比

扩展:同一个脚本做图像检索评估

--retrieval 参数,同一脚本即可转做特征导出:把数据库与查询集的特征向量 dump 成 .npy 数组,再配合 retrieval/ 目录下的 notebook(faiss_nn.ipynb、reranking.ipynb)完成近邻检索与重排序。各嵌套维度下的检索质量对比:

MRL图像检索mAP@10指标:8到2048各嵌套维度下的检索准确率对比

常见问题 FAQ

  1. checkpoint 里带 module. 前缀:训练端用 DDP 保存,utils.pyget_ckpt() 会自动剥掉前 7 个字符,通常无需处理。
  2. --rep_size 没生效:MRL 模型会忽略该参数,--mrl 下默认评估全部 9 个维度。
  3. 加载官方 checkpoint 报错:旧版头结构(MultiHead/SingleHead)必须加 --old_ckpt 才能正确构建。
  4. 只能在 GPU 上跑:脚本内部固定调用 model.cuda(),需 CUDA 环境。
  5. 想改嵌套维度:修改 pytorch_inference.py 顶部的 NESTING_LIST 即可。

小结

  • 一条命令 pytorch_inference.py --mrl 即可输出 8→2048 共 9 个嵌套维度的 Top-1/Top-5 准确率与单图耗时
  • --tta 用水平翻转测试时增强,轻松再提一点准确率
  • --save_* 系列开关保存推理结果,衔接模型分析与 GradCAM 可视化
  • 鲁棒性基准(V2/A/R/sketch)与图像检索评估复用同一套脚本,参数切换即可

【免费下载链接】MRL Code repository for the paper - "Matryoshka Representation Learning" 【免费下载链接】MRL 项目地址: https://gitcode.com/gh_mirrors/mrl/MRL

更多推荐