수학

MATH / 중급 24번

softmax 유도: 점수 벡터를 확률로 바꾸는 함수는 왜 그 모양인가

분류기의 마지막 줄에도, 어텐션 가중치에도, MoE 라우터에도 같은 함수 하나가 들어갑니다. 실수 벡터를 확률로 보내는 함수가 지켜야 할 조건 넷을 적으면 exp가 강제로 튀어나오고, 그때 남는 상수가 정확히 온도라는 것까지 유도합니다.

PALDYN Team18 MIN READ

트랜스포머 블록 하나를 읽다 보면 같은 함수를 세 번 만납니다. 어텐션이 QK⊤/dkQK^\top/\sqrt{d_k} 를 계산한 뒤 한 번, MoE 라우터가 어느 전문가로 보낼지 정할 때 한 번, 마지막 층이 어휘 5만 개 중 다음 토큰을 고를 때 한 번. 전부 torch.softmax 한 줄입니다.

softmax(z)i=ezi∑jezj\mathrm{softmax}(\mathbf{z})_i = \frac{e^{z_i}}{\sum_j e^{z_j}}

식 자체는 외우기 쉽습니다. 그런데 왜 하필 ee 인가라는 질문에는 대개 「양수로 만들어야 하니까」라는 답이 돌아옵니다. 그러면 제곱은 왜 안 되나요. 절댓값은요. 2z2^{z} 는요.

지난 글까지가 분포를 «읽는» 쪽이었다면 이 글은 분포를 «만드는» 쪽입니다. 조건 몇 개를 적어 두면 exp 말고는 답이 없다는 것을 직접 유도해 보겠습니다.

문제를 정확히 적기

모델이 내놓는 것은 실수 벡터 z=(z1,…,zn)\mathbf{z} = (z_1, \dots, z_n) 입니다. 어휘 5만 개짜리 언어모델이라면 nn 이 5만이고, 각 ziz_i 는 «ii 번 토큰이 얼마나 그럴듯한가»를 나타내는 점수입니다. 이 점수를 로짓이라고 부릅니다 — 왜 그렇게 부르는지는 글 뒤쪽에서 유도로 답합니다.

우리가 원하는 것은 이 벡터를 확률분포로 보내는 함수입니다.

σ:Rn→Δn−1,Δn−1={p:pi>0, ∑ipi=1}\sigma : \mathbb{R}^n \to \Delta^{n-1}, \qquad \Delta^{n-1} = \Big\{ \mathbf{p} : p_i > 0,\ \textstyle\sum_i p_i = 1 \Big\}

Δn−1\Delta^{n-1} 은 «성분이 전부 양수이고 합이 1인 벡터들의 모임»이고 확률단체(simplex)라고 부릅니다. 이제 이 함수가 지켜야 할 것을 적습니다.

조건 뜻
① 양수·정규화 pi>0p_i > 0 이고 ∑pi=1\sum p_i = 1
② 좌표마다 같은 취급 pi∝f(zi)p_i \propto f(z_i) — 어느 좌표든 같은 함수 ff 를 거친다
③ 순서 보존 zi>zjz_i > z_j 이면 pi>pjp_i > p_j
④ 평행이동 불변 모든 점수에 같은 수를 더해도 p\mathbf{p} 가 변하지 않는다

①과 ③은 설명이 필요 없습니다. ②는 «ii 번 자리라서 특별대우» 같은 것이 없다는 뜻으로, 어휘의 3번 토큰과 4만 번 토큰을 다른 규칙으로 다루면 곤란하니 당연히 요구됩니다.

④가 이 글의 열쇠입니다. 점수는 그 자체로 의미가 있는 값이 아니라 서로 비교하라고 있는 값입니다. 모든 로짓이 10씩 크다는 것은 아무 정보도 아닙니다 — 마지막 층의 편향에 10을 더하면 그렇게 되니까요. 그러니 의미를 갖는 것은 차이뿐이고, 확률은 그 차이만으로 정해져야 합니다.

exp가 강제되는 과정

②와 ①을 합치면 함수의 모양이 이미 절반 정해집니다.

pi=f(zi)∑jf(zj),f>0p_i = \frac{f(z_i)}{\sum_j f(z_j)}, \qquad f > 0

이제 ④를 붙입니다. 두 좌표의 확률비를 보면

pipj=f(zi)f(zj)\frac{p_i}{p_j} = \frac{f(z_i)}{f(z_j)}

인데, ④에 따라 ziz_i 와 zjz_j 에 같은 수를 더해도 이 비가 변하면 안 됩니다. 즉 이 비는 zi−zjz_i - z_j 에만 의존해야 합니다.

