수학

MATH / 중급 37번

softmax와 교차엔트로피의 기울기가 정확히 p − y인 이유

분류 모델의 마지막 층에서 실제로 흐르는 기울기는 «예측 확률 빼기 정답» 한 줄입니다. softmax의 야코비안 diag(p) − ppᵀ 를 유도하고, 교차엔트로피와 합성했을 때 그 행렬이 어떻게 통째로 소거되는지 끝까지 계산합니다. 그리고 두 함수를 왜 반드시 붙여서 구현하는지까지.

PALDYN Team37 MIN READ

분류 모델의 마지막 두 줄은 어디서나 같습니다.

logits = model(x)                          # (B, C) 짜리 실수 점수
loss = F.cross_entropy(logits, labels)     # softmax + 교차엔트로피

loss.backward()를 부르면 logits에 기울기가 채워집니다. 그 값이 무엇일까요.

softmax는 지수함수와 나눗셈이 얽힌 함수이고, 교차엔트로피에는 로그가 들어 있습니다. 둘을 합성해 미분하면 지저분한 식이 나올 것 같은데, 실제로 나오는 것은

∂L∂z=p−y\frac{\partial L}{\partial z} = p - y

예측 확률에서 정답을 뺀 것 한 줄입니다. 이 글은 그 식을 끝까지 유도하고, 왜 이렇게 깨끗해지는지와 그 사실이 구현에 무엇을 강제하는지를 봅니다.

지난 글까지 연쇄법칙·야코비안·VJP·행렬 미분 규약이 모였으니, 네 도구를 한 번에 쓰는 자리입니다.

softmax의 야코비안

softmax 유도 글에서 세운 정의로 시작합니다. 로짓 z∈RCz \in \mathbb{R}^C 를 확률로 보내는 함수입니다.

pi=ezi∑kezkp_i = \frac{e^{z_i}}{\sum_{k} e^{z_k}}

입력도 출력도 CC 차원이므로 야코비안은 C×CC \times C 입니다.

성분별 미분

성분 ∂pi/∂zj\partial p_i/\partial z_j 를 구합니다. 분모 S=∑kezkS = \sum_k e^{z_k} 에 모든 zjz_j 가 들어 있다는 점이 요점이라, 몫의 미분을 그대로 적용합니다.

∂pi∂zj=∂ezi∂zj⋅S−ezi⋅∂S∂zjS2\frac{\partial p_i}{\partial z_j} = \frac{\dfrac{\partial e^{z_i}}{\partial z_j}\cdot S - e^{z_i}\cdot \dfrac{\partial S}{\partial z_j}}{S^2}

∂S/∂zj=ezj\partial S/\partial z_j = e^{z_j} 입니다. 분자의 첫 항은 i=ji = j 일 때만 ezie^{z_i} 이고 아니면 0이므로 경우가 갈립니다.

i=ji = j 일 때

∂pi∂zi=eziS−ezieziS2=eziS−(eziS)2=pi−pi2=pi(1−pi)\frac{\partial p_i}{\partial z_i} = \frac{e^{z_i}S - e^{z_i}e^{z_i}}{S^2} = \frac{e^{z_i}}{S} - \left(\frac{e^{z_i}}{S}\right)^2 = p_i - p_i^2 = p_i(1 - p_i)

i≠ji \neq j 일 때

∂pi∂zj=0−eziezjS2=−pipj\frac{\partial p_i}{\partial z_j} = \frac{0 - e^{z_i}e^{z_j}}{S^2} = -p_ip_j

두 경우를 한 줄로 묶습니다. δij\delta_{ij} 를 i=ji = j 이면 1, 아니면 0인 기호(크로네커 델타)라 하면

∂pi∂zj=pi(δij−pj)\frac{\partial p_i}{\partial z_j} = p_i(\delta_{ij} - p_j)

이고, 이것을 C×CC \times C 행렬로 묶으면

softmax의 야코비안. J=diag⁡(p)−p pTJ = \operatorname{diag}(p) - p\,p^{\mathsf T}

입니다. 앞 항은 대각선에 pip_i 를 놓은 행렬이고, 뒤 항은 지난 글에서 본 외적입니다.

softmax의 야코비안은 대각행렬에서 외적을 뺀 것이다

대칭성과 영공간

z=(2, 1, 0.1)z = (2,\, 1,\, 0.1) 로 수를 넣어 봅니다. p=(0.659001, 0.242433, 0.098566)p = (0.659001,\, 0.242433,\, 0.098566) 이므로

J=[0.224719−0.159764−0.064955−0.1597640.183659−0.023896−0.064955−0.0238960.088851]J = \begin{bmatrix} 0.224719 & -0.159764 & -0.064955 \\ -0.159764 & 0.183659 & -0.023896 \\ -0.064955 & -0.023896 & 0.088851 \end{bmatrix}

눈에 띄는 것이 셋입니다.

대칭입니다. diag⁡(p)\operatorname{diag}(p) 도 ppTpp^{\mathsf T} 도 대칭이니 당연합니다. 이 한 줄이 아래 둘을 편하게 만듭니다.

