InstaNAS
AnjieCheng/InstaNAS加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下:一个图像分类AI,在识别一张简单猫咪照片时,调用轻量的计算路径;遇到一张背景杂乱、光照复杂的猫咪照片时,自动切换到更深层的神经网络。这种「因图制宜」的智能,正是 InstaNAS 的核心思想——让同一个模型里的不同子网络,对应不同难度的输入样本。
传统神经架构搜索(NAS)的研究范式,是从庞大的搜索空间中找到一个「最优架构」,然后用这个架构处理所有输入。但这个思路存在明显的天花板:简单样本浪费了复杂模型的算力,复杂样本又超出了简单模型的容量上限。 2018年,来自台湾国立清华大学和 Google AI 的研究团队(核心作者 An-Chieh Cheng 和 Hubert Lin)在 arXiv 发布论文,提出了一种革命性的思路:不搜索单一架构,而是搜索「架构的分布」(distribution of architectures)。其核心论文《InstaNAS: Instance-aware Neural Architecture Search》先后发表于 ICML 2019 AutoML Workshop 和 AAAI 2020。 论文发表至今,该 GitHub 仓库已获得 93 颗星、10 次 Fork,在学术复现项目中属于相当不错的传播度。团队成员后续也活跃于 Google DeepMind 等顶级 AI 机构。
InstaNAS 的搜索空间设计极为精妙。整体架构由两个核心组件构成: 1. Meta-graph(预训练权重共享网络) 研究团队设计了一个超网络(super-network),包含多个可选的计算路径——每个残差块内部存在多个候选操作(如不同膨胀率的卷积、不同通道数的子分支)。这些候选路径通过权重共享的方式共存于同一网络中,形成一个「一网络多路径」的搜索空间。 预训练阶段,团队提供了在 CIFAR-10 和 ImageNet 上的预训练权重,学习共享参数和基础的图像表征能力。 2. 策略网络(Policy Network / Controller) 核心创新在于:训练一个 LSTM 控制器,为每个输入样本预测一条穿过超网络的路径(即每个残差块中应该激活哪些候选操作)。这是一个典型的强化学习问题:
# 延迟稀疏奖励:越接近目标延迟区间,奖励越高
highest_point = - (lb - ub)*(ub - lb)/4
sparse_reward = -1 * (elasped.cuda().data - ub) * (elasped.cuda().data - lb) / highest_point
sparse_reward = torch.clamp(sparse_reward, min=0.)
# 分类正确 × 正权重,错误 × 负权重
reward[match] *= args.pos_w # 正确+低延迟 → 大奖
reward[match==0] = args.neg_w # 错误 → 重罚
这种基于梯度的强化学习方法(类似 DARTS),让控制器能够端到端地学习样本难度与架构路径之间的映射关系。实验证明,控制器学到的难度估计与人类直觉高度吻合:杂乱背景、高类内变异、光照复杂等图片会被路由到更深层的架构。
InstaNAS 的核心评测基于 MobileNetv2 的搜索空间,目标是同时优化精度和推理延迟。研究团队在 CIFAR-10、CIFAR-100、Tiny-ImageNet 和 ImageNet 四个数据集上进行了广泛实验。
图1:InstaNAS 在 MobileNetv2 精度-延迟权衡曲线上的表现,所有变体(A-E 和 A-C)均在一个搜索周期内获得
实验结果极具说服力:在 ImageNet 上,InstaNAS-D 将 MobileNetv2 1.0 的精度从 72.0% 提升至 74.0%,同时保持相近的延迟;在 CIFAR-100 上,InstaNAS-A 相比基线提升了超过 3 个百分点的准确率。
更重要的是,搜索本身非常高效——整个搜索过程在单块 NVIDIA GPU 上仅需数小时,而不像早期 NAS 方法需要数百 GPU 日。
仓库代码分为两个核心阶段:
Pretrain 阶段(pretrain/ 目录):
main.py:预训练 Meta-graph 权重,支持 CIFAR/ImageNet 多数据集models/instanas.py:核心模型定义,InstaNas 类实现多路径前向传播dataloader.py、utils.py:数据加载和训练辅助函数args.py:统一的命令行参数解析
Search 阶段(search/ 目录):search.py:策略网络训练主脚本,包含完整的 REINFORCE 类奖励计算models/controller.py:ResNet 和 ResNet32 控制器模型finetune.py:搜索完成后对选定架构进行微调test.py:测试脚本
代码整体风格偏向学术研究用途,结构清晰但不追求工程化封装。依赖 Python 3.6 + PyTorch 0.4.1(较老的版本),数据集需要用户自行下载准备。项目另一个亮点是对搜索结果的解读性分析。通过 UMAP 将搜索到的架构分布投影到二维空间,可以清晰看到:简单样本对应的架构倾向于使用更少的残差块(浅层路径),而困难样本则被路由到更深、更宽的计算路径。这种可视化为理解 NAS 的内部决策提供了直观的窗口。
图2:搜索到的架构分布 UMAP 可视化,展示了不同难度样本的架构差异
InstaNAS 作为 2018-2019 年的研究工作,存在一些局限性:
torch.autograd.Variable 在新版本已移除)。在现代环境中运行需要降级 PyTorch 或做一定的代码适配。InstaNAS 的核心贡献——「为不同难度的输入样本搜索不同的计算路径」——对后续研究产生了深远影响。这一思想直接启发了: