지식 증류를 대규모로 실행하기에 충분히 저렴하게 만드는 방법
요약
본 논문은 지식 증류 과정에서 발생하는 막대한 VRAM 요구 사항 문제를 해결하는 방법을 제시합니다. 교사 모델의 Top-K 로짓 캐싱과 메모리 효율적인 KL-발산 손실을 도입하여, 대규모 LLM 훈련 비용을 크게 절감했습니다.
핵심 포인트
- Top-K 로짓 캐싱으로 교사 모델 상주 필요성 제거
- 메모리 효율적 KL-손실로 VRAM 사용량 대폭 감소
- 단일 GPU에서 긴 컨텍스트 치유가 가능해짐
- 대규모 LLM 훈련을 실용적인 수준으로 만듦
지식 증류(Knowledge distillation)는 더 큰 교사 모델(teacher model)의 성능을 작은 학생 모델(student model)이 따라 하도록 훈련하는, 머신러닝에서 잘 알려진 기술입니다. gpt-oss, Qwen, GLM 또는 Kimi와 같은 최근 오픈 소스 대규모 언어 모델(LLMs)의 물결과 함께, 이는 다시 주류 연구 주제가 되었습니다. 이러한 매우 큰 모델을 배포하는 것은 비용이 많이 듭니다. 예를 들어, 최근의 Kimi-K3 모델은 2조 8천억 개의 파라미터를 가지고 있으며 로드하는 데만 약 3TB의 VRAM이 필요합니다. 따라서 지식 증류를 통해 이들을 더 작은 모델로 압축하고 원래의 기능을 복구하는 것이 표준 관행이 되었으며, Nvidia(Nemotron 3 Puzzle 75B)나 Multiverse Computing(Hypernova 60B) 같은 회사들이 최근 고품질의 압축된 모델을 출시했습니다.
증류 단계가 최종 품질의 대부분을 결정하지만, 일반적으로 파이프라인에서 가장 비용이 많이 드는 부분입니다. 교사와 학생 모두를 로드하고 모든 토큰에 대해 전체 어휘 사전을 걸친 확률 분포를 생성하는 것은 엄청난 양의 VRAM을 필요로 하며, 일반적으로 수백 개의 GPU와 신중한 텐서 병렬 처리(tensor-parallelism) 전략으로만 실현 가능합니다. 저희의 최신 논문인 'Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss'는 이 문제를 두 가지 시스템 변경 사항으로 다룹니다. 첫째, 교사의 top-K 로짓(logits)을 한 번 캐싱하여 교사가 학생과 함께 메모리에 상주할 필요가 없게 만들고, 둘째, 전체 어휘 크기 × 시퀀스 길이 행렬을 절대 생성하지 않는 새롭고 메모리 효율적인 KL-발산 손실(KL-divergence loss)을 도입했습니다. 이는 PyTorch나 NVIDIA Megatron-Bridge와 같은 라이브러리의 기본 구현이 달성하는 VRAM 사용량보다 훨씬 적은 양으로 줄입니다. 이 두 가지 변경 사항 덕분에 훈련 비용이 충분히 절감되어 단일 GPU에서 긴 컨텍스트 치유(long-context healing)가 가능해졌고, 대규모 실험을 실용적으로 만들 만큼 저렴해졌습니다.
표준 설정인 Kullback-Leibler divergence loss (KL loss)를 사용한 온라인 증류(online distillation)는 교사 모델(teacher)과 학생 모델(student)을 동시에 로드해야 합니다. 매 훈련 단계마다 교사는 전체 순전파(forward pass)를 실행하여 출력 분포를 생성하고, 학생은 이 분포와 일치하도록 훈련됩니다. 이는 전체 교사 분포가 사용 가능하기 때문에 가장 표현력이 풍부한 설정이지만, 메모리와 컴퓨팅 자원 측면에서도 가장 많은 자원을 소모합니다. 토큰 위치당 두 개의 전체 어휘(full-vocabulary) 텐서를 유지해야 하며, 교사의 동작이 훈련 실행 전반에 걸쳐 변하지 않음에도 불구하고 매 단계마다 교사를 재계산해야 합니다.
실제 예로, gpt-oss-120b는 201,088개의 토큰 어휘를 가지고 있습니다. 시퀀스 길이(sequence length)가 32K이고 배치 크기(batch size)가 4일 경우, 교사 확률 텐서(teacher-probability tensor)만 해도 4 × 201,088 × 32,768의 형태를 가집니다.
; bfloat16 기준으로 이는 이미 단일 텐서 하나에 약 50GB의 VRAM을 차지합니다. 여기에 기울기(gradients), 활성화값(activations), 모델 가중치(model weights), 최적화기 상태(optimizer states)를 더하면, 증류의 단일 훈련 반복은 최대 약 250GB의 VRAM에 도달할 수 있으며, 이는 H200이나 B200 GPU가 제공하는 용량보다도 많습니다. 본 게시물에서는 KL loss를 재구성하여 데이터를 청크(chunks) 단위로 처리함으로써 이 비용을 거의 '제로' 수준으로 줄일 수 있음을 보여줍니다.