각 열의 합이 정확히 0입니다. ∑ipi(δij−pj)=pj−pj∑ipi=pj−pj=0\sum_i p_i(\delta_{ij} - p_j) = p_j - p_j\sum_i p_i = p_j - p_j = 0 이기 때문이고, 뜻은 「로짓 하나를 흔들어도 확률의 총합은 1로 남는다」입니다. 그리고 대칭이므로 각 행의 합도 0입니다. 행 쪽 등식은 J1=0J\mathbf{1} = 0 으로 적히는데, 이것은 성분이 전부 1인 벡터 1\mathbf{1} 이 고윳값 0에 딸린 고유벡터라는 말입니다.

반양정치입니다. 아무 벡터 vv 를 가져와도 vTJv≥0v^{\mathsf T}Jv \ge 0 이라는 뜻이고, 계산하면 한 줄로 보입니다.

vTJv=∑ipivi2−(∑ipivi)2=Ep[v2]−Ep[v]2=Var⁡p(v)v^{\mathsf T}Jv = \sum_i p_i v_i^2 - \left(\sum_i p_i v_i\right)^2 = \mathbb{E}_p[v^2] - \mathbb{E}_p[v]^2 = \operatorname{Var}_p(v)

pp 를 확률분포로 보고 vv 를 그 위의 값으로 보면 야코비안이 재는 것은 vv 의 분산입니다. 분산은 음수가 될 수 없으므로 반양정치이고, 0이 되는 것은 vv 가 모든 자리에서 같은 값일 때뿐입니다. 곧 고윳값 0에 딸린 방향은 1\mathbf{1} 방향 하나뿐이고, 그 방향은 로짓 전체를 같은 값만큼 옮기는 상수 이동입니다. softmax가 상수 이동에 안 변한다는 성질이 야코비안 쪽에서 「그 방향으로는 분산이 0」으로 다시 나타난 것입니다.

쏠림과 포화

pp 가 한쪽으로 쏠리면 야코비안의 성분이 전부 0에 가까워집니다. z=(10, 1, 0.1)z = (10,\, 1,\, 0.1) 로 재 봅니다. 확률은 p=(0.999826, 0.000123, 0.000050)p = (0.999826,\, 0.000123,\, 0.000050) 이고 야코비안은

J=[1.735×10−4−1.234×10−4−5.016×10−5−1.234×10−41.234×10−4−6.190×10−9−5.016×10−5−6.190×10−95.016×10−5]J = \begin{bmatrix} 1.735\times10^{-4} & -1.234\times10^{-4} & -5.016\times10^{-5} \\ -1.234\times10^{-4} & 1.234\times10^{-4} & -6.190\times10^{-9} \\ -5.016\times10^{-5} & -6.190\times10^{-9} & 5.016\times10^{-5} \end{bmatrix}

입니다. 가장 큰 성분이 1.735×10−41.735\times10^{-4} 로, 앞의 z=(2,1,0.1)z=(2,1,0.1) 에서 0.2247190.224719 였던 것보다 1,300배쯤 작습니다. 위의 분산 읽기로 보면 이유가 분명합니다 — 확률이 한 자리에 몰리면 그 분포 위의 분산이 0으로 내려앉기 때문입니다. 이것을 포화라고 부릅니다.

포화가 위험한 이유는 softmax 뒤에 다른 것을 붙였을 때 그 기울기가 그대로 죽는다는 데 있습니다. 그런데 교차엔트로피를 붙이면 그 일이 안 일어납니다. 다음 절이 그 이야기입니다.

합성에서 소거되는 야코비안

이제 손실을 붙입니다. 교차엔트로피 글의 정의 그대로입니다.

L=−∑iyilog⁡piL = -\sum_i y_i \log p_i

여기서 yy 는 목표 분포이고, 아래 유도에서 필요한 성질은 ∑iyi=1\sum_i y_i = 1 하나뿐입니다. 원-핫이어도 되고 라벨 스무딩을 먹인 부드러운 분포여도 됩니다.

약분되는 자리

먼저 pp 에 대한 미분입니다. pip_i 는 항 하나에만 들어가므로

∂L∂pi=−yipi\frac{\partial L}{\partial p_i} = -\frac{y_i}{p_i}

입니다. 분모에 pip_i 가 있다는 것을 기억해 둡니다. 이제 VJP로 야코비안을 통과시킵니다. ui=∂L/∂piu_i = \partial L/\partial p_i 라 두면

∂L∂zj=∑iui∂pi∂zj=∑i(−yipi)pi(δij−pj)\frac{\partial L}{\partial z_j} = \sum_i u_i \frac{\partial p_i}{\partial z_j} = \sum_i \left(-\frac{y_i}{p_i}\right) p_i(\delta_{ij} - p_j)

여기가 이 글의 결정적인 자리입니다. 교차엔트로피가 내놓은 1/pi1/p_i 와 야코비안이 들고 있던 pip_i 가 만나 약분됩니다.

=−∑iyi(δij−pj)= -\sum_i y_i(\delta_{ij} - p_j)

1/pᵢ 와 pᵢ 가 약분되면서 행렬이 통째로 사라진다

남은 것을 두 항으로 나눠 정리합니다. 첫 항의 합에서는 i=ji = j 인 것만 살아남고, 둘째 항에서는 pjp_j 가 합 밖으로 나옵니다.

=−∑iyiδij+pj∑iyi=−yj+pj⋅1= -\sum_i y_i\delta_{ij} + p_j\sum_i y_i = -y_j + p_j \cdot 1

∑iyi=1\sum_i y_i = 1 을 쓴 자리가 마지막 등호입니다. 그러니

