训练越久,模型反而越会泛化吗?理解 Grokking 现象
Swift Lv6

Grokking: Generalization Beyond Overfitting on Small Algorithmic Datasets这篇论文中,神经网络在训练集上早已达到接近 100% 的准确率,但继续训练很长时间后,验证集准确率才突然从接近随机水平上升到接近完美,作者把这种现象称为 grokking

grok

一、这篇文章到底研究了什么?

我们通常希望神经网络学习数据背后的规律,而不是简单地记住训练样本。

例如,假设我们给模型一些这样的题目:

输入 输出
$a \circ b$ $c$
$a \circ d$ $e$
$b \circ c$ $f$

模型需要根据已经看到的部分答案,推断没有见过的组合。

它可以采用两种方式:

  1. 记忆训练样本:见过的输入直接查表,没见过的输入不会做;
  2. 学习底层规则:发现 $a \circ b$ 背后存在一种统一的运算规律,然后推断新的输入。

这篇论文研究的正是:模型什么时候会从第一种方式转向第二种方式。

作者构造了许多小型的算法数据集,包括模 $97$ 的加法、减法、除法,以及抽象群 $S_5$ 中的排列运算。每条数据都可以写成:

其中 $a$、$b$ 和 $c$ 都被表示成没有内部结构的离散符号。模型不能直接看到“这是数字 3”或“这是某个排列”,只能通过符号之间的关系自己发现规律。Datasets

二、什么是 Grokking?

Grokking 的典型过程可以分成两个阶段。

第一阶段:模型先学会“背答案”

训练刚开始时,模型很快就能把训练集做对。因为模型参数很多,它可以给每个训练样本单独存储一个答案。

于是会出现这样的情况:

  • 训练准确率:接近 100%
  • 验证准确率:接近随机水平
  • 模型:已经记住训练数据,但还不会处理新数据

在论文的模 $97$ 除法实验中,训练准确率不到 $10^3$ 步就接近完美,但验证准确率可能要到接近 $10^6$ 步才达到同样水平。也就是说,模型在很长时间内都处于“训练集全会,验证集不会”的状态。Grokking Example

第二阶段:模型逐渐找到统一规则

如果继续训练,模型的内部参数还会持续变化。它可能逐渐从大量零散的记忆,转向一种更加统一、简单、可泛化的表示。

此时,验证准确率可能长期没有明显变化,然后突然快速上升:

从外部观察,这就像模型突然“顿悟”了。

但这里的“顿悟”不是模型产生了人类式的意识,也不是某一步突然写入了一个“理解模块”。更准确地说,它是模型参数经过长时间优化后,逐渐从一个复杂的记忆解转移到了一个更简单的规则解。

三、模型为什么会从记忆转向规则?

关键在于:能够拟合训练集的模型解并不只有一种。

假设训练集中有 $100$ 个样本,那么模型至少可以有两种解:

  • 一个复杂的解:专门记住这 $100$ 个答案;
  • 一个简单的解:学习背后的规则,并因此正确回答更多未见样本。

这两种解在训练集上的表现可能完全一样,都是 100% 准确率。但它们的泛化能力不同。

可以把它类比成学生做数学题。

学生拿到一组题目后,可能先把题目和答案背下来:

看到题目 A,回答 7;看到题目 B,回答 3。

这样,他可以把练习册上的题全部答对,但遇到新题就不会。

如果他继续思考,可能最终发现:

原来这些题都遵循同一个公式。

这时,他不再需要逐题查答案,而是可以通过公式解决新的题目。Grokking 在模型身上呈现出来的现象,与这个过程有些相似。

四、为什么训练更久,有时反而有利于泛化?

这和“模型偏好什么样的解”有关。

在训练集上,复杂的记忆解和简单的规则解都可以得到零训练误差。但在某些训练条件下,优化过程可能逐渐偏向参数更简单、结构更规则的解。

尤其是权重衰减可能发挥重要作用。权重衰减会惩罚过大的参数,使模型不太容易维持一个极其复杂的记忆方案。论文发现,在这些算法任务中,权重衰减对数据效率的提升非常明显,所需样本量相比多数其他干预甚至减少了一半以上。Weight Decay

此外,优化过程中的噪声也可能帮助模型找到泛化更好的解。小批量随机梯度下降,或者对梯度和参数加入噪声,有时会促使模型进入更平坦的损失区域,而较平坦的解通常对参数的小变化不那么敏感。Optimization Noise

因此,模型可能经历下面的过程:

flowchart LR
    A[开始训练] --> B[快速拟合训练集]
    B --> C[训练集接近100%]
    C --> D[继续优化与参数调整]
    D --> E[形成更简单的内部表示]
    E --> F[验证集性能突然提升]

五、这是否意味着训练越久越好?

不是。

这正是我在理解这篇论文时产生的疑问:按照通常的机器学习经验,训练太久容易过拟合。既然如此,为什么这篇论文里训练更久反而提升了泛化?

答案是:传统过拟合和 Grokking 是两种可能出现的训练动态。

普通过拟合

普通过拟合通常表现为:

模型逐渐学会训练数据中的噪声和偶然细节,因此在训练集上更好,在验证集上更差。

Grokking

Grokking 则可能表现为:

论文观察到,验证损失在某些实验中会出现第二次下降。作者认为,这种现象可能和经典的 double descent 有关联,但也强调 Grokking 发生在训练集已经被拟合很久之后,因此可能是不同的现象。Double Descent

