Grokking: Delayed Generalization¶
개요¶
| 항목 | 내용 |
|---|---|
| 분류 | Deep Learning / Generalization Theory |
| 원논문 | "Grokking: Generalization Beyond Overfitting on Small Algorithmic Datasets" (OpenAI, 2022) |
| 저자 | Alethea Power, Yuri Burda, Harri Edwards, Igor Babuschkin, Vedant Misra |
| 핵심 개념 | 과적합 이후 갑자기 발생하는 지연된 일반화 현상 |
| 관련 분야 | Mechanistic Interpretability, Neural Network Theory, Optimization |
정의¶
Grokking은 신경망이 학습 데이터에 완전히 과적합(overfitting)한 후, 추가적인 학습 반복을 통해 갑자기 테스트 데이터에 대해 완벽한 일반화(generalization)를 달성하는 현상이다. 일반적인 학습 패턴(훈련/검증 성능이 함께 향상)과 달리, grokking에서는 두 성능이 분리되어 움직인다.
용어 유래¶
"Grok"은 Robert Heinlein의 SF 소설 "Stranger in a Strange Land"(1961)에서 유래한 단어로, "깊이 이해하다"라는 의미를 가진다.
현상의 특징¶
학습 곡선 패턴¶
정확도
^
| 훈련 정확도 ━━━━━━━━━━━━━━━━━━━━━━━━━━
| / ╭━━━━ 테스트 정확도 (grokking)
| / /
| / /
| / /
| /______________________/ <-- 긴 정체기
| /
+------+---------------------------------> 에폭
과적합 시작 Grokking 발생
핵심 특성¶
| 특성 | 설명 |
|---|---|
| 지연된 일반화 | 훈련 손실이 수렴한 후에도 테스트 성능은 오랫동안 정체 |
| 급격한 전환 | 테스트 정확도가 랜덤 수준에서 완벽한 일반화로 급등 |
| 작은 데이터셋 | 주로 알고리즘 생성 데이터셋(모듈러 연산 등)에서 관찰 |
| 과파라미터화 | 데이터 대비 과도하게 큰 모델에서 발생 |
발생 조건¶
필수 조건¶
- Weight Decay (L2 정규화)
- 정규화 없이는 grokking이 발생하지 않거나 매우 늦게 발생
-
큰 가중치를 가진 암기 솔루션보다 작은 가중치의 일반화 솔루션 선호
-
충분한 학습 시간
- 일반적인 학습보다 10~100배 이상의 에폭 필요
-
데이터셋이 작을수록 더 긴 학습 시간 필요
-
적절한 데이터셋 크기
- 너무 크면 일반적인 학습 패턴
- 너무 작으면 grokking 불가능
- "임계" 크기에서 grokking 발생
실험적으로 관찰된 조건¶
| 요인 | 영향 |
|---|---|
| 초기화 스케일 | 작은 초기화 -> grokking 가속 |
| Optimizer | Adam이 SGD보다 grokking 유발 경향 |
| 학습률 | 적절한 학습률 범위에서만 발생 |
| 배치 크기 | 작은 배치 -> grokking 지연 |
이론적 설명¶
1. Phase Transition 관점¶
Grokking은 학습 과정에서의 상전이(phase transition)로 이해할 수 있다:
- Phase 1 (암기): 모델이 훈련 데이터를 단순 암기
- Phase 2 (전환): 내부 표현이 재구조화
- Phase 3 (일반화): 일반적인 알고리즘 학습 완료
2. Lazy-to-Rich Training Dynamics¶
Lazy Regime Rich Regime
(초기화 근처 유지) ------> (task-relevant 방향으로 이동)
| |
v v
암기 솔루션 일반화 솔루션
(고차원, 큰 가중치) (저차원, 작은 가중치)
Kumar et al. (2023)의 연구에 따르면: - Lazy regime: 가중치가 초기화에서 크게 벗어나지 않음, 암기 발생 - Rich regime: 가중치가 task-relevant 방향으로 이동, 일반화 달성 - Grokking은 이 두 regime 사이의 전환
3. Circuit Efficiency 관점¶
Varma et al. (2023) - "Explaining grokking through circuit efficiency":
- 모델 내부에 두 가지 "회로"가 경쟁:
- 암기 회로: 빠르게 형성, 큰 가중치 필요
- 일반화 회로: 느리게 형성, 작은 가중치로 효율적
- Weight decay가 점차 암기 회로를 억제하고 일반화 회로 선호
4. Complexity Phase Transition¶
DeMoss et al. (2025) - Physica D에 발표:
- Grokking 중 모델의 복잡도(Kolmogorov complexity)가 급격히 변화
- 암기 -> 일반화 전환 시 표현의 복잡도가 감소
주요 후속 연구¶
Omnigrok (ICLR 2023)¶
Liu, Michaud & Tegmark의 연구: - Grokking이 알고리즘 데이터셋에만 국한되지 않음을 증명 - MNIST, 이미지 분류 등에서도 조건에 따라 grokking 발생 - 초기화 스케일이 핵심 요인임을 밝힘
Deep Grokking (2024)¶
Fan, Pascanu & Jaggi의 연구: - 깊은 신경망에서도 grokking 관찰 - 깊이가 증가할수록 grokking 패턴 변화 - 더 깊은 네트워크가 더 나은 일반화 가능
Grokfast (2024)¶
Lee et al.의 연구: - Grokking 시간을 획기적으로 단축하는 기법 - Slow gradient 성분을 증폭하여 일반화 가속 - 훈련 시간 최대 50배 단축
Unifying Grokking and Double Descent¶
Davies, Langosco & Krueger (2023): - Grokking과 Double Descent를 pattern learning speeds 프레임워크로 통합 - Epoch-wise (시간에 따른)와 Model-wise (크기에 따른) grokking 구분
Mechanistic Interpretability와의 연결¶
Grokking은 Mechanistic Interpretability 연구의 중요한 테스트베드:
| 연구 주제 | Grokking 활용 |
|---|---|
| 회로 발견 | Grokking 전후 회로 구조 비교 |
| 표현 학습 | 암기 vs 일반화 표현 분석 |
| 학습 역학 | 내부 활성화 변화 추적 |
| Progress measures | Grokking 예측 지표 개발 |
실용적 의미¶
시사점¶
- 훈련 조기 종료 주의: 과적합처럼 보여도 더 학습하면 일반화 가능
- 정규화의 중요성: Weight decay가 일반화 솔루션 발견에 핵심
- 작은 데이터셋 학습: 충분한 시간을 주면 일반화 가능
- 모델 해석: 내부 회로 분석으로 학습 과정 이해
한계 및 주의사항¶
- 대규모 실제 데이터셋에서는 명확한 grokking 관찰 어려움
- 계산 비용이 높음 (긴 학습 시간 필요)
- 모든 문제에서 grokking이 발생하지는 않음
Python 구현 예제¶
모듈러 덧셈 Grokking 실험¶
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
from tqdm import tqdm
# 모듈러 덧셈 데이터셋 생성
def create_modular_addition_dataset(p: int, train_ratio: float = 0.5):
"""
(a + b) mod p 데이터셋 생성
Args:
p: 모듈러 값
train_ratio: 훈련 데이터 비율
"""
# 모든 (a, b) 쌍 생성
data = []
for a in range(p):
for b in range(p):
x = (a, b)
y = (a + b) % p
data.append((x, y))
# 셔플 및 분할
np.random.shuffle(data)
split = int(len(data) * train_ratio)
train_data = data[:split]
test_data = data[split:]
return train_data, test_data
# 간단한 MLP 모델
class GrokkinngMLP(nn.Module):
def __init__(self, p: int, embed_dim: int = 128, hidden_dim: int = 256):
super().__init__()
self.embed_a = nn.Embedding(p, embed_dim)
self.embed_b = nn.Embedding(p, embed_dim)
self.mlp = nn.Sequential(
nn.Linear(embed_dim * 2, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, p)
)
def forward(self, a, b):
ea = self.embed_a(a)
eb = self.embed_b(b)
x = torch.cat([ea, eb], dim=-1)
return self.mlp(x)
def train_grokking_experiment(
p: int = 97,
train_ratio: float = 0.3,
epochs: int = 50000,
lr: float = 1e-3,
weight_decay: float = 1.0,
log_interval: int = 100
):
"""
Grokking 실험 실행
Args:
p: 모듈러 값 (소수 권장)
train_ratio: 훈련 데이터 비율 (작을수록 grokking 명확)
epochs: 총 에폭 수
lr: 학습률
weight_decay: L2 정규화 강도 (grokking 필수!)
log_interval: 로깅 간격
"""
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 데이터 준비
train_data, test_data = create_modular_addition_dataset(p, train_ratio)
train_a = torch.tensor([x[0][0] for x in train_data], device=device)
train_b = torch.tensor([x[0][1] for x in train_data], device=device)
train_y = torch.tensor([x[1] for x in train_data], device=device)
test_a = torch.tensor([x[0][0] for x in test_data], device=device)
test_b = torch.tensor([x[0][1] for x in test_data], device=device)
test_y = torch.tensor([x[1] for x in test_data], device=device)
# 모델 및 옵티마이저
model = GrokkinngMLP(p).to(device)
optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
criterion = nn.CrossEntropyLoss()
history = {'train_loss': [], 'train_acc': [], 'test_acc': []}
for epoch in tqdm(range(epochs), desc="Training"):
# 훈련
model.train()
optimizer.zero_grad()
logits = model(train_a, train_b)
loss = criterion(logits, train_y)
loss.backward()
optimizer.step()
# 평가
if epoch % log_interval == 0:
model.eval()
with torch.no_grad():
# 훈련 정확도
train_pred = logits.argmax(dim=-1)
train_acc = (train_pred == train_y).float().mean().item()
# 테스트 정확도
test_logits = model(test_a, test_b)
test_pred = test_logits.argmax(dim=-1)
test_acc = (test_pred == test_y).float().mean().item()
history['train_loss'].append(loss.item())
history['train_acc'].append(train_acc)
history['test_acc'].append(test_acc)
if epoch % (log_interval * 10) == 0:
print(f"Epoch {epoch}: Loss={loss.item():.4f}, "
f"Train Acc={train_acc:.4f}, Test Acc={test_acc:.4f}")
return model, history
# 실행 예시
if __name__ == "__main__":
model, history = train_grokking_experiment(
p=97,
train_ratio=0.3,
epochs=50000,
weight_decay=1.0
)
# 결과 시각화
import matplotlib.pyplot as plt
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
ax1.plot(history['train_loss'], label='Train Loss')
ax1.set_xlabel('Steps (x100)')
ax1.set_ylabel('Loss')
ax1.set_title('Training Loss')
ax1.legend()
ax2.plot(history['train_acc'], label='Train Acc', color='blue')
ax2.plot(history['test_acc'], label='Test Acc', color='red')
ax2.set_xlabel('Steps (x100)')
ax2.set_ylabel('Accuracy')
ax2.set_title('Grokking: Delayed Generalization')
ax2.legend()
ax2.axhline(y=1.0, color='gray', linestyle='--', alpha=0.5)
plt.tight_layout()
plt.savefig('grokking_result.png', dpi=150)
plt.show()
Grokfast 구현 (가속화)¶
class GrokfastOptimizer:
"""
Grokfast: 느린 gradient 성분을 증폭하여 grokking 가속
Reference: Lee et al., "Grokfast: Accelerated Grokking by
Amplifying Slow Gradients" (2024)
"""
def __init__(self, optimizer, alpha: float = 0.99, lamb: float = 5.0):
"""
Args:
optimizer: 기본 옵티마이저 (Adam, AdamW 등)
alpha: EMA 계수 (느린 성분 추출)
lamb: 증폭 계수
"""
self.optimizer = optimizer
self.alpha = alpha
self.lamb = lamb
self.ema_grads = {}
def step(self):
for group in self.optimizer.param_groups:
for p in group['params']:
if p.grad is None:
continue
grad = p.grad.data
# EMA gradient 업데이트 (느린 성분)
if id(p) not in self.ema_grads:
self.ema_grads[id(p)] = torch.zeros_like(grad)
ema = self.ema_grads[id(p)]
ema.mul_(self.alpha).add_(grad, alpha=1 - self.alpha)
# 느린 성분 증폭
p.grad.data = grad + self.lamb * ema
self.optimizer.step()
def zero_grad(self):
self.optimizer.zero_grad()
# 사용 예시
optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1.0)
grokfast_opt = GrokfastOptimizer(optimizer, alpha=0.99, lamb=5.0)
for epoch in range(epochs):
grokfast_opt.zero_grad()
loss = compute_loss(model, data)
loss.backward()
grokfast_opt.step()
핵심 논문 목록¶
| 연도 | 제목 | 저자 | 학회/저널 |
|---|---|---|---|
| 2022 | Grokking: Generalization Beyond Overfitting | Power et al. | arXiv (OpenAI) |
| 2022 | Towards Understanding Grokking | Liu et al. | NeurIPS 2022 |
| 2023 | Omnigrok: Grokking Beyond Algorithmic Data | Liu et al. | ICLR 2023 |
| 2023 | Explaining grokking through circuit efficiency | Varma et al. | arXiv |
| 2023 | Grokking as the Transition from Lazy to Rich | Kumar et al. | arXiv |
| 2023 | Unifying Grokking and Double Descent | Davies et al. | arXiv |
| 2024 | Grokfast: Accelerated Grokking | Lee et al. | ICML 2024 |
| 2024 | Deep Grokking | Fan et al. | arXiv |
| 2025 | The complexity dynamics of grokking | DeMoss et al. | Physica D |
관련 주제¶
- Neural Scaling Laws - 스케일과 일반화
- Lottery Ticket Hypothesis - 네트워크 pruning과 일반화
- Mechanistic Interpretability - 내부 회로 분석
- Double Descent - 또 다른 비직관적 일반화 현상 (문서 준비 중)