前言
本文是对BN漂移的文末留白回应,介绍一些解决BN漂移问题的工程方法。
当然,如题全失败了,希望大家引以为戒。
问题的由来
上一篇我们确认了一件事,BN 统计量更新带来的 mAP 变化,方向和幅度都不稳定:
- 物体尺度上,小物体 +0.0031、中物体 −0.0037、大物体 +0.0005
- 模型规模上,YOLOv8n +0.0036、YOLOv8l −0.0017
- 数据集上,UAVDT +0.0036、VisDrone 各个 epoch 全是负
作用效果不确定,那么能不能在更新之前,只看图像和特征,就先判断出这次到底安不安全?
这就是 SCTR(Safe BN-Update Router)。测试时手上没有标注,算不出 mAP,所以想法是训练一个估计器,输入图像和特征,输出安全还是危险。
制定规则
动手之前先把评价标准定下来,不然容易乱。
- 目标:不看训练状态,只靠输入预测更新 BN 是否安全
- 难点:测试时没有标注,算不出 mAP
- 门槛:F1 要比平凡基线高出至少 0.05
第三条是关键。平凡基线就是永远回答安全,是一个什么都不做的策略。要是不能比它高0.05,几乎是没生效。
穷举所有候选阈值(>= 和 <= 两个方向都试),取最好的 F1。这是 oracle 上界 ,即上帝视角,也就是已经知道答案了再回头挑最好的那条分数线。
确定指标
因为测试时算不了mAP,所以我们需要自己确定指标。
同一个窗口,跑两个模型
candidate BN-update 之后的模型
base BN-freeze 的模型
两个 mAP 相减得 delta
delta >= 0 safe_bn_update 不掉分
delta > 0 beneficial_bn_update 涨分
结果先验
| 阶段 | 粒度 | 描述子数 | 方法 | SCTR | 平凡 F1 | 超出 | 失败原因 |
|---|---|---|---|---|---|---|---|
| 0A | 序列 | 36 | 频谱特征(谱熵/功率谱) | 0.607 | 0.692 | −0.085 | 真实标签不可预测 |
| 0B | 窗口 | 81 | +多分辨率预测差异 | 0.710 | 0.718 | −0.008 | 退化为永远更新 |
| 0C | 窗口 | 270 | +内部激活特征 | 0.704 | 0.718 | −0.014 | 激活变化 ≠ 安全 |
| 0D | 窗口 | 390 | +分布形状特征 | 0.687 | 0.718 | −0.031 | 跨序列不泛化 |
| 0E | 窗口 | 3 | 预测自一致性(flip) | 0.842 | 0.820 | +0.022 | 信号太弱(<0.05) |
描述子从 36 加到 390,翻了十倍,结果反而掉了。只有 0E 涨,而且它只用了 3 个描述子。
阶段解析
前置
- 数据单元 :UAVDT val 按文件名前缀分成一个个序列,同一段视频算一个序列。0B 开始每 64 张切一个窗口,编号用
idx // 64,比如{seq}_w0001 - 标签:就是前文说的指标
- 打分:穷举阈值取 oracle F1,跟平凡基线比
0A 图像频域统计(36 个)
图像糊不糊,跟 BN 更新安不安全有关系。
将每个序列采 24 张图,每张图转灰度后算下面这些量。
python
x = gray - gray.mean()
fft = np.fft.fftshift(np.fft.fft2(x, norm="ortho"))
power = fft.real**2 + fft.imag**2
hfer = power[radius >= tau].sum() / power.sum()
prob = power / power.sum()
entropy = -(prob * np.log(prob)).sum() / np.log(prob.size)
gray - gray.mean() 先去均值。频谱最中心那一点是 DC,也就是整张图的平均亮度,它的能量通常比别处大好几个量级,不减掉的话后面算占比基本都被它一个人吃掉。
np.fft.fft2 是二维傅里叶变换,把图像从像素值变成配方表,每个位置记一个频率成分有多少。
norm="ortho" 是让正变换和逆变换各除一个根号 N,这样变换完能量总量不变,方便直接比大小。
np.fft.fftshift 是把 DC 从左上角搬到正中间。不搬的时候最高频在正中心,搬完之后离中心多远就对应频率多高,这样才好画圈做掩码。
这几步的推导之前单独写过一篇频域的频域数学推理版,这里就不展开了。
power 是功率谱,复数取模的平方。复数有实部和虚部,模的平方就等于实部平方加虚部平方。
hfer 是高频能量比。radius 是每个频率点离中心的归一化半径,tau 取 0.5,也就是把外面那一圈的能量加起来,除以总能量。
entropy 是谱熵。先把功率谱归一化成一个加起来等于 1 的分布,再套信息熵的公式。它量的是能量分布得均不均匀,完全集中在一点就是 0,完全铺开就是 1。除以 np.log(prob.size) 是压到 0 到 1 之间。
除了这两个频域的,再加两个空域的。
python
edges = cv2.Canny(gray_u8, 80, 160) # 边缘密度
lap_var = cv2.Laplacian(gray, cv2.CV_64F).var() # Laplacian 方差
Canny 是边缘检测,80 和 160 是两个阈值,出来的 edges 是一张只有 0 和 255 的图,统计非零点占比就是边缘密度。
Laplacian 是二阶导之和,对灰度做一次卷积。模糊的图二阶导接近 0,方差就小;清晰的图方差大,所以它能当清晰度指标用。
36 个描述子怎么来的?每张图出 4 个指标,分别是 HFER、谱熵、边缘密度、Laplacian 方差。每个指标在序列的 24 张图上算完后,再取 mean、p50、p90 三个统计量,4 乘 3 等于 12。剩下的是几组变体凑到 36。最终每个序列一行、36 列。
结果 0.607,平凡基线 0.692,低了 0.085。
0B 切窗口(81 个)
序列粒度太粗了,一个序列几百张图才出一个标签,样本太少。所以切成 64 张一个窗口,增大样本容量。
描述子从 36 加到 81,加的是多分辨率下的预测差异------同一张图缩放到不同尺寸分别跑推理,看结果差多少。
这一轮改了两处:
第一处,阈值搜索换成七档分位数网格。
python
quantile_cuts = [0.15, 0.25, 0.35, 0.50, 0.65, 0.75, 0.85]
原来是固定几个阈值,现在是按分位数切。等于不管描述子的量级和分布长什么样,都能保证每一档里有样本。
第二处,交叉验证改成按序列分组的 LOOCV。
python
grouped_loocv_one_threshold(X, y, groups=seq_ids)
LOOCV 是留一序列交叉验证,每次留一个序列当测试集,剩下全当训练集,轮一圈。
分组的意思是,留的单位不是单张图或单个窗口,而是整个序列。这一步不能省:同一段视频前后帧长得很像,如果随机切,测试集里的帧和训练集里的帧很可能来自同一段,等于透题,分数会虚高。按序列分组,测试时整个序列模型都没见过,才量得出真实水平。
结果 0.710,平凡 0.718,差 0.008。失败原因是退化为永远更新------它学到的最好策略,就是不管输入什么都猜安全,跟平凡基线干的是同一件事。
0C 看网络内部激活(270 个)
图像方向不管用,那就看架构。
先看 hook 怎么挂。PyTorch 里想拿到某一层的输入输出,又不想改模型源码,就用 forward_hook,它会在这一层前向跑完之后回调你给的函数。
python
handles = []
for name in ["model.12", "model.15", "model.16",
"model.18", "model.19", "model.21", "model.22"]:
layer = get_module_by_name(model, name)
handles.append(layer.register_forward_hook(make_hook(name)))
这几层正好是上一篇说的 60 个 keys 的更新区,也就是 Neck 加 Head 加 Detect。Backbone 在这套配置下是全程冻结的。
每个 hook 里采 BN 层的输入和输出:
python
pre_mean, pre_var = pre.mean(dim=(2, 3)), pre.var(dim=(2, 3))
post_sparsity = (post.abs() < 1e-3).float().mean(dim=(1, 2, 3))
特征图形状是 [B, C, H, W]。在 dim=(2, 3) 上求均值,就是把每张图的长宽取均值,类似空间注意力,出来形状是 [B, C]。因为BN 本来就是按通道算统计量的,所以我们关心的也是每个通道一个数,空间位置不重要。
post_sparsity 是输出里绝对值小于 1e-3 的比例,也就是有多少激活基本是死的。
核心是下面这组,拿当前这批数据算出来的统计量,跟三个 checkpoint 各自存着的 running stats 比距离。
python
mean_dist = (pre_mean - run_mean).abs() / torch.sqrt(run_var + eps)
var_dist = (torch.sqrt(pre_var + eps)
- torch.sqrt(run_var + eps)).abs() / torch.sqrt(run_var + eps)
run_mean 和 run_var 是模型里存着的那套 running stats,上一篇讲过,推理时归一化用的就是它。除以 torch.sqrt(run_var) 是做归一化。
eps 是个很小的常数,一般 1e-5,防止分母为 0。
270 怎么来的?上面这些特征对 base、freeze、update 三个 checkpoint 各算一遍,7 个 BN 层每个层一组,7 乘 3 再乘若干指标,凑到 270。
结果 0.704,比 0B 还低。问题在这一组距离量的是当前数据和模型记着的差多少,也就是变化的大小。但标签要的是这个变化是好是坏,设计的有点偏了。
0D 再加分布形状(390 个)
0C 用的是均值和方差这两个汇总量,等于把整个激活分布压成两个数,中间的信息全丢了。两个分布可以均值、方差都一样,但形状完全不同。
0D 就补这个,开 include_shape=True,直接从激活值里采点统计形状。
python
sample = post.flatten().float()
if sample.numel() > 4096:
idx = torch.linspace(0, sample.numel() - 1, 4096).long()
sample = sample[idx]
feats = torch.cat([
sample.mean().unsqueeze(0),
sample.std().unsqueeze(0),
torch.quantile(sample, torch.tensor([0.05, 0.25, 0.50, 0.75, 0.95])),
((sample - sample.mean()) ** 3).mean().unsqueeze(0), # 偏度
])
flatten() 是把整张特征图拉成一维,不管原来是 [B, C, H, W] 什么形状。
采样 4096 个点,是因为激活值数量太大,全拿来算分位数太慢。
torch.quantile 是分位数。0.05 那个位置的值、0.25 那个位置的值,依次类推。几个分位数拼起来,就大致描出了分布的轮廓。
最后那个三次方的均值是偏度,量分布歪不歪。对称的分布它接近 0,一边拖长尾巴它就明显偏离 0。
270 加 120 等于 390。
结果 0.687,又降了。失败原因是跨序列不泛化。描述子越细,越容易记住单个序列自己的特点,换个序列就不灵了。这跟 0B 里用分组 LOOCV 测出来的是同一个毛病。
0E 只看预测自不自洽(3 个)
前面四轮一直在加东西,0E 反过来做减法。不看网络内部,只看更新前后模型自己的预测结果一不一致。
同一个窗口,用更新前和更新后的模型各跑一遍,比两边的结果。
python
score = view_jaccard + match_ratio + matched_iou \
- matched_score_abs_delta - count_abs_delta / 10.0
一项一项说。
view_jaccard 是两边所有框并起来的区域,交集除以并集。量的是整体覆盖了同一片地方没有。
match_ratio 是能配对上的框占多少。配对规则是 IoU 超过阈值就算一对。IoU 就是两个框交集面积除以并集面积。
matched_iou 是配好对的那些框,IoU 平均下来多少。
前三项都是两边结果像不像。
matched_score_abs_delta 是配好对的框,置信度差了多少,取绝对值。这是罚分项,减号。
count_abs_delta / 10.0 是框的总数差了多少,除以 10 是把量级缩一下。不然框数动辄几十上百,会盖过前面几项。
逻辑也不复杂:更新之后如果模型自己都跟自己对不上,那这次更新多半不妙。
结果 0.842,平凡 0.820,超出 0.022,五轮里最好的。但离 0.05 还差一截,判失败。
为什么全失败?
- 全局非线性:60 个 BN 层同时变,总效果不是各层单独效果的相加
- 信号太弱:mAP 变化在 0.003 量级,比激活本身的波动小太多,特征里看不见
- 方向不可知:能测出变化的大小,测不出变化是正是负
- 训练状态混杂:统计量不只取决于当前数据,还取决于整个训练过程
第 2 条是最根本的。上一篇量过,BN 漂移带来的 mAP 变化只有 0.003,而激活的自然波动比这大得多。信号就这么点,描述子做得再细也捞不出来。
当然也没有白做实验:
- 反过来说明 BN-freeze 对照是唯一可靠的办法。想知道更新好不好,就两个都跑一遍,别猜
- 证伪过程是完整的。oracle 上界、平凡基线、分组 LOOCV,三个都在,才敢下"失败"这个结论
当然你也可以认为是我实验没做对,思考方向不够全面。我发这篇笔记的原因就是给大家一个思考的提议,希望你能避开错的方向,并且吸取经验总结出更好的实验。
希望各位能以此为戒吧!