수학

MATH / 중급 59번

재매개변수화 트릭과 폐형 forward 과정

x = μ + σz라는 한 줄이 표집을 미분 가능하게 만듭니다. 이 추정량이 로그 미분 트릭보다 왜 분산이 작은지를 차원별로 재 보고, 가우시안을 t번 더한 결과가 다시 가우시안이라는 성질로 q(x_t|x_0)의 폐형을 유도해 임의의 t에서 학습이 가능해지는 이유까지 따라갑니다.

PALDYN Team28 MIN READ

확산 모델의 학습 루프는 이렇게 생겼습니다.

t = torch.randint(0, T, (batch,), device=x0.device)   # 스텝을 아무거나 하나
noise = torch.randn_like(x0)
x_t = sqrt_abar[t, None, None, None] * x0 + sqrt_one_minus_abar[t, None, None, None] * noise
loss = F.mse_loss(model(x_t, t), noise)

두 가지가 이상합니다.

첫째, 잡음을 1,000번 단계적으로 더하는 과정이라고 배웠는데 여기서는 한 번의 곱셈 두 개로 t=743t = 743 짜리 표본이 나옵니다. 743번을 돌리지 않습니다.

둘째, x_t는 무작위로 뽑힌 값인데 이 값이 그대로 모델에 들어가고 손실이 역전파됩니다. 확률변수를 뽑는 연산은 미분할 수 없다고 알고 있는데, 여기서는 아무 일 없이 그래디언트가 흐릅니다.

두 이상함은 같은 장치에서 나옵니다. 이 글은 그 장치가 무엇인지, 그리고 왜 그것이 없으면 확산 모델도 VAE도 학습할 수 없는지를 봅니다. 지난 글이 다른 분포의 표본을 재사용하는 대가를 셌다면, 여기서는 표본을 뽑는 일 자체를 계산 그래프 안으로 끌어들입니다. U-Net 구조와 DDPM 구현은 확산 모델의 기초가 다루고, 이 글은 두 줄의 수식만 맡습니다.

미분 가능한 표집

분포 안의 파라미터

목표는 이런 형태의 기댓값을 파라미터로 미분하는 것입니다.

L(θ)=Ex∼pθ[f(x)]L(\theta) = \mathbb{E}_{x \sim p_\theta}[f(x)]

문제는 θ\theta 가 분포 안에 있다는 것입니다. ff 를 아무리 미분해도 θ\theta 는 나오지 않습니다.

로그 미분 트릭

로그 미분 트릭은 이 문제를 우회로 풉니다. ∇θpθ=pθ∇θlog⁡pθ\nabla_\theta p_\theta = p_\theta \nabla_\theta \log p_\theta 를 쓰면

∇θL=Ex∼pθ ⁣[f(x) ∇θlog⁡pθ(x)]\nabla_\theta L = \mathbb{E}_{x \sim p_\theta}\!\left[f(x)\, \nabla_\theta \log p_\theta(x)\right]

가 되어, ff 를 건드리지 않고도 표본만으로 그래디언트를 얻습니다. ff 가 미분 불가능해도 되고 심지어 블랙박스여도 됩니다.

재매개변수화 트릭

재매개변수화 트릭(reparameterization trick)은 정반대의 길을 갑니다. 우회하지 않고 길을 뚫습니다. 정규분포에서 뽑는 일은 이렇게 다시 쓸 수 있습니다.

x∼N(μ,σ2)⟺x=μ+σz,    z∼N(0,1)x \sim \mathcal{N}(\mu, \sigma^2) \quad \Longleftrightarrow \quad x = \mu + \sigma z, \;\; z \sim \mathcal{N}(0, 1)

오른쪽에서 무작위성은 전부 zz 에 있고, zz 의 분포는 μ\mu 나 σ\sigma 와 아무 상관이 없습니다. 그러니 zz 를 먼저 뽑아 상수로 고정해 두면 xx 는 μ,σ\mu, \sigma 의 평범한 미분 가능한 함수입니다.

로그 미분과 재매개변수화의 계산 그래프 비교

