수학

MATH / 중급 58번

중요도 비율과 클리핑: 예전 정책의 샘플을 재사용하는 대가

PPO 로그의 clip_frac과 approx_kl은 왜 나란히 찍힐까요. 중요도 표집이 왜 불편추정량인지, 두 분포가 멀어질 때 유효 표본 수가 왜 지수적으로 무너지는지를 e^(Δ²) 계산으로 확인하고, clip이 목적함수를 조각별 선형으로 바꿔 한 걸음을 막는 방식까지 따라갑니다.

PALDYN Team17 MIN READ

RLHF 학습을 돌리면 로그에 이런 줄이 찍힙니다.

step 420 | policy_loss -0.0231 | clip_frac 0.18 | approx_kl 0.0094 | ratio_max 1.41

clip_frac은 이번 미니배치에서 클리핑에 걸린 토큰의 비율이고, approx_kl은 지금 정책이 이 배치를 만든 정책에서 얼마나 멀어졌는지입니다. 실무에서는 이 두 숫자가 각각 0.1~0.3과 0.01 근처에 머무르는지를 봅니다. 넘으면 학습률을 내리거나 배치 재사용 횟수를 줄입니다.

왜 하필 이 두 숫자일까요. 손실값도 보상도 아닌 이것을 보는 이유는, 이 둘이 한 가지 거래의 청구서이기 때문입니다. 거래의 내용은 이렇습니다 — 답을 뽑는 일은 비싸니 한 번 뽑은 배치로 여러 걸음을 걷고 싶다. 그러면 두 번째 걸음부터는 지금 정책이 아니라 예전 정책이 뽑은 샘플로 지금 정책을 평가하게 됩니다.

이 글은 그 대가가 정확히 얼마인지를 계산합니다. 지난 글에서 KL 제약이 붙은 목적식을 닫힌 형태로 풀었다면, 여기서는 그 제약을 강제하는 두 장치 — 비율 클리핑과 KL 조기 중단 — 이 무엇을 막고 있는지를 봅니다. 알고리즘 전체의 구현은 PPO가 다루고, 이 글은 추정량의 성질만 맡습니다.

다른 분포에서 뽑은 샘플로 기댓값을 구하기

알고 싶은 것은 지금 정책 pp 아래의 기댓값입니다.

J(p)=Ex∼p[f(x)]J(p) = \mathbb{E}_{x \sim p}[f(x)]

그런데 가진 샘플은 예전 정책 qq 에서 나왔습니다. 다시 뽑을 수 없다면 방법은 하나뿐입니다. 적분 안에 qq 를 곱하고 나눕니다.

Ex∼p[f(x)]=∫p(x)f(x) dx=∫q(x)p(x)q(x)f(x) dx=Ex∼q ⁣[p(x)q(x)f(x)]\begin{aligned} \mathbb{E}_{x \sim p}[f(x)] &= \int p(x) f(x)\,dx \\ &= \int q(x) \frac{p(x)}{q(x)} f(x)\,dx \\ &= \mathbb{E}_{x \sim q}\!\left[\frac{p(x)}{q(x)} f(x)\right] \end{aligned}

중요도 표집(importance sampling)은 이렇게 다른 분포에서 뽑은 샘플에 비율을 곱해 원하는 분포의 기댓값을 얻는 방법입니다. 그 비율 w(x)=p(x)/q(x)w(x) = p(x)/q(x) 를 중요도 비율이라 부릅니다.

예전 분포에서 뽑은 표본에 무게를 곱해 새 분포의 기댓값을 얻는 그림

식 변형이 한 줄이므로 성질도 곧바로 나옵니다. J^=1N∑iw(xi)f(xi)\hat{J} = \frac{1}{N}\sum_i w(x_i) f(x_i) 라 두면

Eq[J^]=Ex∼q ⁣[p(x)q(x)f(x)]=J(p)\mathbb{E}_{q}[\hat{J}] = \mathbb{E}_{x \sim q}\!\left[\frac{p(x)}{q(x)} f(x)\right] = J(p)

