수학

MATH / 중급 25번

범주분포에서 뽑기: 온도·top-k·top-p·Gumbel-max의 수학

확률 벡터 하나에서 토큰 하나를 뽑는 그 한 줄을 끝까지 풉니다. 역변환 표집이 왜 맞는지, top-k와 top-p가 확률질량의 어디를 자르는지, 그리고 로짓에 잡음을 더해 argmax만 취해도 softmax 표집과 정확히 같아진다는 Gumbel-max 트릭을 증명합니다.

PALDYN Team35 MIN READ

언어모델의 디코딩 루프는 두 줄이 전부입니다. 모델이 토큰마다 점수, 곧 로짓을 뱉고, 그것을 확률로 바꿔 하나를 뽑습니다.

probs = torch.softmax(logits / temperature, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)

지난 글이 첫 줄을 맡았습니다. 이 글은 둘째 줄입니다 — 확률 벡터 하나에서 실제로 하나를 뽑는 일입니다. 그 안에 든 것이 생각보다 많습니다. 균등난수 하나로 임의의 이산분포에서 뽑는 방법, 첫 줄에 슬쩍 끼어 있는 temperature가 분포를 얼마나 바꾸는지, top-k와 top-p가 확률질량의 어디를 잘라 내는지, 그리고 잡음을 더한 뒤 최댓값만 골라도 정확히 같은 분포가 나온다는 놀라운 등식까지입니다. 글 내내 같은 여섯 토큰의 분포 하나를 들고 다니며 숫자로 따라갑니다.

역변환 표집

누적확률

우리에게 있는 난수 발생기는 대개 하나뿐입니다. U∼Uniform(0,1)U \sim \mathrm{Uniform}(0,1), 즉 0과 1 사이에서 고르게 하나를 뽑아 주는 것입니다. 이것으로 임의의 확률 p=(p1,…,pn)\mathbf{p} = (p_1, \dots, p_n) 을 따르는 뽑기를 만들어야 합니다.

발상은 확률을 길이로 바꾸는 것입니다. 확률을 순서대로 이어 붙이면 길이 1짜리 자가 하나 생깁니다. 그 자 위에 아무 점이나 고르게 찍으면, 어떤 칸에 떨어질 확률이 그 칸의 길이입니다.

누적확률의 자 위에 균등난수를 찍는 역변환 표집

칸의 경계를 적어 둔 것이 누적확률 Fi=p1+⋯+piF_i = p_1 + \cdots + p_i 입니다. 예를 들어 여섯 토큰의 확률이 0.40, 0.25, 0.15, 0.10, 0.06, 0.040.40,\, 0.25,\, 0.15,\, 0.10,\, 0.06,\, 0.04 이면 누적은 0.40, 0.65, 0.80, 0.90, 0.96, 1.000.40,\, 0.65,\, 0.80,\, 0.90,\, 0.96,\, 1.00 입니다.

역변환 표집. u∼Uniform(0,1)u \sim \mathrm{Uniform}(0,1) 을 뽑고 Fi≥uF_i \geq u 인 가장 작은 ii 를 고른다.

u=0.72u = 0.72 라면 0.650.65 는 넘고 0.800.80 은 안 넘으므로 세 번째 토큰입니다.

맞다는 확인은 한 줄입니다. ii 가 뽑히는 것은 uu 가 구간 (Fi−1,Fi](F_{i-1}, F_i] 에 떨어지는 것과 같은데, 균등분포에서 어떤 구간에 떨어질 확률은 그 구간의 길이이므로

P(i 가 뽑힘)=Fi−Fi−1=piP(i \text{ 가 뽑힘}) = F_i - F_{i-1} = p_i

입니다. 구현도 그대로입니다 — 누적합을 만드는 데 O(n)O(n), 이분탐색으로 자리를 찾는 데 O(log⁡n)O(\log n) 이라, 어휘가 5만 개여도 탐색은 열여섯 번이면 끝납니다(216=655362^{16} = 65536). torch.multinomial이 하는 일이 이것입니다.

별칭법

누적합을 만드는 O(n)O(n) 이 뽑을 때마다 드는 것은 아닙니다. 같은 분포에서 여러 번 뽑는다면 누적합을 한 번 만들어 두고 이분탐색만 되풀이하면 되고, 한 걸음 더 가면 한 번 뽑는 데 O(1)O(1) 까지 내려갑니다. 그 방법이 별칭법(alias method)입니다. 확률 nn 개를 높이 1짜리 기둥 nn 개로 다시 쌓되, 기둥마다 주인이 많아야 둘이 되게 나눠 두는 것입니다.

