Gradient Checkpointing
PyTorch 예제, 트레이드오프 및 실용적인 학습 가이드를 통해 역전파 중에 활성화 값을 재계산하여 그래디언트 체크포인팅이 GPU 메모리를 줄이는 방법을 알아보세요.
그래디언트 체크포인팅(gradient checkpointing)은 순전파(forward pass) 과정에서 선택된 중간 활성화 값(intermediate activation)만 저장하고 나머지는 역전파 중에 다시 계산하여 메모리를 절약하는 학습 기법입니다. **활성화 체크포인트(activation checkpointing)**라고도 하며, 추가적인 연산 비용과 교환하여 최대 피크 메모리 사용량을 낮춥니다. 이름과 달리 이 기법은 파라미터 그래디언트나 모델 파일이 아니라 활성화 값을 체크포인트하므로, 활성화 텐서 때문에 신경망이 사용 가능한 GPU 메모리에 들어가지 않을 때 특히 유용합니다.
Gradient Checkpointing 작동 방식#
표준 순전파 과정에서 신경망은 **활성값(activations)**이라는 중간 텐서를 계산합니다. PyTorch autograd mechanics에서 설명하듯이, 자동 미분 시스템은 나중에 그래디언트를 계산하는 데 필요한 활성값을 유지합니다. 심층 네트워크, 대규모 배치, 고해상도 이미지, 긴 입력 시퀀스로 인해 이러한 저장된 텐서는 상당한 메모리를 소모할 수 있습니다.
Gradient checkpointing은 네트워크를 여러 세그먼트로 나눕니다:
- 순전파는 선택된 세그먼트의 입력 또는 경계 활성값을 저장합니다.
- 해당 세그먼트 내의 다른 중간 활성값은 폐기됩니다.
- 역전파 과정에서 각 체크포인트가 지정된 세그먼트가 다시 순전파를 실행하여 누락된 값을 재구성합니다.
- 재구성된 활성값은 그래디언트를 계산하는 데 즉시 사용됩니다.
PyTorch activation checkpointing API는 torch.utils.checkpoint을 통해 이 동작을 제공합니다. 이와 유사한 개념으로 JAX gradient checkpointing and rematerialization 및 TensorFlow tape checkpointing이 있습니다.
체크포인트 배치는 트레이드오프를 결정합니다. 일반적으로 더 많은 영역을 체크포인트화하면 메모리는 더 많이 절약되지만 연산이 반복됩니다. 활성화가 집중되는 블록 주변에 선택적으로 배치하면 전체 네트워크를 재계산하는 것보다 더 나은 균형을 제공할 수 있습니다.
트레이드오프 및 관련 기법#
Gradient checkpointing은 대개 모델 파라미터, 그래디언트, 옵티마이저 상태를 변경하지 않습니다. 주요 대상은 활성화 메모리이며, 일부 순전파 연산이 두 번 실행되기 때문에 학습이 더 느려집니다. 정확한 결과는 아키텍처, 체크포인트 경계, 배치 형태, 하드웨어에 따라 달라집니다.
이는 몇 가지 관련 기법과 차이가 있습니다:
- **Gradient accumulation**은 모든 샘플을 동시에 로드하지 않고도 모델 가중치를 업데이트하기 전에 여러 마이크로 배치를 처리하여 더 큰 유효 배치를 생성합니다. 반면 gradient checkpointing은 각 마이크로 배치에 대해 유지되는 활성값을 줄입니다.
- **Mixed precision**은 선택된 연산에 대해 더 낮은 정밀도의 데이터 타입을 사용합니다. PyTorch automatic mixed precision workflow는 메모리를 줄이고 호환되는 연산을 가속화할 수 있는 반면, 체크포인팅은 고의로 연산량을 추가합니다.
- Batch size 감소는 더 적은 샘플을 함께 처리하여 메모리를 낮춥니다. 체크포인팅을 사용하면 더 큰 배치나 입력 해상도를 계속 유지할 수 있습니다.
- **학습 체크포인트(Training checkpoints)**는 복구 또는 추론을 위해 가중치와 옵티마이저 상태를 저장합니다. PyTorch model checkpoint guide에서는 활성값 재계산과 무관한 이러한 지속성 메커니즘을 설명합니다.
체크포인트가 적용된 코드는 원본 순전파 실행과 재계산된 순전파 실행 사이에서 기능적으로 일관되어야 합니다. 체크포인트 영역 내의 가변 상태, 디바이스 간 전송, 또는 제어되지 않는 무작위성은 오류나 잘못된 그래디언트를 유발할 수 있습니다. 디코더 스타일 Transformer에서는 key-value cache가 캐시된 추론 상태와 순전파 그래프 재구성이 충돌할 수 있으므로 체크포인트 학습 중에 비활성화해야 할 수도 있습니다.
PyTorch 예제#
다음 예제는 하나의 학습 단계 동안 메모리 집약적인 블록에 체크포인트를 적용합니다:
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
torch.manual_seed(0)
block = nn.Sequential(
nn.Linear(1024, 4096),
nn.ReLU(),
nn.Linear(4096, 1024),
)
optimizer = torch.optim.AdamW(block.parameters())
inputs = torch.randn(8, 1024, requires_grad=True)
targets = torch.zeros_like(inputs)
optimizer.zero_grad(set_to_none=True)
outputs = checkpoint(block, inputs, use_reentrant=False)
loss = nn.functional.mse_loss(outputs, targets)
loss.backward()
optimizer.step()블록의 입력과 필요한 경계 정보만 유지되며, 내부 활성값은 loss.backward() 동안 재구성됩니다. 실제 프로젝트에서는 PyTorch peak allocated GPU memory와 같은 측정 지표를 사용하여 체크포인팅 적용 여부에 따른 실행 결과를 비교해야 합니다.
실제 애플리케이션 사례#
-
고해상도 컴퓨터 비전: 의료 영상 분할, 항공 탐지, 산업 검사에서는 모델 가중치보다 피처 맵이 더 많은 메모리를 소모하는 대형 이미지로 학습을 진행할 수 있습니다. 백본 단계를 선택적으로 체크포인트화하면 입력을 과도하게 다운스케일링하는 대신 이미지 해상도를 보존하거나 배치당 추가 샘플을 허용할 수 있습니다.
-
긴 시퀀스 Transformer 학습: 시퀀스 길이와 레이어 수가 증가함에 따라 활성화 저장 공간이 급격히 증가합니다. Transformer 블록을 재계산하면 동일한 가속기에서 더 긴 컨텍스트나 더 큰 마이크로 배치를 수용할 수 있습니다. NVIDIA activation recomputation guide는 Transformer 레이어에 대한 전체 및 선택적 재계산을 보여줍니다.
실용 가이드#
프로파일링 결과 파라미터나 옵티마이저 상태가 아닌 활성값이 메모리를 지배할 때 gradient checkpointing을 사용하세요. 적용 범위를 넓히기 전에 큰 반복 블록부터 시작하여 피크 메모리, 반복 시간, 검증 동작을 벤치마킹하세요.
문서화된 Ultralytics YOLO training configurations의 경우, 기본적인 메모리 제어 항목에는 자동 혼합 정밀도, 이미지 크기, 물리적 배치 크기가 포함됩니다. Ultralytics AutoBatch reference는 사용 가능한 GPU memory에 기반한 자동 배치 선택 기능을 설명합니다. 로컬 하드웨어가 여전히 부족한 경우, Ultralytics Platform cloud training은 관리형 학습 실행을 위해 설정 가능한 클라우드 GPU를 제공합니다. 사용자 정의 PyTorch 아키텍처에서 더 세밀한 활성화 메모리 관리가 필요할 때 gradient checkpointing이 이러한 제어 기능을 보완할 수 있습니다.






