神经网络训练到最后,权重到底经历了什么?

先说一个你可能没想到的事实。
深度学习里有一个现象叫"grokking"(延迟泛化),一个网络在训练数据上早就已经完美记住答案了,但泛化能力就是迟迟不出现,非要等上几千个训练轮次,突然某一刻,它就"开窍"了,测试准确率从随机猜测直接跳到接近满分。这不是偶发个例,2022年OpenAI的研究者在模运算任务上系统地观察到了这个现象,而且到今天,没人能给出一个完整的机制解释,说清楚训练过程中到底发生了什么,才导致了这个突然的跳跃。
这篇论文想做的事情,就是试图打开这个黑箱。而它给出的答案相当出人意料:用物理学里研究导电网络、森林火灾蔓延的渗流理论,来描述神经网络训练时权重的坍缩过程。
训练到底在"折叠"什么
要理解这篇论文在说什么,得先搞清楚一个背景。
深度神经网络有大量参数,但训练完之后,很多参数其实是"多余"的,它们本可以被压缩、合并,网络照样能干同样的事。这背后有个重要原因是架构对称性:比如一层里的两个神经元,如果作用完全等价,那把它们互换位置,网络的输出不会有任何变化。
架构对称性*:神经网络里由于结构设计(比如多个神经元功能等价)而天然存在的、不影响输出结果的参数变换方式。
这种对称性会在参数空间里制造出大片"平坦"的区域,专业说法叫不变集。一旦训练轨迹掉进这些区域,它就再也跑不出去了,因为这片区域里的梯度指向恰好把轨迹困在原地。
不变集*:一旦参数进入这个区域,无论后续怎么训练,都不会再离开的子空间。
已有研究(Chen等人2023年的工作)证明了SGD(随机梯度下降)确实会把网络推向这些不变集,效果就是网络逐渐收缩成一个更简单的"子网络",好比原本雇了十个人做一份工作,训练结束后发现其中六个人其实一直在做完全一样的事,可以直接合并成一个岗位。
但这项研究只告诉你终点在哪,没告诉你路怎么走的。网络到底是怎么一步步从"十个各自为战的神经元"变成"一个统一的等价类"的?合并过程是平滑渐进,还是像多米诺骨牌一样一批批地倒?这就是这篇论文真正要回答的问题。
用渗流理论给训练"拍X光"
论文的核心思路是把每个子网络看成一个节点,如果两个子网络的参数轨迹被证明会永远绑定在一起(论文里叫吸引子等价),就在它们之间连一条边。随着训练推进,边会越来越多,直到某个时刻,整个网络里绝大部分节点连成了一整块。
这个过程,跟物理学里的渗流现象一模一样。
渗流*:描述一个系统中局部连接如何逐渐扩展成全局连通的数学理论,经典例子是往一堆随机分布的导电颗粒里逐个添加连接,某个临界点上,颗粒突然从互不相通变成整体导电。
论文用一个叫Reeb图的工具,把参数轨迹的合并与分裂过程记录成一张随时间演化的拓扑图,再通过一次"时间尺度分离"的操作,把里面快速抖动的布朗噪声筛掉,只留下缓慢但不可逆的合并趋势。这一步的关键洞察在于,训练存在两种时间尺度:快的是局部抖动,网络在小范围内来回震荡;慢的是宏观趋势,权重实际收缩变小的整体方向。论文证明,只要这两种尺度分得开,局部抖动造成的合并和分裂在统计上会正好抵消,剩下的只有单向的坍缩。
这就好比看一段延时摄影拍下的城市车流。你如果按正常速度看,车子在红绿灯前来回挪动、变道、又倒回来,看起来乱糟糟一片,根本看不出规律。但如果把镜头拉长成一整天的延时,你会发现每天早晚高峰的车流方向其实清晰又固定。论文做的正是这个"拉长镜头"的操作,把训练中那些看似随机的局部波动过滤掉,只留下真正决定网络走向的宏观趋势。如果不做这一步过滤,你会被那些瞬时的、可逆的噪声干扰所迷惑,误以为网络的坍缩是杂乱无章的,看不出背后其实有一条清晰的、不可逆的坍缩轨迹。
不是一个个来,是一批一批地塌
标准渗流理论里,网络的连通性增长通常是平滑的,一条边一条边地加。但这篇论文发现了一个关键的不同点:神经网络的架构对称性会强迫多个子网络同时合并,而不是两两单独合并。
这个"多个一起"背后的机制叫Sn对称群,简单说就是n个功能完全等价的神经元,它们之间存在完整的置换对称性,不是随便两个互换,是这n个中的任意排列组合都不影响输出。
Sn对称群*:描述n个完全等价、可以任意排列而不改变系统行为的对象之间对称关系的数学结构。
这意味着,当这n个子网络要合并进同一个不变集时,它们是同时绑定的,论文管这叫"超凝聚",用数学表达就是组件规模从1直接跳到n,中间没有过渡态。
这就像一场婚礼上突然宣布集体注册结婚,而不是情侣们排队一对一登记。如果登记处只能一对一处理,你会看到队伍缓慢但持续地缩短;但如果规定必须凑齐n对新人一起进场,登记处的队伍长度就会保持不变很长时间,然后突然一下子跳崩,一大批人同时完成手续。这个"一起进场"的规则不是可以选择性忽略的细节,它直接决定了整个系统连通性增长曲线的形状:是平滑曲线,还是带着台阶的阶梯。
论文证明了这种台阶式跳跃是否在网络规模趋于无穷大时依然存在,取决于一个具体条件:合并时那些子网络的规模,相对于整个网络有多大。如果合并的子网络本身已经占了网络相当大的比例(论文用Θ(N)表示,也就是规模随网络整体大小成比例增长),这个跳跃就会一直存在,不会随着网络变大而消失。但如果合并的只是零星几个(论文用O(1)表示,规模跟网络整体大小无关),跳跃幅度就会随着N变大而被"稀释"掉,趋于平滑。
这个区分其实回应了渗流理论里一场老争论。2009年有篇论文(Achlioptas等人)在Science上提出"爆炸性渗流"概念,说某些随机过程会导致连通性突然爆炸式增长;但2011年Riordan和Warnke证明,那类过程在网络规模真正趋于无穷时其实还是连续变化的,之前观察到的"爆炸"只是有限规模下的假象。这篇论文的态度挺诚实:它承认自己的机制在"合并规模小"的情况下,确实会退化成前人说的那种有限规模假象;但只要合并规模足够大,跳跃就是真实存在的,不会消失。这种诚实其实挺难得,很多论文倾向于把自己的发现讲成"绝对成立",这篇论文选择讲清楚边界条件。
怎么"看见"这些跳跃:用方差捕捉瞬间
单独看一次训练的轨迹,你很难精确判断某个合并瞬间到底发生在哪个时刻,因为随机噪声会让不同的训练轮次在稍微不同的时刻发生跳跃,把信号搞得模糊一片。
论文的解决办法是跑一整个训练集合体,也就是用不同的随机种子重复训练很多次,然后计算网络"最大连通组件占比"这个宏观指标(论文称为阶参数)在这个集合体里的相对方差。
阶参数*:衡量网络中已经合并成一体的部分占整体的比例,数值从0(完全没合并)到1(全部合并成一个整体)。
关键的数学直觉是:在合并瞬间发生前后,不同训练轮次会分裂成两群,一群已经跳过去了,一群还没跳。这种"两群并存"的状态叫双峰分布,而双峰分布正好会让方差在这个时间点上出现一个尖锐的峰值。论文严格证明了这个方差峰值的高度,正好和跳跃幅度的平方成正比。
这就好比统计一群人过马路的时间。如果红灯正常倒数,所有人差不多同时开始走,你测出来的"过马路耗时"分布会很集中。但如果有个瞬间规则突然变化,比如警察哨声一响,一部分人已经冲过去了,另一部分人还在等,这时候你测的这批人过马路的用时,会出现两个截然不同的群体,方差自然会猛地拉高。方差峰值本身,就是"规则突变时刻"的信号灯。
论文进一步证明,这些方差峰值出现的位置不是随机散布的,而是严格遵循一个几何级数规律,这个现象在物理学里叫离散标度不变性。
离散标度不变性*:系统只在特定的、成倍数关系的尺度上表现出自相似性,而不是在任意尺度上都自相似,这跟经典的连续标度不变性不同。
具体来说,如果把n个子网络绑定看作一次"放大"操作,放大倍数是nσ(σ是渗流临界指数),那么相邻两次微观转变发生的密度之比,会精确收敛到这个放大倍数的倒数。换句话说,只要你观测到了前两次微观转变发生的时刻,理论上就能推算出整个网络最终彻底崩塌(也就是全局连通)会发生在什么时候。这就像地震学里的前震序列,专业的地震学家能通过分析连续几次小震的间隔和强度规律,去预测下一次更大地震的时间窗口,虽然精度有限,但这种"从局部规律推全局趋势"的思路是共通的。如果没有这条几何级数规律,你观测到的每次方差峰值就只是孤立事件,完全没法拿来预测什么;正是这个精确的倍数关系,才把一堆孤立的峰值串成了一条可以往前推算的轨迹。
这个理论预测在论文的实验里得到了相当干净的验证。在一个简化的三参数玩具模型里,理论预测的放大倍数是2.00,实验测出来的正好是2.00,分毫不差。在一个六参数的、不加约束的SGD训练里,同样测出了约2.00的放大倍数。
Adam和AdamW也逃不过这个规律
理论讲到这儿,其实还有个现实问题没解决:论文前面所有推导都是针对最朴素的SGD优化器的,但现实里几乎没人纯用SGD训练大模型,大家用的是Adam、AdamW这类自适应优化器。
Adam和AdamW跟SGD最大的不同,是它们不光记录当前梯度,还维护两个额外的滑动统计量,一阶动量(梯度的指数加权平均)和二阶动量(梯度平方的指数加权平均),并用二阶动量的平方根去给每个参数的更新步长做个性化缩放。
这个"平方"操作看似不起眼,却会直接破坏之前论文假设的对称性。论文里证明了一个技术细节:一般的正交变换Q(比如任意旋转)跟"逐元素平方"这个操作不兼容,平方会把不同坐标的信息混在一起,破坏原有的对称结构。但坐标置换(也就是单纯交换某几个坐标的位置,不做别的变换)却完全兼容,置换和平方谁先谁后都一样。这意味着Adam和AdamW能保留的对称性类别,比SGD原本假设的要窄一些,只能是置换对称,不能是任意正交变换。好消息是,神经元互换所依赖的那种Sn对称性,恰好就属于这个"窄一些"但依然合法的类别,所以整个理论框架不需要推倒重来,只需要在这个更窄的对称类别里重新证明一遍。
除了对称性收窄,论文还处理了另一个现实问题:真实训练里的梯度噪声往往是重尾分布,极端的梯度尖峰出现的频率,比标准正态分布预测的要高得多,这在使用注意力机制的Transformer架构里已经有专门研究证实过。重尾分布意味着梯度噪声的方差可能是无穷大,这会让Adam维护的二阶动量统计量在数学上变得难以处理。
论文的解决办法是引入一个截断机制,把超过某个阈值的梯度直接砍掉,用截断后的梯度去计算动量统计量,代价是引入一个可以被精确量化的偏差项。这就好比统计一个班级的平均身高,如果有个数据录入错误显示某个学生身高三米,你不会真的把这个数据算进平均值,你会设一个合理的上限,把明显异常的数值截掉,然后接受"截断后的平均值跟真实平均值会有点微小偏差"这个代价,换来整体统计不被极端值搞崩。
在这些调整之上,论文证明了一个"条件性"的结论:只要网络处于某个特定阶段(论文给出了具体可检验的停止时刻条件),Adam和AdamW训练出来的网络,依然会表现出跟SGD一样的离散标度不变性规律,放大倍数的公式形式完全一致。这个证明不是空对空的数学游戏,论文用一个基于AdamW训练的、真实观察到grokking现象的Transformer模型做了初步的诊断检验,确认了理论假设的停止条件在这个实际案例里是满足的。
从三个参数到真实的Transformer
理论说得再漂亮,也得在真实场景里验证。论文的实验分了几个层次,从最简单的可控玩具模型,一路铺到了真实的Transformer训练。
在最基础的三参数玩具实验里,论文人为设计了一个层级式的合并结构,先两个参数按理论预测的时刻合并,再和第三个参数合并。结果测出来的放大倍数正好是2.00,跟理论预测严丝合缝。这个实验的意义在于排除干扰,真实神经网络训练涉及太多其他因素,用一个刻意简化、只保留核心机制的系统去验证理论,能确认这个数学关系本身是站得住的,不是巧合。
接着论文换成了一个六参数、完全不加约束的SGD训练,还在训练中途(第3000步)人为切换了任务目标。有意思的结果出现在这里:任务切换之前,网络确实表现出符合理论预测的合并式坍缩(两个方差峰值,放大倍数约2.00);但任务切换之后,网络出现了反向的碎裂,之前已经绑定在一起的参数重新分裂开,对应了另一个方差峰值。这验证了理论里一个不太直观的推论:这种坍缩不是单向不可逆的物理定律,而是取决于噪声和梯度谁更强,如果外部条件突变,让噪声重新压制住梯度的收敛力量,合并过程是可以逆转的。
这个反转在论文里被称为"反应性碎裂",它跟grokking现象形成了一个有趣的对照:grokking是网络长期停留在一个次优状态,某个瞬间突然跳到更好的状态;而反应性碎裂展示的是相反方向,系统突然从有序坍缩重新变得无序。两者用的是同一套数学框架描述,只是方向相反。
再往上是真实数据集上的实验:UCI心脏病数据集、鲍鱼年龄预测数据集、德国信用数据集这几个表格分类任务,FashionMNIST和MNIST这两个图像分类任务,以及一个基于模运算的Transformer grokking任务。
这里论文遇到了一个测量上的现实问题:真实的大网络里,权重很少会精确塌缩到零,合并更多表现为一种"软性"的???维,一部分参数方向的重要性大幅下降,但不是完全消失。为了捕捉这种软性坍缩,论文换用了谱有效秩这个指标,它通过对权重矩阵做奇异值分解,再计算奇异值分布的香农熵指数,来衡量参数矩阵实际"占用"了多少独立维度。
谱有效秩*:衡量一个矩阵实际有效维度的连续型指标,即使权重没有精确归零,只要奇异值分布变得集中,这个指标也会相应下降,捕捉"软性"的降维现象。
用这个连续指标去套用离散标度不变性理论,论文数学上证明了一个很巧妙的关系:如果n个组件在合并前后只占了整体谱方差的一部分(记为m,取值在0到1之间),那放大倍数会从整数nσ被压缩成分数nmσ。这个理论桥梁让论文可以反过来,从实测的放大倍数去反推,这次合并大概涉及了多少比例的谱方差。
数据表格如下(部分核心结果):
在UCI心脏病数据集上,网络检测到4个方差峰值,拟合出的放大倍数是1.71,拟合优度R?达到0.98,通过随机化对照检验(虚警率仅4.9%,说明这个信号不太可能是噪声假象)。
在鲍鱼数据集上,检测到6个方差峰值,放大倍数1.28,R?为0.97,虚警率15.6%。
在德国信用数据集上,虽然也测出4个峰值和不错的R?(0.86),但虚警率高达80.2%,这意味着这个信号很可能就是普通优化噪声的假象,论文很坦诚地把它当作一个负面对照案例来呈现,而不是选择性隐瞒。
在FashionMNIST图像分类上,测出3个峰值,放大倍数1.57,R?达到1.00,虚警率10.7%。
**最亮眼的结果出现在模运算的Transformer grokking实验上:3个峰值精确排成一条完美的对数线性关系,放大倍数2.11,R?为1.00,虚警率仅0.1%。**
这个虚警率0.1%的数字值得多说一句。论文用了1000次相位随机化的频谱代理数据去做统计检验,简单说,就是保留原始方差曲线的整体频率成分,但打乱其中的具体时间结构,生成1000个"伪造"的对照序列,看看这种伪造序列里,有多大比例会偶然产生出跟真实数据一样漂亮的对数线性模式。0.1%意味着1000次伪造里只有1次会凑巧产生类似结果,这说明观测到的DSI级联结构,基本可以排除是随机噪声凑巧形成的解释。
grokking实验里还有个细节特别值得琢磨:这个3峰级联恰好出现在性能突然跳升之前。论文没有过度声称这个级联就是grokking的唯一原因,它说得很谨慎,只说这个结构"与"性能跳升同时出现,而不是"导致了"性能跳升。这种谨慎其实是研究该有的态度,机制解释和相关现象之间隔着一条不小的证据鸿沟,论文没有急着跨过去。
这套理论跟前人工作是什么关系
这项研究不是凌空出现的,它踩在几条已有的研究脉络上。
最直接的基础是Chen等人2023年发表的"随机坍缩"(Stochastic Collapse)论文,那篇文章首次严格证明了SGD的梯度噪声会把网络推向更简单的子网络结构,本文管这个现象叫"是什么",而这篇论文接过来问"怎么发生的",也就是坍缩过程的时间动态和拓扑结构。
另一条脉络是关于SGD隐式偏置的更早期工作,比如Blanc等人2020年关于隐式正则化的研究,以及一系列关于深度网络"低秩简化偏置"的观察(Huh等人2021年、S,ims,ek等人2021年),这些工作共同确立了一个背景事实:SGD训练出来的网络倾向于收敛到结构更简单、秩更低的表示,这不是偶然,而是训练动力学本身的固有倾向。
再往后看,论文也提到了两条建立在渗流理论本身之上的近期方向。一条是Devlin和Sanders 2025年的工作,用渗流工具去研究dropout训练下的网络连通性;另一条是关于"神经电容"的研究(Li等人2022年),试图通过网络边的动态特性去预测模型性能。这篇论文跟它们的区别在于,边的形成机制不是随机删除或者相关性驱动的,而是完全由架构对称性决定的,这是一个更"硬"的、有数学证明支撑的机制,而不是纯经验观察。
这套框架留下的空白和它带来的启发
读完这篇论文,有几个地方让人真的会停下来想一想。
第一个触动我的地方是那个诚实的负面结果,德国信用数据集那组80.2%的虚警率。很多论文在遇到不理想的结果时,倾向于换个数据集重新调参直到数字好看,或者干脆不呈现。这篇论文选择把它明明白白地放进正文的表格里,当作一个负面对照案例摆出来。这个态度本身,比任何单个的正面结果都更能说明这篇论文的可信度。一个理论框架如果只在精心挑选的案例上生效,那它的解释力就值得怀疑;而这篇论文主动展示了自己什么时候不生效,反而让人更信它什么时候生效是真的。
第二个值得琢磨的地方,是论文对"离散"和"连续"这两个概念关系的处理。经典的物理相变理论里,连续标度不变性是一个近乎信条般的假设,系统在临界点附近的行为,应该在任意尺度上都自相似。这篇论文提出的离散标度不变性,本质上是说:当系统受到某种"硬约束"(比如n个东西必须一起动,不能一个个来)的时候,这种连续对称性会被打断成一个只在特定倍数上生效的、更弱的对称性。这个思路其实可以推广到很多"必须成群结队才能动"的系统里,不只是神经网络,任何受到离散组合约束的复杂系统,理论上都可能表现出类似的级联结构。这让我想到,也许很多我们习惯用连续曲线去拟合的经验规律(比如神经网络的scaling law),背后其实藏着一串被平均掉的离散跳跃,论文的附录里也提到了这个猜测,但没有深入验证。
第三点是关于Sn到底是多少这个问题。论文在附录里做了一个反推:对于grokking实验测出的2.11倍放大系数,如果假设是简单的两两合并(n=2),理论上限是2的1次方,也就是2,可实测的2.11已经超出这个上限。这说明这次坍缩不可能是简单的两两合并,得假设至少是三个一组的合并(n=3),这样反推出来,这次合并大概占了68%的谱方差。这个数字本身没什么特别的,但这个反推的逻辑很有意思:它把一个原本只能定性描述的现象("网络在grokking前发生了某种坍缩"),变成了一个可以定量刻画、甚至能推测出具体涉及多少个子结构、卷入多大比例参数的可检验假说。
这套理论到目前为止,最大的空白应该是"实际网络到底处于哪个渗流区域"这个问题没有答案。论文自己也承认,那个决定跳跃能不能在网络无限大时依然存在的条件,合并的子网络规模是不是跟整体网络规模成比例增长,目前没有给出一个针对真实大模型架构的判定方法。这不是一个小问题,因为它直接决定了这套理论对现实大模型(几十亿甚至上万亿参数)到底有没有解释力,还是只在小规模玩具实验里成立。
Q&A
Q1:什么是离散标度不变性(DSI)在神经网络训练中的具体表现?
A:论文发现,由于神经元置换对称性的限制,网络中的子网络合并不是连续渐进的,而是按照几何级数排列的一串离散跳跃发生的,相邻两次跳跃发生的密度之比会固定收敛到一个放大倍数,这个规律在SGD和Adam/AdamW训练中都被观察到。
Q2:grokking现象和这篇论文的渗流理论有什么关系?
A:论文在一个基于模运算训练的Transformer上观察到,delayed泛化(grokking)出现前,网络的谱有效秩经历了一个精确对数线性排列的3峰方差级联(放大倍数2.11,拟合优度1.00),说明这个突然的泛化跳跃前存在一系列结构性的拓扑坍缩,但论文谨慎地表示这不是唯一驱动因素。
Q3:这套理论是否适用于Adam和AdamW优化器?
A:论文证明了在特定条件下(包括重尾梯度噪声假设和非退化性假设),Adam和AdamW训练出的网络依然会表现出跟SGD一致的离散标度不变性级联结构,放大倍数公式形式相同,并用一个真实的AdamW训练grokking案例做了初步验证。