∂L∂z=p−y\dfrac{\partial L}{\partial z} = p - y

입니다. C×CC \times C 행렬은 어디에도 남지 않았습니다.

앞의 수로 확인합니다. y=(1,0,0)y = (1,0,0) 이면 L=−log⁡0.659001=0.41703L = -\log 0.659001 = 0.41703 이고

∂L∂z=(0.659001−1,  0.242433,  0.098566)=(−0.340999,  0.242433,  0.098566)\frac{\partial L}{\partial z} = (0.659001 - 1,\; 0.242433,\; 0.098566) = (-0.340999,\; 0.242433,\; 0.098566)

입니다. 야코비안을 실제로 만들어 곱해도, 유한차분으로 재도 같은 값이 나옵니다 — 뒤의 코드에서 셋을 나란히 찍습니다.

이 소거는 우연이 아닙니다. softmax는 로짓에 지수를 씌워 정규화하고, 교차엔트로피는 그 결과에 로그를 씌웁니다. 지수와 로그는 서로를 되돌리는 함수이므로 두 함수를 붙이면 안쪽에서 log⁡pc=zc−log⁡∑kezk\log p_c = z_c - \log\sum_k e^{z_k} 라는 로짓에 대해 거의 선형인 식이 드러납니다. 실제로 이 식을 zjz_j 로 곧바로 미분해도 δcj−pj\delta_{cj} - p_j 가 나오고, 부호를 붙이면 그대로 p−yp - y 입니다. 야코비안을 거쳐 약분하는 길과 로그를 먼저 펼치는 길이 같은 곳에 닿습니다.

sigmoid와 BCE

같은 소거가 이진 분류에서도 일어납니다. 클래스가 둘인 경우를 손으로 따라가면 바로 보입니다.

C=2C = 2 이고 로짓이 (z1,z2)(z_1, z_2) 일 때 softmax의 첫 성분은

p1=ez1ez1+ez2=11+e−(z1−z2)=σ(z1−z2)p_1 = \frac{e^{z_1}}{e^{z_1} + e^{z_2}} = \frac{1}{1 + e^{-(z_1 - z_2)}} = \sigma(z_1 - z_2)

입니다. 로짓의 차 하나로 줄어들고 그 함수가 바로 시그모이드입니다. 목표도 y=(y1,1−y1)y = (y_1, 1 - y_1) 로 수 하나이므로, 교차엔트로피는

L=−y1log⁡p1−(1−y1)log⁡(1−p1)L = -y_1\log p_1 - (1-y_1)\log(1-p_1)

이 되는데 이것이 이진 교차엔트로피입니다. 한 변수 s=z1−z2s = z_1 - z_2 로 미분하면

dLds=p1−y1\frac{dL}{ds} = p_1 - y_1

입니다. 시그모이드의 도함수가 σ(1−σ)\sigma(1-\sigma) 이고 손실이 내놓는 1/p11/p_1 과 1/(1−p1)1/(1-p_1) 이 그것과 만나 약분되는, 앞 절과 똑같은 구조입니다. C×CC\times C 행렬이 1×11\times1 로 줄어든 특수 경우라고 읽으면 됩니다. F.binary_cross_entropy_with_logits가 시그모이드를 안에 품고 있는 이유도 F.cross_entropy와 같습니다.

합이 1이 아닐 때

유도에서 쓴 성질이 ∑iyi=1\sum_i y_i = 1 하나였으니, 그것이 깨지면 결론도 깨집니다. 어디서 깨지는지는 위 유도의 마지막 등호가 그대로 알려 줍니다.

∂L∂zj=−yj+pj∑iyi\frac{\partial L}{\partial z_j} = -y_j + p_j\sum_i y_i

여기서 ∑iyi\sum_i y_i 를 1로 바꾼 것이 p−yp-y 였습니다. 합이 1이 아니면 pj∑iyi−yjp_j\sum_i y_i - y_j 가 남습니다.

한 표본이 여러 정답을 가질 수 있는 멀티라벨 문제가 바로 이 경우입니다. 라벨이 둘이면 ∑iyi=2\sum_i y_i = 2 라서 기울기가 2pj−yj2p_j - y_j 가 되고, 정답이 아닌 자리에서도 2pj2p_j 로 두 배가 됩니다. 확률의 합은 1인데 목표의 합은 2라 두 분포를 맞출 방법이 애초에 없습니다. 그래서 멀티라벨에는 softmax + 교차엔트로피를 쓰지 않고, 클래스마다 시그모이드와 이진 교차엔트로피를 따로 답니다.

p − y를 성분마다 읽기

올리는 자리와 내리는 자리

p−yp - y 를 성분마다 읽으면 학습이 무엇을 하는지가 그대로 보입니다.

  • 정답 자리는 pc−1p_c - 1 이라 항상 음수입니다. 경사하강은 기울기의 반대로 가므로 그 로짓을 올립니다.
  • 나머지 자리는 pjp_j 라 항상 양수입니다. 그 로짓들을 내립니다.
  • 크기가 곧 틀린 정도입니다. pc=0.99p_c = 0.99 로 맞히면 기울기가 −0.01-0.01 이라 거의 아무 일도 하지 않고, pc=0.01p_c = 0.01 로 틀리면 −0.99-0.99 로 세게 밀어붙입니다.

기울기의 성분은 정답 로짓을 올리고 나머지를 내린다