그래서 그래디언트가 이렇게 나옵니다.

∇θL=Ez∼N(0,1) ⁣[∇θf(μθ+σθz)]=Ez ⁣[f′(x) ∇θ(μθ+σθz)]\nabla_\theta L = \mathbb{E}_{z \sim \mathcal{N}(0,1)}\!\left[\nabla_\theta f(\mu_\theta + \sigma_\theta z)\right] = \mathbb{E}_z\!\left[f'(x)\, \nabla_\theta(\mu_\theta + \sigma_\theta z)\right]

기댓값 안의 분포가 θ\theta 에서 완전히 떨어져 나갔으므로 미분을 기댓값 안으로 넣을 수 있고, 연쇄법칙이 그대로 작동합니다. 대가는 두 가지입니다 — ff 가 미분 가능해야 하고, 분포를 이렇게 다시 쓸 수 있어야 합니다. 두 번째 대가가 어디까지 허용되는지는 세 번째 절에서 따로 봅니다.

두 추정량의 분산

도함수와 잡음의 곱

두 추정량은 같은 값을 추정합니다. 다르게 흔들릴 뿐입니다. 간단한 예로 재 봅니다 — f(x)=x2f(x) = x^2, x∼N(μ,σ2)x \sim \mathcal{N}(\mu, \sigma^2), 목표는 ∂μE[x2]=2μ\partial_\mu \mathbb{E}[x^2] = 2\mu 입니다.

  • 재매개변수화: ∂μf(μ+σz)=2(μ+σz)\partial_\mu f(\mu + \sigma z) = 2(\mu + \sigma z). 분산은 정확히 4σ24\sigma^2.
  • 로그 미분: x2⋅x−μσ2x^2 \cdot \frac{x - \mu}{\sigma^2}. 세제곱이 들어 있습니다.

μ=σ=1\mu = \sigma = 1 에서 표본 10만 개로 재면 앞쪽 표준편차가 2.00, 뒤쪽이 5.45입니다. 분산으로 7.4배 차이입니다.

차이가 어디서 오는지는 식의 모양이 말해 줍니다. 재매개변수화 추정량은 ff 의 도함수를 쓰고, 로그 미분 추정량은 ff 의 값에 잡음을 곱합니다. 도함수는 ff 가 그 자리에서 어느 쪽으로 기우는지를 직접 알려 주는 정보인데, 뒤쪽은 그것을 ff 값과 잡음의 상관에서 간접적으로 뽑아냅니다. 간접적으로 뽑으면 흔들림이 더 큽니다.

차원별 격차

차원을 올리면 격차가 벌어집니다. f(x)=∥x∥2f(x) = \|x\|^2 로 두고 μ∈Rd\mu \in \mathbb{R}^d 의 첫 성분에 대한 그래디언트를 재 봤습니다.

차원별 두 추정량의 표준편차 비교

import numpy as np
np.random.seed(4)
N = 200_000
for d in [1, 10, 100]:
    mu, sig = np.ones(d), 1.0
    x = mu + sig * np.random.randn(N, d)
    g_rep = 2 * x[:, 0]                                  # 재매개변수화
    g_score = (x**2).sum(axis=1) * (x[:, 0] - mu[0])     # 로그 미분
    print(f"d={d:3d}  재매개변수화 {g_rep.std():6.2f}   로그미분 {g_score.std():7.2f}")
d=  1  재매개변수화   2.00   로그미분    5.45
d= 10  재매개변수화   2.00   로그미분   23.43
d=100  재매개변수화   2.00   로그미분  202.98

재매개변수화 쪽은 차원과 무관하게 2σ2\sigma 에 머뭅니다. ∂μ1∥x∥2=2x1\partial_{\mu_1} \|x\|^2 = 2x_1 이라 다른 성분이 아예 식에 없기 때문입니다. 로그 미분 쪽은 ∥x∥2\|x\|^2 가 dd 개 항의 합이라 그 흔들림을 그대로 받고, 분산이 d=100d = 100 에서 1만 배로 벌어집니다.

VAE의 잠재 벡터가 수백 차원인 것을 생각하면 이 차이는 "조금 나은 정도"가 아닙니다. 로그 미분 트릭으로도 원리상 학습은 되지만, 같은 정확도를 얻으려면 배치를 1만 배로 키워야 합니다. 그래서 잠재변수가 연속이면 재매개변수화를 쓰고, 이산이라 쓸 수 없을 때만 로그 미분으로 돌아갑니다.

이산 잠재변수

위치·척도 족

"이산 분포는 안 된다"는 말을 두 절에 걸쳐 두 번 했으니 경계를 정확히 그어 둡니다. 재매개변수화가 되는 분포는 위치·척도 족(location-scale family)입니다 — 표준 분포 하나를 정해 두고 파라미터가 그것을 옮기고 늘리는 데만 쓰이는 족입니다.

x=μ+σ w,w∼(파라미터와 무관한 표준 분포)x = \mu + \sigma\,w, \qquad w \sim (\text{파라미터와 무관한 표준 분포})

이 꼴로 쓸 수 있으면 ww 를 먼저 뽑아 상수로 고정할 수 있고, 그러면 xx 가 파라미터의 미분 가능한 함수가 됩니다.

분포 되는가 어떻게
정규 N(μ,σ2)\mathcal{N}(\mu,\sigma^2) 된다 μ+σz\mu + \sigma z
균등 U(a,b)U(a,b) 된다 a+(b−a)ua + (b-a)u, u∼U(0,1)u \sim U(0,1)
라플라스 Lap(μ,b)\text{Lap}(\mu,b) 된다 μ+b w\mu + b\,w, ww 는 표준 라플라스
감마 Γ(k,θ)\Gamma(k,\theta) 척도만 된다 θ\theta 는 곱하면 되지만 kk 는 모양 자체를 바꾼다
베르누이 Bern(p)\text{Bern}(p) 안 된다 아래 소절

감마가 경계선을 잘 보여 줍니다. 척도 θ\theta 는 표준 감마 표본에 곱하기만 하면 되지만, 모양 파라미터 kk 는 표준 분포 자체를 다른 모양으로 바꾸므로 고정된 ww 하나로는 만들 수 없습니다. 파라미터가 옮기고 늘리는 일 말고 다른 일을 하면 이 트릭은 거기서 멈춥니다.

베르누이와 계단 함수

이산은 이유가 조금 다릅니다. 베르누이 표본도 경로 자체는 쓸 수 있습니다.

x=1[u<p],u∼U(0,1)x = \mathbb{1}[u < p], \qquad u \sim U(0,1)

무작위성이 uu 에 몰려 있고 pp 는 밖에 나와 있으니 형태는 재매개변수화와 같습니다. 그런데 이 함수는 pp 에 대해 계단입니다. pp 를 아주 조금 키워도 거의 모든 uu 에서 출력이 그대로 0이거나 그대로 1이고, u=pu = p 인 한 점에서만 0에서 1로 통째로 뛰어오릅니다.

그러니 도함수가 거의 어디서나 0이고 한 점에서 정의되지 않습니다. 길은 뚫려 있는데 그 길로 아무 정보도 안 흐릅니다. 연속 분포에서는 파라미터를 조금 바꾸면 표본이 조금 움직였는데, 이산에서는 표본이 움직일 자리가 없고 어느 값이 나올 확률만 바뀝니다. 확률이 바뀌는 것은 표본 하나를 보고는 알 수 없는 변화라, 로그 미분 트릭처럼 여러 표본의 상관에서 간접적으로 읽어 내는 수밖에 없습니다.

Gumbel-softmax와 straight-through

그 자리를 메우는 방법이 둘 있습니다. 둘 다 정확한 추정을 포기하고 편향을 받는 대신 그래디언트를 얻습니다.

Gumbel-softmax(concrete 분포라고도 합니다)는 계단을 매끄러운 것으로 갈아 끼웁니다. 범주분포에서 뽑는 일은 원래 arg⁡max⁡i(log⁡πi+gi)\arg\max_i(\log \pi_i + g_i) 로 쓸 수 있는데(gig_i 는 검벨 잡음 −log⁡(−log⁡ui)-\log(-\log u_i) 입니다), 여기서 arg⁡max⁡\arg\max 를 softmax로 바꿉니다.

yi=exp⁡((log⁡πi+gi)/τ)∑jexp⁡((log⁡πj+gj)/τ)y_i = \frac{\exp\big((\log \pi_i + g_i)/\tau\big)}{\sum_j \exp\big((\log \pi_j + g_j)/\tau\big)}

온도 τ\tau 가 0으로 가면 원래의 이산 표본에 가까워지고, 크면 매끄럽지만 원래 분포에서 멀어집니다. 잡음이 π\pi 밖에 나와 있는 형태라 π\pi 로 미분이 됩니다 — 이산 분포를 못 쓰는 대신 그것에 가까운 연속 분포를 쓰는 것입니다.

straight-through는 더 노골적입니다. 순전파에서는 이산 값을 그대로 쓰고, 역전파에서는 그 자리에 매끄러운 함수가 있었던 셈 치고 도함수를 흘려보냅니다. 앞뒤가 다른 함수를 쓰는 셈이라 편향된 추정량이지만, 실제로 잘 동작해서 VQ-VAE의 코드북처럼 이산 토큰을 쓰는 자리에서 표준으로 쓰입니다.

정리하면 이렇습니다. 연속이면 재매개변수화, 이산이면 로그 미분이거나 위 두 방법으로 근사입니다. 확산 모델이 픽셀 같은 연속 값을 다루는 한 이 고민은 없습니다 — 그래서 다음 절의 유도가 깔끔합니다.

폐형 forward 과정

forward 사슬

이제 두 번째 이상함으로 갑니다. 확산의 forward 과정은 정의부터 사슬입니다.

xk=1−βk xk−1+βk zk,zk∼N(0,I)x_k = \sqrt{1 - \beta_k}\, x_{k-1} + \sqrt{\beta_k}\, z_k, \qquad z_k \sim \mathcal{N}(0, I)

한 스텝마다 원래 신호를 조금 줄이고 새 잡음을 조금 섞습니다. αk=1−βk\alpha_k = 1 - \beta_k 로 줄여 쓰면 xk=αk xk−1+1−αk zkx_k = \sqrt{\alpha_k}\,x_{k-1} + \sqrt{1-\alpha_k}\,z_k 입니다.

정의를 그대로 따르면 t=743t = 743 짜리 표본을 얻으려면 743번을 돌려야 합니다. 배치마다 무작위 tt 를 뽑는 학습에서는 감당할 수 없는 비용입니다.

t번의 사슬과 한 번의 점프

두 스텝의 합성

빠져나갈 구멍은 정규분포의 성질 하나입니다 — 독립인 정규분포 둘을 더하면 다시 정규분포이고, 분산은 그냥 더해집니다. 두 스텝을 손으로 이어 봅니다.

xt=αt xt−1+1−αt zt=αt(αt−1 xt−2+1−αt−1 zt−1)+1−αt zt=αtαt−1 xt−2+αt(1−αt−1) zt−1+1−αt zt⏟독립인 정규분포 둘의 합\begin{aligned} x_t &= \sqrt{\alpha_t}\,x_{t-1} + \sqrt{1-\alpha_t}\,z_t \\ &= \sqrt{\alpha_t}\left(\sqrt{\alpha_{t-1}}\,x_{t-2} + \sqrt{1-\alpha_{t-1}}\,z_{t-1}\right) + \sqrt{1-\alpha_t}\,z_t \\ &= \sqrt{\alpha_t \alpha_{t-1}}\,x_{t-2} + \underbrace{\sqrt{\alpha_t(1-\alpha_{t-1})}\,z_{t-1} + \sqrt{1-\alpha_t}\,z_t}_{\text{독립인 정규분포 둘의 합}} \end{aligned}

밑줄 친 부분의 분산을 더합니다.

αt(1−αt−1)+(1−αt)=αt−αtαt−1+1−αt=1−αtαt−1\alpha_t(1 - \alpha_{t-1}) + (1 - \alpha_t) = \alpha_t - \alpha_t\alpha_{t-1} + 1 - \alpha_t = 1 - \alpha_t\alpha_{t-1}

αt\alpha_t 가 깔끔하게 상쇄되면서 1−αtαt−11 - \alpha_t \alpha_{t-1} 만 남습니다. 그러니 두 스텝은 한 스텝과 같은 모양입니다.

xt=αtαt−1 xt−2+1−αtαt−1 zˉx_t = \sqrt{\alpha_t\alpha_{t-1}}\,x_{t-2} + \sqrt{1 - \alpha_t\alpha_{t-1}}\,\bar{z}

같은 계산을 끝까지 반복하면, αˉt=∏k=1tαk\bar{\alpha}_t = \prod_{k=1}^{t}\alpha_k 라 두고

  q(xt∣x0)=N ⁣(αˉt x0,  (1−αˉt)I),xt=αˉt x0+1−αˉt z  \boxed{\;q(x_t \mid x_0) = \mathcal{N}\!\left(\sqrt{\bar{\alpha}_t}\,x_0,\; (1 - \bar{\alpha}_t) I\right), \qquad x_t = \sqrt{\bar{\alpha}_t}\,x_0 + \sqrt{1-\bar{\alpha}_t}\,z\;}

를 얻습니다. 이것이 폐형 forward 과정입니다 — 반복 없이 곧바로 값이 나오는 식을 폐형(closed form)이라 부릅니다. 두 계수의 제곱합이 αˉt+(1−αˉt)=1\bar{\alpha}_t + (1-\bar{\alpha}_t) = 1 이라 xtx_t 의 크기가 tt 와 무관하게 일정한 것도 여기서 함께 나옵니다. 신호가 줄어든 만큼 정확히 잡음이 채워 들어오는 셈이라, 모델이 받는 입력의 크기가 tt 에 따라 들쭉날쭉하지 않습니다.

스케줄과 남은 신호

신호 계수와 잡음 계수의 스케줄 곡선

β\beta 를 0.0001에서 0.02까지 선형으로 늘린 1,000스텝 스케줄이면 이렇게 됩니다.

tt 10 100 500 1000
αˉt\bar{\alpha}_t 0.9981 0.8970 0.0786 0.00004
αˉt\sqrt{\bar{\alpha}_t} · 남은 신호 0.9991 0.9471 0.2803 0.0064
1−αˉt\sqrt{1-\bar{\alpha}_t} · 섞인 잡음 0.0435 0.3209 0.9599 1.0000

식이 맞는지 사슬을 실제로 돌려 확인합니다. x0=2x_0 = 2 에서 시작해 20만 개를 1,000스텝까지 굴렸습니다.

tt 사슬 평균 폐형 평균 사슬 표준편차 폐형 표준편차
10 1.9982 1.9981 0.0435 0.0435
100 1.8943 1.8942 0.3212 0.3209
500 0.5625 0.5607 0.9559 0.9599
1000 0.0086 0.0127 1.0016 1.0000

1,000번을 돌린 것과 곱셈 두 번이 소수 셋째 자리까지 같습니다.

β의 크기와 코사인 스케줄

β\beta 를 왜 하필 0.0001에서 0.02 사이에 두는지는 위 표의 양 끝이 답입니다. 끝에서 αˉT\bar\alpha_T 가 0에 붙어야 한다는 것이 유일한 실질적 조건입니다 — 역과정은 순수한 잡음 N(0,I)\mathcal{N}(0,I) 에서 출발하므로, xTx_T 에 원본이 조금이라도 남아 있으면 출발점의 분포가 어긋납니다.

αˉT=∏(1−βk)\bar\alpha_T = \prod(1-\beta_k) 이므로 β\beta 가 크면 빨리 0에 닿고 작으면 스텝이 더 필요합니다. 위 스케줄은 평균 β≈0.01\beta \approx 0.01 로 1,000스텝을 도는데, αˉ1000=0.00004\bar\alpha_{1000} = 0.00004 이니 조건을 딱 맞춰 놓은 값입니다. 위쪽 0.02는 그 조건을 만족시키는 하한이고, 아래쪽 0.0001은 초반 스텝이 원본을 거의 안 건드리게 두려는 값입니다 — 초반에 크게 망가뜨리면 역과정에서 되돌릴 정보가 남지 않습니다.

문제는 중간입니다. 선형 스케줄은 αˉt\bar\alpha_t 가 너무 일찍 0으로 꺼집니다.

선형 스케줄과 코사인 스케줄의 ᾱ 곡선 비교

t=500t = 500 에서 이미 αˉt=0.079\bar\alpha_t = 0.079 라 남은 신호가 0.28밖에 안 됩니다. 뒤쪽 절반은 잡음에 잡음을 더하는 구간이라 모델이 배울 것이 거의 없는데도 학습 예산의 절반이 거기로 갑니다.

코사인 스케줄은 αˉt\bar\alpha_t 를 직접 정의해 이 자리를 고칩니다.

αˉt=h(t/T)h(0),h(u)=cos⁡2 ⁣(u+s1+s⋅π2)\bar\alpha_t = \frac{h(t/T)}{h(0)}, \qquad h(u) = \cos^2\!\left(\frac{u+s}{1+s}\cdot\frac{\pi}{2}\right)

ss 는 0.008 정도의 작은 수로, tt 가 0 근처에서 β\beta 가 지나치게 작아지지 않게 하는 보정입니다. 같은 t=500t=500 에서 αˉt=0.494\bar\alpha_t = 0.494 라 중반이 완만하게 지나가고, 대신 끝에서 가파르게 떨어져 αˉT≈0\bar\alpha_T \approx 0 이라는 조건은 그대로 지킵니다. 고치는 것은 시작도 끝도 아니라 가운데를 쓸 수 있게 만드는 것입니다.

확산과 VAE

학습 코드의 한 줄

이제 처음의 코드 세 줄을 다시 읽을 수 있습니다.

x_t = sqrt_abar[t] * x0 + sqrt_one_minus_abar[t] * noise

이 한 줄이 폐형이면서 동시에 재매개변수화입니다. 폐형이라서 tt 를 무작위로 골라도 비용이 같고, 재매개변수화 형태라서 xtx_t 를 통과하는 그래디언트가 살아 있습니다. 둘 중 하나만 있었다면 학습 루프가 이 모양이 될 수 없었습니다.

  • 폐형이 없으면 배치마다 수백 스텝을 굴려야 하니 무작위 tt 학습이 불가능합니다.
  • 재매개변수화 형태가 아니면 xtx_t 가 그래디언트의 벽이 됩니다.

VAE 인코더

VAE의 인코더도 정확히 같은 줄을 씁니다.

z = mu + torch.exp(0.5 * logvar) * torch.randn_like(mu)   # VAE

x_t와 z는 하는 일이 다릅니다 — 하나는 데이터를 망가뜨린 결과이고 하나는 데이터를 압축한 결과입니다. 그런데 표집을 계산 그래프 안에 남기는 방식은 글자 그대로 같습니다. 확산을 "잠재변수가 아주 많은 계층적 VAE"로 읽는 관점이 여기서 시작합니다. 두 모델이 공유하는 것은 손실의 모양이 아니라 이 한 줄이고, 손실 쪽에서 둘이 어떻게 만나는지는 다음 글에서 봅니다.

연습 문제

답은 문항을 눌러 펼칩니다. 연습 4는 표준정규의 적률 E[z]=0\mathbb{E}[z]=0, E[z2]=1\mathbb{E}[z^2]=1, E[z3]=0\mathbb{E}[z^3]=0, E[z4]=3\mathbb{E}[z^4]=3, E[z5]=0\mathbb{E}[z^5]=0, E[z6]=15\mathbb{E}[z^6]=15 를 씁니다.

연습 1 — 균등·라플라스 분포

  1. 균등분포 U(a,b)U(a,b) 에서 뽑는 일을 u∼U(0,1)u \sim U(0,1) 하나로 다시 쓰고, aa 와 bb 로 각각 미분하세요.
    x=a+(b−a)ux = a + (b-a)u 입니다. uu 의 분포가 a,ba,b 와 무관하므로 고정할 수 있고, ∂x∂a=1−u\dfrac{\partial x}{\partial a} = 1-u, ∂x∂b=u\dfrac{\partial x}{\partial b} = u 입니다. 둘을 더하면 1이 되는데, aa 와 bb 를 같은 만큼 옮기면 표본도 그만큼 옮겨진다는 뜻입니다.
  2. 라플라스 분포 Lap(μ,b)\text{Lap}(\mu, b) 를 같은 꼴로 쓰고, 이것이 위치·척도 족인 이유를 한 줄로 적으세요.
    표준 라플라스 w∼Lap(0,1)w \sim \text{Lap}(0,1) 을 뽑아 x=μ+b wx = \mu + b\,w 로 씁니다. 밀도가 12bexp⁡ ⁣(−∣x−μ∣b)\dfrac{1}{2b}\exp\!\left(-\dfrac{|x-\mu|}{b}\right) 로 x−μb\dfrac{x-\mu}{b} 에만 의존하므로, 파라미터가 하는 일이 옮기기와 늘리기뿐입니다.

연습 2 — 두 스텝 합성

  1. αt=0.99\alpha_t = 0.99, αt−1=0.98\alpha_{t-1} = 0.98 일 때 xt−2x_{t-2} 에서 xtx_t 로 한 번에 가는 식의 두 계수를 구하세요.
    신호 계수는 αtαt−1=0.9702≈0.985\sqrt{\alpha_t\alpha_{t-1}} = \sqrt{0.9702} \approx 0.985, 잡음 계수는 1−0.9702=0.0298≈0.173\sqrt{1-0.9702} = \sqrt{0.0298} \approx 0.173 입니다.
  2. 그 두 계수의 제곱을 더하면 얼마이고, 왜 그런지 적으세요.
    0.9702+0.0298=10.9702 + 0.0298 = 1 입니다. 폐형이 αˉt\sqrt{\bar\alpha_t} 와 1−αˉt\sqrt{1-\bar\alpha_t} 를 계수로 쓰므로 제곱합은 언제나 αˉt+(1−αˉt)=1\bar\alpha_t + (1-\bar\alpha_t) = 1 이고, 그래서 xtx_t 의 크기가 tt 와 무관하게 유지됩니다.

연습 3 — 스케줄 어림

  1. 모든 kk 에서 βk=0.01\beta_k = 0.01 로 같다면 αˉt=0.99t\bar\alpha_t = 0.99^t 입니다. t=100t = 100 에서 남은 신호 계수와 섞인 잡음 계수를 구하세요.
    αˉ100=0.99100≈0.366\bar\alpha_{100} = 0.99^{100} \approx 0.366 이므로 신호 계수는 0.366≈0.605\sqrt{0.366} \approx 0.605, 잡음 계수는 0.634≈0.796\sqrt{0.634} \approx 0.796 입니다. 100스텝만에 잡음 쪽이 더 커집니다.
  2. 같은 β\beta 로 αˉT<10−4\bar\alpha_T < 10^{-4} 가 되려면 TT 를 얼마로 두어야 하는지 어림하세요.
    Tln⁡0.99<ln⁡10−4T\ln 0.99 < \ln 10^{-4} 에서 T>9.210.01005≈916T > \dfrac{9.21}{0.01005} \approx 916 이라 917스텝쯤 필요합니다. 본문의 1,000스텝 스케줄이 평균 β≈0.01\beta \approx 0.01 인 것과 맞아떨어집니다.

연습 4 — 로그 미분 추정량의 분산

  1. d=1d=1, μ=σ=1\mu = \sigma = 1 에서 로그 미분 추정량 g=x2x−μσ2g = x^2\dfrac{x-\mu}{\sigma^2} 의 평균과 분산을 손으로 구하세요.
    x=1+zx = 1 + z 로 놓으면 g=(1+z)2z=z+2z2+z3g = (1+z)^2 z = z + 2z^2 + z^3 입니다. 평균은 E[g]=0+2+0=2\mathbb{E}[g] = 0 + 2 + 0 = 2 로 목표값 2μ=22\mu = 2 와 같아 불편추정량입니다. 제곱을 펼치면 g2=z2+4z4+z6+4z3+2z4+4z5g^2 = z^2 + 4z^4 + z^6 + 4z^3 + 2z^4 + 4z^5 이므로 E[g2]=1+12+15+0+6+0=34\mathbb{E}[g^2] = 1 + 12 + 15 + 0 + 6 + 0 = 34 이고 분산은 34−4=3034 - 4 = 30, 표준편차는 30≈5.48\sqrt{30} \approx 5.48 입니다. 본문에서 표본 10만 개로 잰 5.45가 이 값이었습니다. 재매개변수화 쪽은 2(1+z)2(1+z) 라 분산이 4, 표준편차가 2입니다.

정리

  • 재매개변수화 트릭은 x∼N(μ,σ2)x \sim \mathcal{N}(\mu, \sigma^2) 을 x=μ+σzx = \mu + \sigma z 로 다시 써서 무작위성을 파라미터 밖으로 뺀다. zz 를 고정하면 xx 는 평범한 미분 가능한 함수가 된다.
  • 로그 미분 트릭은 ff 를 미분하지 않고 우회하고, 재매개변수화는 f′f' 를 실제로 쓴다. 대신 ff 가 미분 가능해야 하고 분포를 위치·척도 형태로 다시 쓸 수 있어야 한다.
  • 분산이 다르다. f(x)=∥x∥2f(x) = \|x\|^2 에서 재매개변수화 추정량의 표준편차는 차원과 무관하게 2인데, 로그 미분 쪽은 d=100d = 100 에서 203이다 — 분산으로 1만 배다.
  • 그래서 잠재변수가 연속이면 재매개변수화를 쓰고, 이산이라 쓸 수 없을 때만 로그 미분으로 돌아간다.
  • 되는 경계는 위치·척도 족이다. 파라미터가 옮기고 늘리는 일만 하면 되고(정규·균등·라플라스), 감마의 모양 파라미터처럼 분포 모양 자체를 바꾸면 안 된다.
  • 이산은 경로가 있어도 계단이라 도함수가 거의 어디서나 0이다. Gumbel-softmax는 argmax를 온도 τ\tau 의 softmax로 갈아 끼우고, straight-through는 순전파만 이산으로 두고 역전파는 매끄러운 함수의 도함수를 흘린다. 둘 다 편향을 받는 대가로 그래디언트를 얻는다.
  • 확산의 forward 과정은 정의상 tt 번의 사슬이지만, 독립인 정규분포의 합이 다시 정규분포라는 성질을 쓰면 αt\alpha_t 가 상쇄되면서 q(xt∣x0)=N(αˉtx0,(1−αˉt)I)q(x_t \mid x_0) = \mathcal{N}(\sqrt{\bar{\alpha}_t}x_0, (1-\bar{\alpha}_t)I) 라는 폐형이 나온다.
  • 두 계수의 제곱합이 1이라 xtx_t 의 크기가 tt 와 무관하게 일정하다. 20만 표본으로 사슬을 1,000스텝 굴린 결과가 폐형과 소수 셋째 자리까지 일치했다.
  • β\beta 의 범위를 정하는 조건은 끝에서 αˉT\bar\alpha_T 가 0에 붙을 것 하나다. 선형 스케줄은 그 조건을 지키는 대신 t=500t=500 에서 이미 αˉt=0.079\bar\alpha_t = 0.079 라 뒤쪽 절반을 낭비하고, 코사인 스케줄은 같은 자리를 0.494로 두어 가운데를 되살린다.
  • 학습 코드의 한 줄이 폐형이자 재매개변수화다. 폐형이라 임의의 tt 를 공짜로 고를 수 있고, 재매개변수화라 그 값을 통과해 그래디언트가 흐른다.

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

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