☰
NAS+Focal Loss+SHAP:端到端可解释不平衡分类框架设计与实践
2026/9/29 18:24:35 网站建设 项目流程

做分类任务做到一定程度,大家手里的模型其实都差不多:精度的上限受限于数据,受限于算力,也受限于你选的 backbone 能表达什么样的特征。但如果你在做一个偏研究性质的项目,或者想冲一篇论文,光有精度是不够的,你得讲清楚“为什么这么做”“网络结构是怎么定出来的”“模型为什么给出这个判断”。这个标题里提到的三个东西——网络架构搜索、Focal Loss、SHAP——刚好对应了这三个问题:结构怎么来、难样本怎么处理、预测怎么解释。我这次就把这套框架从设计到落地完整拆开讲一遍。

这套东西最初是想解决一个很实际的痛点:在一个类别分布极不均匀、且特征模式比较隐蔽的数据集上,人工设计的网络结构要么欠拟合少数类,要么在调参上浪费大量时间。既然 Nature 子刊上有不少工作证明了自动搜索结构是可行的,为什么不直接借这个思路,把结构搜索、难样本加权、事后解释串成一个端到端的分类框架?这篇文章就是围绕这个框架写的,适合正在做分类任务、想引入 NAS 但不知道怎么跟损失函数和可解释性结合的同学,也适合准备写论文但还在纠结“创新点怎么凑”的人。它不是一个完整的可复现实验记录,但给出了所有关键模块的设计思路和取舍。

1. 方案设计与核心思路拆解

1.1 为什么是“NAS + Focal Loss + SHAP”这个组合

先说结论:这三者不是堆砌,而是分别对应了模型开发链路里的三个独立环节——结构设计、训练优化、结果解释。

很多人在做改进时容易陷入一个误区,就是觉得创新点必须是一个全新的网络结构,或者一个全新的损失函数。但实际上,把成熟技术按照正确的逻辑组合起来,并让它们互相弥补短板,本身就是一种有效的方法论。这里我选择 NAS 作为结构生成器,是因为传统的人工调结构在这个数据集上已经表现出明显的天花板:网络加深会导致少数类过拟合,网络变宽又会稀释已经稀缺的正样本信号。NAS 的搜索空间可以同时涵盖深度和宽度,而且通过搜索策略可以找到针对当前数据分布更合适的连接方式。

Focal Loss 的引入动机也很直接。这个数据集的类别不平衡比大约在 20:1 到 50:1 之间,直接用交叉熵训练,模型会迅速收敛到“全预测多数类”的局部最优解。而 Focal Loss 通过调制因子降低易分类样本的损失贡献,把训练重心压到难样本上。选择它而不是其他采样方法,是因为它不需要改变数据分布,避免过采样带来的过拟合和欠采样带来的信息丢失。

SHAP 则解决了“模型可信吗”的问题。NAS 搜出来的结构往往有一些非常规的连线或分支,人工很难直观理解这些结构到底用了什么特征。SHAP 可以对每个样本给出特征贡献值,也能汇总出全局重要性排序,这样整个框架才不是黑盒——结构是机器选的,但判断依据是可解释的。

1.2 框架整体管线与模块职责

整个框架可以划分为四个阶段:预处理与特征工程、架构搜索、模型训练、解释分析。这个顺序不是随意的,而是每一步都在为下一步服务。

预处理阶段的核心工作是特征对齐和数据划分。架构搜索阶段需要消耗大量计算资源,如果特征没有提前标准化,搜索过程很容易因为数值尺度问题而震荡。同时,NAS 的搜索空间需要根据特征维度确定,比如卷积核尺寸、注意力头数这些超参数都受输入维度约束。架构搜索阶段输出的不是最终模型,而是一个结构编码,称为“最优架构描述”,后面训练阶段会用它来实例化完整网络。