세 번째 줄이 앞 절의 포화와 정반대라는 점이 중요합니다. softmax만 놓고 보면 쏠린 자리에서 야코비안이 죽었는데, 교차엔트로피를 붙이면 크게 틀린 표본일수록 기울기가 커집니다. 두 표본을 나란히 세우면 차이가 한눈에 보입니다.

잘 맞힌 표본과 크게 틀린 표본의 p − y 막대 비교

성분의 합은 언제나 0입니다. ∑j(pj−yj)=1−1=0\sum_j (p_j - y_j) = 1 - 1 = 0 이기 때문이고, 뜻은 「로짓 전체를 같은 값만큼 올려도 손실이 변하지 않는다」입니다. 앞 절에서 본 1\mathbf{1} 방향의 고윳값 0이 여기서 다시 나타난 것이고, 그래서 로짓은 절대적인 크기가 아니라 서로의 차이만 의미를 갖습니다.

목표 분포가 부드러워도 결론이 그대로라는 점도 확인해 둡니다. y=(0.9, 0.05, 0.05)y = (0.9,\, 0.05,\, 0.05) 로 라벨 스무딩을 먹이면 기울기는

p−y=(−0.240999,  0.192433,  0.048566)p - y = (-0.240999,\; 0.192433,\; 0.048566)

입니다. 유도에서 쓴 것이 ∑yi=1\sum y_i = 1 뿐이었으므로 당연한 결과이고, 정답 자리를 미는 힘이 −0.341-0.341 에서 −0.241-0.241 로 줄어든 것이 스무딩이 하는 일 전부입니다.

배치와 패딩

지금까지는 표본 하나였습니다. 실제로 흐르는 값에는 배치가 한 겹 더 얹힙니다.

F.cross_entropy의 기본값인 reduction='mean' 은 표본마다의 손실을 평균합니다.

L=1B∑bLb⟹∂L∂z(b)=p(b)−y(b)BL = \frac1B\sum_{b} L_b \quad\Longrightarrow\quad \frac{\partial L}{\partial z^{(b)}} = \frac{p^{(b)} - y^{(b)}}{B}

앞에 1/B1/B 가 붙으므로 배치 크기를 두 배로 늘리면 표본 하나가 미는 힘은 절반이 됩니다. 학습률을 그대로 두고 배치만 키웠을 때 학습이 느려지는 자리가 여기이고, reduction='sum' 으로 바꾸면 이 분모가 사라져 p−yp - y 그대로 흐릅니다.

언어모델처럼 길이가 다른 시퀀스를 한 배치에 모으면 한 겹이 더 붙습니다. 짧은 문장은 뒤를 채움 토큰으로 메우는데, 이 패딩 자리는 정답이 없으므로 손실에서 빼야 합니다. ignore_index 가 하는 일이 정확히 그것이고, 결과는 두 가지입니다.

  • 패딩 자리의 기울기는 0입니다. 그 자리는 손실에 아무 항도 안 보탰으니 당연합니다.
  • 분모의 개수에서도 빠집니다. 평균을 내는 BB 가 전체 토큰 수가 아니라 패딩을 뺀 실제 토큰 수가 됩니다.

둘째 줄을 빠뜨리면 조용히 틀립니다. 패딩까지 세어 나누면 패딩이 많은 배치일수록 기울기가 작아져, 배치를 어떻게 묶느냐가 학습 속도를 바꿉니다. 손실 값 자체도 배치마다 비교할 수 없게 됩니다.

마지막 층의 세 줄

한 걸음 더 가면 마지막 층의 가중치까지 닿습니다. 마지막 은닉 표현을 hh 라 하고 z=Wh+bz = Wh + b 이면, 지난 글의 공식이 그대로 적용됩니다.

∂L∂W=(p−y) hT,∂L∂b=p−y,∂L∂h=WT(p−y)\frac{\partial L}{\partial W} = (p - y)\,h^{\mathsf T}, \qquad \frac{\partial L}{\partial b} = p - y, \qquad \frac{\partial L}{\partial h} = W^{\mathsf T}(p - y)

분류 모델의 마지막 층 전체가 이 세 줄입니다. 첫 줄은 「틀린 만큼을 표현 방향으로 뿌린 외적」이고, 둘째 줄은 편향이 잔차를 그대로 받는다는 말이고, 셋째 줄이 그 아래 층들로 흘러 들어가는 값입니다.

셋째 줄을 한 번 더 읽어 둘 만합니다. WW 의 각 행이 클래스 하나의 방향이므로, WT(p−y)W^{\mathsf T}(p-y) 는 정답 클래스의 방향으로 hh 를 당기고 나머지 클래스의 방향에서 밀어내는 합입니다. 분류 손실이 표현 공간을 어떻게 정리하는지가 이 한 줄에 들어 있습니다.

온도가 바꾸는 눈금

로짓을 softmax에 넣기 전에 상수로 나누는 일이 흔합니다. 그 상수를 온도라 하고 TT 로 적습니다.

p(T)=softmax⁡(z/T)p^{(T)} = \operatorname{softmax}(z/T)

TT 가 크면 분포가 평평해지고 작으면 뾰족해집니다. 기울기 쪽에서 무슨 일이 일어나는지를 봅니다.

T로 나눈 로짓의 야코비안

