LLM WikiAccess-protected knowledge portal
← 스터디 홈
163편 · 약 14분

LLM 지식 증류 효율화: Offline Top-K Logits와 Fused Chunked KL Loss로 교사를 메모리에서 꺼내는 방법 (arXiv:2608.03796)

요약

대형 언어 모델(LLM)을 소형 모델로 압축할 때 가장 효과적인 방법 중 하나가 지식 증류(Knowledge Distillation, KD)다. 교사 모델의 출력 분포를 학습 신호로 삼아 학생 모델을 훈련하면, 레이블만으로 학습할 때보다 훨씬 좋은 성능을 얻을 수 있다.

문제는 비용이다. 표준 온라인 KD는 매 학습 스텝마다 교사 모델을 메모리에 올리고 순전파를 실행해야 한다. 교사가 70B 모델이라면, 학생 7B와 교사 70B가 동시에 GPU에 올라가야 하므로 실제 연구실이 아닌 곳에서는 거의 불가능하다.

arXiv:2608.03796(2026년 8월 4일 제출)은 이 병목을 두 가지 방법으로 동시에 해결한다.

  1. Offline Top-K Logits: 교사의 상위 K개 로짓을 미리 한 번만 계산해 캐시하고, 학습 내내 캐시만 참조한다. 교사가 메모리에서 사라진다.
  2. Fused Chunked KL Loss: 어휘(vocabulary) 전체 크기의 로짓 텐서를 한 번도 완전히 생성하지 않아 피크 메모리를 시퀀스 길이에 선형으로 줄인다.

결과: 온라인 KD 대비 훈련 손실이 거의 동일하면서 반복당 29% 빠르고, H200 GPU에서 처리량이 최대 41% 높다.


배경: 왜 지식 증류가 어려운가

표준 KD의 손실 함수

교사 분포 $P_T$와 학생 분포 $P_S$ 사이의 KL 발산을 최소화한다.

L_KD = KL(P_T || P_S)
      = Σ_{v ∈ V} P_T(v) · log[P_T(v) / P_S(v)]

여기서 $V$는 어휘 집합(vocabulary). 현대 LLM의 어휘 크기는 32,000~128,256 토큰이다.

두 가지 메모리 병목

메모리 병목 1: 교사 모델 상주 비용

온라인 KD에서 교사는 매 스텝 순전파를 실행해야 한다. 훈련 루프 안에 교사와 학생이 동시에 올라간다.

모델 조합교사 메모리학생 메모리합산
교사 7B → 학생 1B~14 GB (BF16)~2 GB~16 GB
교사 70B → 학생 7B~140 GB (BF16)~14 GB~154 GB
교사 405B → 학생 70B~810 GB (BF16)~140 GB~950 GB

교사가 클수록 단일 노드에서 훈련은 불가능에 가까워진다.

메모리 병목 2: 전체 어휘 로짓 텐서

KL 손실을 계산하려면 교사·학생 양쪽의 로짓 텐서(batch × seq_len × |V|)를 생성해야 한다. 예를 들어:

batch=8, seq_len=4096, |V|=128,256
→ 텐서 하나의 크기: 8 × 4096 × 128256 × 2 bytes(BF16) ≈ 8.4 GB
→ 교사·학생·그라디언트 합산: 25+ GB

이 텐서가 피크 메모리의 주요 원인이 된다.


핵심 아이디어 1: Offline Top-K Logits

동작 원리

훈련 전에 교사를 한 번만 실행해 각 토큰 위치에서 상위 K개 로짓만 저장한다.

온라인 KD (기존)
교사 모델 (메모리 상주)
↓ 매 스텝 순전파
전체 어휘 로짓 (|V| × seq)
↓ KL 손실
학생 모델 업데이트
오프라인 KD (논문 방법)
교사 사전 실행 (1회)
↓ Top-K 로짓 캐시 저장
캐시 (K × seq, K=100)
↓ Fused Chunked KL 손실
학생 모델 업데이트
오프라인: 훈련 중 교사 메모리 0 · 피크 메모리 선형 감소
온라인 KD vs 오프라인 KD 흐름 비교

왜 Top-K면 충분한가

소프트맥스 분포에서 실제로 확률 질량이 집중되는 위치는 상위 소수 토큰이다. 어휘 전체의 99.5% 이상의 확률은 상위 100~500개 토큰에 집중된다는 것이 경험적으로 알려져 있다. 나머지 토큰의 로짓은 학습 신호로서의 효용이 낮다.

저장 효율:

전체 어휘: |V| = 128,256 개 로짓/토큰
Top-K:    K   =     100 개 로짓/토큰
압축률: 99.92%

캐시는 교사가 생성한 (token_index, logit_value) 쌍 K개만 저장하므로 디스크 공간도 크게 줄어든다.


핵심 아이디어 2: Fused Chunked KL Loss

전통적인 KL 손실의 문제

PyTorch나 NVIDIA Megatron의 기본 KL 손실은 다음 두 텐서를 전체 어휘 크기로 생성한다.

# 기존: 두 텐서 모두 (batch × seq_len × |V|) 크기
log_p_T = teacher_logits.log_softmax(dim=-1)  # 거대한 텐서
log_p_S = student_logits.log_softmax(dim=-1)  # 또 다른 거대한 텐서
loss = F.kl_div(log_p_S, log_p_T.exp(), reduction='batchmean')

