SCTR 五次失败的安全 BN 路由器

前言

本文是对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_meanrun_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 还差一截,判失败。


为什么全失败?

  1. 全局非线性:60 个 BN 层同时变,总效果不是各层单独效果的相加
  2. 信号太弱:mAP 变化在 0.003 量级,比激活本身的波动小太多,特征里看不见
  3. 方向不可知:能测出变化的大小,测不出变化是正是负
  4. 训练状态混杂:统计量不只取决于当前数据,还取决于整个训练过程

第 2 条是最根本的。上一篇量过,BN 漂移带来的 mAP 变化只有 0.003,而激活的自然波动比这大得多。信号就这么点,描述子做得再细也捞不出来。

当然也没有白做实验:

  • 反过来说明 BN-freeze 对照是唯一可靠的办法。想知道更新好不好,就两个都跑一遍,别猜
  • 证伪过程是完整的。oracle 上界、平凡基线、分组 LOOCV,三个都在,才敢下"失败"这个结论

当然你也可以认为是我实验没做对,思考方向不够全面。我发这篇笔记的原因就是给大家一个思考的提议,希望你能避开错的方向,并且吸取经验总结出更好的实验。

希望各位能以此为戒吧!

相关推荐
AgentMaster1 小时前
元数据、血缘、质量、安全四大模块能力拆解,数据治理方案对比:4 种技术路线深度评测
大数据·数据库·数据仓库·人工智能·原型模式
霸道流氓气质1 小时前
Spring AI vs Spring AI Alibaba:技术选型与平滑迁移策略
java·人工智能·spring
影视飓风TIM1 小时前
C++11 核心新特性完整梳理
数据结构·c++·算法
艾莉丝努力练剑1 小时前
【AI大模型接入SDK】Ollama本地大语言模型部署
c++·人工智能·语言模型·自然语言处理·面试
晴天的雨.9921 小时前
[C++]算法双指针 复写0
数据结构·c++·算法
Raas1001 小时前
AI网关和OpenRouter区别在哪?MAI Gateway(魔芋企业级AI网关)统一治理方案深度解析
大数据·人工智能·gateway·ai网关·mai gateway
天赐范式1 小时前
天赐范式第159天:让广播走出本机——HTTP通道发出第一声
python·broadcast·urllib·http.server·天赐范式·动态运行时·http通道
JJJennie7771 小时前
MAI Gateway能力解析:大模型网关支持本地模型吗?AI网关核心功能详解
人工智能