NAS-Bench-201 快速上手指南:在固定搜索空间中查询 15,625 个神经单元的训练指标
【免费下载链接】NAS-Bench-201NAS-Bench-201 API and Instruction项目地址: https://gitcode.com/gh_mirrors/na/NAS-Bench-201
NAS-Bench-201 是一个面向神经架构搜索(NAS)的可复现基准:固定的 cell 搜索空间中共有 15,625 个候选架构,每个架构都附带预训练好的指标。本指南带你在几分钟内完成 NAS-Bench-201 API 的安装与初始化,并查得任意一个架构的训练损失与精度。
项目定位:它解决什么问题,不适合什么场景
它解决的是“不同 NAS 算法各自使用不同搜索空间与数据切分,评测结果无法横向比较”的问题。项目把搜索空间固定为 4 节点 cell、5 种操作,并预先记录每个架构在 CIFAR-10、CIFAR-100、ImageNet16-120 上的训练曲线、评测指标与计算开销,你实现或对比算法时无需重新训练即可获得数据。
适合:正在实现、评估或对比 NAS 算法,以及需要逐架构指标做分析的人。不适合:想用其中的结构直接上线跑业务,或需要加载完整模型权重的人(权重归档文件体量达数百 GB,详情以仓库 README 为准)。另外仓库在 README 开头已标注归档状态,作者的后续工作已转向 NATS-Bench;如果你要开展新的研究,可先评估是否改用后者,但仅用本基准做实验的话,仓库内内容是自洽完整的。
快速准备:安装 API 并拿到数据文件
📦 环境要求不高:Python ≥ 3.6.0、PyTorch ≥ 1.2.0。API 有两种获取方式,装包或装源码二选一:
pip install nas-bench-201 # 或克隆源码后以开发模式安装,便于阅读与修改实现 # git clone https://gitcode.com/gh_mirrors/na/NAS-Bench-201 # cd NAS-Bench-201 && pip install -e .数据文件方面,官方推荐NAS-Bench-201-v1_1-096897.pth(约 4.7 GB),它包含更多 trial 以及全部数据集上 12 epoch 的训练结果;另有较早的 v1_0 文件,体量较小但覆盖更少。下载渠道列在 README.md 中,获取后放到任意目录均可,建议直接放入$TORCH_HOME(默认~/.torch/),因为这是 API 的默认查找位置。若你要自行重新生成数据集或训练模型,还需要 CIFAR-10/100 与 ImageNet16-120 的原始训练评测数据,链接同样在 README 中。
最小可用入口:两行代码创建 API 实例
仓库里没有可执行的主程序,也没有命令行入口,真正的起点是实例化NASBench201API类。构造时只需传入 .pth 文件路径;文件若已放在 TORCH_HOME 下,可直接传None走默认路径。加上verbose=False可去掉过程日志:
from nas_201_api import NASBench201API as API api = API('<path_to_data>/NAS-Bench-201-v1_1-096897.pth', verbose=False) print(len(api), api[1]) info = api.query_meta_info_by_index(1) print(info.get_metrics('cifar10', 'train'))验证是否装好:len(api)应返回 15625;api[i]返回第 i 个架构的字符串编码;api.show(1)会打印该架构在各数据集上的指标摘要与 FLOPs、参数量、延迟等开销。
关键数据与参数:一个 .pth 承载全部状态
.pth 是唯一的数据载体
基准文件是一个 torch 字典,核心键为meta_archs(架构字符串列表)、arch2infos(逐架构的 trial 结果,按 12/200 两档超参分开存放)、evaluated_indexes(已被训练过的架构索引集合)。API 构造时会加载并整理这份字典,运行过程不会再产生额外状态文件。
参数全部通过函数实参传递
项目没有独立的配置文件,行为完全由你传入的参数决定,常见的几组取值:
- dataset:
cifar10、cifar100、ImageNet16-120,另有cifar10-valid,它是 CIFAR-10 的另一种切分,不要与前两者混用。 - hp:
'12'或'200',对应“训练到 12 epoch”和“训练到 200 epoch”两套超参。 - iepoch:指定查看某一 epoch 的指标,传
None则取最后一个 epoch。 - is_random:
True随机返回一个 trial,False返回全部 trial 的均值。
get_more_info、query_by_index等查询函数都由这些参数组合而成,各参数的精确含义见 API 实现 中的函数注释。
关键文件地图:先看哪几个文件
- nas_201_api/api_201.py:
NASBench201API的实现,查询、模拟训练、架构编码转换都在这里。 - nas_201_api/api_utils.py:
ArchResults(单架构全部 trial)与ResultsCount(单个 trial)两个数据类,以及抽象基类。 - nas_201_api/init.py:API 版本号,附带一个
test_api自检函数,可直接跑通验证数据文件。 - README.md:数据下载链接、完整 API 示例与引用信息。
- setup.py:打包元数据,MIT 协议。
常见坑与下一步建议
⚠️ 最常踩的两处:一是 .pth 路径不存在时构造函数会直接断言报错,放进 TORCH_HOME 再调API(None)最稳妥;二是数据集名称必须精确匹配,cifar10与cifar10-valid是两个不同的切分,用错名字会直接抛异常。
下一步建议:先完整运行一次 nas_201_api/init.py 里的test_api做全量自检;如果后续要加载某个架构的真实权重,再看reload方法与 README 中权重归档的说明。
【免费下载链接】NAS-Bench-201NAS-Bench-201 API and Instruction项目地址: https://gitcode.com/gh_mirrors/na/NAS-Bench-201
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考