세 토큰 0.5, 0.3, 0.20.5,\, 0.3,\, 0.2 로 해 봅니다. 각각에 n=3n = 3 을 곱하면 1.5, 0.9, 0.61.5,\, 0.9,\, 0.6 입니다. 둘째 기둥은 0.90.9 라 0.10.1 이 모자라니 첫째에게서 0.10.1 을 빌리고, 셋째 기둥은 0.40.4 를 빌립니다. 첫째에게 남은 것은 1.5−0.1−0.4=1.01.5 - 0.1 - 0.4 = 1.0 이라 제 기둥을 혼자 채웁니다. 뽑을 때는 기둥 하나를 1/31/3 씩 고르게 고르고, 그 기둥 안에서 균등난수 하나로 주인과 빌려 준 쪽을 가립니다. 첫째 토큰이 나올 확률은 (1.0+0.1+0.4)/3=0.5(1.0 + 0.1 + 0.4)/3 = 0.5 로 맞아떨어집니다. 난수 둘과 비교 한 번이라 어휘 크기와 상관없는 O(1)O(1) 입니다.

확률 0.5·0.3·0.2를 높이 1짜리 기둥 셋으로 다시 쌓는 별칭법

다만 언어모델의 디코딩에서는 이 이득이 없습니다. 토큰을 하나 뽑을 때마다 모델이 새 로짓을 내므로 분포가 매번 바뀌고, 한 번 쓰고 버릴 표를 만드느라 O(n)O(n) 을 치르게 됩니다. 별칭법은 추천 시스템의 음성 표본처럼 고정된 분포에서 수백만 번 뽑는 자리에서 값어치가 있습니다.

누적합의 끝

수학에서는 Fn=1F_n = 1 이지만 컴퓨터에서는 아닐 수 있습니다. 표준편차 3인 로짓 5만 개를 float32로 softmax한 뒤 차례로 더해 보면, 한 번 해 본 결과 끝이 0.99995390.9999539 에서 멈췄습니다. 작은 수를 큰 누적에 더할 때마다 반올림이 조금씩 새기 때문입니다. 이러면 uu 가 0.99995390.9999539 보다 클 때 Fi≥uF_i \ge u 인 ii 가 없어 이분탐색이 배열 밖의 자리를 돌려줍니다. 확률로는 약 2만 번에 한 번이라, 토큰 수백만 개를 뽑는 서버에서는 하루에도 여러 번 터집니다.

고치는 법은 둘입니다. 균등난수를 [0,Fn)[0, F_n) 에서 뽑도록 uu 에 FnF_n 을 곱하거나, 마지막 칸의 경계를 강제로 1로 덮어씁니다. 앞쪽은 모든 칸을 같은 비율로 늘리는 셈이라 분포를 해치지 않고, 뒤쪽은 어긋난 몫을 마지막 토큰에 몰아줍니다. 아래 코드는 앞쪽을 씁니다.

float32 언더플로

더 조용한 고장도 있습니다. float32가 나타낼 수 있는 가장 작은 양수는 약 1.4×10−451.4 \times 10^{-45}, 곧 e−103.3e^{-103.3} 쯤입니다. 그래서 1등보다 로짓이 약 103.3 넘게 낮은 토큰은 softmax 안에서 ezi−zmax⁡e^{z_i - z_{\max}} 가 정확히 0이 됩니다. 이렇게 너무 작은 수가 0으로 내려앉는 일을 언더플로라고 합니다. 그 토큰은 이론상 확률이 있는데 영영 안 뽑힙니다.

누적합에서는 문턱이 훨씬 높습니다. float32는 0.9 근처에서 이웃한 두 수의 간격이 6×10−86 \times 10^{-8} 이라, 누적이 0.9에 이른 뒤에 오는 10−810^{-8} 짜리 확률은 더해도 합이 안 바뀝니다. 칸의 폭이 0이 되니 역시 안 뽑힙니다. 대부분은 해가 없고 오히려 반가운 절단이지만, 꼬리의 확률을 정확히 재야 하는 평가에서는 로그 확률로 계산하고 float64로 누적해야 합니다.

온도

확률의 비

첫 줄의 temperature는 로짓을 TT 로 나눕니다. 지난 글에서 유도했듯 pi∝ezi/Tp_i \propto e^{z_i/T} 이고, 이 TT 가 온도입니다. 온도가 무엇을 하는지는 확률 둘의 비를 보면 한 줄로 드러납니다.

pi(T)pj(T)=e(zi−zj)/T=(pi(1)pj(1))1/T\frac{p_i(T)}{p_j(T)} = e^{(z_i - z_j)/T} = \left(\frac{p_i(1)}{p_j(1)}\right)^{1/T}

