enas
通过参数共享机制大幅降低神经架构搜索的计算成本,让AI自动设计高性能神经网络
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
通过参数共享机制大幅降低神经架构搜索的计算成本,让AI自动设计高性能神经网络
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下:你是一位建筑设计师,要从零设计一栋摩天大楼。但传统做法是——每设计一个方案,就要从头建一个完整模型来测试抗风性、地震响应、能耗效率。一个方案测试三个月,成本惊人。这个过程重复成百上千次,才能找到那个"最优解"。
这正是2018年以前深度学习工程师的日常:设计新的神经网络架构,需要在庞大的搜索空间中逐一训练、评估。这种方法叫神经架构搜索(Neural Architecture Search, NAS),但计算成本极高——通常需要数百到数千个GPU运行数周,普通人根本无法承担。
2018年,Google Brain的工程师们提出了一个精妙的解法,发表在ICML顶会上——Efficient Neural Architecture Search via Parameter Sharing(ENAS),中文译作"通过参数共享实现高效神经架构搜索"。它的核心思想一句话概括:不是每次都从零训练模型,而是让所有候选架构共享同一个"超级模型"的权重,新架构只需继承微调,大幅削减计算成本。

图1:ENAS论文中的RNN Cell结构示意,不同颜色代表不同操作节点。
本文分析的仓库是ENAS论文的官方TensorFlow实现,在GitHub上获得1579颗星、383次Fork,至今仍是AutoML领域的重要参考。
在2017年之前,ResNet、DenseNet、VGG等经典网络架构均由人工设计。工程师凭借经验和直觉,通过堆叠卷积层、设置跳跃连接、调整超参数来提升性能。这种方式效率低、依赖专家知识,且容易陷入局部最优。
2016-2017年,Google相继提出NASNet、MetaQNN等自动化架构搜索方法,通过强化学习控制器(RNN)来生成架构描述,并用验证集准确率作为奖励信号训练控制器。理论上这是"让AI设计AI",但代价是极高的计算开销——NASNet用800块GPU搜索了28天。
ENAS的核心洞察是:NAS的计算瓶颈来自反复从头训练子模型,但这些子模型其实有大量重叠的子结构。
想象一棵决策树:多个分支的前几层可能完全相同,只有后面的节点不同。神经网络也一样——不同架构的早期层(边缘特征提取)往往相似,差异主要在中后期。ENAS的解决方案是为整个搜索空间构建一个超图(Supernetwork),所有子架构共享超图的权重。控制器(RNN)负责决定超图中哪些边被激活——激活哪些边,就构成一个子架构。
这样一来,评估一个新架构不再是"从头训练",而是"继承权重+快速验证",计算量从O(N)降到O(1)。
ENAS的控制器是一个LSTM网络,每一步的输出决定网络的一个决策:
对于图像任务(CIFAR-10),控制器决定:
对于NLP任务(PTB语言模型),控制器决定:
控制器用**策略梯度(Policy Gradient)**训练,最大化子架构在验证集上的期望奖励。训练过程中,控制器逐渐学会采样高奖励的架构。
仓库中实现了两种搜索空间:
Macro搜索空间:直接搜索整个网络的端到端结构。每一层从6个候选操作中选择(conv_3x3、sep_conv_3x3、conv_5x5、sep_conv_5x5、avg_pool、max_pool),并决定与哪些前序层建立跳跃连接。架构用一串数字编码,例如"0 1 0 1 1 0 1 0 ..."。
Micro搜索空间:搜索"Cell"(基本单元),然后堆叠多个相同Cell构成网络。这是后来NASNet采用的方法,更高效、更通用。每个Cell内有B个块(Block),每个块由(index_1, op_1, index_2, op_2)定义,其中index选择两个输入节点,op选择操作类型。
权重共享通过GeneralChild和MicroChild两个类实现。它们继承自Model基类,在前向传播时根据控制器采样出的架构动态构建计算图:
GeneralChild:宏搜索空间实现,支持任意Skip连接MicroChild:微搜索空间实现,基于预定义的Cell模板关键代码在src/common_ops.py中实现了自定义LSTM(lstm()和stack_lstm()函数),用于RNN架构搜索,而非使用TF内置的tf.nn.rnn_cell。
src/
├── utils.py # 参数解析、日志、训练优化器
├── common_ops.py # LSTM、权重初始化等通用操作
├── controller.py # 控制器基类
├── cifar10/
│ ├── controller.py # CIFAR-10卷积网络控制器
│ ├── micro_controller.py # 微搜索空间控制器
│ ├── general_controller.py # 宏搜索空间控制器
│ ├── general_child.py # 宏搜索空间子模型(~27KB,最核心)
│ ├── micro_child.py # 微搜索空间子模型(~31KB,最核心)
│ ├── models.py # 模型基类、训练循环
│ ├── image_ops.py # 卷积、BN、ReLU等图像操作
│ ├── data_utils.py # CIFAR-10数据加载
│ └── main.py # 入口脚本,定义所有搜索参数
└── ptb/
├── ptb_enas_child.py # PTB语言模型子模型(~21KB)
├── ptb_enas_controller.py # PTB RNN控制器
├── ptb_ops.py # PTB相关操作
├── data_utils.py # PTB数据预处理
└── main.py # PTB入口脚本
数据目录data/包含PTB语料库(train/valid/test,已预处理为二进制格式,~5MB each)和CIFAR-10原始数据集(需用户自行下载)。
搜索阶段(约12小时,单GPU):
# 宏搜索空间
./scripts/cifar10_macro_search.sh
# 微搜索空间
./scripts/cifar10_micro_search.sh
控制器会采样数千个子架构,通过共享权重快速评估,输出最优架构的编码。
最终训练阶段(约3天,单GPU):
# 固定最优架构,从头训练验证性能
./scripts/cifar10_macro_final.sh
./scripts/cifar10_micro_final.sh
论文报告的最终结果:
./scripts/ptb_search.sh # 搜索最优RNN Cell
./scripts/ptb_final.sh # 固定架构后完整训练
⚠️ 重要勘误:README中明确标注,PTB语言模型实现存在错误,正确实现在google-research/google-research/enas_lm,请勿使用本仓库的PTB代码进行正式实验。
所有参数在main.py中通过TF App Flags定义,核心超参数包括:
| 参数 | 含义 | 典型值 |
|---|---|---|
child_num_layers | 网络层数(Macro) | 5 |
child_num_cells | Cell数量(Micro) | 5 |
child_out_filters | 初始通道数 | 48 |
controller_lr | 控制器学习率 | 1e-3 |
controller_bl_dec | Baseline衰减系数 | 0.99 |
child_lr | 子模型学习率 | 0.1 |
child_grad_bound | 梯度裁剪阈值 | 5.0 |
| 组件 | 技术选型 | 说明 |
|---|---|---|
| 深度学习框架 | TensorFlow 1.x | 代码使用Python 2.x语法(print语句),依赖TF 1.x API |
| 优化器 | SGD + Momentum / Adam | 子模型用SGD,控制器用Adam |
| 强化学习 | REINFORCE策略梯度 | 控制器奖励机制 |
| 数据集 | CIFAR-10、PTB | 图像分类+语言建模 |
| 编程语言 | Python 2.x | 2018年代码,已有Python 3兼容性问题 |
print "..."语句,需Python 2.7或2to3转换tf.contrib等API已在TF 2中移除ENAS的意义远超论文本身,它开创的参数共享范式深刻影响了后续研究:
ENAS还推动了**权重继承(Weight Inheritance)**成为NAS的标配技术,极大降低了AutoML的门槛。如今,任何人用一块消费级GPU,也能在数天内完成一次完整的神经架构搜索。
本仓库为纯研究代码,不具备生产部署条件:
推荐替代方案:如果想在现代环境(Python 3 + TF2/PyTorch)中复现ENAS思想,可参考:
melodyguan/enas 是神经架构搜索领域的里程碑式工作,其参数共享机制从根本上解决了NAS的计算瓶颈问题。虽然代码基于2018年的TensorFlow 1.x+Python 2.x环境编写,已略显过时,但其核心思想——用控制器指导架构采样、用权重共享降低评估成本——至今仍是AutoML研究的重要基石。
对于AI爱好者和开发者,这个仓库是理解NAS从"蛮力时代"到"高效时代"演进过程的绝佳教材,也是复现AutoML里程碑工作的起点。