推理阶段的计算图一旦固定,结构上的冗余就会被固化成延迟和功耗。神经架构搜索(NAS)把设计网络结构从经验活变成可自动求解的组合优化问题,其目标不是在训练集上无限提高精度,而是在给定硬件和延迟预算内找到最能兼顾精度与推理效率的那一个结构。

一、推理架构寻优与传统NAS的区别
传统NAS大多以最终验证精度作为选择标准,搜索得到的网络可能参数量小,但在特定硬件上推理并不快。例如深度可分离卷积的FLOPs较低,但在部分GPU上访存开销大,实际延迟未必低于普通卷积。推理架构寻优需要把延迟、内存峰值、功耗等部署指标纳入评估,而不是只看浮点运算数。因此搜索目标常写成精度与惩罚项的加权和,例如目标值等于验证精度减去延迟惩罚系数乘以实测延迟或代理延迟。
这种差异会直接影响搜索空间设计。面向推理的搜索空间通常限制在可部署的算子集合内,包括标准卷积、深度可分离卷积、倒残差模块、注意力模块等,并且每个候选操作会预先标注其在不同输入形状下的推理成本。通过查询成本表或使用小型延迟预测器,搜索算法可以在不真正编译运行每个候选网络的情况下过滤掉明显超预算的结构,从而把有限的计算资源用在更有希望的候选上。
另一个关键变化是约束条件的表达。传统NAS往往只约束参数总量或FLOPs,而推理架构搜索会引入多级约束,例如移动端NPU要求卷积核尺寸不能超过5×5,DSP对混合精度支持有限,或者内存带宽不允许通道数过大的特征图。搜索策略需要在满足这些硬约束的可行域内搜索,否则产出的架构即使精度再高也无法落地。
二、推理代价如何进入搜索目标
最常见的做法是代理指标法。训练一个轻量回归模型,输入候选网络的结构编码,输出预测的推理延迟或功耗。结构编码可以包含每层的算子类型、输入通道、输出通道、卷积核大小、步长和激活函数。延迟预测器在数百个随机采样的真实测量值上训练,误差通常能控制在5%以内,之后搜索过程中无需反复实测,速度可提升数十倍。
更贴近硬件的方法是查表累加法。把网络每一层在不同输入尺寸、不同通道数下的延迟预先测量并存储为查找表,搜索时按层累加得到总延迟。这种方法假设层与层之间不存在明显调度重叠,对串行推理的移动端CPU和DSP比较适用。对于GPU等并行度较高的硬件,层间融合和内核启动开销不能忽略,查表结果可能偏乐观,此时需要引入端到端实测校准。
如果部署环境支持模型级二进制搜索,还可以把推理成本直接放进进化算法的适应度函数中。例如每次变异得到新候选结构后,先进行一次低精度快速训练,再在目标设备上跑50次推理取平均延迟,最后按精度与延迟的比值选择下一代。虽然单次评估更慢,但搜索过程能够感知真实调度、算子融合和内存分配带来的影响,适合硬件确定、需要一次搜索长期复用的场景。
三、权重共享与进化搜索的组合实践
权重共享是降低搜索成本的核心思路。构建一个超网络,把所有候选操作和候选通道都包含在内,搜索时从超网络中抽取子网络并继承对应权重,只需少量微调即可评估子网络表现。与从头训练每个候选相比,权重共享能让评估成本降低两到三个数量级。但权重共享也带来了子网络评估排名不准的问题,因为共享权重会互相干扰,搜索后期可能偏向参数更多、更容易被充分训练的结构。
进化搜索在推理架构寻优中常与权重共享配合。初始化一组候选结构作为种群,每一代根据验证精度和推理延迟计算适应度,淘汰表现差的个体,再通过变异和交叉生成新候选。变异操作可以改变某一层的算子类型、通道数或卷积核大小,交叉操作则交换两个候选的连续层片段。由于子网络从超网络继承权重,评估速度快,种群可以维持较大规模,从而增加搜索多样性。
实践中还可以加入帕累托前沿保存机制。将每一代中精度和延迟都不被其他候选支配的个体单独保存,最终从帕累托前沿上根据部署预算选择一个或多个结构。相比固定加权求和,帕累托方法能够为不同硬件平台提供不同选择,例如给服务器GPU选高精度版本,给移动端NPU选低延迟版本,而不需要重新搜索。
四、用PyTorch实现一个轻量推理NAS示例
下面示例定义一个只搜索通道宽度和卷积核大小的简化空间,用于说明推理约束如何参与选择。代码中构建了一个超网络,随机采样多个子网络,用验证损失与延迟惩罚的加权和作为分数,并用简单进化策略更新候选配置。
import random
import torch
import torch.nn as nn
class SearchableBlock(nn.Module):
def __init__(self, in_channels, out_choices, kernel_choices):
super().__init__()
self.in_channels = in_channels
self.out_choices = out_choices
self.kernel_choices = kernel_choices
def forward(self, x, out_channels, kernel_size):
conv = nn.Conv2d(self.in_channels, out_channels, kernel_size,
padding=kernel_size // 2, bias=False)
bn = nn.BatchNorm2d(out_channels)
relu = nn.ReLU(inplace=True)
return relu(bn(conv(x)))
class SearchableNet(nn.Module):
def __init__(self, blocks_config, num_classes=10):
super().__init__()
self.blocks = nn.ModuleList()
for cfg in blocks_config:
self.blocks.append(SearchableBlock(cfg[0], cfg[1], cfg[2]))
self.pool = nn.AdaptiveAvgPool2d((1, 1))
self.fc = nn.Linear(blocks_config[-1][1][-1], num_classes)
def forward(self, x, config):
for i, (blk, cfg) in enumerate(zip(self.blocks, config)):
out_ch = cfg[0]
kernel = cfg[1]
x = blk(x, out_ch, kernel)
x = nn.functional.max_pool2d(x, 2)
x = self.pool(x)
x = torch.flatten(x, 1)
x = self.fc(x)
return x
def latency_penalty(config, coeff=0.01):
cost = 0.0
for out_ch, kernel in config:
cost += out_ch * kernel * kernel
return coeff * cost
def search_step(model, configs, val_loader, device):
scores = []
for config in configs:
model.eval()
total, correct = 0, 0
with torch.no_grad():
for images, labels in val_loader:
images, labels = images.to(device), labels.to(device)
logits = model(images, config)
total += labels.size(0)
correct += (logits.argmax(1) == labels).sum().item()
acc = correct / total
score = acc - latency_penalty(config)
scores.append(score)
ranked = sorted(zip(scores, configs), key=lambda item: item[0], reverse=True)
survivors = [cfg for _, cfg in ranked[:5]]
next_generation = []
for _ in range(20):
parent = random.choice(survivors)
child = []
for out_ch, kernel in parent:
if random.random() < 0.3:
out_ch = random.choice([16, 32, 64])
if random.random() < 0.3:
kernel = random.choice([1, 3, 5])
child.append((out_ch, kernel))
next_generation.append(child)
return next_generation
这段代码把每个候选配置的延迟近似为输出通道数乘以卷积核面积,实际项目中可以替换成查表或延迟预测器。搜索步函数先计算验证精度,再减去延迟惩罚,选择排名靠前的候选作为父代,通过变异生成新一代。为简化示例,卷积层在每次前向时重新创建,实战中建议把不同输出的卷积层预建在超网络中,使用mask选择通道,以减少显存和训练时间。
该示例没有处理层间通道匹配问题,实际搜索时相邻块的输出通道与输入通道必须一致,或者引入可学习的通道维数变换。较完整的实现会为每个搜索块建立独立的候选卷积算子,并通过argmax或Gumbel-Softmax进行可微松弛,使超网络能够端到端梯度更新。
五、落地时容易忽略的三个问题
第一个问题是搜索空间过大导致的不可行性。如果每一层都有十几种算子、十几种通道宽度,组合数会膨胀到无法搜索。可以先用灵敏度分析删除对精度影响较小的层,或者固定浅层与深层的结构,只在关键层搜索。搜索空间一旦缩小,权重共享的排名也会更可靠。
第二个问题是代理指标与真实硬件不一致。FLOPs低的模型可能在目标设备上因为内存非线性访问而变慢。解决方式是阶段式评估:先用FLOPs过滤大部分候选,再用延迟预测器筛选,最后对排名前十的结构做真实设备测试。不要只用一个代理指标从搜索开始到结束。
第三个问题是搜索结束后的重训练不足。搜索过程为了快速评估通常只训练少量轮次,选出的最终结构如果直接继承搜索阶段权重,部署精度往往偏低。需要对最终架构进行完整的数据增强、学习率调度和正则化训练,再与搜索阶段的分数进行对照,防止搜索排名与实际重训练结果发生反转。