입니다. 표본 수가 몇이든 기댓값이 참값과 같은 추정량을 불편추정량(unbiased estimator)이라 하고, 중요도 표집 추정량은 불편추정량입니다. q(x)>0q(x) > 0 인 곳에서 p(x)>0p(x) > 0 이기만 하면 그렇습니다.

여기까지만 보면 공짜입니다. 편향이 없으니 샘플을 재사용해도 되는 것처럼 보입니다. 대가는 분산에 있습니다.

무게의 분산은 거리의 제곱의 지수로 자란다

비율 자체의 분산을 봅니다. Eq[w]=∫q⋅pq=1\mathbb{E}_q[w] = \int q \cdot \frac{p}{q} = 1 이므로

Varq(w)=Eq[w2]−1=∫p(x)2q(x)dx−1\mathrm{Var}_q(w) = \mathbb{E}_q[w^2] - 1 = \int \frac{p(x)^2}{q(x)}dx - 1

적분 안이 p2/qp^2/q 라는 것이 핵심입니다. qq 가 분모에 있으므로 qq 가 얇은데 pp 가 두꺼운 자리가 하나라도 있으면 이 값이 크게 튑니다.

정규분포 두 개로 계산하면 정확한 값이 나옵니다. q=N(0,1)q = \mathcal{N}(0, 1), p=N(Δ,1)p = \mathcal{N}(\Delta, 1) 이라 두면 지수를 정리해서

Eq[w2]=eΔ2,Varq(w)=eΔ2−1\mathbb{E}_q[w^2] = e^{\Delta^2}, \qquad \mathrm{Var}_q(w) = e^{\Delta^2} - 1

을 얻습니다. 한편 두 분포의 KL 발산은 KL(p∥q)=Δ2/2\mathrm{KL}(p\|q) = \Delta^2/2 입니다. KL은 Δ2\Delta^2 에 비례해 자라는데 무게의 분산은 그 지수로 자랍니다.

이 차이를 눈에 보이게 하는 값이 유효 표본 수(effective sample size, ESS)입니다. 무게가 제각각인 NN 개의 표본이 실제로는 몇 개의 균등한 표본만큼 일하는지를 재는 값으로, 이렇게 정의합니다.

ESS=(∑iwi)2∑iwi2\mathrm{ESS} = \frac{\left(\sum_i w_i\right)^2}{\sum_i w_i^2}

무게가 전부 같으면 ESS=N\mathrm{ESS} = N 이고, 한 표본이 무게를 독차지하면 ESS→1\mathrm{ESS} \to 1 입니다. 위의 정규분포 예에서는 ESS/N→e−Δ2\mathrm{ESS}/N \to e^{-\Delta^2} 로 갑니다.

평균 차이에 따라 유효 표본 수가 무너지는 표

숫자로 확인해 봅니다. 표본 20만 개로 잰 값입니다.

import numpy as np
np.random.seed(1)
N = 200_000
for D in [0.5, 1.0, 2.0, 3.0]:
    x = np.random.randn(N)              # q = N(0,1)에서 뽑는다
    w = np.exp(D * x - D**2 / 2)        # p = N(D,1)에 대한 비율
    ess = w.sum()**2 / (w**2).sum() / N
    print(f"D={D}  KL={D*D/2:.3f}  Var(w)={w.var():9.2f}  ESS/N={ess:.5f}")
D=0.5  KL=0.125  Var(w)=     0.28  ESS/N=0.77938
D=1.0  KL=0.500  Var(w)=     1.70  ESS/N=0.37055
D=2.0  KL=2.000  Var(w)=    61.25  ESS/N=0.01637
D=3.0  KL=4.500  Var(w)=  1378.60  ESS/N=0.00071

KL이 0.5일 때 이미 표본의 63%가 사라집니다. KL이 2가 되면 1,000개를 뽑아도 실제로 일하는 것은 열여섯 개 남짓입니다(이론값은 열여덟 개).