所以,“训练越久越好”并不是一般规律。更准确的说法是:

在某些任务上,模型可能需要很长时间才能从记忆解转向规则解;但在另一些任务上,继续训练只会让模型越来越严重地过拟合。

六、什么类型的任务更容易出现 Grokking?

我的进一步理解是:任务类型确实非常重要。

如果一个任务存在明确、统一、可压缩的规则,那么模型才有可能从训练样本中发现这个规则。比如:

  • 模运算
  • 排列和群运算
  • 简单程序
  • 符号推理
  • 数学公式
  • 某些人工构造的算法任务

这篇论文中的任务都具有明确的生成规则,而且训练数据量很小、标签没有明显噪声,模型又足够大,可以同时表达“记忆解”和“规则解”。

在这种情况下,长时间训练有可能让模型逐渐找到更具泛化能力的规律。

可以粗略地表示为:

不过,有规则并不代表一定会 Grokking

论文中也有一些看似有规律的任务,例如:

在允许的训练预算内,模型并没有成功泛化,基本上只是记住了训练数据。Failed Generalization

这说明还需要满足几个条件:

  1. 规则必须足够简单,或至少对当前模型来说足够容易发现
  2. 模型必须能够表达这种规则
  3. 优化过程必须有机会找到这个规则
  4. 训练数据不能包含太多随机噪声
  5. 训练步数和超参数要合适

七、噪声会不会阻止 Grokking?

通常会。

作者还专门研究了训练集中加入随机标签的情况。结果显示,少量异常样本对泛化的影响可能不明显,但大量异常样本会显著降低模型最终成功泛化的范围。Outliers

这也很好理解。

如果训练集中的每个样本都遵循同一个规则,模型有机会发现统一规律。但如果训练集中混入大量错误答案,模型可能不得不学习:

此时,纯粹学习规则已经无法完全拟合训练集,模型更容易被迫保留记忆和例外。

因此,Grokking 更容易出现在标签干净、规律明确的数据集上,而不是噪声很多的现实数据中。

八、模型真的“理解”了吗?

这里需要谨慎使用“理解”这个词。

从行为上看,模型确实从“只能回答见过的问题”变成了“可以回答没见过的问题”。这说明它学到了一些能够支持泛化的内部结构。

作者还对模型的 embedding 进行了可视化,发现模型有时会形成与底层数学对象对应的结构。例如,模加法中的 embedding 可能呈现类似圆环的拓扑,排列运算中的 embedding 则可能形成与群结构有关的簇。Embedding Structure

但这并不等于模型像人一样理解了数学概念。我们更稳妥的说法是:

模型形成了一个能够复现底层规则、并支持新样本预测的内部表示。

至于模型为什么一定会形成这个表示,作者在这篇文章中并没有给出完整答案。

论文提供了一个重要线索:在 $S_5$ 实验中,验证准确率和 sharpness 的 Spearman 相关系数为 $-0.79548$。这意味着验证集表现更好的模型,往往位于更平坦的损失区域。但这仍然是相关性,并不能证明“平坦性”就是 Grokking 的唯一原因。Sharpness

九、实际训练时要不要刻意训练久一点?

结论:

不要无条件地延长训练;但对于规则型、小数据、低噪声任务,值得专门测试更长的训练周期。

对于一般的图像分类、文本分类或业务预测任务,建议仍然观察验证集:

  • 训练集性能上升、验证集性能也上升:可以继续训练;
  • 训练集性能上升、验证集性能下降:可能正在普通过拟合;
  • 训练集已经很好、验证集长期停滞:可以尝试更长训练,但要同时调整正则化和学习率;
  • 验证集性能在长时间停滞后重新上升:这可能是 Grokking 的迹象。

实际操作中,可以保存不同训练阶段的 checkpoint:

1
2
3
4
训练 10,000 步  → 保存模型并评估验证集
训练 100,000 步 → 保存模型并评估验证集
训练 1,000,000 步 → 保存模型并评估验证集
选择验证集表现最好的 checkpoint

不要简单地认为“最后一个模型一定最好”,也不要因为训练准确率已经达到 100% 就立刻停止。

十、最佳实践

该论文复现较为容易,对数据集、模型、算力要求不高。详细复现的pipeline见:grokking_reproduction

十一、总结

我对这篇论文的理解可以总结为:

模型训练过程不一定是从“不会”逐渐变成“会”,而可能先从“记住训练样本”开始,再经过很长时间的优化,转向“发现数据背后的规则”。

因此,训练集准确率达到 100% 并不总意味着模型已经学会了任务。它可能只是找到了一个能够记忆训练数据的复杂解。

如果任务背后存在明确、简单、可表达的规律,那么继续训练可能让模型逐渐偏向更简单、更统一的规则解,从而出现 Grokking。论文中的权重衰减和优化噪声实验,也说明正则化和优化路径可能帮助模型摆脱单纯的记忆。Main Findings

但这并不推翻“训练过久会过拟合”的常识。更准确的经验应该是:

普通任务要依靠验证集早停;规则明确的小数据任务,则要警惕过早早停,并尝试检查验证性能是否会在长时间停滞后重新上升。

这也是 Grokking 最有价值的地方:它提醒我们,过拟合并不总是学习过程的终点,有时只是模型从记忆走向泛化之前的中间阶段。

Powered by Hexo & Theme Keep
This site is deployed on
Unique Visitor Page View