ff 가 양수이므로 로그를 씌워도 됩니다. g=log⁡fg = \log f 라 두고, 「g(a)−g(b)g(a) - g(b) 가 a−ba-b 에만 의존한다」를 식으로 적으면 어떤 함수 hh 가 있어

g(a)−g(b)=h(a−b)g(a) - g(b) = h(a - b)

입니다. 여기서 b=0b = 0 을 넣으면 h(a)=g(a)−g(0)h(a) = g(a) - g(0) 이고, 이것을 원래 식에 다시 넣으면

h(a−b)=(h(a)+g(0))−(h(b)+g(0))=h(a)−h(b)h(a-b) = \big(h(a) + g(0)\big) - \big(h(b) + g(0)\big) = h(a) - h(b)

이 남습니다. x=a−bx = a-b, y=by = b 로 이름을 바꾸면 이렇습니다.

h(x+y)=h(x)+h(y)h(x + y) = h(x) + h(y)

더하기를 더하기로 옮기는 함수는 무엇인가. 이 조건을 코시 함수방정식이라고 하고, 답이 h(x)=kxh(x) = kx 하나뿐이라는 것이 알려져 있습니다 — 단, 병리적인 답을 배제하려면 단조성이나 연속성 같은 조건이 하나 붙어야 하는데 우리에게는 ③(순서 보존)이 이미 있습니다.

되돌아가면

g(z)=kz+g(0)⟹f(z)=Cekz,C=eg(0)>0g(z) = kz + g(0) \quad \Longrightarrow \quad f(z) = C e^{kz}, \qquad C = e^{g(0)} > 0

입니다. 그리고 CC 는 분자와 분모에 똑같이 곱해져 약분되어 사라집니다.

pi=Cekzi∑jCekzj=ekzi∑jekzjp_i = \frac{C e^{k z_i}}{\sum_j C e^{k z_j}} = \frac{e^{k z_i}}{\sum_j e^{k z_j}}

k>0k > 0 인 것은 ③에서 나옵니다. kk 가 음수면 큰 점수가 작은 확률을 받아 순서가 뒤집힙니다.

조건 넷을 적었을 뿐인데 답이 하나로 좁혀졌습니다. 제곱도, 절댓값도 ④를 못 지킵니다. f(z)=z2f(z)=z^2 이면 (1,2)(1,2) 는 (0.2,0.8)(0.2, 0.8) 을, 여기에 10을 더한 (11,12)(11,12) 는 (0.457,0.543)(0.457, 0.543) 을 줍니다 — 아무것도 안 바뀌었는데 확률이 달라졌습니다. 2z2^z 만은 통과하는데, 그것은 2z=ezln⁡22^z = e^{z \ln 2} 라 이미 k=ln⁡2k = \ln 2 인 같은 식이기 때문입니다.

softmax의 세 걸음: 점수, exp, 정규화

남은 상수 kk 가 곧 온도다

유도가 끝나고 자유롭게 남은 것이 kk 하나입니다. 이 값은 없앨 수 없습니다 — 점수의 «단위»를 정하는 값이기 때문입니다. 「점수 차이 1이 확률비 몇 배를 뜻하는가」에 답하는 것이 kk 이고, 우리가 정해 줘야 합니다.

k=1k = 1 로 두면 점수 차이 1이 확률비 e≈2.718e \approx 2.718 배를 뜻합니다. 이것이 기본형입니다. 그리고 관례상 kk 대신 그 역수를 쓰고 온도라고 부릅니다.

T=1k,pi=ezi/T∑jezj/TT = \frac{1}{k}, \qquad p_i = \frac{e^{z_i/T}}{\sum_j e^{z_j/T}}

즉 온도는 나중에 덧붙인 손잡이가 아니라 유도가 끝나고 반드시 남는 자유도입니다. 이름이 온도인 것은 통계물리에서 같은 모양의 식(볼츠만 분포)이 나오고 거기서 TT 가 실제 온도이기 때문입니다.

TT 를 바꾸면 무엇이 달라지는지는 두 극단을 보면 분명합니다.

  • T→0T \to 0 이면 zi/Tz_i/T 의 차이가 무한히 벌어져 가장 큰 점수 하나가 확률 1을 독차지합니다. argmax와 같아집니다.
  • T→∞T \to \infty 이면 zi/Tz_i/T 가 전부 0으로 몰려 균등분포가 됩니다. 점수를 아예 안 본 것과 같습니다.

z=(2,1,−1)\mathbf{z} = (2, 1, -1) 로 세 값을 계산해 보겠습니다.

온도에 따라 뾰족해지고 평평해지는 분포