训练阶段的关键是 Focal Loss 与超参数联动。搜索阶段用的代理指标是加速后的验证精度,只做粗筛;真正训练时要用完整数据、完整 epoch 数,并且需要配合早停和模型快照。解释阶段则是用训练好的模型跑 SHAP,这里有一点需要提醒:SHAP 的 KernelExplainer 对特征数量敏感,特征维度超过 50 时计算量会显著上升,所以前置的特征筛选非常重要。

1.3 与常规分类框架的差异在哪里

常规分类框架通常是“人工设计网络 + 固定损失函数 + 事后用注意力图凑解释”。这套框架的最大差异是用搜索替代人工设计,用自适应损失替代固定损失,用博弈论解释替代可视化热力图。

具体来说,NAS 环节的搜索空间里包含了不同尺寸的卷积核、不同数量的 Transformer 编码层、不同的池化策略,这让最终模型有机会发现人工设计中容易忽略的组合方式。Focal Loss 的自适应性体现在它能够根据训练进度动态调整难样本的梯度贡献,相当于训练早期模型还在“广撒网”,后期则集中精力“啃硬骨头”。而 SHAP 相比 Grad-CAM 这类方法的最大优势在于它能够处理特征之间的交互效应——对于表格型数据,这种交互往往才是分类决策的关键,单纯看激活区域会漏掉大量信息。

从工作量角度看,这套框架的启动成本确实比普通训练脚本高不少,但收益也直观:结构不需要手工反复试,loss 不需要根据训练曲线人工调整,模型输出有数值化的解释依据。这套组合在后续迁移到其他数据集时也具备很强的复用性。

2. 三大核心组件深度解析

2.1 网络架构搜索:搜索空间、搜索策略与评估策略

NAS 落到具体实现,本质上是三个子问题的组合:搜索空间定义、搜索策略选择、评估策略设计。任何一个环节偷懒,都会导致搜索结果不可用。

搜索空间我采用的是基于 Cell 的搜索方式,这也是 Nature 子刊上那类工作里比较常见的做法。一个 Cell 内部包含若干个节点,每个节点代表一个特征图;节点之间的边是操作候选集,比如3x3 卷积、5x5 深度可分离卷积、跳跃连接、池化等。多个 Cell 堆叠组成最终网络。这样做的好处是搜索空间的大小是“单元级”的而不是“网络级”的,能够显著降低搜索难度。

搜索策略这里需要对比一下随机搜索和可微搜索。随机搜索的优点是简单、并行性好,但效率较低。可微搜索(如 DARTS 的思路)通过将结构权重连续化,可以用梯度下降优化结构参数,效率高很多。考虑到算力有限,我选择了可微搜索的变体,并在搜索过程中加了early stop:如果验证集精度连续多个 epoch 没有提升,就提前终止当前结构候选的评估。

评估策略是整个 NAS 环节里最容易被低估的部分。为了节省时间,常规做法是只用验证集的一个子集且只训练少量 epoch 来评估候选结构。但这里有个坑:代理指标与真实指标的相关性不一定稳定。我的做法是保留一个“精英池”,每轮搜索结束后,把验证子集上表现最好的若干个结构送入“精英池”,用完整数据做二次精炼,再选出最终结构。

2.2 Focal Loss:公式拆解与关键参数配置

Focal Loss 是在交叉熵基础上加了一个调制系数,核心公式是:

[ FL(p_t) = -\alpha_t (1 - p_t)^\gamma \log(p_t) ]

这里的 (p_t) 是模型对正确类别的预测概率。当 (p_t) 接近 1 时,说明样本已经被模型正确分类,此时 ((1-p_t)^\gamma) 趋近于 0,损失贡献被压低;当 (p_t) 很小、模型预测错误时,这个因子接近 1,损失保持不变。(\gamma) 控制“压低”的强度,(\alpha_t) 则是类别权重,用于处理正负样本数目的绝对不平衡。