이 과정에서 중간 텐서가 피크 메모리를 지배한다. 어휘 크기가 클수록 문제가 심해진다.

Fused Chunked KL: 어휘 차원을 청크로 나누기

논문은 어휘 차원을 여러 청크로 나눠 각 청크만큼씩 계산하고 합산한다. 전체 어휘 텐서를 한 번도 완전히 생성하지 않는다.

피크 메모리: O(batch × seq_len × chunk_size)
             (기존: O(batch × seq_len × |V|))

"Fused"는 softmax, log, KL 계산을 단일 CUDA 커널로 합쳐 청크 단위의 연산 오버헤드를 최소화했다는 뜻이다.

오프라인 Top-K 캐시와 결합하면 교사 로짓을 청크 단위로 복원해 KL을 계산할 수 있어 메모리 절감 효과가 배가된다.


성능 결과

처리량과 속도 (H200 기준)

방법반복 속도처리량
온라인 KD (기준선)1.0×1.0×
오프라인 Top-K KD1.29×1.41×
  • 반복당 29% 빠름: 교사 순전파가 사라지므로 GPU 계산 시간이 줄어든다.
  • 처리량 최대 41% 향상: 교사 메모리가 해제되어 더 큰 배치를 돌릴 수 있다.

훈련 손실 동등성

오프라인 KD가 온라인 KD 대비 훈련 손실이 거의 동일하다는 것이 핵심이다. 상위 K=100개 로짓만으로도 교사 분포의 학습 신호를 충분히 전달한다.

오픈 질문으로 남는 부분: K 값 선택이 도메인별로 달라질 수 있다. 수학·코드처럼 엔트로피가 낮은 분포에는 K=50도 충분할 수 있고, 창의적 텍스트 생성처럼 엔트로피가 높은 분포에는 더 큰 K가 필요할 수 있다.


운영 적용 가이드

사전 계산(Pre-computation) 파이프라인

1. 교사 모델 로드 (추론 전용, 그라디언트 불필요)
   → 메모리: 교사 가중치만, 옵티마이저 없음

2. 훈련 데이터셋 전체 순전파
   → 각 토큰 위치에서 Top-K 로짓과 인덱스 저장

3. 캐시 저장: {token_id: [(vocab_idx, logit_val), ...] × K}
   → 포맷: int16(인덱스) + bf16(값) 권장
   → 1B 토큰 데이터셋, K=100 → 약 340 GB (BF16 기준)

4. 교사 모델 언로드, 학생 훈련 시작
   → 교사 GPU 메모리 0

캐시 크기 추정

캐시 크기 = 토큰 수 × K × (2 bytes_idx + 2 bytes_val)
예: 1B 토큰, K=100 → 400 GB

실용 팁:
- int16 인덱스: |V| < 65536이면 충분
- NVMe SSD에서 sequential read → 학습 병목이 아님
- DataLoader에서 prefetch 로딩 권장

기존 라이브러리 통합

Megatron-LM, LlamaFactory, Axolotl 등 주요 훈련 프레임워크에서 손실 함수만 교체하면 적용 가능하다. 교사 사전 계산 스크립트와 Fused KL 커널이 독립적으로 동작하므로 훈련 루프 수정은 최소화된다.


제한과 주의사항

항목내용
교사 업데이트교사 가중치 변경 시 캐시를 전체 재생성해야 한다
디스크 공간대규모 데이터셋에서 캐시가 수 TB에 달할 수 있다
K 값 민감도극단적으로 작은 K(≤10)에서는 손실 동등성이 보장되지 않는다
교사와 학생 어휘 불일치교사와 학생이 동일한 토크나이저를 사용해야 한다
평가 지표논문은 훈련 손실을 주 지표로 사용했다. 하위 태스크 성능 검증은 추가 실험이 필요하다

요점 정리

온라인 KD의 두 가지 메모리 병목 — 교사 모델 상주와 전체 어휘 로짓 텐서 — 을 각각 오프라인 캐싱청크 단위 Fused 커널로 해결했다.

훈련 손실의 동등성을 유지하면서 속도와 처리량을 동시에 개선한 점이 중요하다. 이는 KD를 실험실 밖의 실제 운영 환경에서 쓸 수 있게 만드는 첫 번째 실용적 단계다.

소형 LLM을 온프레미스나 비용 제약 환경에서 전문화해야 하는 조직이라면, 이 두 기법을 조합해 70B 교사의 지식을 7B 학생에게 단일 H100 노드 안에서 증류하는 파이프라인을 구성할 수 있다.


References

  • arXiv:2608.03796 — "Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss" (2026.08.04): https://arxiv.org/abs/2608.03796
  • Hugging Face Papers 페이지: https://huggingface.co/papers/2608.03796
  • 지식 증류 서베이 arXiv:2402.13116 — "A Survey on Knowledge Distillation of Large Language Models" (2024.02)
  • 강화 학습 인식 지식 증류 arXiv:2602.22495 — "Reinforcement-aware Knowledge Distillation for LLM Reasoning" (2026.02)
  • Offline Top-K 구현 해설 (Hugging Face Blog): https://huggingface.co/blog/MultiverseComputingCAI/efficient-knowledge-distillation