TL;DR
사전학습 데이터의 출처 구성과 혼합 전략이 LLM 성능과 학습 효율을 결정한다. HDS는 데이터 혼합을 에이전트가 실시간으로 최적화하도록 하여 학습 스텝과 연산 비용을 크게 줄이고 downstream 성능을 동시에 향상시켰다.
왜 중요한가
사전학습 데이터의 출처 구성과 혼합 전략이 LLM 성능과 학습 효율을 결정한다. HDS는 데이터 혼합을 에이전트가 실시간으로 최적화하도록 하여 학습 스텝과 연산 비용을 크게 줄이고 downstream 성능을 동시에 향상시켰다.
핵심 기여
Holistic Data Scheduler (HDS)
데이터 품질(data-driven), 도메인 간 영향(inter-domain/gradient alignment), 모델 안정성(model-driven weight-norm) 세 가지 보상 성분을 결합한 다목적 보상으로 온라인 데이터 믹싱을 수행하는 프레임워크를 제시한다. 제어 공간을 연속 확률 심플렉스로 설정하고 Soft Actor-Critic(SAC)을 사용해 안정적으로 정책을 학습한다.
학습 효율 및 최종 성능 개선
Pythia-1B를 The Pile(50B tokens)로 학습한 결과, HDS는 TPW 대비 동일한 최종 perplexity를 달성하는 데 필요한 학습 스텝을 약 57% 절감했고, AC-ODM 대비 약 44% 절감했다. 최종적으로 검증 perplexity에서 TPW 대비 13.6% 개선을 보였고 MMLU 0-shot 정확도가 0.26915(약 +7.2% 상대)였다.
설계 및 구현 가이드
가벼운 Transformer 기반 actor/critic 아키텍처(약 5M 파라미터, 에이전트 총 26.5M)를 제시하고, 상태(state)·보상 구성, 레이어 선택, 에이전트 크기(≈0.5% LLM 파라미터) 등 실무적 하이퍼파라미터 튜닝 가이드를 제공한다.
보상 구성의 분해(ablations)
r_align(gradient alignment), r_diversity(MTLD 기반 스케줄), r_stability(weight-norm 변화 억제) 세 보상 성분의 기여도를 실험적으로 분리하여 r_align의 기여가 가장 크고, r_diversity가 조기 학습 수렴에 중요하며 r_stability가 안정화에 기여함을 확인했다.
스케일 검증
Pythia-12B 사전학습(25B tokens)에서도 HDS가 ODM 대비 모든 체크포인트에서 더 낮은 검증 perplexity를 기록했고, 마지막 체크포인트에서는 약 33% 상대 개선을 보였다.
핵심 아이디어 이해하기
출발점과 기존 한계: LLM 사전학습에서는 서로 다른 도메인들의 혼합비가 학습 효율과 최종 성능에 직접적인 영향을 준다. 기존의 온라인 데이터 믹싱(ODM, AC-ODM 등)은 단일 관점(예: 도메인별 손실 또는 gradient alignment)에 기반한 보상 신호에 의존해 도메인 간 상호작용이나 학습 단계별 데이터 난이도 변화를 충분히 고려하지 못했다. 이로 인해 학습 초기 단계에서 지나치게 어려운 데이터가 과다 샘플링되거나, 상호 보완적 도메인을 활용하지 못해 수렴 속도가 떨어졌다.
해결 원리: HDS는 데이터 혼합을 연속 제어 문제로 모델링하고, 에이전트가 도메인 가중치(확률 심플렉스)를 직접 출력하도록 한다. 상태는 도메인별 누적 샘플 수, 도메인별 검증 손실과 손실 변화, 선택된 레이어의 weight L2 norm 및 변화 등으로 구성되어 에이전트가 모델의 현재 학습 단계와 내부 안정성을 파악하도록 한다. 보상은 세 축으로 구성된다: (1) gradient alignment — 입력: 도메인 i의 gradient g_i와 다른 도메인들의 gradient 합 Σ_{j≠i} g_j → 연산: 내적 ⟨g_i, Σ_{j≠i} g_j⟩ → 결과: 스칼라값 → 의미: 양수이면 도메인 i의 업데이트가 다른 도메인에 유익하다는 신호, (2) scheduled lexical diversity — 입력: 배치의 MTLD → 연산: 정규화 및 t'와 결합(r_diversity = t' / (MTLD_norm + ε)) → 결과: 시간에 따른 curriculum 보상 → 의미: 초기에는 낮은 MTLD(단순 텍스트) 보상, 후반에는 높은 MTLD(복잡 텍스트) 보상, (3) model stability — 입력: 선택된 레이어들의 L2-norm 변화 Δ||ω||_2 → 연산: 역수 1/(|Δ||ω||_2|+ε), 상한치 적용 → 결과: 스칼라 보상 → 의미: 급격한 파라미터 변화를 억제해 안정적 수렴 유도.
달라지는 점(구체적 변화): 이 구성은 에이전트가 '어떤 도메인을, 언제 더 많이 샘플링할지'를 학습 단계별로 조정하게 해 학습 초기에는 단순·반복 데이터로 기초를 다지고, 중·후반에는 고다양성 도메인으로 전환하는 커리큘럼을 자동 형성한다. 실험에서 이 접근은 TPW 대비 동일 perplexity 달성에 필요한 스텝을 약 57% 줄였고, downstream(MMLU 0-shot) 성능을 0.26915로 끌어올렸다.
방법론
전체 접근 방식: HDS는 온라인 데이터 믹싱을 Markov Decision Process로 정식화하여 에이전트가 연속 행동 공간(확률 심플렉스)에서 도메인 가중치 a^{t+1}을 출력하도록 한다. 환경은 LLM의 학습 프로세스 자체이며, 에이전트는 현재 상태 s^t(샘플 수, 스텝, 도메인별 검증 손실 및 델타, 선택된 레이어의 L2 norm 및 그 변화)를 관찰해 stochastic policy π_{θ_A}(a|s)를 통해 샘플링 확률을 결정한다. 이 행동으로 구성된 배치 B^t로 LLM을 업데이트하고 보상 r^t를 계산해 transition을 리플레이 버퍼에 쌓아 SAC 에이전트를 업데이트한다.
핵심 메커니즘/알고리즘 상세: 보상은 세 성분의 가중합으로 구성된다. (1) Inter-domain influence r_align: 도메인 i의 gradient g_i = ∇{θ_M} ℒ(θ_M, B_i)와 다른 도메인들의 gradient 합의 내적 ⟨g_i, Σ{j≠i} g_j⟩을 계산해 도메인의 cross-domain 유용성을 측정한다(입력: gradient 벡터들 → 연산: 내적 → 결과: 스칼라 alignment score → 의미: 양수면 긍정적 전이). (2) Scheduled lexical diversity r_diversity: 배치별 MTLD를 MTLD_norm으로 정규화하고 현재 상대 스텝 t' = t/T_total을 이용해 r_diversity = t'/(MTLD_norm + ε)로 계산한다(입력: MTLD → 연산: 정규화 및 역수 결합 → 결과: 시간 의존적 보상 → 의미: 초기에는 단순 텍스트 보상, 후반에는 복잡 텍스트 보상). (3) Model stability r_stability: 선택된 레이어들의 L2-norm 변화량 | ||ω^t||_2 − ||ω^{t-1}||_2 |을 계산해 역수를 취하고 상한치로 캡하여 안정성 보상을 얻는다(입력: L2-norm 변화 → 연산: 역수 및 상한 적용 → 결과: 스칼라 → 의미: 급격한 파라미터 변화를 억제).
학습·구현 세부사항: SAC 에이전트는 twin critics와 actor, learnable entropy temperature α_ent을 사용해 업데이트된다. critic 목표값은 y = r + γ(min_k Q'{Ck}(s',a') − α_ent log π(a'|s'))이며, critic은 MSE로, actor는 α_ent log π − min_k Q를 최소화해 갱신된다. 네트워크 아키텍처는 입력 선형투영(256) → 8개의 Transformer encoder block(hidden=512) → 4-layer MLP 출력 구조로, 각 네트워크 약 5M 파라미터이며 에이전트 전체 파라미터는 약 26.5M이다. 하이퍼파라미터 예: w_align=1, w_diversity=10, w_stability=10, r_stability 상한 5, N_ac=2. 데이터 샘플링: action a^{t+1}에 따라 도메인 D_i를 확률적으로 선택하고 UNIFORM(D_i)에서 미니배치를 샘플링해 인스턴스 수준 분포 P{a^t}를 구성한다.
관련 Figure

이 다이어그램은 에이전트가 상태를 관찰해 도메인 가중치를 결정하고, 해당 가중치로 샘플링한 배치로 LLM을 업데이트하며 보상이 계산되어 에이전트가 학습되는 전체 루프를 시각적으로 요약한다. 방법론(정식화된 MDP, 입력·출력 관계)을 이해하는 데 핵심적인 참고자료다.
HDS의 전체 프레임워크 다이어그램: SAC 에이전트(Actor/Critic), LLM 환경, 리플레이 버퍼, 데이터셋 간 상호작용을 보여준다.
주요 결과
메인 벤치마크 결과: Pythia-1B를 The Pile(50B tokens, 41,667 steps)로 학습한 실험에서 HDS는 validation perplexity 감소 곡선이 가장 가파르게 나타났다. HDS는 AC-ODM 대비 동일한 최종 perplexity를 약 44% 적은 스텝으로 달성했고, TPW 대비 동일한 perplexity 달성에 필요한 스텝을 약 57% 절감했다(실제 Steps: TPW 41667 → HDS 17917). 최종 스텝(41,667)에서 HDS는 TPW 대비 검증 perplexity 13.6% 감소, ODM 대비 8.9% 감소, AC-ODM 대비 5.3% 감소를 기록했다. Downstream MMLU 성능은 0-shot 0.26915, 5-shot 0.31064로 보고되었다(Table 1), 0-shot에서 AC-ODM(0.25146) 대비 절대 +0.01769 포인트 개선을 달성했다.
Ablation study 결과: 보상 성분 제거 실험에서 r_align 제거가 가장 큰 성능 저하를 초래했고(r_align은 핵심 요인), r_diversity 제거가 수렴 속도에 큰 영향을 미쳤다. r_stability 제거 시에도 성능 저하가 관찰되었으나 영향도는 상대적으로 작았다(그럼에도 통합 모델이 최저 perplexity를 기록).
효율성/속도 분석: 에이전트 통합에 따른 시간 오버헤드는 미미했다(시간/스텝 2.47s → 2.49s, <1% 증가). HDS는 TPW 기준 동일 perplexity 달성 시 총 학습 시간에서 2.21x의 speedup을 보였다(Table 3). 모델 스케일 실험(Pythia-12B, 25B tokens)에서도 HDS가 ODM 대비 모든 체크포인트에서 더 낮은 perplexity를 기록했고, 마지막 체크포인트(20,832 steps)에서 HDS 4.89 vs ODM 7.32로 약 33% 상대 개선을 보였다(Table 5).
관련 Figure

이 Figure는 HDS가 학습 초기부터 더 빠르게 perplexity를 낮추며 최종 단계까지 우위를 유지함을 보여준다. 그래프의 수치 지표(예: 동일 최종 perplexity 달성 시 필요한 스텝 감소)는 논문의 효율성 및 수렴 가속 주장을 직접 뒷받침한다.
HDS와 baselines(TPW, ODM, AC-ODM)의 validation perplexity(학습 스텝에 따른 감소)를 비교한 그래프.

이 Figure는 HDS가 약 20k 스텝 이후부터 baselines를 앞서며 downstream 성능 향상이 LM perplexity 개선으로 이어졌음을 보여준다. MMLU 0-shot의 수치(예: HDS 0.26915)는 테이블 결과와 일치한다.
학습 스텝에 따른 MMLU 0-shot 테스트 정확도 변화(각 방법의 성능 추이)를 나타내는 선그래프.

도메인별 성능 분해는 HDS가 특정 소규모 또는 전문 도메인에서 우수한 성능을 내며 전역적으로 균형 잡힌 개선을 달성했음을 보여 준다. 이 정보는 HDS가 일부 도메인에만 특화된 것이 아니라 전반적 데이터 분포에서 이득을 얻었음을 뒷받침한다.
The Pile의 22개 도메인별 테스트 perplexity를 각 방법별로 비교한 막대그래프.

이 Figure는 r_align 제거 시 성능 저하가 가장 크고, r_diversity가 조기 수렴에 중요하며 r_stability가 안정성에 기여함을 시각적으로 확인시킨다. 보상 성분들의 상보적 역할을 정량적으로 검증하는 근거다.
r_align, r_diversity, r_stability 보상 성분을 각각 제거한 경우와 전체 모델(HDS)의 validation perplexity 비교(ablations).

이 Figure는 HDS 정책의 탐색-수렴 과정을 보여 준다: 초기 탐색기 동안 가중치가 급변하고 이후 안정화되어 Book3 같은 고다양성 도메인으로 가중치가 증가하는 반면 Ubuntu IRC는 감소하는 등 'simple-to-complex' 커리큘럼 효과와 도메인별 역할 분화가 관찰된다.
학습 스텝에 따른 선택된 도메인들(Book3, Arxiv, Pile-CC, DM Mathematics, Ubuntu IRC)의 샘플링 가중치 변화(시간에 따른 정책 적응)를 보여주는 시계열 그래프.
기술 상세
아키텍처 구조: 에이전트(Actor/Critic)는 입력 선형 투영(dim=256) → 8개의 Transformer encoder block(hidden=512) → 4-layer MLP로 구성되어 상태의 고차원 상호의존성을 모델링한다. 각 네트워크는 약 5M 파라미터이며 전체 에이전트 파라미터는 약 26.5M이다. LLM은 Pythia 계열의 decoder-only Transformer(실험 주요: Pythia-1B, 추가 스케일 실험: Pythia-12B)이다.
핵심 메커니즘의 수학적/알고리즘적 기반: 문제는 MDP(𝒮, 𝒜, 𝒫, ℛ, γ)로 정식화된다. 상태 s^t = (n^t, t, l^t, Δl^t, ||ω^t||2, ||Δω^t||2)로 구성된다. 행동 a^t는 도메인 가중치 벡터 a ∈ Δ^K(확률 심플렉스). r_align 계산: 입력 g_i^t(도메인 i의 gradient)와 Σ{j≠i} g_j^t → 연산: 내적 ⟨g_i^t, Σ{j≠i} g_j^t⟩ → 결과: 스칼라 alignment score → 의미: 양수는 cross-domain 전이가 긍정적임을 나타낸다. r_diversity 계산: 입력 MTLD(B_i^t) → 연산: 정규화 MTLD_norm 및 t' = t/T_total 결합 → 결과: r_diversity = t'/(MTLD_norm + ε) → 의미: 시간 가중 커리큘럼. r_stability 계산: 입력 ||ω^t||_2 변화 → 연산: 역수 및 상한 적용 → 결과: 안정성 보상.
Prior work 대비 차별점: 기존 ODM은 도메인별 손실을 보상으로 사용하거나 MAB 관점으로 접근했고, AC-ODM는 gradient alignment를 도입했으나 단일 관점 보상이 주가 되었다. HDS는 세 가지 관점의 보상을 통합해 도메인 선택을 시간에 따라 조율하는 커리큘럼 효과와 교차 도메인 시너지를 동시에 취한다.
구현 및 학습 세부사항: 트레이닝 하드웨어는 Intel Xeon Platinum 8468 CPU와 NVIDIA H800 80GB GPU 8대. 시퀀스 길이 1024, 글로벌 배치 1152(micro-batch 8 per GPU, grad accumulation 18). Warm-up 833 steps 동안 The Pile 원래 가중치를 소음 N(0,0.02)로 사용 후 HDS 활성화. SAC 하이퍼파라미터: w_align=1, w_diversity=10, w_stability=10, r_stability cap=5, N_ac=2. 액터·크리틱 업데이트는 SAC 표준 손실을 사용하며 타깃 네트워크는 Polyak averaging으로 갱신된다.
스케일 가이드: 에이전트 파라미터는 LLM 파라미터의 약 0.3%~1.5% 범위 권장(≈0.5%에서 비용 대비 성능 균형이 우수). r_align 계산 시 더 깊은(후반) 레이어 선택이 약간 더 유리한 것으로 보고되었다.
관련 Figure

이 도식은 상태 인코딩 방식과 actor/critic이 어떻게 Transformer 기반 블록을 사용해 고차원 상태를 처리하는지, 그리고 액션(도메인 가중치)과 LLM 가중치 정보가 네트워크로 유입되는 방식을 명확히 한다. 네트워크 설계(선형 투영 256 → 8 Transformer → 4-layer MLP)의 근거를 연결해 준다.
Actor와 Critic 네트워크 아키텍처(입력 처리, Transformer encoder blocks, concatenation 후 MLP)를 단순화해 보여주는 도식.
실무 활용
HDS는 학습 중 데이터 샘플링 비율을 실시간으로 조정해 학습 효율을 개선하므로, 대규모 LLM 사전학습 또는 재학습 파이프라인에서 비용과 시간을 절감하면서 성능을 향상시킬 수 있다. 에이전트 파라미터가 LLM 파라미터의 약 0.3–1.5% 범위이면 실무적 오버헤드는 작다.
- 대규모 LLM 사전학습(예: The Pile 기반)에서 데이터 믹싱 정책을 자동화하여 총 토큰/스텝 절감
- 도메인 혼합이 중요한 파이프라인에서 특정 도메인의 과적합이나 편향을 억제하면서 전반적 일반화 성능 개선
- 사전학습 비용이 큰 환경에서 연산·전력 소비 절감을 목적으로 한 학습 효율화
코드 공개 여부: 공개
코드 저장소 보기키워드
용어 해설
- Multi-Objective Reinforcement Learning
- — 여러 목표(보상)를 동시에 최적화하기 위해 강화학습 문제를 구성하는 접근법이다. 이 논문에서는 data-quality, inter-domain influence, model-stability 세 가지 보상 성분을 가중합하여 에이전트가 데이터 믹싱 정책을 학습하도록 한다. 다목적 보상은 서로 상충하는 신호들 사이의 균형을 찾아 전반적 학습 효율을 개선한다.
- Gradient Alignment
- — 한 도메인에서 계산된 gradient와 다른 도메인들의 gradient 합과의 내적을 계산하여 양(positive) 혹은 음(negative) 영향력을 평가하는 지표다. 입력: 도메인별 gradient 벡터 → 연산: 내적(⟨g_i, Σ_{j≠i} g_j⟩) → 결과: 스칼라값 → 의미: 양수면 해당 도메인 업데이트가 다른 도메인 학습에 유리하다는 신호다.
- MTLD
- — 문서의 어휘 다양성(텍스트 복잡도)을 수치화한 지표다. 논문에서는 배치별 MTLD를 [0,1] 범위로 정규화한 후 시간 스텝 t에 따라 보상으로 사용해 'simple-to-complex' 커리큘럼을 형성한다. 입력: 배치 텍스트 → 연산: MTLD 계산 및 정규화 → 결과: 정규화 점수 → 의미: 낮은 값은 단순한 텍스트, 높은 값은 복잡한 텍스트를 의미한다.
- Probability Simplex
- — 도메인 가중치 벡터 a∈Δ^K가 속하는 공간으로, 모든 요소가 0 이상이며 합이 1이 되는 K차원 공간이다. 이 논문에서 에이전트의 Action은 이 심플렉스 상의 샘플링 확률 분포이다. 입력: 실수 벡터 → 연산: 비음수 및 합=1 제약 → 결과: 도메인 샘플링 확률 → 의미: 다음 배치에서 각 도메인이 선택될 확률을 규정한다.
코드 예제
Input : Corpus D = { D1, … , DK }, hyperparameters
Initialize : Actor π_{θ_A}, critics Q_{θ_{C1},C2}, target critics θ' ← θ, replay buffer B, Na_c, entropy temperature α_ent, LLM θ_M.
for t = 1, …, T do
Observe state s^t from the LLM environment;
Sample action a^t;
Sample data batch B^t according to P_{a^{t+1}};
Update LLM parameters: θ_M^{t+1} ← θ_M^t − η_M ∑_{i=1}^K a_i^{t+1} ∇ ℒ(θ_M^t, B_i^t);
Compute reward vector r^t using Eq. (1)-(5);
Observe next state s^{t+1};
Store transition (s^t, a^t, r^t, s^{t+1}) in B;
for n_ac = 1, …, N_ac do
Sample a minibatch {(s^j, a^j, r^j, s^{j+1})} from B;
Compute target y^j for each transition;
Update critics by minimizing L(θ_{C1}) and L(θ_{C2});
Update actor by minimizing L(θ_A);
Update entropy temperature by minimizing L(α_ent);
Update target networks using Polyak averaging;
end
end
Algorithm 1 The Holistic Data Scheduler (HDS)HDS의 의사코드로, 에이전트가 상태를 관찰해 도메인 가중치(action)를 샘플링하고 LLM을 업데이트한 뒤 리플레이 버퍼로부터 에이전트 파라미터를 갱신하는 전체 루프를 보여준다.
AI 요약 · 북마크 · 개인 피드 설정 — 무료
출처 · 인용 안내
인용 시 "요약 출처: AI Trends (aitrends.kr)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.