TT p1p_1 p2p_2 p3p_3
0.5 0.879 0.119 0.002
1 0.705 0.259 0.035
2 0.547 0.331 0.122

T=0.5T=0.5 일 때 3등의 확률이 0.002까지 내려간 것이 눈에 띕니다. 온도를 낮추는 것은 꼬리를 죽이는 일입니다.

이름이 조금 잘못 붙었다

T→0T \to 0 에서 argmax가 된다는 것을 보고 나면 이름이 이상하게 들립니다. softmax가 근사하는 것은 최댓값 max⁡\max 가 아니라 최댓값의 «위치»를 알려 주는 argmax이기 때문입니다. 실제로 이 함수의 출력은 수 하나가 아니라 벡터이고, TT 가 작아질수록 1등 자리만 1인 원-핫 벡터에 가까워집니다. 정확한 이름은 soft-argmax인 셈입니다.

부드러운 max⁡\max 에 해당하는 함수는 따로 있습니다 — 바로 앞에서 상수로 넘긴 log⁡∑jezj\log \sum_j e^{z_j} 입니다. 이 값은 언제나 max⁡jzj\max_j z_j 보다 조금 크고, 점수 차이가 벌어질수록 그 «조금»이 0으로 줄어듭니다. 이름이 엇갈린 채로 굳어졌을 뿐, 두 함수는 한 식의 분자와 분모로 붙어 있습니다.

상수를 더해도 같다 — 그것이 곧 구현의 안전장치

④는 유도의 재료였지만 코드에서는 그 자체로 쓸모가 있습니다. 임의의 상수 cc 에 대해

ezi−c∑jezj−c=e−cezie−c∑jezj=ezi∑jezj\frac{e^{z_i - c}}{\sum_j e^{z_j - c}} = \frac{e^{-c} e^{z_i}}{e^{-c}\sum_j e^{z_j}} = \frac{e^{z_i}}{\sum_j e^{z_j}}

이므로 무엇을 빼도 답이 같습니다. 그러면 가장 편한 cc 를 고르면 됩니다. 그 편한 값이 c=max⁡jzjc = \max_j z_j 입니다.

이유는 지수와 로그를 다룬 글에서 본 것과 짝을 이룹니다. 거기서는 확률의 곱이 0으로 죽는 언더플로가 문제였는데, 여기서는 반대쪽입니다. float64가 담을 수 있는 가장 큰 수가 대략 1.8×103081.8 \times 10^{308} 이라 eze^{z} 는 zz 가 710 정도만 넘어도 무한대가 됩니다. 로짓이 그만큼 커지는 일이 실제로 있고, 그러면 ∞/∞\infty/\infty 라 결과가 nan입니다.

최댓값을 빼고 나면 가장 큰 지수가 정확히 e0=1e^0 = 1 이고 나머지는 그보다 작으므로 넘칠 수가 없습니다.

최댓값을 빼도 확률이 같고, 그래야 넘치지 않는다

(1000,999,997)(1000, 999, 997) 에서 1000을 빼면 (0,−1,−3)(0, -1, -3) 인데, 이것은 (2,1,−1)(2,1,-1) 에서 2를 뺀 것과 같은 벡터입니다. 두 로짓 벡터는 정확히 같은 확률을 줍니다 — 차이가 같으니까요.

로짓은 로그 확률이다

앞의 유도를 뒤집으면 이름의 유래가 나옵니다. k=1k=1 일 때 양변에 로그를 씌우면

log⁡pi=zi−log⁡∑jezj\log p_i = z_i - \log \sum_j e^{z_j}

인데 오른쪽 둘째 항은 ii 와 무관한 상수입니다. 이 항이 지수와 로그 글에서 본 log-sum-exp이고, 그러니 로짓은 «상수 하나 차이로» 로그 확률입니다.

zi=log⁡pi+constz_i = \log p_i + \text{const}

이 관점에서 보면 조건 ④가 왜 그렇게 자연스러웠는지도 분명해집니다. 로그 확률에서 상수는 정규화를 맡을 뿐이니 처음부터 정보가 아니었습니다. 그리고 두 좌표를 비교하면 상수가 사라져

log⁡pipj=zi−zj\log \frac{p_i}{p_j} = z_i - z_j

가 됩니다. 점수의 차이가 곧 로그 확률비이고, «로그 확률비」를 뜻하는 통계 용어 logit이 이름의 출처입니다.

선택지가 둘이면 sigmoid가 된다

n=2n=2 를 넣어 봅시다.

p1=ez1ez1+ez0p_1 = \frac{e^{z_1}}{e^{z_1} + e^{z_0}}

분자와 분모를 ez1e^{z_1} 으로 나누면