그리고 마지막 줄에 함정이 하나 더 있습니다. Δ=3\Delta = 3 에서 이론값은 Var(w)=e9−1≈8,102\mathrm{Var}(w) = e^9 - 1 \approx 8{,}102 인데 표본에서 잰 값은 1,379입니다. 분산이 큰 상황에서는 분산 추정값 자체가 아래로 치우칩니다. 분산을 지배하는 것은 아주 드물게 나오는 거대한 무게인데 20만 개 안에 그것이 안 들어왔기 때문입니다. 실제 학습에서 이것은 "지표가 조용한데 학습이 무너진다"로 나타납니다.

기댓값 추정의 정확도로 바꿔 보면 이렇습니다. 표본 1,000개로 Ep[x]=Δ\mathbb{E}_p[x] = \Delta 를 추정할 때

평균 차이 중요도 표집 표준편차 직접 뽑았을 때
Δ=1\Delta = 1 0.117 0.032
Δ=2\Delta = 2 0.912 0.032
Δ=3\Delta = 3 5.99 0.032

편향은 없습니다. 다만 Δ=3\Delta = 3 에서 추정값의 흔들림이 참값(3)의 두 배입니다. 불편하다는 말은 무한히 반복하면 맞는다는 뜻이지, 한 번의 미니배치가 쓸 만하다는 뜻이 아닙니다.

한 걸음을 어떻게 막을 것인가

정책 경사법의 목적식을 비율로 다시 씁니다. 이점 함수 AA 를 붙이면

LIS(θ)=Eq ⁣[r(θ) A],r(θ)=πθ(a∣s)πold(a∣s)L^{\text{IS}}(\theta) = \mathbb{E}_{q}\!\left[r(\theta)\,A\right], \qquad r(\theta) = \frac{\pi_\theta(a \mid s)}{\pi_{\text{old}}(a \mid s)}

입니다. 이 식을 그냥 최대화하면 앞 절이 말한 자리로 걸어 들어갑니다 — A>0A > 0 인 행동의 rr 을 끝없이 키우는 것이 목적식을 계속 키우기 때문입니다. 그런데 그 방향으로 멀리 갈수록 rr 이 근거하는 추정 자체가 못 믿을 것이 됩니다.

신뢰 영역(trust region)은 근사식을 믿을 수 있는 범위를 정해 두고 한 걸음을 그 안에 가두는 발상입니다. 원래는 KL에 명시적인 제약을 걸어 풀었는데, 그러려면 매 걸음 제약 있는 최적화를 풀어야 합니다. PPO의 선택은 제약을 푸는 대신 목적함수의 모양을 바꾸는 것입니다.

LCLIP(θ)=E[min⁡(r(θ)A,  clip(r(θ),1−ϵ,1+ϵ) A)]L^{\text{CLIP}}(\theta) = \mathbb{E}\Big[\min\big(r(\theta)A,\; \mathrm{clip}(r(\theta), 1-\epsilon, 1+\epsilon)\,A\big)\Big]

클리핑(clipping)은 값을 정해진 구간 밖으로 못 나가게 잘라 내는 것이고, clip(r,1−ϵ,1+ϵ)\mathrm{clip}(r, 1-\epsilon, 1+\epsilon) 은 rr 을 [0.8,1.2][0.8, 1.2] 같은 띠 안에 가둡니다. 여기에 min⁡\min 이 하나 더 붙어 있는 것이 이 식의 전부인데, 그 min⁡\min 이 하는 일은 부호에 따라 갈립니다.

A의 부호에 따라 어느 쪽이 평평해지는지 보이는 두 패널

A>0A > 0 일 때 목적값을 rr 의 함수로 적어 보면

rr 0.5 0.7 0.9 1.1 1.2 1.4 1.8
목적값 0.50 0.70 0.90 1.10 1.20 1.20 1.20
기울기 1 1 1 1 — 0 0