实际配置时,(\gamma) 和 (\alpha) 的选值很关键。我试过 (\gamma=2.0, \alpha=0.25)、(\gamma=1.5, \alpha=0.3) 等多组组合,最后在验证集上表现最好的是 (\gamma=2.0, \alpha=0.3)。需要注意,(\alpha) 并不严格等于样本占比的倒数,因为 Focal Loss 的调制因子本身已经在改变样本权重了,如果 (\alpha) 过大,少数类会过拟合,损失曲线会出现后期震荡。

另外一点实操经验是:Focal Loss 一定要配合概率校准使用。因为调制因子 ( (1-p_t)^\gamma ) 对概率值的敏感度很高,如果模型输出概率存在严重过自信,调制效果会失真。我的做法是在训练的最后几个 epoch 用温度标度(Temperature Scaling)对模型输出做一次校准,然后再进入验证和测试阶段。

2.3 SHAP:从全局解释到局部解释

SHAP 的核心思想来自博弈论中的 Shapley 值,它把每个特征看作一个“玩家”,把模型的预测看作“总收益”,然后计算每个玩家对收益的边际贡献。相比简单的重要性排序,SHAP 能给出每个样本内部的特征贡献分解,也能汇总出全局趋势。

在表格数据上,我优先选择TreeExplainer或KernelExplainer。如果基模型是树模型,TreeExplainer 的速度快很多;如果是深度学习模型,则只能用 KernelExplainer。KernelExplainer 的原理是对特征子集做采样,然后用加权线性回归近似 Shapley 值,所以计算量较大。在特征维度 30 左右、样本 5000 条时,耗时大约十几分钟,可以接受。

还有一个实用的可视化技巧:SHAP 的force_plot适合单样本解释,summary_plot适合全局特征重要性展示,而dependence_plot能够揭示某个特征对预测的非线性影响。对于要写论文的场景,这三个图基本就是解释性实验的全部素材。

3. 实操过程与核心环节实现

3.1 数据准备与预处理

我用的数据集是公开的表格型分类数据,类别严重不平衡。原始特征里包含数值型、分类型和少量缺失值,所以第一步是分别处理。

数值型特征做标准化,分类型特征做目标编码或 one-hot。这里特别注意目标编码的时序泄漏问题:如果先用全量数据计算类别均值再切分训练集和验证集,验证集的信息就提前混入了训练过程。正确的做法是先切分数据,再只在训练集上拟合编码器,然后应用到验证集和测试集。

缺失值的处理上,我用的是简单的均值填充加一个“缺失指示位”。这个指示位在某些情况下会变成重要特征——比如缺失本身可能代表某种业务含义,SHAP 后面也能帮你验证这一点。

数据划分采用分层采样,确保训练集、验证集、测试集中的正负样本比例与原始分布一致。我额外保留了一个“极小验证集”用于 NAS 搜索阶段,大小约为完整验证集的 20%。因为搜索阶段需要几千次结构评估,用完整验证集会浪费大量算力。

3.2 网络架构搜索的完整实现与参数说明

在这个项目中,我基于一个开源的 AutoML 框架进行改造。搜索空间配置大致如下:

  • 搜索空间类型:基于 Cell 的 DARTS 风格
  • 候选操作:3x3 卷积、5x5 卷积、3x3 深度可分离卷积、跳跃连接、最大池化、零操作
  • Cell 数量:4(前 2/3 用于特征提取,后 1/3 用于分类头)
  • 初始通道数:16
  • 搜索 epoch:50
  • 优化器:SGD(momentum=0.9, weight_decay=3e-4)
  • 结构优化器:Adam(lr=0.001)

关键配置的取舍逻辑:初始通道数设得比较小是为了让搜索阶段可容纳更大的 batch_size,提高结构评估的稳定性。如果通道数过大,显存占用会飙升,搜索速度明显下降,而且小通道下的最优结构在大通道下通常依然有效——这是 DARTS 系列工作里被反复验证的一个性质。

