如何用 pytorch-image-models 的 validate.py 在 ImageNet 验证集上评估 timm 模型并导出 CSV 结果?
【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 & V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models
在 timm(pytorch-image-models)里,validate.py是仓库根目录自带的验证脚本:它在一个按 ImageNet 结构组织的验证集上加载一个 timm 模型,跑完整个验证集后输出 Top-1 / Top-5 准确率,并把汇总结果写入 CSV(或 JSON)文件。本文的目标就一条:从仓库中取出validate.py,在 ImageNet 验证集上评估一个已带预训练权重的模型,并把结果落盘成 CSV。
前提:脚本不在 pip 发行包里,需要仓库根目录
pip install timm安装的是库本身,但训练/验证/推理脚本并不随 pip 包分发(见 Scripts 文档 开头说明:"Scripts are not currently packaged in the pip release")。所以要同时拿到validate.py和可用的timm库,最省事的路径是从源码安装,这样仓库根目录里的脚本和库一起就位(参考 Installation 的 From Source 一节):
git clone <仓库地址> pytorch-image-models # 换成你使用的 pytorch-image-models 仓库地址 cd pytorch-image-models pip install -e .安装完成后,可用 Installation 给出的检查命令确认库能正常加载:
python -c "from timm import list_models; print(list_models(pretrained=True)[:5])"能打印出一批模型名即表示timm安装成功。注意validate.py里的模型名、数据加载和权重都依赖timm库,两者必须来自同一套代码,这也是要一起准备的原因。
准备 ImageNet 验证集目录
validate.py的--data-dir指向的是验证图片所在文件夹本身,而不是训练时那种包含train/validation两个子目录的根目录(Scripts 文档 明确:"Specify the folder containing validation images, not the base as in training script")。以标准 ImageNet-1k 为例,应指向 5 万张验证图片的目录,例如/imagenet/validation/。
- 默认数据集读取方式是 ImageFolder(
--dataset留空时的默认值,见 validate.py 的--dataset说明)。 - 如果你用的是 tar 打包的验证集等其它组织方式,可通过
--dataset "<type>/<name>"指定数据集类型与名称。 - 默认加载 4 个 worker(
-j)、batch size 256(-b),可在显存不足时调小。
运行单模型验证(主路径)
下面这条命令用模型的预训练权重在验证集上评估seresnext26_32x4d(该示例命令直接取自 Scripts 文档),并把结果写入val.csv:
python validate.py --data-dir /imagenet/validation/ --model seresnext26_32x4d --pretrained --results-file val.csv各参数与适用条件:
--data-dir:验证图片目录,必填。--model / -m:模型架构名,默认dpn92。名字可用通配符,也可指向一个包含模型名的文本文件(见 validate.py 的main()逻辑)。--pretrained:使用模型自带的预训练权重。不加该参数、又未提供--checkpoint时,脚本内部会把pretrained置为True(validate.py 中args.pretrained = args.pretrained or not args.checkpoint)。--results-file:结果输出文件名,留空则不写文件(只打印到终端)。--results-format:输出格式,csv(默认)或json(见 validate.py 的--results-format)。
脚本的默认设备是cuda(--device默认值cuda);需要换到别的加速器时显式传--device。要开启混合精度推理可加--amp(默认float16,可用--amp-dtype改为bfloat16,见 validate.py)。
从训练 checkpoint 评估
如果要评估的是自己训练的 checkpoint 而不是官方预训练权重,改用--checkpoint指定权重文件(该示例命令取自 Scripts 文档):
python validate.py --data-dir /imagenet/validation/ --model mobilenetv3_large_100 --checkpoint ./output/train/model_best.pth.tar--checkpoint指向单个.pth.tar/.pth文件。当--checkpoint指向一个目录时,脚本会批量验证该目录下同架构的所有 checkpoint(glob匹配*.pth.tar和*.pth,见 validate.py 的main()),结果会按 Top-1 排序。若 checkpoint 里带有 EMA 权重,可加--use-ema选用 EMA 版本。
验证输出与 CSV 结果
运行结束时,脚本会在日志里打印一行汇总,例如:
* Acc@1 xxx.xxx (xxx.xxx) Acc@5 xxx.xxx (xxx.xxx)(以上xxx为占位,实际数值取决于你的模型与验证集,文档未给出固定预期值。)同时脚本会把完整结果以 JSON 形式打印到标准输出,并用--result作为分隔标记,方便上层脚本解析(见 validate.py 末尾的print)。
当指定了--results-file时,write_results会把汇总写成 CSV。CSV 的列来自结果字典的键,包含model、top1、top1_err、top5、top5_err、param_count、img_size、crop_pct、interpolation等字段(validate.py 的resultsOrderedDict)。仓库里已有一份真实结果样例 results-imagenet.csv,其表头为:
model,img_size,top1,top1_err,top5,top5_err,param_count,crop_pct,interpolation打开你生成的val.csv,确认出现对应模型的一行记录、且 Top-1/Top-5 为合理数值,就说明这次验证与导出成功了。这些数值本身取决于模型与验证集,不是固定目标,不要与样例文件里的数字做绝对对比。
可选分支:批量验证多个模型
如果要在一次运行里评估一批模型并汇成一个 CSV,可用 bulk_runner.py 作为外层驱动,它会对每个模型在独立进程里调用validate.py。文档给出的示例(见 bulk_runner.py 顶部)用通配符筛选了模型列表:
python bulk_runner.py --model-list all --results-file val.csv --pretrained validate.py --data-dir /imagenet/validation/ --amp -b 512 --retry--model-list可换成vit*这类筛选表达式,只跑匹配的子集,避免一次跑全部模型。--retry启用 batch size 衰减与重试(显存吃紧时更稳,见 validate.py 的_try_run)。- 该批量入口仍只做"验证 + 汇总 CSV"这一件事,不会顺带执行训练等其它任务,因此可以作为当前场景的可选路径。
限制与注意
- ImageNet-1k 的 5 万张验证集在训练时也被用来选模型,所以它不是真正的测试集(见 results/README.md)。需要衡量泛化时,仓库还提供了 ImageNetV2、ImageNet-Sketch 等额外测试集的 CSV(见 results/README.md 的 Datasets 一节)。
- 默认设备为
cuda、默认 batch 256;显存不足时先调小-b,或加--retry让脚本自动降 batch 重试。 - 需要额外指标(precision / recall / F1)时加
--metrics-avg {micro,macro,weighted},该功能依赖 scikit-learn,未安装时脚本会给出pip install scikit-learn的提示并跳过(见 validate.py)。 --results-file留空时不会生成 CSV,只有终端日志与 stdout 的 JSON;要落盘必须显式传该参数。
【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 & V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考