긴 context, 큰 vocab로 LLM을 학습할 때 가장 먼저 OOM(Out-Of-Memory)을 유발하는 지점이 의외의 곳에 있습니다. 모델 weight도 아니고, attention도 아닙니다. LM head + cross entropy loss 단계가 메모리 폭탄의 주범입니다. 16B 모델을 128K context로 학습하면서 실제로 만난 OOM에서 출발해, cross entropy의 forward/backward가 왜 그렇게 무거운지, 그리고 어떻게 메모리 사용량을 줄이는지를 정리합니다.

128K Context에서 LM Head가 만드는 메모리 블랙홀
128K 토큰, vocab size 128K, hidden dim 4096 모델 기준으로 한 번 계산해 봅시다.
- logits 텐서 크기:
128K tokens × 128K vocab × 4 bytes(float32)= 약 64GB - gradient 텐서: forward와 동일한 크기로 backward 시 추가 = +64GB
- 중간 activation 저장용 (log-softmax 결과 등): 수 GB
이거 하나가 H100 80GB GPU 한두 장을 거덜 내기에 충분합니다. 더 큰 모델, 더 긴 시퀀스, 더 넓은 vocab에서는 선형적으로 증가합니다.
왜 Cross Entropy가 이렇게 무거운가
Cross entropy loss의 forward는 본질적으로 두 단계입니다.
logits = lm_head(hidden_states)— vocab_size 만큼의 fully connected projectionlog_softmax(logits) → nll_loss(log_probs, targets)— 클래스 수 만큼의 log-softmax + negative log likelihood
nll_loss 자체는 가벼워 보이지만, log_softmax가 vocab 차원 전체에서 exp/sum/log를 수행하면서 큰 activation을 남깁니다. backward 시에는 softmax의 미분(=exp - sum·identity)을 다시 계산해야 하므로, forward activation을 모두 저장해 두어야 합니다.
해결책은 이 giant tensor를 GPU 메모리에 절대 안 올리는 것입니다.
3가지 메모리 절약 전략
전략 1 — Chunked Cross Entropy (CE 분할 계산)
logits를 vocab 차원 기준으로 청크로 잘라서 하나씩 처리하는 방법입니다. PyTorch의 torch.nn.functional.cross_entropy는 내부적으로 모든 청크를 단일 텐서로 만들기 때문에, 직접 chunked 구현을 해야 합니다.
def chunked_cross_entropy(hidden, lm_head_weight, targets, chunk_size=8192):
# hidden: [num_tokens, hidden_dim]
# lm_head_weight: [hidden_dim, vocab_size]
# targets: [num_tokens]
total_loss = 0.0
total_count = 0
for i in range(0, hidden.size(0), chunk_size):
h_chunk = hidden[i:i+chunk_size] # [chunk, hidden]
t_chunk = targets[i:i+chunk_chunk] # [chunk]
logits_chunk = h_chunk @ lm_head_weight # [chunk, vocab]
loss_chunk = F.cross_entropy(logits_chunk, t_chunk, reduction='sum')
total_loss += loss_chunk.item()
total_count += t_chunk.size(0)
del logits_chunk, h_chunk
return total_loss / total_count
메모리 사용량이 O(chunk_size × vocab)로 제한되어 청크 크기만 적절히 잡으면 수 GB 수준으로 떨어집니다.
전략 2 — Gradient Checkpointing on LM Head
LM head의 projection을 gradient checkpointing 안에 넣어서, backward 시 projection을 재계산하는 방법입니다. Forward는 logits을 저장하지 않고, backward에서 hidden_states를 다시 lm_head에 통과시켜 logits을 복원합니다. Activation 메모리는 zero에 가깝게 줄지만, 재계산 비용 30% 추가가 듭니다.
전략 3 — Fused Loss (Triton / Megablocks)

Triton이나 Megablocks의 fused loss 커널을 사용하면, logits 텐서를 material하지 않고 forward+backward를 하나의 fused kernel에서 처리합니다. A100/H100에서는 이게 가장 빠르면서도 메모리 효율적입니다.
실전 비교: 같은 16B 모델, 128K Context
| 전략 | 피크 메모리 | 학습 throughput | 구현 난이도 |
|------|------------|----------------|------------|
| Naive (전체 logits) | 110GB+ | 1.0x | 쉬움 |
| Chunked CE (8K) | 14GB | 0.85x | 중간 |
| Grad checkpointing | 4GB | 0.70x | 중간 |
| Fused Triton loss | 6GB | 1.05x | 어려움 |
128K 시퀀스 길이에서는 chunked + fused loss 조합이 throughput과 메모리 모두에서 가장 균형이 좋습니다.
즉시 적용할 수 있는 팁
- vocab_size가 32K 미만이라면 chunked CE만으로 충분합니다.
- H100 80GB 한 장으로 70B 모델 8K context 학습이 가능합니다 (fused loss 사용 시).
- gradient accumulation step을 늘리면 micro batch는 작게 유지하면서 effective batch는 키울 수 있습니다.
- flash attention과 함께 쓰면 attention activation도 줄이므로 시너지가 큽니다.
- bf16이 아니라 fp32로 loss를 계산해야 numerical stability가 보장됩니다. logits만 bf16, loss accumulation은 fp32 권장.
전망 — Longer Context가 만드는 새로운 표준
Context length가 200K, 500K, 1M으로 늘어나면서 cross entropy 최적화는 단순한 trick이 아니라 필수 인프라가 되고 있습니다. Megablocks, TransformerEngine, Flash-Attention 3 같은 라이브러리들이 fused loss를 표준으로 채택하기 시작했고, custom kernel 없이도 메모리 절반 이하로 줄이는 것이 가능해지고 있습니다.
앞으로는 "어떤 chunk size로 cross entropy를 짤 것인가"가 모델 학습의 기본 옵션이 될 가능성이 높습니다.
요약
- 128K context + 큰 vocab에서 cross entropy는 logits 텐서가 메모리 폭탄의 주범 (40GB~)
- 해결책: chunked cross entropy, gradient checkpointing, fused loss kernel
- 16B + 128K 기준, fused Triton loss는 피크 메모리 6GB, throughput 1.05x로 베스트
- bf16 logits + fp32 loss accumulation이 numerical stability의 기본
- longer context trend에서 cross entropy 최적화는 필수 인프라로 자리 잡는 중
'AI 뉴스' 카테고리의 다른 글
| 덜한 것이 더 낫다, 대체로 — AI 시대 제품 설계의 본질 (0) | 2026.07.06 |
|---|---|
| AI 시대, Figma를 다시 생각하다 — 캔버스가 코드로 자라나는 방식 (0) | 2026.07.06 |
| dbtrail — MySQL을 위한 타임머신, 모든 행 변경을 기억하고 되돌리는 MySQL Time Machine (0) | 2026.07.06 |
| OpenTag — Slack용 Claude Tag의 오픈소스 대안, 그리고 셀프 호스팅 AI 에이전트의 진짜 가치 (1) | 2026.07.06 |
| 더 나은 모델, 더 나빠진 도구 — LLM 도구 호출 역설과 하네스의 미래 (0) | 2026.07.06 |