nas-without-training
无需训练即可预测网络性能,Gram矩阵对数行列式为NAS搜索加速(ICML 2021)
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
无需训练即可预测网络性能,Gram矩阵对数行列式为NAS搜索加速(ICML 2021)
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下,你是一位 AI 研究员,手头有一个绝妙的神经网络架构想法,但不知道它是否真的比现有的更好。传统做法是:花几天甚至几周时间训练它,再花同样的时间训练其他候选架构,比较效果,再重复。这个过程枯燥、昂贵,而且效率低得令人发指。
这就是神经架构搜索(Neural Architecture Search, NAS)的痛点:搜索空间巨大,评估每个架构的成本极高。NAS-Bench-101 有 423k 个候选网络,NAS-Bench-201 有 15625 个。每个都训练一遍?光是电费就让人头疼。
BayesWatch 团队(来自爱丁堡大学)提出了一个大胆的假设:也许我们根本不需要训练,就能预测架构的最终性能。论文发表于 ICML 2021。
这个项目的核心方法叫 "NAS without Training"(NASWOT),其核心思想极为优雅:
关键假设:神经网络的训练过程可以被其**权重曲面的局部曲率(local curvature)所预测。具体来说,研究者发现 Gram 矩阵的对数行列式(log-determinant)**与架构的最终准确率存在强相关。
具体操作步骤:
这个分数被称为 "hook log-det score"(因为用到了 PyTorch 的 forward hook 机制来捕获中间层激活)。
scores.py — 评分函数接口
只暴露两个函数,极其干净。hooklogdet 计算 Gram 矩阵的对数行列式,get_score_func 是工厂函数,支持扩展新的评分指标。框架设计简洁,可轻松扩展新的评分方法。
score_networks.py — 网络评分入口
捕获 Jacobian 矩阵,配合 hook 机制提取中间层激活,构建 Gram 矩阵。核心函数 get_batch_jacobian 调用 net(x) 两次(一次 forward 记录激活,一次 backward 计算 Jacobian),这是 NASWOT 方法的标准实现。
nasspace.py — 搜索空间抽象层
通过统一抽象层支持 NAS-Bench-101(423k 候选)、NAS-Bench-201(15625 候选)和 NDS(Facebook 的自然搜索空间)。每个搜索空间都实现了 __len__、get_network、__iter__ 接口,评分脚本只需调用统一的 API。
search.py — 架构搜索逻辑
支持按「分数最高的 K 个架构」进行筛选,并可结合 余弦相似度(cosine similarity) 对分数接近的架构做聚类去重,避免在相似架构上浪费评估次数。
env.yml — 环境依赖
依赖 2020 年的旧版 PyTorch(1.6.0)+ CUDA 9.2,GPU 支持为 Volta/Turing 系列(V100、RTX 20系)。新版 GPU(如 RTX 30/40 系、A100/H100)需要更新 CUDA 版本。
1. 数据集偏差:方法在 NAS-Bench-101/201 上效果显著,但这些 benchmark 本身是「精心构造」的搜索空间。在自然搜索空间(如 NDS)上的效果尚未充分验证。
2. 分数噪声:log-det 分数对小批量大小(batch_size)和随机种子敏感,需要多次重复取最大(maxofn 参数)来降噪。
3. 缺乏验证:作者未提供完整的单元测试,主要通过 scorehook.sh 脚本重现论文实验结果,代码可复现性依赖第三方 benchmark API。
4. 无法预测训练不稳定架构:log-det 分数只能衡量「可训练性」,无法预测训练过程中的梯度爆炸/消失问题。
| 维度 | 评估 |
|---|---|
| 容器化 | 无 Dockerfile |
| Web UI | 纯 CLI 工具 |
| GPU 需求 | 必须(log-det 计算量与网络规模成正比) |
| 环境复杂度 | 高(需 conda + 第三方数据下载) |
| 适用用户 | NAS 研究者、AutoML 工程师 |
这篇论文(ICML 2021)的最大贡献不在于提出了一个完美的评分函数,而在于重新定义了 NAS 的评估范式:能否在「不训练」的情况下,对架构进行快速筛选?
后续工作(如 ZeroCost-NAS 系列、TE-NAS)沿着这条路不断深化。其中最直接的继承者是 ZeroCostNAS(ECCV 2022),将相同方法扩展到更多视觉任务,并开源了 4M+ 网络的大规模预计算分数。
该项目代码质量适中,但工程设计清晰——尤其是搜索空间抽象层的设计,被后续多个 NAS 项目直接借鉴引用。