LLM 강화학습의 학습-추론 엔진 불일치를 log-odds 변위에서 교정하는 CIS
Rethinking Training-Inference Mismatch in LLM Reinforcement Learning: Where It Arises and How to Correct It
무엇인가
RLVR(검증 가능한 보상 기반 강화학습)로 LLM의 추론 능력을 올리는 파이프라인은 처리량을 위해 롤아웃 생성과 그래디언트 계산을 분리한다. 롤아웃은 vLLM·SGLang 같은 추론 엔진이, 그래디언트는 FSDP·Megatron 같은 학습 엔진이 담당한다. 두 엔진은 같은 파라미터를 로드하지만 커널 구현, reduction 순서, 수치 정밀도가 달라 같은 토큰에 다른 확률을 부여한다. 명목상 on-policy인 학습이 실제로는 off-policy가 되는 것이다. 이 논문이 겨냥하는 지점은 이 불일치가 특히 MoE 모델에서 심각하다는 사실이다. 게이팅 네트워크의 top-k 선택은 이산 결정이라 작은 수치 차이만으로 두 엔진이 서로 다른 expert를 활성화하고, 그 차이가 이후 모든 레이어를 통과하며 누적된다. 저자들은 이것이 학습 불안정이나 붕괴로 이어진다고 본다.
어떻게 동작하나
표준적 대응은 중요도 샘플링이다. 학습 엔진 확률 p와 추론 엔진 확률 q의 비 k=p/q로 각 토큰의 그래디언트를 재가중한다. 문제는 두 엔진이 특정 토큰에서 날카롭게 어긋날 때 k가 매우 커진다는 점이다. 정확 보정(exact ratio)은 편향이 없지만 분산이 통제 불가능해진다. 기존 방법들은 그래서 k 자체에 규칙을 건다. TIS는 k를 고정 상한 C로 자르고, IcePop은 고정 구간 밖의 토큰을 마스킹하며, KPop은 두 엔진 확률의 binary KL이 임계값을 넘는 토큰을 마스킹한다. 시퀀스 레벨에서는 응답 전체의 비율을 자르거나(Seq-TIS, Seq-MIS) 시퀀스 목적함수로 바꾸고(GSPO), 수치 정밀도를 FP16으로 올리는 접근도 있다. 논문의 비판은 이 규칙들이 모두 k 또는 확률에서 직접 계산한 양에 작용한다는 것이다. k는 두 가지를 뒤섞는다. 엔진 간 실제 불일치와 토큰 자체의 신뢰도다. 같은 크기의 불일치도 신뢰도가 낮은 토큰에서는 비율을 1에서 멀리 밀어내지만 신뢰도가 높은 토큰에서는 거의 움직이지 않는다. 결과적으로 고정 임계값은 주로 저신뢰 토큰에 작용하고, 절단 편향이 그쪽에 몰린다.
무엇과 다른가
CIS(Calibrated Importance Sampling)의 출발점은 Lemma 3.1의 비율-변위 항등식이다. k_t = p_t + (1-p_t)·exp(ε_t), k_t - 1 = (1-p_t)(exp(ε_t)-1)이며, 여기서 ε_t = -δ_{y_t} + log Σ_{j≠y_t} w_j exp(δ_j)다. δ는 softmax 이전 logit 벡터의 per-logit 섭동이고, ε_t는 두 엔진이 샘플된 토큰에 부여한 log-odds의 차이, 즉 logit p_t - logit q_t와 정확히 같다. 즉 불일치는 log-odds 좌표에서 순수한 덧셈 변위로 나타나며, 두 엔진이 기록한 log-prob만으로 직접 계산된다. 실측에서 dense 모델의 δ는 bf16 반올림 오차 수준으로 0 근처에 집중되지만, MoE에서는 몸통은 여전히 0 근처인데 꼬리가 두드러지게 두껍다. 반올림 오차는 몸통만 설명하고, 두 엔진의 라우팅 불일치가 꼬리를 증폭하는 주범이며, 서로 다른 expert를 선택하는 레이어가 많을수록 꼬리가 두꺼워진다. 이 꼬리는 ε_t로 그대로 전이된다.
어떻게 쓰나
이론적 동기는 명확하다. Theorem 3.2에 따르면 정확 보정 추정량의 MSE 상한은 E[e^{2ε_t}] 항에 지배된다. 학습된 MoE에서 측정한 raw moment E[e^{αε}]는 α=1에서 7.07이지만 α=2에서 9.2×10^6까지 치솟고, (1-p_t)^α로 압축된 가중 모멘트도 0.16에서 22.9로 뛴다. 보정하지 않은 경우의 기대값이 1이므로 상한이 20배 이상 커지는 셈이다. 중요한 관찰은 분산 문제가 본질적으로 단방향이라는 점이다. exp(ε_t) ≤ 1이면 p_t < k_t ≤ 1이라 k_t^2 ≤ 1로 유계다. exp(ε_t) > 1일 때만 비율이 1을 넘어 무한정 커질 수 있다. 그래서 CIS는 큰 양의 변위만 자른다. 하방을 건드리는 것은 분산의 원인을 해결하지 못한 채 편향만 더한다. 절단 위치는 실측이 정한다. ε_t를 1-p_t에 대해 그려보면 불확실성이 6자릿수에 걸쳐 변하는 동안 중앙값은 0에 머물고 퍼짐은 2배 이내로 변하며 Spearman 상관은 -0.028에 불과하다. 같은 크기의 변위는 신뢰도와 무관하게 비슷하게 이상하다는 뜻이므로, 변위 좌표에 단일 상수 임계값 exp(ε_t) ≤ 1+λ를 놓는다. 이를 비율 공간으로 되돌리면 f_λ(k_t, p_t) = min{k_t, 1+λ(1-p_t)}가 된다. 상한에 붙은 (1-p_t)는 설계 선택이 아니라 항등식이 변위를 비율로 옮기는 계수다. 결과적으로 신뢰도가 높을수록 상한이 더 타이트해진다. 비교하자면 고정 비율 상한 k_t ≤ C는 exp(ε_t) ≤ 1+(C-1)/(1-p_t)에 대응하므로, 신뢰도가 커질수록 변위 임계값이 점점 관대해진다. Theorem 3.3은 CIS가 무계 second moment를 (1+λ)^2로 유계인 항으로 바꾸고, 그 대가로 잘린 초과분이 제어하는 편향을 낸다는 것을 보인다. λ→∞면 정확 보정으로 복귀한다.
전제와 한계
구현에는 실무적인 함정이 하나 있다. 1-p_t가 아주 작아지면 상한 1+λ(1-p_t)가 1에 가까워져 log-prob 저장 해상도보다 좁아진다. 그러면 절단이 실제 불일치가 아니라 반올림 오차에 작용한다. 저자들은 1-p_t를 저장 해상도 κ=5×10^-3으로 floor 처리한다(φ_t = max(1-p_t, κ)). 이 처리를 빼면 평균 정확도가 34.78에서 29.02로 떨어진다. CIS는 elementwise 패스 하나만 추가할 뿐 순전파·역전파 계산을 늘리지 않으며, 추론 측 log-prob은 샘플링 시점에 기록한 raw 값이어야 하고 가중치는 detach해야 한다.
실험은 MoE 모델 3종(Qwen1.5-MoE-A2.7B-Chat, DeepSeek-V2-Lite, Qwen3-30B-A3B)에서 GSM8K로 학습하고, GSM8K·MATH500·SVAMP·Minerva-Math·OlympiadBench 5개 수학 추론 벤치마크에서 greedy pass@1로 평가했다. 학습 시드 3개 평균이다. Qwen1.5-MoE에서 CIS는 5개 벤치마크 평균 34.78로, 보정 없음 30.99, Exact Ratio 32.48, KPop† 33.50, IcePop 34.18, TIS 31.40을 모두 앞섰다. 학습 도메인 GSM8K에서는 67.88을 기록했다. DeepSeek-V2-Lite에서는 37.16, Qwen3-30B-A3B에서는 69.88로 두 블록 모두 최고 평균이다. 다만 Qwen3-30B-A3B는 베이스 모델이 이미 GSM8K 95.22%에 도달해 학습 과제에서 남은 여지가 거의 없고, 저자들 스스로 IcePop 대비 일부 전이 벤치마크의 마진이 시드 간 변동보다 작아 개별 열을 통계적으로 분리된 것으로 해석하지 않는다고 밝힌다. RL 단계에서 쓰지 않은 지식 3개·코드 2개 벤치마크에서도 CIS가 통합 정확도 최고였지만 IcePop과의 차이는 시드 퍼짐 이내다.
진단 분석은 설계 선택의 근거를 수치로 뒷받침한다. 학습된 MoE의 샘플 토큰에서 k_t > 1인 토큰이 Var(k_t)의 99.96%를 차지하고 k_t < 1은 0.04%만 기여한다. 하방까지 함께 자르는 양방향 가중(w_t = max{min{k_t, 1+λφ_t}, (1+p_t)/2})은 평균 정확도를 34.78에서 6.01로 떨어뜨린다. 신뢰도 의존 상한을 고정 TIS 상한으로 바꾸면 31.40, CIS와 비슷한 비율의 토큰(4.3% 대 3.8%)을 자르는 C=1.13을 쓰면 8.88로 무너진다. 같은 양을 잘라도 고정 상한이 그것을 저신뢰 토큰에 몰아넣으면 학습이 망가진다는 뜻이다. λ=0(TIS C=1에 해당)은 8.75로 붕괴한다. 편향이 어디로 가는지도 측정됐다. p_t < 0.3 구간에서 CIS가 유발하는 편향은 TIS의 절반 이하이고, 대신 고신뢰 구간에서 더 많은 편향을 진다. 전체 편향은 0.0064로 TIS의 0.0078보다 낮고, 추가 분산은 0.0081로 Exact Ratio의 0.0132보다 훨씬 작다. 하이퍼파라미터 민감도는 관대하다. λ는 기본값 2.3에서 34.78로 최고였고, 양수인 모든 임계값이 보정 없음(30.99)보다 높았으며 그 범위는 32.53~34.08이었다. floor는 κ=5×10^-3과 2×10^-2에서 34.78, 34.39였고, 10^-3은 1.4점 손실, 제거 시 29.02, 0.1로 올리면 32.17이었다.
개발자 관점에서 이 논문의 실용적 메시지는 세 가지다. 첫째, 분리형 RL 파이프라인에서 학습-추론 불일치는 선택이 아니라 구조적 조건이고, MoE에서는 라우팅 불일치 때문에 꼬리가 두꺼워져 정확 보정이 위험해진다. 둘째, 기존 TIS/IcePop식 고정 임계값은 편향을 저신뢰 토큰에 몰아넣으므로, 같은 절단률이라도 어디를 자르는지가 성능을 가른다. CIS는 log-prob 두 개의 차이만 있으면 되고 elementwise 한 번이라 도입 비용이 낮다. 셋째, 도입 시 반드시 확인할 것은 추론 측 log-prob이 샘플링 시점의 raw 값인지, 가중치가 detach됐는지, 그리고 1-p_t floor가 저장 정밀도에 맞게 설정됐는지다. floor를 빼면 정확도가 보정 없음보다도 낮아진다.
저자들이 밝힌 전제와 한계도 분명하다. 이론 분석은 ‖ℓ'_t‖₂ ≤ b와 토큰 기여의 독립성을 가정하며, 상관이 있으면 n을 n_eff로 바꿔야 한다고 명시한다. 또한 보정 대상인 k_t는 관측된 prefix가 주어진 조건부 next-token 분포만 교정하고 prefix 자체의 분포는 재가중하지 않는다. k_t가 θ에 의존하지 않는다는 가정 위에서 그래디언트 추정량을 세운다. λ는 하이퍼파라미터로 남아 있고, floor κ도 이론에서 유도된 값이 아니라 저장 해상도에 맞춘 실무적 장치다. 실험은 MoE 3종과 수학 추론 벤치마크에 한정되며, 개별 벤치마크 마진이 시드 변동보다 작은 경우가 있어 저자들 스스로 통계적 분리로 해석하지 않는다.