w=z/Tw = z/T 로 두면 p(T)=softmax⁡(w)p^{(T)} = \operatorname{softmax}(w) 이므로, 앞에서 구한 야코비안이 ww 에 대해 그대로 성립합니다. 여기에 ∂w/∂z=(1/T)I\partial w/\partial z = (1/T)I 를 연쇄법칙으로 이으면

∂p(T)∂z=1T(diag⁡(p(T))−p(T)p(T)T)\frac{\partial p^{(T)}}{\partial z} = \frac1T\left(\operatorname{diag}(p^{(T)}) - p^{(T)}{p^{(T)}}^{\mathsf T}\right)

입니다. 모양은 그대로이고 눈금만 1/T1/T 배로 바뀌었습니다. 교차엔트로피를 붙였을 때의 소거도 그대로 일어나므로 결론은 한 줄입니다.

∂L∂z=p(T)−yT\frac{\partial L}{\partial z} = \frac{p^{(T)} - y}{T}

z=(2,1,0.1)z = (2,1,0.1) 과 y=(1,0,0)y=(1,0,0) 으로 재 봅니다.

TT p(T)p^{(T)} 기울기 정답 자리의 크기
1 (0.6590, 0.2424, 0.0986)(0.6590,\ 0.2424,\ 0.0986) (−0.3410, 0.2424, 0.0986)(-0.3410,\ 0.2424,\ 0.0986) 0.3410
2 (0.5017, 0.3043, 0.1940)(0.5017,\ 0.3043,\ 0.1940) (−0.2492, 0.1521, 0.0970)(-0.2492,\ 0.1521,\ 0.0970) 0.2492
4 (0.4165, 0.3244, 0.2590)(0.4165,\ 0.3244,\ 0.2590) (−0.1459, 0.0811, 0.0648)(-0.1459,\ 0.0811,\ 0.0648) 0.1459

TT 가 커질수록 로짓으로 흘러 들어가는 신호가 작아집니다. 줄어드는 이유가 둘이라는 점이 중요합니다 — 앞에 붙은 1/T1/T 가 한 번 줄이고, 분포가 평평해지면서 p(T)−yp^{(T)} - y 자체도 줄어듭니다. TT 가 클 때 이 둘이 겹쳐 전체가 대략 1/T21/T^2 로 작아집니다.

반대쪽도 봐 둡니다. T=0.5T = 0.5 로 내리면 정답 자리의 크기가 0.27250.2725 로 T=1T=1 보다 오히려 작습니다. 분포가 뾰족해져 pcp_c 가 1에 가까워지는 쪽, 곧 앞 절의 포화가 1/T1/T 보다 빨리 먹기 때문입니다. 온도는 양쪽 끝에서 모두 신호를 줄입니다.

증류 손실의 T²

큰 모델의 출력 분포를 작은 모델에 옮기는 지식 증류에서 이 1/T1/T 가 문제가 됩니다. 부드러운 목표를 쓰려고 TT 를 4나 8로 올리는데, 그러면 위의 이유로 기울기가 T2T^2 분의 1쯤으로 줄어듭니다. 같은 학습률로는 부드러운 목표 쪽이 거의 아무 일도 못 합니다.

그래서 관행이 하나 있습니다. 부드러운 목표에서 나온 손실에 T2T^2 을 곱합니다. 줄어든 만큼을 도로 곱해 눈금을 되돌리는 것이고, 이렇게 해 두면 TT 를 바꿔도 학습률을 다시 맞출 필요가 없습니다. 정답 라벨에서 나온 손실은 T=1T=1 로 계산하므로 안 건드리고, 둘을 섞는 비율만 따로 정합니다.

L=α T2 Lsoft+(1−α) LhardL = \alpha\,T^2\,L_{\text{soft}} + (1-\alpha)\,L_{\text{hard}}

T2T^2 이 어디서 왔는지 모르면 이 식이 임의의 숫자로 보입니다. 앞 소절의 1/T1/T 두 겹을 되돌리는 값이라고 읽으면 외울 것이 없습니다.

어텐션의 1/√d_h

같은 자리가 어텐션에도 있습니다. 점수를 만들 때 내적을 dh\sqrt{d_h} 로 나눕니다.

A=softmax⁡ ⁣(QKTdh)A = \operatorname{softmax}\!\left(\frac{QK^{\mathsf T}}{\sqrt{d_h}}\right)

이 분모는 온도 T=dhT = \sqrt{d_h} 와 정확히 같은 자리입니다. 이유도 온도와 같습니다 — dhd_h 가 크면 내적의 크기가 대략 dh\sqrt{d_h} 에 비례해 커지므로, 나눠 두지 않으면 점수가 커지면서 softmax가 한 자리에 쏠리고 앞 절의 포화가 그대로 일어납니다. 나눗셈 하나가 야코비안이 죽는 것을 막는 장치입니다.

다만 어텐션에는 교차엔트로피가 안 붙습니다. 그래서 소거가 일어나지 않고, 그 자리가 다음 절의 마지막 소절입니다.

붙여 구현하는 세 이유

F.cross_entropy가 softmax와 교차엔트로피를 따로 부르지 않고 한 연산으로 묶여 있는 데는 이유가 셋 있습니다.

야코비안을 안 만든다

따로 구현하면 softmax의 backward가 C×CC \times C 행렬을 만들어야 합니다. 어휘 크기가 128,000인 언어모델이면