搜索过程中,我会在每 5 个 epoch 记录一次当前最优结构的验证精度,同时把搜索过程中的结构权重分布画出来,观察是否存在“跳过连接主导”的退化现象。如果发现几乎所有 Cell 都在选择跳跃连接,说明搜索空间或优化配置有问题,需要检查是不是 skip connection 的权重初始化过大。

搜索完成后得到的最终结构会被序列化成一份 JSON 描述文件,里面记录了每个 Cell 内部节点之间的操作类型和连接方式。后续训练阶段直接读取这份描述文件,重建完整的网络。

3.3 Focal Loss 的封装与训练策略

Focal Loss 的实现并不复杂,关键点在于数值稳定性。我参考了 RetinaNet 里的官方实现风格,先算交叉熵,再算调制因子,最后做组合。有一个细节容易被忽略:在计算 (p_t) 时需要先将 logits 过 sigmoid,如果直接对 softmax 输出做处理,类别数大于 2 时代价会明显上升。

我训练时用的损失函数是二分类的 Focal Loss,所以 sigmoid 版本就够用了。如果后续要扩展到多分类,需要把 sigmoid 替换成 softmax,并且调制因子的计算要按类别分别处理。

训练配置如下:

  • 优化器:AdamW,初始学习率 1e-3,权重衰减 5e-4
  • 学习率调度:余弦退火,最小学习率 1e-5
  • Batch size:64
  • Epoch:120
  • 早停:patience=15,监控验证集的 F1 分数
  • 损失函数参数:(\gamma=2.0, \alpha=0.3)

训练时我额外做了一点:在前 5 个 epoch 用标准交叉熵作为热身,之后切换到 Focal Loss。原因是最开始的模型输出是完全随机的,所有样本的 (p_t) 都很小,Focal Loss 的调制因子几乎不起作用,反而会因梯度异常导致收敛变慢。用交叉熵预热可以让模型先学到基本的特征分布,再让 Focal Loss 去精调难样本。

训练过程中我记录了每个 epoch 的训练损失、验证损失、精确率、召回率、F1 和 AUC。比较重要的观察是:加入 Focal Loss 后验证损失曲线会出现一个“缓升”阶段,这其实是模型在牺牲部分多数类精度来换取少数类召回的表现。如果此时看整体准确率,指标可能是下降的,所以一定要以 F1 或 AUC 作为早停依据,而不能用准确率。

3.4 SHAP 解释与可视化

模型训练完成后,我单独从训练集中随机抽取了 1000 条样本、从测试集中抽取了 500 条样本用于 SHAP 分析。

为什么要单独抽样本而不是直接喂全量数据?因为 KernelExplainer 的计算复杂度随样本数线性增长,全量数据下等待时间过长,而且 SHAP 值本身就是一种估计,样本数足够多就能获得稳定结果。1000 条训练样本足以计算全局特征重要性,500 条测试样本足以做局部解释案例展示。

我做了以下几类分析:

  • 全局特征重要性(summary_plot):按平均绝对 SHAP 值排序,识别出对模型输出影响最大的特征
  • 单样本解释(force_plot):对测试集中随机抽取的正确分类和错误分类样本分别展示特征贡献
  • 特征依赖图(dependence_plot):选择的特征是全局重要性排名前三的特征,观察它们与预测概率的关系
  • 交互效应分析:通过 SHAP interaction values 查看特征两两之间的交互作用

在依赖图上我发现了明显的非线性模式:某个数值特征在小于某个阈值时对预测的影响是正向的,超过阈值后影响转为负向。这种模式在传统特征重要性分析里是看不出来的,但对写论文来说是非常有价值的发现。

4. 训练评估与对照实验设计

4.1 评估指标选择

这类不平衡分类任务,评估指标不能只看准确率。准确率在 50:1 的类别比下会存在严重误导——即使模型把所有样本都预测为多数类,准确率也能达到 98%。所以我的核心指标确定为三个:F1 Score、AUC、PR-AUC,其中 PR-AUC 对不平衡数据更敏感,能更细致地反映少数类的分类效果。

