Gradient checkpointing: N개의 활성화 함수만 유지하고 나머지는 재계산 — 추가적인 순전파(forward pass) 1회로
요약
심층 신경망 학습 시 VRAM 사용량을 줄이기 위한 Gradient Checkpointing 기법을 설명합니다. 모든 활성화 함수를 저장하는 대신 일부 체크포인트만 유지하고 필요할 때 재계산하여 메모리 효율을 극대화하는 트레이드오프를 다룹니다.
핵심 포인트
- 활성화 함수 저장 대신 재계산을 통해 VRAM 사용량 절감
- 수학적 결과는 동일하며 오직 메모리 스케줄만 변경됨
- 세그먼트 경계의 체크포인트만 저장하여 피크 메모리 관리
- 연산량(Compute)과 메모리(Memory) 사이의 최적 트레이드오프 활용
심층 네트워크(deep network)를 학습시킬 때 VRAM을 폭발적으로 사용하는 주범은 대개 가중치(weights)가 아니라 바로 _활성화 함수(activations)_입니다. 역전파(Backprop)는 해당 레이어의 그래디언트(gradient)를 계산하기 위해 모든 레이어의 순전파(forward) 출력을 필요로 합니다. 따라서 단순한(naive) 학습 단계에서는 이 모든 것을 저장하며, 활성화 메모리는 깊이(depth)에 따라 O(N)으로 증가합니다(배치 크기(batch size) 및 시퀀스 길이(sequence length)에 따라서도 마찬가지입니다). Gradient checkpointing은 약간의 연산량을 대가로 많은 양의 메모리를 확보합니다. 즉, 소수의 활성화 함수만 유지하고 나머지는 역전파(backward pass) 중에 필요할 때마다 _재계산(recompute)_하는 방식입니다. 이는 수학적인 결과에는 아무런 영향을 미치지 않습니다. 그래디언트(gradients)와 학습된 가중치(trained weights)는 비트 단위로 완전히 동일하게 산출되며, 오직 메모리 스케줄(memory schedule)만 바뀔 뿐입니다. 저는 이 트레이드오프(trade-off)를 실시간으로 계산하는 데모를 만들었습니다. 그 아이디어는 다음과 같습니다.
역전파(Backprop)는 순전파 활성화 함수(forward activations)를 필요로 합니다
레이어의 가중치 그래디언트(weight gradient)에 대한 연쇄 법칙(chain rule)에는 해당 레이어의 입력값이 포함됩니다. y = W·x인 경우, dL/dW = dL/dy · xᵀ가 성립하며, 즉 역전파(backward) 시점에 문자 그대로 x가 필요합니다. 따라서 단순한(naive) 패스는 모든 활성화 함수의 테이프(tape)를 보관합니다:
acts = []
h = x
for layer in layers:
...
트레이드오프: 버리고, 필요할 때 재계산하기
네트워크를 세그먼트(segments)로 나누고 세그먼트의 경계인 _체크포인트(checkpoints)_만 저장합니다. 내부 순전파(interior forward)는 no_grad 상태로 실행하여 아무것도 유지되지 않도록 하고, 이후 역전파(backward pass) 중에 해당 값들이 실제로 필요할 때 그 세그먼트의 순전파(forward)를 다시 실행합니다.
with torch.no_grad(): # 내부 활성화 함수를 저장하지 않는 순전파(forward)
h = segment(x)
# ...나중에, 역전파(backward)에서 값이 필요할 때:
...
연산(Compute)은 저렴하고 병렬적이지만, 메모리(memory)는 희소합니다. 따라서 전체 스텝 동안 텐서(tensor)를 확보하기 위해 몇 마이크로초(microseconds)의 재계산을 사용하는 것은 이득입니다.
세그먼트 메커니즘(The segment mechanism)
역전파 루프(backward loop)는 세그먼트를 역순으로 훑으며, 저장된 체크포인트로부터 한 번에 하나씩 재계산합니다. 오직 현재 세그먼트의 활성화 함수(activations)만 실체화(materialised)되기 때문에, 한 번에 하나의 세그먼트 분량만 살아있게 되며, 이것이 바로 전체적인 절약의 핵심입니다.
for k in reversed(range(len(segs))):
h = checkpoints[k].detach().requires_grad_()
out = segs[k](h) # 이 세그먼트를 재계산 (이제 그래프가 구축됨)
...
왜 √N인가 — s + N/s 최소화
어느 순간이든 당신은 저장된 s개의 체크포인트와 재계산 중인 세그먼트의 N/s개 내부 활성화 함수(activations)를 보유하게 됩니다. 따라서 피크 메모리(peak memory)는 s + N/s이며, 이를 최소화하는 것은 하나의 미분 문제입니다:
# d/ds (s + N/s) = 1 − N/s² = 0 -> s = √N -> peak ≈ 2√N
peak(64, 8) # 16 sqrt(64) 세그먼트 -> O(√N), 4배 절감
peak(64, 64) # 65 모두 저장 -> O(N)
...
두 극단적인 경우 모두 약 N의 비용이 들지만, 중간 지점에서만 급격히 감소합니다. N=64일 때 64에서 16 유닛으로 감소하며, N=10,000일 때는 10,000에서 약 200으로 감소합니다. 이는 50배의 절감이며, 네트워크가 깊어질수록 이득은 더 커집니다. 비용은 한 번의 추가적인 순전파(forward pass)입니다. 재계산은 순전파 작업을 대략 두 배로 늘리며, 역전파(backward)가 이미 순전파의 약 2배 비용이 들기 때문에, 전체 단계는 약 3에서 약 4의 순전파 상당치로 늘어납니다. 이는 실제 실행 시간(wall-clock) 기준으로 약 +30%의 증가를 의미합니다 (Chen et al., 2016).
실제로 적용하는 방법
역방향 루프를 직접 구현할 필요는 없습니다. PyTorch는 checkpoint를 통해 블록을 감싸고, checkpoint_sequential을 통해 Sequential을 약 √N개의 세그먼트로 분할하며, HuggingFace는 model.gradient_checkpointing_enable()을 제공합니다. 단 하나의 정확성 규칙은 다음과 같습니다: 재계산은 순전파를 정확하게 재현해야 하므로, 드롭아웃(dropout)이나 배치 정규화(batch-norm) 통계와 같은 무작위 요소는 시드(seed)를 고정하거나 동결해야 하며, 그렇지 않으면 use_reentrant=False를 전달해야 합니다.
이는 실제로 그래디언트(gradient)의 크기를 재조정(rescale)하여 업데이트를 변경하는 그래디언트 클리핑(gradient clipping)과 대조해 볼 가치가 있습니다. 체크포인팅은 결과에서 측정 가능한 그 어떤 것도 변경하지 않으며, 오직 메모리가 사용되는 위치만 바꿉니다. 이와 유사한 기술들은 더 나아가기도 합니다: 활성화 오프로딩(activation offloading)은 활성화 함수를 CPU RAM으로 이동시키고, 가역 레이어(reversible layers)는 거의 제로에 가까운 저장 공간을 위해 입력을 재구성하며, FlashAttention은 어텐션(attention) 자체에 동일한 재계산 트릭을 적용합니다.
깊이와 세그먼트 수를 설정하고 피크 메모리가 2√N에서 최저점을 찍는 것을 확인해 보세요:
https://dev48v.infy.uk/dl/day48-gradient-checkpointing.html
AI 자동 생성 콘텐츠
본 콘텐츠는 Dev.to AI tag의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기