T=1T = 1 일 때의 비를 1/T1/T 제곱한 것이 온도 TT 에서의 비입니다. 우리 분포에서 1등과 2등의 비는 0.40/0.25=1.60.40 / 0.25 = 1.6 입니다. T=0.5T = 0.5 로 낮추면 제곱이 되어 2.562.56, T=2T = 2 로 높이면 제곱근이 되어 1.2651.265 입니다. 1등과 6등은 더 극적입니다. 원래 1010 배이던 차이가 T=0.5T = 0.5 에서 100100 배로 벌어지고 T=2T = 2 에서 3.163.16 배로 좁혀집니다. 온도를 절반으로 내리는 것은 비를 조금 키우는 것이 아니라 지수를 두 배로 만드는 일이라, 차이가 큰 쌍일수록 더 크게 벌어집니다.

세 온도의 비교

로짓을 zi=ln⁡piz_i = \ln p_i 로 두면 T=1T = 1 이 원래 분포를 그대로 돌려줍니다. 여기에 셋을 넣은 결과가 아래 그림입니다.

같은 로짓에 온도 0.5·1·2를 넣어 만든 확률 막대 세 벌

T=0.5T = 0.5 에서는 0.615, 0.240, 0.086, 0.038, 0.014, 0.0060.615,\, 0.240,\, 0.086,\, 0.038,\, 0.014,\, 0.006 이고 T=2T = 2 에서는 0.277, 0.219, 0.170, 0.139, 0.107, 0.0880.277,\, 0.219,\, 0.170,\, 0.139,\, 0.107,\, 0.088 입니다. 「뾰족해진다」를 한 숫자로 재는 양이 엔트로피 H=−∑ipilog⁡2piH = -\sum_i p_i \log_2 p_i 로, 분포가 얼마나 넓게 퍼져 있는지를 비트로 잽니다. 다음 글의 주제라 여기서는 값만 봅니다. 셋은 차례로 1.542, 2.201, 2.4761.542,\, 2.201,\, 2.476 비트이고, 여섯 개가 똑같을 때의 상한은 log⁡26=2.585\log_2 6 = 2.585 비트입니다.

꼬리가 쥔 몫으로 봐도 같습니다. 상위 다섯의 누적은 0.994, 0.960, 0.9120.994,\, 0.960,\, 0.912 라, 6등이 가진 몫이 T=0.5T = 0.5 에서는 0.6%0.6\% 인데 T=2T = 2 에서는 8.8%8.8\% 입니다. 열한 번에 한 번꼴로 가장 그럴듯하지 않은 토큰이 나온다는 뜻입니다.

두 극한

TT 를 0으로 보내면 1등과 나머지의 비 (p1/pj)1/T(p_1/p_j)^{1/T} 가 끝없이 커지므로 확률이 전부 1등에게 몰립니다. T=0.1T = 0.1 에서 이미 0.991, 0.009, 0.0001,…0.991,\, 0.009,\, 0.0001, \dots 입니다. 극한은 argmax, 곧 탐욕 디코딩입니다. 그래서 구현들은 대개 temperature=0을 받으면 0으로 나누는 대신 argmax로 따로 처리합니다.

반대로 T→∞T \to \infty 이면 모든 비가 11 로 가서 균등분포가 됩니다. T=10T = 10 에서 0.1870.187 부터 0.1480.148 까지로 이미 1/6≈0.1671/6 \approx 0.167 언저리에 모입니다. 이때 모델이 배운 것은 거의 다 지워집니다.

순위 불변

x↦x/Tx \mapsto x/T 는 T>0T > 0 이면 순서를 지키는 함수이므로 로짓의 순위가 그대로이고, 그 뒤의 softmax도 순서를 지키므로 확률의 순위도 그대로입니다. 1등은 어느 온도에서나 1등입니다. 또 ez/Te^{z/T} 는 언제나 양수이므로 어떤 토큰도 0이 되지 않습니다. 온도는 모든 토큰의 확률을 조금씩 옮길 뿐 순서도 안 바꾸고 누구도 지우지 않습니다. 만 번에 한 번쯤 엉뚱한 토큰이 튀어나와 문장이 무너지는 일을 확실히 막으려면 꼬리를 실제로 잘라 내는 다른 조작이 필요합니다.

후보 자르기

top-k

top-k 표집. 확률이 큰 순으로 kk 개만 남기고 나머지를 0으로 만든 뒤, 남은 것의 합으로 나눠 다시 정규화한다.

