TL;DR
Kanana-2 시리즈는 온디바이스 제약을 겨냥해 Dense 3B 모델을 시작점으로 단계적 프루닝과 Teacher 기반 디스틸레이션을 결합하여 3B·1.3B·0.9B SLM을 생성했고, 히든 차원 프루닝에 PCA 회전을 적용해 주요 표현을 보존하면서 구조를 효율적으로 축소했다. 온디바이스 추론 병목을 줄이기 위해 Sliding Window Attention 레이어와 Full Attention 레이어를 3:1로 교차 배치한 하이브리드 구조를 적용해 32K 문맥에서 KV Cache를 최대 72.7% 절감했고 Window size 1024는 Full Attention과 대등한 성능을 보였다. Post-training 단계에서는 Staged SFT로 도메인별 Expert를 얻고 SCE 방식으로 병합한 뒤 DAPO 기반 RL을 적용해 응답 다양성과 안정성을 확보했으며, 임베딩은 Bidirectional+Mean Pool 초기화와 HN-only 손실, Temperature=0.02 설정으로 한국어·영어에서 우수한 성능과 차원 효율을 달성했다. 전체 파이프라인은 Proxy Token Scale 기반의 LR 스케일링, PCA 기반 프루닝, Cascade/Elastic 훈련 옵션을 통해 성능과 비용의 균형을 맞추는 방향으로 설계되어 실무적 재현 가능성과 온디바이스 실용성을 동시에 목표로 한다.
빠른 이해
새로운 점
PCA 기반 히든 차원 프루닝과 SWA 하이브리드 배치의 결합으로 소형 온디바이스 모델에서 성능 손실을 최소화하면서 KV Cache를 크게 절감한 점
핵심 메커니즘
입력으로는 Dense 3B Base 모델과 Teacher rollout·캘리브레이션 데이터가 주어지고, 처리 과정에서는 PCA 회전을 통한 히든 차원 축소와 단계적(Cascade) 프루닝·디스틸레이션, SWA 하이브리드 배치로 KV Cache를 제한하며 Post-training에서 DAPO 기반 RL로 정책을 정교화한 뒤 출력으로는 3B·1.3B·0.9B SLM과 고효율 임베딩 모델이 산출된다.
핵심 수치
- Pre-training 토큰량: Stage-1 7.5T + Stage-2 2T
- SWA 하이브리드(Kanana-2-1.3B, 32K): KV Cache 최대 -72.7%- SWA:Full=3:1 대비 Full-only
- 소형 모델 Distillation 토큰: 각 단계별 약 300B- 2B→1.3B→0.9B 각 단계에 대해 300B 토큰으로 Distillation 수행
섹션별 상세
프로젝트 개요와 목표
사전학습 전략과 하이퍼파라미터 탐색
- 사전학습은 TPU 기반으로 Stage-1 7.5T 토큰, Stage-2 2T 토큰의 2단계로 진행되었고 Proxy D_proxy=100B와 β=0.32를 사용해 학습률을 스케일링했다. — TPU 기반 Pre-Training 섹션의 토큰 수와 Token Horizon LR 스케일링 공식
Teacher 기반 Distillation과 Long Context 학습
프루닝 및 구조 축소 설계
PCA 기반 히든 차원 프루닝 파이프라인
def prune(model, target_config):
# Activations are cached during calibration forward passes.
# Norm inputs are collected before rotation, while attention and MLP
# activations are collected after rotation.
# Hidden dimension pruning
norm_inputs = collect_norm_inputs(model)
rotation = PCA( norm_inputs, n_components=target_config.hidden_dim, )
model = apply_rotation( model, rotation, target_dim=target_config.hidden_dim, )
# Attention head pruning
head_importance = sum( score_query_heads(layer.attn) for layer in model.layers )
heads_to_keep = select_topk_heads_per_kv_group( head_importance, k=target_config.query_heads, )
for layer in model.layers:
layer.attn = prune_query_heads( layer.attn, heads_to_keep, )
# MLP intermediate dimension pruning
for layer in model.layers:
intermediate = ( silu(layer.mlp.gate_proj.output_x) * layer.mlp.up_proj.output_x )
importance = aggregate_intermediate_importance(intermediate)
dims_to_keep = topk( importance, k=target_config.intermediate_dim, )
layer.mlp = prune_intermediate_dims( layer.mlp, dims_to_keep, )
return model이 코드는 PCA 기반 회전을 적용한 뒤 Attention 헤드와 MLP 중간 차원을 순차적으로 프루닝하는 파이프라인의 핵심 로직을 예시한 파이썬 형태의 의사코드이다.
온디바이스 최적화와 Sliding Window Attention
Post-Training, RL 및 임베딩 설계
- SWA:Full=3:1 하이브리드 구조는 Kanana-2-1.3B 모델에서 32K 문맥 기준 KV Cache를 최대 72.7% 절감했다. — KV Cache 크기 비교 표(Architecture별 Length별 MiB 수치)
용어 해설
- PCA 기반 히든 차원 프루닝(PCA-based Hidden Dimension Pruning)
- — 레이어별 활성화 통계를 모아 전역 회전 행렬을 계산하고, 토큰 임베딩·어텐션·MLP 투영에 일관되게 회전을 적용한 뒤 차원을 축소하여 주요 표현을 보존하면서 파라미터를 줄이는 방식으로, 프루닝으로 인한 정보 손실을 줄이고 압축 모델의 성능 회복 속도를 개선한다.
- 슬라이딩 윈도우 어텐션(Sliding Window Attention)
- — 각 토큰이 고정된 윈도우 범위 내에서만 어텐션을 수행하도록 하여 KV Cache 길이를 윈도우 크기로 제한함으로써 토큰당 메모리 읽기량을 시퀀스 길이와 무관하게 일정하게 유지해 온디바이스 추론 시 메모리와 레이턴시 병목을 완화하는 지역적 어텐션 기법이다.
- 단계적(카스케이드) 프루닝 및 디스틸레이션(Cascade Pruning & Distillation)
- — 대형 모델에서 작은 모델로 한 번에 압축하지 않고 중간 규모의 서브모델을 차례로 생성하면서 각 단계에서 프루닝과 교사 모델 기반 디스틸레이션을 반복하여 최종 소형 모델의 초기 성능을 높이고 학습 안정성을 확보하는 압축 파이프라인이다.
- Decoupled Clip and Dynamic Sampling Policy Optimization(DAPO)
- — 동적 샘플링으로 유효 샘플만 배치에 포함시키고 클립 범위를 확장해 정책 변화로 인한 엔트로피 붕괴를 방지하며 토큰 단위의 정책 그래디언트와 과긴 출력에 대한 패널티를 결합해 RL 학습의 효율과 안정성을 개선하는 최신 강화학습 최적화 기법이다.
- 토큰 지평선 학습률 스케일링(Token Horizon LR Scaling)
- — 작은 토큰 규모(Proxy)에서 탐색한 최적 학습률을 타깃 토큰 규모로 변환하기 위한 스케일링 법칙으로, LR*(D_target) ≈ LR*(D_proxy) · (D_target / D_proxy)^{-β} 공식을 사용해 긴 토큰 지평선에서 안정적인 학습률 설정을 확보한다.
기술
- TPU v5e
- Megatron-LM
- YaRN
- PCA
- Sliding Window Attention
- DAPO
언급된 리소스
AI 요약 · 북마크 · 개인 피드 설정 — 무료
출처 · 인용 안내
인용 시 "요약 출처: AI Trends (aitrends.kr)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.