MRL推理与评估实战如何用pytorch_inference.py测出每个嵌套维度的准确率附TTA技巧【免费下载链接】MRLCode repository for the paper - Matryoshka Representation Learning项目地址: https://gitcode.com/gh_mirrors/mrl/MRLMRLMatryoshka Representation Learning套娃表示学习是同名论文的官方开源实现其评估脚本pytorch_inference.py可以一条命令测出同一个 ResNet50 模型在 82048 各嵌套维度下的 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路径--datasetV11K 验证集/ 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参数即可开启 TTATest-Time Augmentation。实现非常简洁原图与水平翻转图各推理一次两组 logits 相加融合# utils.py 中 TTA 的核心逻辑 logits model(img_input) logits model(torch.flip(img_input, dims[3]))torch.flip(..., dims[3])生成镜像图两组 logits 求和后取 softmax2 倍缩放不影响 argmax 排序从而稳定带来小幅准确率提升。⚠️ 注意官方论文报告的分类结果不含 TTATTA 主要服务于自适应分类的模型级联场景可参考 model_analysis/ 目录。进阶一鲁棒性基准一键切换--dataset支持四个鲁棒性测试集命令结构完全一致V2ImageNetV2经 imagenetv2_pytorch 自动下载无需本地路径AImageNet-ARImageNet-RsketchImageNet-Sketch后三者需放在ROOT目录默认../../IMAGENET/。ImageNet-A/R 只有 200 个类脚本通过 imagenet_id.py 中的索引映射自动选取对应 logits 列无需任何手工处理。进阶二保存推理结果衔接模型分析--save_logits、--save_softmax、--save_gt、--save_predictions四个开关会把每张图的 logits、概率、真值标签与预测结果存成 .pth 文件命名自动组合为mrl1_efficient0_datasetV1_ttaFalse_logits.pth形式供 model_analysis/ 下的分析 notebook 直接使用GradCAM可视化各嵌套维度的注意力可见小维度更易混淆同一超类内的类别Custom SuperClass基于 WordNet 层级的 30 个超类性能分析Oracle Upper Bound为每张图片寻找最优维度计算理论上限Model Cascades模型级联策略扩展同一个脚本做图像检索评估加--retrieval参数同一脚本即可转做特征导出把数据库与查询集的特征向量 dump 成 .npy 数组再配合 retrieval/ 目录下的 notebookfaiss_nn.ipynb、reranking.ipynb完成近邻检索与重排序。各嵌套维度下的检索质量对比常见问题 FAQcheckpoint 里带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与图像检索评估复用同一套脚本参数切换即可【免费下载链接】MRLCode repository for the paper - Matryoshka Representation Learning项目地址: https://gitcode.com/gh_mirrors/mrl/MRL创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考