r>1.2r > 1.2 에서 평평합니다. 좋았던 행동의 확률을 이미 20% 넘게 올렸으면 더 올리라는 힘이 사라집니다. 반대로 아래쪽은 자르지 않습니다 — rr 이 0.5로 내려가 있으면 되돌아오라는 힘이 그대로 살아 있어야 하니까요.

A<0A < 0 이면 평평해지는 쪽이 뒤집힙니다.

rr 0.5 0.7 0.8 1.0 1.4 1.8
목적값 −0.80 −0.80 −0.80 −1.00 −1.40 −1.80
기울기 0 0 — −1 −1 −1

나빴던 행동을 이미 20% 넘게 줄였으면 더 줄이라는 힘이 없어지고, 위쪽은 열려 있습니다. 두 경우 모두 한 걸음에 밀어낼 수 있는 거리를 목적함수의 기울기로 막습니다. 제약을 푸는 대신 그래디언트를 0으로 만드는 방식이라 구현이 몇 줄로 끝납니다.

ratio = torch.exp(logp_new - logp_old)          # r(θ)
unclipped = ratio * adv
clipped = torch.clamp(ratio, 1 - eps, 1 + eps) * adv
loss = -torch.min(unclipped, clipped).mean()

clip_frac은 이 두 항 중 clipped 쪽이 선택된 비율입니다. 그 값이 0이면 띠가 너무 넓어 아무것도 막지 못하는 것이고, 0.5를 넘으면 절반이 기울기 0인 자리에 있다는 뜻이라 배치의 절반이 학습에 기여하지 못합니다.

clip이 못 막는 것

여기서 두 번째 지표가 필요해집니다. clip은 한 걸음의 크기를 막지 걸음의 횟수를 막지 못합니다.

같은 배치를 K번 다시 쓰는 동안 rr 은 매번 다시 계산됩니다. 한 번 걸을 때마다 띠 안쪽에 머물러도, 걸음이 쌓이면 정책은 계속 밀려 나갑니다. 게다가 클리핑에 걸린 토큰은 그래디언트가 0이지만 걸리지 않은 나머지가 계속 정책을 옮기므로, 걸린 토큰들의 rr 은 다음 반복에서 더 멀어져 있습니다.

한 배치를 여러 번 재사용하는 루프와 조기 중단 위치

그래서 반복마다 KL을 재고 예산을 넘으면 남은 반복을 버립니다. 이것이 KL 조기 중단(early stopping)입니다. approx_kl은 그 KL의 값싼 추정값으로, 보통 E[(r−1)−log⁡r]\mathbb{E}[(r - 1) - \log r] 을 씁니다. 이 양은 항상 0 이상이고 r≈1r \approx 1 근처에서 KL\mathrm{KL} 과 이차 근사로 일치합니다.

두 장치가 막으려는 붕괴가 정확히 무엇인지 이제 말할 수 있습니다. 비율이 커지는 것 자체가 문제가 아니라, 비율이 커진 만큼 그 비율을 곱해서 만든 그래디언트 추정이 못 믿을 것이 된다는 것이 문제입니다. 앞 절의 숫자로 하면 KL이 0.01일 때는 유효 표본이 거의 그대로지만 KL이 0.5로 가면 3분의 1로 줍니다. 표본이 그만큼 줄었는데도 걸음 크기는 그대로라면, 정책은 잡음을 신호로 알고 한쪽으로 달려갑니다. 그 뒤로는 뽑히는 답이 한 가지로 몰리고 보상은 올라가는데 문장은 무너지는, 흔히 보상 해킹으로 부르는 모양이 됩니다.

ratio_max가 1.41이었던 앞의 로그로 돌아가 봅니다. ϵ=0.2\epsilon = 0.2 이면 그 토큰은 클리핑에 걸려 기울기가 0입니다. clip_frac 0.18은 배치의 18%가 그 상태라는 뜻이고, approx_kl 0.0094는 아직 예산 0.01 안쪽이라 이번 반복을 마저 돌려도 된다는 뜻입니다. 세 숫자가 한 문장이 됩니다 — 몇몇 토큰은 이미 띠 밖으로 나갔지만 정책 전체는 아직 표본이 떠받칠 수 있는 거리 안에 있다.