pi′={pi∑j∈Skpji∈Sk0그 밖p'_i = \begin{cases} \dfrac{p_i}{\sum_{j \in S_k} p_j} & i \in S_k \\[4pt] 0 & \text{그 밖} \end{cases}

SkS_k 는 확률 상위 kk 개의 자리입니다. 위 예에서 k=3k=3 이면 0.40,0.25,0.150.40, 0.25, 0.15 만 남고 합이 0.800.80 이므로 각각을 0.800.80 으로 나눠 0.5, 0.3125, 0.18750.5,\, 0.3125,\, 0.1875 가 됩니다.

남은 것의 합으로 나누는 이 걸음이 재정규화이고, 잘려 나간 0.200.20 만큼의 확률질량이 남은 셋에게 비율 그대로 나눠진 셈입니다. 1등과 2등의 비는 자르기 전에도 뒤에도 1.61.6 입니다.

top-p

top-k의 불편한 점은 kk 가 고정이라는 것입니다. 모델이 확신에 차 있어 1등이 0.95를 가진 자리에서도 3등까지 살려 두고, 반대로 정말 애매해서 스무 개가 고만고만한 자리에서도 셋만 남깁니다.

top-p 표집(핵 표집, nucleus sampling). 확률이 큰 순으로 더해 가다가 누적이 처음으로 pp 이상이 되는 자리까지만 남기고 재정규화한다.

같은 분포에서 p=0.9p = 0.9 라면 누적이 0.40→0.65→0.80→0.900.40 \to 0.65 \to 0.80 \to 0.90 이므로 네 개가 남습니다. 그런데 분포가 0.22, 0.20, 0.19, 0.14, 0.13, 0.120.22,\, 0.20,\, 0.19,\, 0.14,\, 0.13,\, 0.12 처럼 평평하면 누적이 0.880.88 까지 가도 모자라 여섯 개 전부가 남습니다.

top-k와 top-p가 남기는 후보의 개수가 달라지는 모습

같은 설정으로도 top-p는 남기는 개수를 스스로 바꿉니다. 모델이 확신할 때는 좁게, 헷갈릴 때는 넓게 — 그래서 top-p가 기본값 자리를 차지했습니다.

min-p

top-p에도 약점이 있습니다. 1등이 0.400.40 인 분포와 0.950.95 인 분포에서 「누적 0.9」는 전혀 다른 깊이까지 파 내려갑니다. 앞쪽은 넷을, 뒤쪽은 1등 하나를 남기는데, 온도를 높여 분포가 평평해지면 누적이 느리게 차서 꼬리까지 딸려 들어옵니다.

min-p 표집은 문턱을 1등의 확률에 묶습니다. 1등 확률에 비율 mm 을 곱한 값보다 작은 토큰을 버리는 것입니다. m=0.2m = 0.2 로 우리 분포에 걸면 문턱이 0.2×0.40=0.080.2 \times 0.40 = 0.08 이라 0.100.10 까지 넷이 남습니다. 평평한 분포에서는 문턱이 0.2×0.22=0.0440.2 \times 0.22 = 0.044 로 내려가 여섯이 다 남습니다. 모델이 확신할수록 1등이 커지고 문턱도 따라 올라가니, 확신의 정도를 문턱이 직접 읽습니다.

적용 순서

조작을 겹쳐 쓰면 순서가 결과를 바꿉니다. 흔한 순서는 로짓에 손대는 조작을 먼저, 확률을 자르는 조작을 나중에 두는 것입니다.

반복 벌점은 이미 나온 토큰의 로짓을 깎는 조작입니다. 확률이 아니라 점수에 손을 댄다는 것이 요점입니다. 1등 토큰이 이미 나왔다고 로짓에서 11 을 빼면 그 가중치가 e−1e^{-1} 배가 되어 0.40→0.1470.40 \to 0.147 이고, 전체 합 0.7470.747 로 다시 나누면 0.1970.197 로 떨어져 2등(0.3350.335)에게 자리를 내줍니다. 로짓에서 빼는 것은 확률에 곱하는 것과 같으므로 나머지 토큰 사이의 비는 하나도 안 변합니다.

그다음이 온도, 그다음이 자르기입니다. 온도는 순위를 안 바꾸므로 top-k가 고르는 집합은 온도와 상관없이 같습니다. 그러나 질량을 옮기므로 top-p가 남기는 개수는 달라집니다. 앞의 누적을 다시 보면 T=0.5T = 0.5 에서 0.615→0.855→0.9420.615 \to 0.855 \to 0.942 라 셋, T=1T = 1 에서 넷, T=2T = 2 에서 다섯입니다. min-p도 마찬가지로 m=0.2m = 0.2 에서 둘, 넷, 여섯입니다. 온도를 올려 다양하게 하려다 top-p가 꼬리를 더 많이 들여보내 효과가 두 배가 되는 셈이라, 둘을 함께 조정할 때는 이 겹침을 셈에 넣어야 합니다. 마지막에 한 번 재정규화하고 역변환 표집으로 뽑습니다.

Gumbel-max 트릭

여기까지는 확률을 만들고 자로 재서 뽑는 한 갈래였습니다. 전혀 다른 길이 하나 있습니다.

Gumbel-max 트릭. 로짓 ziz_i 마다 독립인 Gi∼Gumbel(0,1)G_i \sim \mathrm{Gumbel}(0,1) 을 더하고 가장 큰 것의 자리를 고르면, 그 자리는 정확히 softmax(z)\mathrm{softmax}(\mathbf{z}) 에서 뽑은 것과 같은 분포를 따른다.

arg⁡max⁡i (zi+Gi) ∼ Categorical(softmax(z))\arg\max_i \,(z_i + G_i) \ \sim \ \mathrm{Categorical}\big(\mathrm{softmax}(\mathbf{z})\big)

확률로 바꾸지도, 누적합을 만들지도 않았는데 결과가 정확히 같다는 주장입니다.

Gumbel 분포

정의. Gumbel 분포 Gumbel(0,1)\mathrm{Gumbel}(0,1) 의 누적분포함수는 F(g)=exp⁡(−e−g)F(g) = \exp(-e^{-g}) 이고 밀도는 f(g)=e−gexp⁡(−e−g)f(g) = e^{-g}\exp(-e^{-g}) 이다.

이름은 홍수 수위처럼 「해마다 가장 큰 값」을 연구한 통계학자 에밀 굼벨에게서 왔습니다. 꼬리가 지수적으로 줄어드는 변수를 많이 모아 최댓값을 취하면, 그 최댓값이 이 분포로 모입니다. 태생이 최댓값의 분포이니 argmax와 맞물리는 것이 우연이 아닙니다.

뽑는 것은 한 줄입니다. U∼Uniform(0,1)U \sim \mathrm{Uniform}(0,1) 에 대해 G=−log⁡(−log⁡U)G = -\log(-\log U) 로 두면 됩니다. 확인해 봅시다.

P(G≤g)=P(−log⁡(−log⁡U)≤g)=P(−log⁡U≥e−g)=P(U≤e−e−g)=exp⁡(−e−g)P(G \le g) = P\big(-\log(-\log U) \le g\big) = P\big(-\log U \ge e^{-g}\big) = P\big(U \le e^{-e^{-g}}\big) = \exp(-e^{-g})

가운데에서 부등호가 뒤집힌 것은 양변에 음수를 곱했기 때문이고, 마지막은 균등분포의 누적분포가 P(U≤t)=tP(U \le t) = t 이기 때문입니다. 역변환 표집을 연속분포에 그대로 쓴 것입니다.

로짓에 Gumbel 잡음을 더하고 argmax를 취하는 절차

증명

ii 번이 이길 확률을 구합니다. GiG_i 의 값을 gg 로 고정해 두고 나머지가 전부 그보다 작을 확률을 곱한 뒤, gg 에 대해 적분하는 순서입니다.

P(i 가 최대)=∫−∞∞f(g)∏j≠iP(zj+Gj<zi+g) dgP(i \text{ 가 최대}) = \int_{-\infty}^{\infty} f(g) \prod_{j \neq i} P\big(z_j + G_j < z_i + g\big)\, dg

P(Gj<g+zi−zj)=F(g+zi−zj)=exp⁡ ⁣(−e−gezj−zi)P(G_j < g + z_i - z_j) = F(g + z_i - z_j) = \exp\!\big(-e^{-g} e^{z_j - z_i}\big)

이므로 곱은 지수의 합으로 모입니다.

∏j≠iF(g+zi−zj)=exp⁡ ⁣(−e−g∑j≠iezj−zi)\prod_{j \neq i} F(g + z_i - z_j) = \exp\!\Big(-e^{-g} \sum_{j \neq i} e^{z_j - z_i}\Big)

여기에 f(g)=e−gexp⁡(−e−g)f(g) = e^{-g}\exp(-e^{-g}) 를 곱하면, 앞의 exp⁡(−e−g)\exp(-e^{-g}) 가 j=ij = i 항을 채워 줍니다 — ezi−zi=1e^{z_i - z_i} = 1 이니까요. 그래서 합의 범위가 j≠ij \neq i 에서 전체로 넓어집니다.

S=∑jezj−zi라 두면P(i 가 최대)=∫−∞∞e−gexp⁡(−Se−g) dgS = \sum_{j} e^{z_j - z_i} \quad \text{라 두면} \quad P(i \text{ 가 최대}) = \int_{-\infty}^{\infty} e^{-g} \exp(-S e^{-g})\, dg

이제 t=e−gt = e^{-g} 로 치환합니다. dt=−e−g dgdt = -e^{-g}\,dg 이므로 e−g dg=−dte^{-g}\,dg = -dt 이고, gg 가 −∞-\infty 에서 ∞\infty 로 갈 때 tt 는 ∞\infty 에서 00 으로 갑니다.

∫∞0e−St (−dt)=∫0∞e−St dt=1S\int_{\infty}^{0} e^{-St}\,(-dt) = \int_{0}^{\infty} e^{-St}\,dt = \frac{1}{S}

마지막으로 SS 를 풀어 쓰면 끝입니다.

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

정확히 softmax입니다. 근사가 아니라 등식입니다. 지수적으로 생긴 Gumbel 분포의 꼬리가 softmax의 eze^{z} 와 딱 맞물려 적분이 깔끔하게 떨어지는 것이 이 트릭의 전부입니다.

최댓값의 분포

이긴 자리뿐 아니라 이긴 값 M=max⁡i(zi+Gi)M = \max_i (z_i + G_i) 도 분포를 갖습니다. 위와 같은 곱을 하면 P(M≤m)=exp⁡ ⁣(−e−(m−LSE))P(M \le m) = \exp\!\big(-e^{-(m - \mathrm{LSE})}\big) 이 나옵니다. 여기서 LSE=log⁡∑jezj\mathrm{LSE} = \log \sum_j e^{z_j} 는 softmax의 분모에 로그를 씌운 값, 곧 logsumexp입니다. 최댓값이 다시 Gumbel 분포이고 위치만 LSE로 옮겨 간 것입니다. 그래서 평균은 LSE+γ\mathrm{LSE} + \gamma 이고 γ≈0.5772\gamma \approx 0.5772 는 오일러 상수입니다. 우리 로짓 zi=ln⁡piz_i = \ln p_i 는 ∑jezj=1\sum_j e^{z_j} = 1 이라 LSE가 0이므로, 최댓값의 평균이 0.57720.5772 여야 합니다. 40만 번 뽑아 보면 0.57530.5753 이 나옵니다. 정규화 상수를 모르는 채 그 로그를 잡음으로 추정하는 길이 여기서 열립니다.

모든 로짓에 같은 상수 cc 를 더하면 모든 zi+Giz_i + G_i 가 똑같이 cc 만큼 올라가므로 1등의 자리는 그대로이고 최댓값만 cc 만큼 옮겨 갑니다. softmax가 상수 이동에 불변이던 성질이 여기서도 그대로 섭니다. 위 실험에서 로짓에 5를 더하면 뽑힌 자리는 40만 번 모두 같았고 최댓값의 평균만 5.57535.5753 이 됐습니다. 온도도 같은 눈으로 읽힙니다. arg⁡max⁡(zi/T+Gi)=arg⁡max⁡(zi+TGi)\arg\max (z_i/T + G_i) = \arg\max (z_i + T G_i) 이므로, 온도를 올리는 것은 로짓은 그대로 두고 잡음을 TT 배로 키우는 일입니다. 잡음을 고정해 두고 TT 만 바꿔 보면 1등의 로짓 우위가 잡음 차이보다 작아지는 자리에서 결과가 갈립니다.

Gumbel top-k

잡음을 한 번 더한 뒤 argmax 하나가 아니라 상위 kk 개를 취하면, 그 kk 개는 한 번 뽑은 것을 빼고 다시 뽑기를 kk 번 되풀이한 것과 같은 분포를 따릅니다. 한 번 뽑은 것을 다시 넣지 않는 이 방식이 비복원 추출이고, 잡음 한 번으로 그것을 얻는 방법이 Gumbel top-k입니다.

우리 분포에서 1등이 먼저, 2등이 그다음에 나올 확률은 비복원 추출로 0.40×0.25/0.60=0.16670.40 \times 0.25/0.60 = 0.1667 입니다. 첫 뽑기 뒤 남은 질량이 0.600.60 이라 2등의 몫이 0.25/0.600.25/0.60 으로 커진 것입니다. 복원 추출이라면 0.40×0.25=0.100.40 \times 0.25 = 0.10 이었을 것입니다. 20만 번 세어 보면 0.16630.1663 이 나옵니다. 빔 서치에서 후보를 겹치지 않게 여럿 뽑거나 확률적 빔 서치를 할 때 이 성질을 씁니다.

쓰임은 둘이 더 있습니다. 로짓만 있으면 되고 ∑jezj\sum_j e^{z_j} 를 계산할 필요가 없으니 후보를 여러 기계에 나눠 둔 자리에서 합을 모으는 비용을 건너뜁니다. 그리고 잡음을 미리 뽑아 두면 같은 로짓에 대해 언제나 같은 토큰이 나오므로 뽑기가 재현됩니다. 무작위성을 어느 자리에서 흔들 것인가의 문제로 옮겨 놓은 셈입니다.

Gumbel-softmax

연속 완화

argmax에는 치명적인 결함이 있습니다. 기울기가 흐르지 않습니다. 로짓을 조금 바꿔도 1등이 그대로면 출력이 하나도 안 변하고, 어느 순간 1등이 바뀌면 출력이 툭 튑니다. 미분이 거의 모든 곳에서 0이고 나머지에서 정의되지 않으니 역전파가 여기서 끊깁니다.

지난 글에서 softmax가 사실은 soft-argmax이고 온도를 낮추면 argmax에 가까워진다고 했습니다. 이 글의 온도 절에서 본 T→0T \to 0 극한이 바로 그것입니다. 그 성질을 그대로 씁니다.

Gumbel-softmax. arg⁡max⁡\arg\max 대신 y=softmax((z+G)/τ)\mathbf{y} = \mathrm{softmax}\big((\mathbf{z} + \mathbf{G})/\tau\big) 를 쓴다.

τ→0\tau \to 0 이면 y\mathbf{y} 는 원-핫 벡터로 수렴하므로 극한에서 정확히 Gumbel-max 트릭, 즉 정확한 표집입니다. 그리고 τ>0\tau > 0 인 동안에는 softmax가 매끄러운 함수라 기울기가 z\mathbf{z} 까지 흘러갑니다. 이산 선택을 매끄러운 함수로 바꿔 끼우는 이런 조작을 완화라고 부릅니다.

온도 τ에 따라 원-핫에 가까워지는 Gumbel-softmax

직통 추정량

τ>0\tau > 0 인 y\mathbf{y} 는 원-핫이 아니라 (0.7,0.2,0.1)(0.7, 0.2, 0.1) 같은 섞인 벡터입니다. 다음 층이 정말 토큰 하나를 받아야 한다면 이것으로는 안 됩니다. 그래서 순전파에서는 y\mathbf{y} 를 원-핫으로 반올림해 진짜 이산 선택을 하고, 역전파에서는 그 반올림이 없었던 것처럼 y\mathbf{y} 의 기울기를 그대로 흘립니다. 이 조합이 직통 추정량(straight-through estimator)입니다. 코드로는 y_hard - y.detach() + y 한 줄입니다 — 값은 y_hard이고 기울기는 y의 것입니다.

대가는 기울기가 틀린다는 것입니다. 순전파가 계산한 함수는 반올림된 원-핫인데 역전파는 반올림 전의 함수를 미분하므로, 흘러가는 기울기는 실제 손실의 기울기와 평균적으로도 어긋납니다. 이런 기울기를 편향된 추정량이라 합니다. 그래도 방향이 대체로 맞아 VQ-VAE의 코드북 선택이나 이산 잠재변수를 학습하는 자리에서 널리 씁니다.

τ 스케줄

τ\tau 를 고르는 일은 저울질입니다. 크면 기울기는 잘 흐르지만 y\mathbf{y} 가 원-핫과 멀어 실제 이산 선택과 딴판이고, 작으면 그 반대입니다. 그래서 학습 초반에 1 언저리로 두었다가 지수적으로 줄여 0.5 안팎에서 멈추는 스케줄을 흔히 씁니다.

너무 빨리 줄이면 기울기가 죽습니다. softmax의 미분은 yi(δij−yj)/τy_i(\delta_{ij} - y_j)/\tau 꼴인데, τ\tau 가 작아 y\mathbf{y} 가 원-핫에 붙으면 yi(1−yi)y_i(1 - y_i) 가 거의 0이 되어 대부분의 걸음에서 기울기가 사라집니다. 어쩌다 두 후보가 비등한 걸음에서만 1/τ1/\tau 배로 크게 튀므로 분산이 커지고 학습이 흔들립니다. 온도 절의 두 극한이 여기서는 두 고장이 됩니다.

잡음 없는 완화

표집 자체가 필요 없는 경우도 많습니다. 다음 계산이 선택의 기댓값만 쓴다면 잡음 없이 softmax(z/τ)\mathrm{softmax}(\mathbf{z}/\tau) 로 가중평균을 내면 됩니다. 어텐션이 정확히 이 길입니다. 키 하나를 고르는 대신 모든 키를 softmax 가중치로 섞습니다. 결정적이라 분산이 없고 기울기도 편향되지 않지만, 한 번의 순전파가 여러 후보를 섞어 버리므로 「정말로 하나만 골랐을 때」의 행동은 학습하지 못합니다. 이산 선택 자체가 중요할 때 Gumbel-softmax가, 섞어도 될 때 이쪽이 맞습니다.

코드로 확인하기

역변환과 Gumbel-max

import numpy as np

rng = np.random.default_rng(0)
p = np.array([0.40, 0.25, 0.15, 0.10, 0.06, 0.04])

# ① 역변환 표집 — 누적합과 이분탐색
def inverse_sample(p, n, rng):
    F = np.cumsum(p)
    u = rng.random(n) * F[-1]             # 끝이 1이 아니어도 안전
    return np.searchsorted(F, u)          # F_i >= u 인 가장 작은 i

s = inverse_sample(p, 200_000, rng)
print(np.bincount(s, minlength=6) / 200_000)
# [0.400135 0.2504   0.15146  0.099295 0.059275 0.039435]

z = np.log(p)                             # 로짓은 상수 차이로 로그 확률
# ③ Gumbel-max — 정규화도 누적합도 없이 argmax 하나
g = -np.log(-np.log(rng.random((200_000, 6))))
picks = (z + g).argmax(axis=1)
print(np.bincount(picks, minlength=6) / 200_000)
# [0.39863  0.24971  0.150585 0.099625 0.0614   0.04005 ]

두 결과가 같은 자리에 떨어졌습니다. 확률로 바꾸지도 누적합을 만들지도 않고 잡음을 더해 최댓값만 골랐는데도 그렇습니다. 둘째 줄의 0.06140.0614 가 0.060.06 과 조금 벌어진 것은 표본 20만 개의 흔들림입니다. 그 칸의 표준오차가 약 0.00050.0005 라 세 배 안쪽입니다.

온도와 Gumbel top-k

# ② 온도 셋 — 확률, 엔트로피, top-p(0.9)가 남기는 개수
for T in (0.5, 1.0, 2.0):
    q = np.exp(z / T); q /= q.sum()
    H = -(q * np.log2(q)).sum()
    kept = np.searchsorted(np.cumsum(q), 0.9 - 1e-12) + 1
    print(T, q.round(3), round(H, 3), kept)
# 0.5 [0.615 0.24  0.086 0.038 0.014 0.006] 1.542 3
# 1.0 [0.4  0.25 0.15 0.1  0.06 0.04] 2.201 4
# 2.0 [0.277 0.219 0.17  0.139 0.107 0.088] 2.476 5

# ④ Gumbel top-k — 상위 둘의 순서쌍을 세기
top2 = np.argsort(-(z + g), axis=1)[:, :2]
print(np.mean((top2[:, 0] == 0) & (top2[:, 1] == 1)))   # 0.4 * 0.25/0.6
# 0.16633

온도 셋의 확률·엔트로피·남는 개수가 그림과 같고, Gumbel top-k의 순서쌍 빈도가 비복원 추출의 0.16670.1667 과 맞습니다. top-p 줄의 0.9 - 1e-12는 누적이 부동소수점으로 0.90.9 를 아주 조금 밑돌아 한 칸 더 가는 일을 막으려는 것으로, 앞의 누적합의 끝과 같은 고장입니다.

디코딩 루프의 두 줄로 돌아가 봅시다. 첫 줄이 점수를 확률로, 둘째 줄이 확률에서 하나를 뽑습니다. 사용자가 만지는 temperature·top_k·top_p·min_p와 반복 벌점은 전부 그 사이에 끼어 분포의 모양을 바꾸는 조작이었습니다. 온도는 비를 1/T1/T 제곱해 순서를 지킨 채 뾰족하게 하거나 평평하게 하고, 자르기는 꼬리를 실제로 0으로 만들고, 벌점은 점수 쪽에서 특정 토큰을 누릅니다. 어느 것도 뽑는 방법 자체를 바꾸지는 않았고, 뽑는 방법은 누적확률의 자 위에 난수를 찍든 로짓에 Gumbel 잡음을 더해 argmax를 취하든 같은 분포를 줍니다.

그런데 이 확률들이 좋은지는 무엇으로 재는가. 온도 절에서 값만 보고 넘어간 엔트로피가 그 답의 절반입니다. 학습 로그에 찍히는 loss와 ppl은 사실 하나의 서로 다른 표현이고, 다음 글은 그 하나가 무엇인지 — 분포의 펼쳐진 정도를 재는 양을 정의하는 데서 시작합니다.


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

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