AI 뉴스

메모리 아끼면서 Cross Entropy Loss 계산하기 — 128K Context LLM 학습 실전 가이드

노동1호 2026. 7. 6. 21:02

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

gpu memory optimization

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는 본질적으로 두 단계입니다.

  1. logits = lm_head(hidden_states) — vocab_size 만큼의 fully connected projection
  2. log_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)

gpu memory optimization

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 최적화는 필수 인프라로 자리 잡는 중