정리

  • 중요도 표집은 Ep[f]=Eq[pqf]\mathbb{E}_p[f] = \mathbb{E}_q[\frac{p}{q}f] 라는 한 줄의 항등식이고, 표본 수와 무관하게 불편추정량이다. q>0q > 0 인 곳에서만 p>0p > 0 이면 된다.
  • 대가는 분산이다. Varq(w)=∫p2/q−1\mathrm{Var}_q(w) = \int p^2/q - 1 이고, 단위 분산 정규분포 두 개가 Δ\Delta 만큼 떨어져 있으면 정확히 eΔ2−1e^{\Delta^2} - 1 이다.
  • 유효 표본 수 (∑w)2/∑w2(\sum w)^2 / \sum w^2 는 그 분산을 표본 개수로 번역한 값이다. 같은 예에서 ESS/N=e−Δ2\mathrm{ESS}/N = e^{-\Delta^2} 로, KL이 2일 때 1,000개 중 열여덟 개만 남았다.
  • KL은 Δ2/2\Delta^2/2 로 자라는데 분산은 eΔ2e^{\Delta^2} 로 자란다. KL을 조금 넘겼다는 말과 분산이 조금 늘었다는 말은 같은 뜻이 아니다.
  • 분산이 큰 구간에서는 표본으로 잰 분산이 아래로 치우친다 — Δ=3\Delta = 3 에서 이론값 8,102에 대해 20만 표본이 1,379를 내놨다. 지표가 조용해도 안전하다는 근거가 되지 못한다.
  • clip은 목적함수를 조각별 선형으로 바꿔 한 걸음의 크기를 막는다. A>0A > 0 이면 위쪽이, A<0A < 0 이면 아래쪽이 기울기 0이 되고, 되돌아오는 방향은 언제나 열려 있다.
  • KL 조기 중단은 걸음의 횟수를 막는다. clip은 매 걸음 안쪽에 머물러도 걸음이 쌓여 밀려 나가는 것을 못 잡기 때문이다.
  • 두 장치가 막는 붕괴는 "비율이 커진 것"이 아니라 "그 비율로 만든 그래디언트 추정을 떠받칠 표본이 남아 있지 않은 것"이다.

읽어주셔서 감사합니다. 😊

LATEST

수학의 최신 글

수학2026.09.07

양자화 오차: 격자 사상, 오차 분산, 이상치 채널

실수를 2^b개 격자에 사상할 때 오차의 분산이 왜 Δ²/12인지 유도하고, 그것이 비트당 6.02dB라는 SNR로 번역되는 과정을 실측과 대조했습니다. 이상치 하나가 나머지 값의 유효 비트를 어떻게 먹는지, 그리고 int4에서 성능이 무너지는 지점을 오차 예산으로 미리 계산하는 법까지.

중급18 MIN
수학2026.09.07

수치적으로 안정한 계산 패턴 모음

최댓값 빼기, 로그 공간, log1p·expm1, 분산의 두 공식, 정규화의 ε, fp32 누산, 역행렬 대신 solve — 프레임워크가 몰래 해 주는 일곱 가지를 하나씩 꺼내 각각 어떤 고장을 막는지 직접 재 봤습니다. 수식을 그대로 옮긴 코드가 왜 라이브러리보다 나쁜지에 대한 목록입니다.

중급22 MIN
수학2026.09.07

부동소수점은 어디서 새는가: 반올림, 상쇄, 더하는 순서

0.1 + 0.2가 0.3이 아닌 이유부터 시작해 머신 엡실론을 유도하고, 같은 16비트인데 fp16과 bf16이 서로 다른 지점에서 터지는 이유, 비슷한 수를 뺄 때 유효자리가 사라지는 파괴적 상쇄, 그리고 1,000만 개를 순서만 바꿔 더했을 때 오차가 백만 배 갈리는 실험까지 직접 재 봤습니다.

중급23 MIN