콘텐츠로 이동
Data Prep
상세

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 발생

핵심 특성

특성 설명
지연된 일반화 훈련 손실이 수렴한 후에도 테스트 성능은 오랫동안 정체
급격한 전환 테스트 정확도가 랜덤 수준에서 완벽한 일반화로 급등
작은 데이터셋 주로 알고리즘 생성 데이터셋(모듈러 연산 등)에서 관찰
과파라미터화 데이터 대비 과도하게 큰 모델에서 발생

발생 조건

필수 조건

  1. Weight Decay (L2 정규화)
  2. 정규화 없이는 grokking이 발생하지 않거나 매우 늦게 발생
  3. 큰 가중치를 가진 암기 솔루션보다 작은 가중치의 일반화 솔루션 선호

  4. 충분한 학습 시간

  5. 일반적인 학습보다 10~100배 이상의 에폭 필요
  6. 데이터셋이 작을수록 더 긴 학습 시간 필요

  7. 적절한 데이터셋 크기

  8. 너무 크면 일반적인 학습 패턴
  9. 너무 작으면 grokking 불가능
  10. "임계" 크기에서 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 예측 지표 개발

실용적 의미

시사점

  1. 훈련 조기 종료 주의: 과적합처럼 보여도 더 학습하면 일반화 가능
  2. 정규화의 중요성: Weight decay가 일반화 솔루션 발견에 핵심
  3. 작은 데이터셋 학습: 충분한 시간을 주면 일반화 가능
  4. 모델 해석: 내부 회로 분석으로 학습 과정 이해

한계 및 주의사항

  • 대규모 실제 데이터셋에서는 명확한 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

관련 주제