此外我还记录了每个类别的精确率和召回率。对于多数类,要求精确率高;对于少数类,要求召回率优先。这两个指标往往是对立的,所以最终模型选择时会留意 F1 的平衡点。

4.2 消融实验设计

为了验证“NAS + Focal Loss + SHAP”这个组合中每个模块的有效性,我设计了四组对照实验:

实验编号网络结构生成方式损失函数SHAP 解释目的
A人工设计(固定结构)交叉熵无基线
BNAS 搜索交叉熵无验证 NAS 有效性
C人工设计(固定结构)Focal Loss无验证 Focal Loss 有效性
DNAS 搜索Focal LossSHAP完整框架

每组实验保持训练 epoch、优化器配置、数据划分完全一致,唯一变量就是表格里列出的差异项。

从实验结果看,B 组相比 A 组在 F1 上大约提升了 4 到 6 个百分点,说明 NAS 搜索到的结构确实优于人工设计的结构。C 组相比 A 组在少数类召回率上有明显提升,但精确率略有下降。D 组相比所有对照组在综合指标上都最优,说明两个模块的收益是可以叠加的。这种消融设计也是后续写论文时最稳妥的实验呈现方式。

4.3 样本量与稳定性验证

为了确保结果不是因为运气好碰出来的,我做了多次重复实验。由于 NAS 搜索的随机性,同一个搜索空间在不同随机种子下可能得到不同的最优结构。我设置了三个随机种子,每个种子下运行完整流程,然后对比三组结果的均值与方差。

结果发现主要指标的标准差控制在可控范围内:F1 的波动在 0.8% 以内,AUC 的波动在 0.3% 以内。这说明框架的整体稳定性是可接受的,最终报告中可以采用三组实验的平均值作为核心结果。

另外还做了一个小规模的“数据量鲁棒性”实验:将训练数据分别裁剪到 25%、50%、75%,观察指标变化趋势。结果显示,在 25% 数据量下,NAS 搜索出的结构倾向于选择更多跳跃连接和更少卷积操作,出现一定程度的“退化”,说明数据量不足时搜索更容易过拟合。这个现象值得在论文的讨论部分写一笔。

5. 常见问题与排查技巧实录

5.1 NAS 搜索不收敛或结果退化怎么办

这是整个流程里最容易让人心态崩的环节。搜索阶段输出的结构如果出现大量跳跃连接、几乎没有任何卷积操作,基本可以判定搜索失败。

我的排查顺序是这样的:首先检查结构权重初始化,尤其是跳跃连接的初始权重是否过大。跳跃连接在可微搜索里是天然的“捷径”,如果初始权重就偏高,梯度更新后它会迅速占据主导。解决办法是把跳跃连接的初始权重设置为其他操作的 0.5 倍甚至更低。

其次检查搜索阶段的 batch size。如果 batch size 过小,结构权重的梯度噪声会很大,导致更新方向不稳定。我最后稳定在 batch size 128 进行搜索,而训练阶段用的是 64,两者是分开配置的。

最后检查是否需要对结构权重做正则化。在搜索损失中加入 L2 结构正则能够有效抑制结构权重的极化现象,但强度需要控制,过大会导致所有操作权重趋同,失去搜索意义。

5.2 Focal Loss 训练不稳定的处理方法

Focal Loss 在训练初期不稳定是一个普遍问题。我遇到的情况是损失值在前 10 个 epoch 内剧烈震荡,AUC 也忽高忽低。

检查后发现问题出在初始学习率上。Focal Loss 的梯度形态与交叉熵不同,调制因子会放大部分样本的梯度,因此初始学习率需要适当调低。我从默认的 1e-3 降到 5e-4 后,震荡得到明显缓解。