128,0002=1.64×1010개×4바이트=65.5 GB128{,}000^2 = 1.64 \times 10^{10} \text{개} \times 4\text{바이트} = 65.5\,\text{GB}

입니다. 토큰 하나마다 그렇습니다 — 애초에 만들 수 없습니다. 붙여 구현하면 뺄셈 한 번이라 길이 128,000짜리 벡터 하나로 끝납니다.

오버플로와 log 0

두 자리에서 터집니다.

  • ezie^{z_i} 가 넘칩니다. float64에서 e710e^{710} 부터 무한대이고, 로짓이 그 근처로 가는 일은 드물지 않습니다. 해법은 softmax의 상수 이동 불변성을 써서 zz 에서 최댓값을 빼는 것이고, 실제 구현이 언제나 그렇게 합니다.
  • log⁡pc\log p_c 가 무한대가 됩니다. pcp_c 가 매우 작으면 실수 표현에서 0으로 내려앉고, 그러면 log⁡0=−∞\log 0 = -\infty 입니다. 여기서는 이동만으로 부족합니다.

둘째 문제의 해법이 붙여 구현하는 진짜 이유입니다. 정답 자리의 로그확률을 pp 를 거치지 않고 곧바로 계산할 수 있습니다.

log⁡pc=zc−log⁡∑kezk\log p_c = z_c - \log\sum_k e^{z_k}

오른쪽에는 나눗셈도 로그의 0도 없습니다. 그리고 log⁡∑kezk\log\sum_k e^{z_k} 는 최댓값 MM 을 빼서 M+log⁡∑kezk−MM + \log\sum_k e^{z_k - M} 으로 안전하게 계산합니다 — 이 조합을 log-sum-exp 요령이라고 부릅니다. 확률 pp 를 한 번도 만들지 않고 손실이 나옵니다.

나눴다 곱하는 왕복

따로 계산하면 pip_i 로 나눴다가 다시 pip_i 를 곱하는 왕복이 실제로 일어납니다. 수학적으로는 1이지만 부동소수점에서는 그렇지 않고, pip_i 가 작을수록 오차가 커집니다. 붙여 구현하면 그 왕복 자체가 없습니다.

나눠 구현하면 만들어야 하는 것과, 붙여 구현하면 남는 것

소거가 없는 어텐션

앞의 셋은 뒤에 교차엔트로피가 붙어 있을 때의 이야기입니다. 어텐션의 backward에는 −log⁡-\log 가 없어 diag⁡(p)−ppT\operatorname{diag}(p) - pp^{\mathsf T} 가 소거되지 않고 그대로 남습니다. 어텐션 확률 AA 는 손실이 아니라 값 행렬 VV 와 곱해져 다음 층으로 가므로, 위에서 내려오는 u=∂L/∂Au = \partial L/\partial A 가 임의의 벡터입니다. 약분해 줄 1/pi1/p_i 가 없습니다.

그래도 행렬을 만들 필요는 없습니다. 야코비안이 대칭이라는 것을 쓰면 VJP가 한 줄로 적힙니다.

JTu=Ju=diag⁡(p)u−p(pTu)=p⊙(u−p⋅u)J^{\mathsf T}u = Ju = \operatorname{diag}(p)u - p(p^{\mathsf T}u) = p \odot \bigl(u - p\cdot u\bigr)

p⋅up \cdot u 는 스칼라 하나이고 ⊙\odot 는 성분별 곱입니다. 곱셈 두 번과 뺄셈 한 번이면 끝나고, 필요한 메모리는 길이 CC 짜리 벡터뿐입니다. 행렬을 세우는 길과 결과가 같은데 비용만 다릅니다.

비용이 얼마나 다른지는 원소 수로 세면 분명합니다. 어텐션 확률은 (B,H,T,T)(B, H, T, T) 모양이고, 그중 한 행마다 T×TT \times T 야코비안이 붙습니다.

무엇 T=4,096T = 4{,}096 에서
확률 텐서 (B,H,T,T)(B,H,T,T), B=8B=8 · H=32H=32 17.2 GB
한 행의 야코비안 T×TT \times T 67.1 MB
모든 행의 야코비안 (B,H,T,T,T)(B,H,T,T,T) 70.4 TB

마지막 줄은 어느 장비에도 안 올라갑니다. VJP를 한 줄로 적을 줄 아는 것이 선택이 아니라 조건인 자리가 여기입니다. 앞 절의 1/dh1/\sqrt{d_h} 와 이 VJP가 어텐션에서 softmax를 실제로 쓸 수 있게 만드는 두 장치입니다.

코드로 확인하기

야코비안과 p − y 맞대기

import math

def softmax(z):
    m = max(z)                                  # 상수 이동 — 오버플로 방지
    e = [math.exp(v - m) for v in z]
    s = sum(e)
    return [v / s for v in e]

z = [2.0, 1.0, 0.1]
y = [1.0, 0.0, 0.0]
p = softmax(z)
print([round(v, 6) for v in p])       # [0.659001, 0.242433, 0.098566]

# ① 야코비안 diag(p) − ppᵀ 를 실제로 만들어 본다
J = [[p[i] * ((i == j) - p[j]) for j in range(3)] for i in range(3)]
for row in J:
    print([round(v, 6) for v in row])