p1=11+e−(z1−z0)=σ(z1−z0)p_1 = \frac{1}{1 + e^{-(z_1 - z_0)}} = \sigma(z_1 - z_0)

여기서 σ(d)=1/(1+e−d)\sigma(d) = 1/(1+e^{-d}) 가 sigmoid 함수입니다. 즉 sigmoid는 softmax의 특수한 경우이고, 이진 분류에서 출력을 하나만 두는 이유도 여기 있습니다 — 어차피 차이 하나만 쓰이므로 점수를 둘 둘 필요가 없습니다.

둘일 때 softmax가 sigmoid로 줄어드는 모습

d=0d = 0 이면 정확히 0.5이고, d=1d=1 이면 1/(1+e−1)=0.7311/(1+e^{-1}) = 0.731 입니다. 곡선이 1에 한없이 다가가지만 결코 닿지 않는 것도 식에서 그대로 읽힙니다 — 분모의 e−de^{-d} 가 0이 되지는 않으니까요. softmax가 확률 0이나 1을 정확히 내놓지 못한다는 것은 그래서 특수한 사정이 아니라 조건 ①(모든 pip_i 가 양수)을 지킨 결과입니다.

코드로 확인하기

import numpy as np

def softmax(z, T=1.0):
    z = np.asarray(z, dtype=np.float64) / T
    z = z - z.max()              # ④ — 답은 그대로, 넘침만 막는다
    e = np.exp(z)
    return e / e.sum()

z = np.array([2.0, 1.0, -1.0])
print(softmax(z).round(4))            # [0.7054 0.2595 0.0351]
print(softmax(z + 1000).round(4))     # [0.7054 0.2595 0.0351]  ← 평행이동 불변
print(softmax(z, T=0.5).round(4))     # [0.8789 0.1189 0.0022]
print(softmax(z, T=2.0).round(4))     # [0.5465 0.3315 0.122 ]

# 로짓 차이 = 로그 확률비
p = softmax(z)
print(np.log(p[0] / p[1]).round(4), (z[0] - z[1]).round(4))   # 1.0 1.0

# 최댓값을 안 빼면 이렇게 된다
bad = np.exp(z + 1000); print(bad / bad.sum())   # [nan nan nan]

# 둘일 때는 sigmoid
d = 1.0
print(softmax([d, 0.0])[0].round(6), (1 / (1 + np.exp(-d))).round(6))   # 0.731059 0.731059

정리

  • 문제는 «실수 벡터를 확률분포로 보내기» 하나고, 조건은 양수·정규화, 좌표마다 같은 취급, 순서 보존, 평행이동 불변 넷이다.
  • 넷을 적으면 f(z)=Cekzf(z) = Ce^{kz} 가 강제된다. 확률비가 차이에만 의존한다는 조건이 코시 함수방정식으로 넘어가고, 단조성이 답을 h(x)=kxh(x)=kx 로 못 박는다.
  • CC 는 약분되어 사라지고 kk 만 남는다. 그 역수 T=1/kT=1/k 가 온도이고, 유도가 남긴 유일한 자유도다.
  • T→0T \to 0 은 argmax, T→∞T \to \infty 는 균등분포다. 온도를 낮추면 꼬리부터 죽는다.
  • 최댓값을 빼는 구현은 조건 ④ 그 자체다. 답은 안 변하고 오버플로만 막힌다.
  • 로짓은 상수 하나 차이로 로그 확률이고, 그 상수가 log-sum-exp다. 그래서 점수 차이가 로그 확률비다.
  • 선택지가 둘이면 sigmoid로 줄어든다. 이진 분류의 출력이 하나뿐인 이유다.

블록 하나에서 세 번 만난다고 했던 그 함수로 돌아가 봅시다. 어텐션이 softmax를 쓰는 것은 «각 위치에 얼마씩 주의를 나눠 줄 것인가»가 합이 1인 배분 문제이기 때문이고, dk\sqrt{d_k} 로 나누는 것은 점수의 폭을 조절하는 일이니 사실상 온도를 정하는 조작입니다. MoE 라우터도 전문가들 사이의 배분이라 같은 함수를 씁니다. 셋 다 조건 넷을 그대로 요구하니 셋 다 같은 함수가 나온 것입니다.

다음 글은 이 확률분포에서 실제로 하나를 뽑는 일을 다룹니다. 온도로 분포를 바꾸는 것까지가 이 글이었다면, top-k와 top-p가 확률질량의 어디를 잘라 내는지, 그리고 로짓에 잡음을 더해 argmax만 취해도 정확히 이 분포에서 뽑은 것과 같아지는 트릭이 이어집니다.


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

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