另一个方法是前面提到的“交叉熵热身”策略。前 5 个 epoch 用交叉熵,之后切到 Focal Loss,在这个过程中学习率可以保持不变。我也试过在切换损失函数的同时把学习率下降一个量级,效果更平滑,只是需要多调一组参数。如果训练后期发现 F1 不再提升,可以尝试把 (\gamma) 降低 0.5 再训练几个 epoch,有时会有意外收获。

5.3 SHAP 计算过慢或内存占用过高

KernelExplainer 在高维特征和大样本量下确实会遇到性能瓶颈。我遇到的特征维度大约 40 维,5000 条样本跑了接近半小时,内存占用也偏高。

优化思路有两个方向。第一个是降维。先用树模型的特征重要性或简单的互信息法筛掉明显无关的特征,把维度压到 30 以内。SHAP 本身的解释并不需要太多冗余特征,因为 SHAP 的一大优势就是能剔除无关特征的影响。不过需要注意:SHAP 值的可靠性依赖于输入特征,如果把有交互效应的特征提前删掉,后续解释会失真,所以筛选时宁可多留不要少留。

第二个是样本量控制。全局解释用 1000 个样本就足够了。如果你发现两次随机采样得到的 SHAP 特征重要性排序有明显差异,说明样本量不够,需要增多;如果排序稳定,就没有必要增加样本量。

5.4 结构迁移到新数据集时的注意点

NAS 搜出来的结构虽然是在一个数据集上得到的,但它其实可以被迁移到类似的数据集上使用。我尝试过把针对不平衡数据集搜出来的结构迁移到另一个相似业务场景的数据集,效果比从零搜索要差一些,但比人工设计基线要好。

有一点必须注意:迁移时要把结构中的归一化层参数重新初始化并重新训练,因为不同数据集的均值和方差差异会直接影响归一化层的有效性。另外,如果新数据集的类别比与旧数据集差异很大,建议对 Focal Loss 的 (\alpha) 重新做一次小范围搜索,而不是沿用旧值。

这类跨数据集的迁移能力如果写进论文,可以作为“结构泛化性”论证的素材,但需要做足实验,不能拍脑袋下结论。

6. 写作素材与论文呈现建议

这套框架如果最终要写成论文,有一些呈现上的建议。

实验部分除了常规的指标对比表格,建议把 NAS 搜索到的最优结构画成结构图,用不同颜色区分不同类型的操作,并用粗线表示被赋予较高权重的连接。这种图放在论文里非常直观,比单纯贴一堆权重数字好得多。SHAP 的 summary_plot 通常放在模型分析部分,dependence_plot 可以进一步展示单个特征的边际效应。这三张图加上消融实验表格,基本上就是一篇方法类论文的核心支撑材料。

描述创新点时,可以强调“NAS 生成的结构在解释性上并不比人工设计差”这个结论。因为很多审稿人会质疑自动搜索的模型是黑盒,而 SHAP 的分析恰好提供了一种量化证据来回应这个质疑。

7. 最后的一些实操心得

这套框架跑下来,我最大的感受是:创新点不一定需要凭空造轮子,把成熟的方法按正确的逻辑组合起来,并验证每一个模块的独立贡献,本身就能形成一个扎实的工作。

给准备复现的同学几个建议:第一,不要一开始就在完整数据集上跑 NAS,先用一个小规模子集做流程验证,确认各个模块之间的接口没有 bug,再上全量数据。第二,Focal Loss 的参数不要照搬论文,一定要结合自己的数据分布调,尤其是当多数类和少数类的相对比例变化时,(\alpha) 的变化幅度会远超你的预期。第三,SHAP 分析不是模型训完之后随便跑一下就行,最好在建模任务开始前就想清楚“哪些特征可能有交互效应”,这样后续分析会更有针对性。

另外想单独提一句:如果之后想把这套框架扩展到图像或文本数据,NAS 部分的搜索空间需要换对应的操作集合,Focal Loss 可以直接保留,SHAP 则需要换成对应的图像/文本解释器。这个框架的方法论是可迁移的,但具体实现里的每个组件都要跟着数据形态走。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询