# [0.224719, -0.159764, -0.064955]
# [-0.159764, 0.183659, -0.023896]
# [-0.064955, -0.023896, 0.088851]
print([round(sum(J[i][j] for i in range(3)), 12) for j in range(3)])   # [0.0, 0.0, 0.0]
print([round(sum(J[i][j] for j in range(3)), 12) for i in range(3)])   # [0.0, 0.0, 0.0]

# ② 야코비안을 통과시킨 값과 p − y 를 맞대 본다
dL_dp = [-y[i] / p[i] for i in range(3)]
chained = [sum(dL_dp[i] * J[i][j] for i in range(3)) for j in range(3)]
print([round(v, 6) for v in chained])            # [-0.340999, 0.242433, 0.098566]
print([round(p[i] - y[i], 6) for i in range(3)]) # [-0.340999, 0.242433, 0.098566]
print(round(sum(chained), 12))                   # 0.0   성분의 합은 언제나 0

# ③ vᵀJv 가 정말 p 위에서 잰 v 의 분산인가
v = [1.5, -0.4, 2.0]
quad = sum(v[i] * J[i][j] * v[j] for i in range(3) for j in range(3))
var = sum(p[i] * v[i] ** 2 for i in range(3)) - sum(p[i] * v[i] for i in range(3)) ** 2
print(round(quad, 12), round(var, 12))           # 0.730624 0.730624
one = [1.0, 1.0, 1.0]                            # 상수 방향은 분산이 0이다
print(round(sum(one[i] * J[i][j] * one[j] for i in range(3) for j in range(3)), 12))  # 0.0

②가 이 글의 유도를 확인한 자리입니다. 야코비안을 만들어 곱한 값과 손으로 유도한 p−yp - y 가 소수점 여섯 자리까지 같습니다. ③은 반양정치를 수로 본 것이고, 마지막 줄에서 1\mathbf{1} 방향의 값이 정확히 0으로 떨어집니다.

온도 셋과 유한차분

def loss(zz, yy=y, T=1.0):
    pp = softmax([v / T for v in zz])
    return -sum(yy[i] * math.log(pp[i]) for i in range(3))

def numeric_grad(zz, T=1.0, eps=1e-6):
    return [round((loss([zz[k] + eps * (k == j) for k in range(3)], T=T)
                 - loss([zz[k] - eps * (k == j) for k in range(3)], T=T)) / (2 * eps), 5)
            for j in range(3)]

for T in (0.5, 1.0, 2.0, 4.0):
    pT = softmax([v / T for v in z])
    closed = [round((pT[i] - y[i]) / T, 5) for i in range(3)]
    print(T, closed, numeric_grad(z, T))
# 0.5 [-0.27245, 0.2338, 0.03865]   [-0.27245, 0.2338, 0.03865]
# 1.0 [-0.341, 0.24243, 0.09857]    [-0.341, 0.24243, 0.09857]
# 2.0 [-0.24916, 0.15214, 0.09701]  [-0.24916, 0.15214, 0.09701]
# 4.0 [-0.14586, 0.0811, 0.06476]   [-0.14586, 0.0811, 0.06476]

# ④ 포화 — z 가 쏠리면 야코비안 성분이 통째로 작아진다
zs = [10.0, 1.0, 0.1]
ps = softmax(zs)
Js = [[ps[i] * ((i == j) - ps[j]) for j in range(3)] for i in range(3)]
print(max(abs(v) for row in Js for v in row))   # 0.00017352423868...
print(max(abs(v) for row in J for v in row))    # 0.22471863783296...
print([round(ps[i] - y[i], 6) for i in range(3)])  # [-0.000174, 0.000123, 5e-05]

(p(T)−y)/T(p^{(T)} - y)/T 와 유한차분이 네 온도에서 모두 같습니다. ④의 마지막 줄이 이 글에서 가장 아픈 자리입니다 — 포화된 자리에서는 p−yp-y 도 함께 작아지므로, 소거가 기울기를 살리는 것은 틀렸을 때뿐입니다. 이미 맞힌 표본은 어차피 밀 이유가 없으니 그래도 됩니다.

합이 1이 아닐 때

# ⑤ 목표 분포가 부드러워도 p − y 그대로다
ys = [0.9, 0.05, 0.05]
print([round(p[i] - ys[i], 6) for i in range(3)])   # [-0.240999, 0.192433, 0.048566]

# ⑥ 합이 1이 아니면 p·Σy − y 가 남는다 (멀티라벨)
ym = [1.0, 1.0, 0.0]                                # 정답이 둘
s = sum(ym)
dL_dp = [-ym[i] / p[i] for i in range(3)]
chained = [sum(dL_dp[i] * J[i][j] for i in range(3)) for j in range(3)]
print([round(v, 6) for v in chained])               # [0.318002, -0.515134, 0.197132]
print([round(p[j] * s - ym[j], 6) for j in range(3)])  # [0.318002, -0.515134, 0.197132]
print([round(p[j] - ym[j], 6) for j in range(3)])      # [-0.340999, -0.757567, 0.098566]

# ⑦ 나눠 구현하면 터지는 두 자리
try:
    math.exp(1000)
except OverflowError as e:
    print("exp 오버플로:", e)                        # exp 오버플로: math range error

big = [1000.0, 0.0, -1000.0]
print(softmax(big))                                 # [1.0, 0.0, 0.0]   이동 덕분에 산다
try:
    math.log(softmax(big)[2])
except ValueError as e:
    print("log(0):", e)                             # log(0): math domain error