Dense KL은 약 250GB에 달하며, 이는 단일 H200의 141GB 용량을 초과합니다. 통합된 청크 손실(fused chunked loss)은 그러한 급증을 만들지 않으며 약 128GB에서 최고치를 기록합니다. 출처: 논문 Figure 1.
오프라인 증류(Offline distillation). 매 단계마다 교사를 재계산하는 대신, 우리는 교사의 출력을 한 번 계산하고 위치별로 가장 가능성이 높은 상위 100개 토큰을 캐시(cache)하며, 학생은 이 캐시를 기반으로 훈련됩니다. 교사는 훈련 중 메모리에 머무를 필요가 없으며, 캐시가 존재하는 한 다시 실행될 필요도 없으므로, 동일한 캐시를 여러 어블레이션(ablations)에 걸쳐 재사용할 수 있습니다.
융합된 청크 KL 손실(A fused, chunked KL loss). 이 손실 자체가 왜 비싼지 이해하려면, 실제로 무엇을 구축하는지 상상해 보세요. 시퀀스의 모든 토큰 위치와 어휘집의 모든 단어에 대해, 학생 모델의 예측이 교사 모델의 예측과 얼마나 불일치하는지를 설명하는 숫자가 필요합니다. 이것을 격자 형태로 펼쳐 놓으면, 어휘 항목당 한 행, 시퀀스 위치당 한 열이 되며, 100K개 이상의 단어를 가진 어휘집과 긴 시퀀스의 경우 이 격자는 엄청나게 크고, KL 손실을 계산하는 기본 방식은 하나의 숫자를 산출하기 전에 전체를 구축해야 합니다.
저희는 수학적으로 동등한 세 가지 방식으로 이 동일한 손실을 계산하는 방법을 비교합니다:
Dense KL은 교과서적인 접근 방식입니다. 캐시된 top-100 로짓(logits)으로부터 완전하고 밀집된(dense) 교사 확률 격자를 재구축하여 학생 모델 자체의 밀집된 로그 확률 그리드와 비교합니다. 이는 온라인 증류(online distillation)가 이미 작동하는 방식과 가장 유사하기 때문에, 저희는 이를 정확성 기준선(correctness baseline)으로 사용하지만, 전체 어휘집 × 시퀀스 격자 전체를 메모리에 두 번 저장해야 합니다.Forward-chunked KL은 교사 모델을 희소하게 유지하고(각 위치별 캐시된 top-100 로짓만 사용하며, 밀집된 그리드로 확장하지 않음) 손실을 조각별로, 즉 시퀀스 위치 슬라이스 단위로 계산합니다. 이는 밀집된 교사와 밀집된 비교를 제거하며, 저희 벤치마크에서 세 가지 방법 중 가장 빠르다는 결과가 나왔습니다. 하지만 여전히 하나의 사각지대가 있습니다. 모델의 출력 레이어에 의해 생성되는 학생 자체의 로짓 그리드는 여전히 전체가 계산되어 역전파(backward pass)를 위해 유지되므로, 메모리가 시퀀스 길이에 따라 가파르게 증가합니다.Fused chunked KL은 저희의 주요 기여로, 한 단계 더 나아가 모델의 출력 투영(output projection)을 손실 계산에 직접 융합합니다. 학생의 전체 로짓 그리드를 전혀 생성하지 않습니다. 대신, 시퀀스의 청크를 하나씩 처음부터 끝까지 처리하면서, 해당 청크에 대해 은닉 상태(hidden states)를 로짓으로 투영하고, 그 결과를 누적되는 손실에 접어 넣은 다음, 다음 청크로 넘어가기 전에 해당 청크는 폐기합니다.
역방향 전파(backward pass)는 청크를 저장하는 대신 실시간으로 재계산합니다. 비용은 해당 투영(projection)을 두 번 수행한다는 점입니다. 한 번은 순방향(forward), 다른 한 번은 역방향이지만, 그 대가로 피크 메모리 사용량은 전체 어휘 크기 × 시퀀스 크기로 급증하는 대신 시퀀스 길이에 따라 선형적으로만 증가합니다.
아래 GIF는 밀집 청크 방식(dense)과 융합 청크 방식(fused-chunked)의 차이를 보여줍니다. 전자는 전체 비교 그리드(comparison grid)를 구축하고 모두 유지하는 반면, 후자는 한 번에 하나의 슬라이스만 구축하고 폐기하므로 메모리가 단일 청크를 초과하여 증가하지 않습니다.
저희는 청크 손실 구현(chunked-loss implementation)을 오픈소스로 공개했습니다: github.com/CompactifAI/Full-Chunked-KL-Loss
아래 표에는 네 가지 설정, 즉 온라인 증류(online distillation)와 방금 설명한 세 가지 오프라인 손실 구현이 나란히 비교되어 있습니다. H200 GPU 단일 장치에서 Llama 3.1 8B Instruct를 교사 모델(teacher)로, 3.2B Llama 모델을 학생 모델(student)으로 사용하여 8K 토큰 컨텍스트 길이에서 비교했을 때, 네 가지 모두 거의 동일한 학습 손실에 도달합니다. 비록 오프라인 실행은 토큰당 캐시된 상위 100개 로짓(top-100 logits)만을 가지고 학습하지만 말입니다.
| 방법 (8K 컨텍스트, H200 단일 장치) | 피크 메모리 | 반복 시간 | 처리량 |
|---|---|---|---|
| 온라인 증류 | 102.8 GB | 25.9 s | 237 TFLOP/s |
| ... |
The 손실 곡선은 네 가지 방법 모두에서 거의 정확하게 겹치며, 상위 100개 캐시된 로짓을 사용한 오프라인 증류가 온라인 증류에 대해 무손실(lossless)임을 확인시켜 줍니다. 출처: 논문 Figure 2. 이 시퀀스 길이에서는 융합 청크 손실이 아직 가장 빠른 옵션은 아니며, 추가적인 역방향 전파 투영 비용 때문에 속도가 약간 떨어지지만, 실제 장점은 컨텍스트 길이가 증가함에 따라 나타나며, 다음 섹션에서 이를 입증합니다.
확장 패턴을 더 명확하게 보기 위해, 저희는 장난감 출력 투영 네트워크(트랜스포머 본체 없이 손실 커널만 사용)에 대해 격리된 벤치마크를 실행했습니다. 32K 토큰에서 피크 메모리는 밀집 손실(dense loss)을 사용할 경우 85.2 GiB였으나, 완전히 청킹된 버전(fully chunked version)에서는 5.45 GiB로 줄어들어 15.6배 감소했으며, 밀집 손실은 64K 토큰부터는 아예 실패했습니다. 256K 토큰에서는 완전히 청킹된 손실이 이전 최적의 청킹 변형(next-best chunked variant)에 비해 11.6 GiB를 사용하여 134.2 GiB였던 메모리를 절약했고, 해당 길이에서 반복당 속도 또한 약 3.3배 빨랐습니다.
GPT-OSS 20B 모델을 32,768 토큰 컨텍스트로 증류(distilling)하는 과정에서, 결합 손실(fused loss)이 확보한 메모리 덕분에 설정 규모가 네 개의 GPU 노드에서 단 하나의 노드로 줄어들었습니다. 단계별 시간은 57.0초에서 12.23초로 약 5배 빨라졌고, GPU당 처리량(throughput per GPU)은 74.2 TFLOP/s에서 345.7 TFLOP/s로 증가했습니다.
이러한 효율적인 오프라인 설정 덕분에 대규모 증류 캠페인이 애초에 저렴해질 수 있었습니다. 그 결과, Llama 3.1 8B Instruct 모델을 약 3.2B 파라미터까지 증류하여 얻은 작고 간결한 학생 모델(student)은 BoolQ와 HellaSwag에서 교사 모델(teacher)의 정확도를 대부분 유지하며, MMLU에서는 파라미터 수가 절반 이하임에도 불구하고 약 9점 차이로 그 성능을 유지합니다.

학생 모델이 크기의 절반 이하에서도 교사 모델의 짧은 컨텍스트 정확도를 대부분 유지합니다. 출처: paper Figure 6.
본 연구는 Multiverse Computing이 증류 및 치유(healing)를 단순히 일회성 레시피가 아니라, 팀들이 저렴하게 반복할 수 있는 방식으로 대규모로 실행 가능하도록 만드는 지속적인 연구의 일부입니다. 이 논문은 또한 손실 함수 선택과 시퀀스 패킹(sequence packing)이 복구 품질에 미치는 영향 등 추가적인 제거 실험(ablations)도 다루고 있습니다.
결합 청킹 손실 뒤에 숨겨진 폐쇄형 근사도 기울기(closed-form gradient)와 전체 훈련 구성 정보를 포함한 모든 기술적 세부 사항을 알고 싶으신가요? 전체 논문을 읽거나, 저희 팀에게 연락하여 귀하의 증류 파이프라인에 적용하는 것에 대해 이야기해 보세요.
또한 청킹 손실 구현체도 오픈 소스화했습니다: github.com/CompactifAI/Full-Chunked-KL-Loss
AI 자동 생성 콘텐츠
본 콘텐츠는 Hugging Face Blog의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기