def log_softmax(z):
    m = max(z)
    lse = m + math.log(sum(math.exp(v - m) for v in z))
    return [v - lse for v in z]

print([round(v, 3) for v in log_softmax(big)])      # [0.0, -1000.0, -2000.0]

# ⑧ VJP 한 줄과 야코비안 곱이 같은가
u = [0.3, -1.2, 0.7]
full = [sum(u[i] * J[i][j] for i in range(3)) for j in range(3)]
dot = sum(p[i] * u[i] for i in range(3))
vjp = [p[j] * (u[j] - dot) for j in range(3)]
print([round(v, 9) for v in full], [round(v, 9) for v in vjp])   # 두 줄이 같다

⑥의 마지막 두 줄을 나란히 보세요. 라벨이 둘일 때 진짜 기울기는 p∑y−yp\sum y - y 이고, p−yp - y 라고 믿으면 첫 성분의 부호가 반대로 나옵니다 — 진짜 값은 +0.318+0.318 이라 정답 로짓을 내리는데 p−yp-y 는 −0.341-0.341 로 올리라고 말합니다. ⑧은 어텐션 쪽 계산으로, 길이 CC 짜리 벡터 셋만 가지고 야코비안 곱과 같은 값을 얻습니다.

정리

  • softmax의 야코비안은 diag⁡(p)−ppT\operatorname{diag}(p) - pp^{\mathsf T} 다. i=ji = j 에서 pi(1−pi)p_i(1-p_i), 아니면 −pipj-p_ip_j 를 한 줄로 묶은 것이다.
  • 그 행렬은 대칭이고 행과 열의 합이 모두 0이다. vTJvv^{\mathsf T}Jv 가 pp 위에서 잰 vv 의 분산이라 반양정치이고, 고윳값 0의 방향은 상수 이동 하나뿐이다.
  • pp 가 한쪽으로 쏠리면 성분이 전부 0에 가까워진다(포화). z=(10,1,0.1)z=(10,1,0.1) 에서 최대 성분이 1.7×10−41.7\times10^{-4} 로 1,300배 작아진다.
  • 교차엔트로피는 ∂L/∂pi=−yi/pi\partial L/\partial p_i = -y_i/p_i 를 내놓는다. 그 1/pi1/p_i 가 야코비안의 pip_i 와 약분되면서 행렬이 통째로 사라진다. 지수와 로그가 서로를 되돌리기 때문이라 우연이 아니다.
  • 남는 것은 ∂L/∂z=p−y\partial L/\partial z = p - y 한 줄이다. 유도에 쓴 성질은 ∑iyi=1\sum_i y_i = 1 뿐이라 라벨 스무딩에도 그대로 성립하고, C=2C=2 로 줄이면 sigmoid + BCE의 p1−y1p_1 - y_1 이 된다.
  • 합이 1이 아니면 pj∑iyi−yjp_j\sum_i y_i - y_j 가 남는다. 멀티라벨에 softmax를 안 쓰는 이유가 이것이다.
  • 배치에서 reduction='mean' 이면 실제로 흐르는 값은 (p−y)/B(p-y)/B 다. 패딩 자리는 기울기가 0이고 분모의 개수에서도 빠진다.
  • 마지막 층 전체는 세 줄이다 — ∂L/∂W=(p−y)hT\partial L/\partial W = (p-y)h^{\mathsf T}, ∂L/∂b=p−y\partial L/\partial b = p-y, ∂L/∂h=WT(p−y)\partial L/\partial h = W^{\mathsf T}(p-y).
  • 온도 TT 로 나누면 야코비안은 모양이 같고 눈금만 1/T1/T 가 되어 기울기가 (p(T)−y)/T(p^{(T)}-y)/T 다. 증류에서 T2T^2 을 곱하는 관행이 이 1/T1/T 두 겹을 되돌리는 값이고, 어텐션의 1/dh1/\sqrt{d_h} 도 같은 자리의 온도다.
  • 붙여 구현하는 이유는 셋이다 — C×CC \times C 행렬을 안 만들고(어휘 12만 8천이면 65.5GB), eze^z 오버플로와 log⁡0\log 0 을 피하고(log-sum-exp), 나눴다 곱하는 왕복의 오차를 없앤다.
  • 어텐션에는 −log⁡-\log 가 없어 소거가 안 일어난다. 그래도 VJP를 p⊙(u−p⋅u)p \odot (u - p\cdot u) 한 줄로 적으면 벡터만으로 끝난다 — 행렬을 세우면 T=4,096T=4{,}096 에서 70.4 TB다.

F.cross_entropy의 backward가 하는 일은 이제 한 줄로 적을 수 있습니다 — 예측 확률에서 정답 분포를 빼는 것입니다. 지수함수와 로그와 C×CC \times C 행렬이 모두 종이 위에서 지워지고 뺄셈 하나만 남았습니다.

여기까지가 6단원입니다. 도함수를 민감도로 다시 읽는 것에서 시작해 연쇄법칙 하나를 세우고, 그것을 행렬로 올려 야코비안을 얻고, 곱하는 순서를 세어 역전파를 강제한 뒤, 규약을 정해 종이 위에 적고, 마지막으로 실제 층 하나를 끝까지 유도했습니다. 이제 「어떻게 최소화하는가」의 계산 쪽은 닫혔습니다.


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

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