지난 글까지 어텐션의 순전파와 역전파를 종이 위에서 닫았습니다. 그런데 그 식을 그대로 코드로 옮기면 터집니다.
언어 모델이 다음 토큰 하나를 고르는 장면을 떠올려 봅니다. 마지막 층이 어휘의 낱말마다 점수 하나씩, GPT-2의 어휘라면 50,257개의 로짓을 내놓습니다. 로짓은 softmax에 넣기 전의 날점수이고, softmax가 그것을 확률로 바꿉니다. 학습이 진행되며 정답 자리의 로짓이 수백까지 커지는 일은 흔한데, 그 점수를 곧이곧대로 지수에 넣는 순간 이렇게 됩니다.
p = [math.exp(v) for v in scores] # OverflowError: math range error
는 float64에서 부터 무한대이고, float32에서는 88.72부터입니다. 표현할 수 있는 가장 큰 수를 넘어 무한대가 되는 것을 넘침이라 합니다. 반대쪽도 마찬가지라 가 그만큼 작으면 0으로 잘리고, 모든 항이 0이 되면 분모까지 0이라 NaN이 나옵니다. NaN은 「수가 아님」이라는 부동소수점의 표시이고, 한 번 생기면 그 뒤의 모든 덧셈과 곱셈으로 번집니다.
이 글은 그 자리를 세 갈래로 정리합니다 — 최댓값을 빼는 것이 왜 근사가 아닌지, 마스킹을 왜 0을 곱하지 않고 를 더해서 하는지, 그리고 점수를 한꺼번에 못 보는 상황에서도 정확한 softmax를 얻는 점화식입니다. 마지막 것이 FlashAttention의 밑에 깔린 대수이고, 어텐션 비용 줄이기가 IO 병목과 타일링을 다루므로 여기서는 대수만 맡습니다. 그 사이에 이 값을 「매끄러운 최댓값」으로 보는 눈과, 역전파가 이 값 하나로 확률 행렬 전체를 되살리는 계산을 끼워 넣습니다.
최댓값 빼기
LSE 항등식
먼저 이름을 붙입니다. log-sum-exp(줄여서 LSE)는 지수의 합에 로그를 씌운 값입니다.
softmax의 분모가 정확히 이고, 교차엔트로피 손실이 이므로 실제 학습에서 계산되는 것은 확률이 아니라 이 값입니다. 앞 장면의 50,257개 로짓도 손실로 가는 길에 정확히 한 번 LSE를 지납니다.
이제 아무 상수 을 잡습니다. 지수법칙 을 넣고 을 합 밖으로 뽑습니다.
log-sum-exp 항등식. 모든 실수 에 대해
근사가 아닙니다. 로그의 곱셈 법칙 한 번을 쓴 항등식이라 이 무엇이든 정확히 성립합니다. 그러니 계산에 가장 편한 을 고르면 되고, 그것이 입니다.
을 최댓값으로 잡으면 모든 이므로 입니다. 가장 큰 항은 정확히 1이고 나머지는 그 아래입니다.
- 넘칠 수 없습니다. 지수의 인자가 0 이하라 결과가 1을 못 넘습니다.
- 분모가 0이 될 수 없습니다. 적어도 한 항이 1이므로 합이 항상 1 이상입니다.
- 작은 항이 0으로 잘리는 것은 여전히 일어나지만, 그때 그 항의 진짜 값도 아래라 합에 기여하지 않습니다.
로 확인합니다. 그냥 계산하면 첫 항에서 OverflowError인데, 800을 빼면 이라 합이 입니다.
같은 값들을 softmax에 넣으면 입니다. 점수가 800이든 0이든 확률은 차이만으로 정해지므로, 최댓값을 빼는 것은 답을 바꾸지 않고 계산만 살립니다.
기준점 선택
항등식이 모든 에서 성립한다면 최댓값 말고 평균을 빼도 되지 않을까요. 평균은 한 번 훑으면 나오고 최댓값도 한 번 훑으면 나오니 비용 차이는 없습니다. 차이는 빼고 난 뒤의 폭에 있습니다.
의 평균은 798이고, 빼면 이라 지수가 입니다. 1을 넘는 항이 생겼지만 넘칠 걱정은 없습니다. 점수가 좁게 모여 있으면 평균도 최댓값도 잘 듭니다.
점수가 넓게 퍼지면 다릅니다. 의 평균은 301이고, 빼면 첫 항이 입니다. float64로는 아직 셀 수 있지만 float32의 한계 88.72를 한참 넘어 무한대가 됩니다. 경계는 「최댓값에서 기준점을 뺀 값이 88.72보다 작은가」 하나이고, 기준점을 최댓값으로 잡으면 그 값이 언제나 0이라 점수가 어떻게 퍼져 있든 조건이 저절로 채워집니다. 최댓값을 고르는 까닭은 편해서가 아니라 조건을 입력과 무관하게 만드는 유일한 선택이라서입니다.
두 항의 LSE
항이 둘뿐이면 식을 한 번 더 다듬을 수 있습니다. 큰 쪽을 으로 뽑으면 남는 항은 1과 입니다.
두 점수의 차이가 커지면 둘째 항이 아주 작아집니다. 차이가 20이면 인데, float32는 1 근처에서 보다 가는 차이를 못 적으므로 이 그냥 1이 되고 로그가 0을 냅니다. 작은 항이 통째로 사라진 것입니다. 그래서 라이브러리는 를 1을 더하지 않고 곧바로 계산하는 log1p를 따로 두고, 그러면 같은 자리에서 이 살아남습니다.
쪽에 이미 수백이 서 있으면 이 차이는 결과의 마지막 자릿수 아래라 보통은 드러나지 않습니다. 드러나는 것은 LSE에서 최댓값을 뺀 나머지 자체가 필요한 경우입니다. 대표적인 것이 활성화 함수 softplus로, 이라 정확히 두 항의 LSE입니다.
로그확률
확률을 만들지 않고 로그확률을 곧바로 얻을 수도 있습니다.
오른쪽에는 나눗셈도 도 없습니다. log_softmax가 하는 일이 이 한 줄이고, 37번에서 softmax와 교차엔트로피를 붙여 구현하는 셋째 이유가 이것이었습니다. 확률 가 이면 float32에서는 0으로 잘려 로그가 가 되지만, 로그확률 쪽으로 곧장 가면 이라는 멀쩡한 수가 나옵니다.
매끄러운 최댓값
두 줄 부등식
LSE는 이름과 달리 최댓값과 아주 가까운 수입니다. 합 안의 항 가운데 가장 큰 것이 이고 항이 개이므로, 합은 그 한 항보다 크거나 같고 그 한 항의 배보다 작거나 같습니다. 양쪽에 로그를 씌우면 두 줄이 나옵니다.
에서 확인하면 입니다. 아래쪽이 등호가 되는 것은 나머지 항이 전부 0일 때이고, 위쪽이 등호가 되는 것은 개가 모두 같을 때입니다.
그래서 LSE를 매끄러운 최댓값이라고 부릅니다. 최댓값처럼 가장 큰 항을 따라가되, 꺾이는 자리 없이 둥글게 따라가는 함수라는 뜻입니다. 앞 장면의 어휘 50,257개라면 위아래 폭이 입니다. 로짓이 수백 단위로 움직이는 동안 LSE는 최댓값에서 11 안쪽을 벗어나지 않습니다.
점수 둘 으로 그리면 그 모양이 드러납니다. 은 에서 꺾이는 선이고, 은 그 위를 둥글게 지나갑니다. 두 선의 틈은 꺾인 자리에서 가장 벌어져 이 되고, 에서는 로 틈이 0.127, 에서는 로 틈이 거의 사라집니다. 앞 절의 두 항 공식에서 가 바로 이 틈입니다.
기울기
둥글다는 것은 어디서나 미분할 수 있다는 뜻입니다. 로 미분하면 바깥의 로그가 분모를 만들고, 안쪽 합에서는 한 항만 살아남습니다.
LSE의 기울기가 정확히 softmax입니다. 분모를 미분하면 분자가 튀어나오는 구조라 이렇게 됩니다. 최댓값의 기울기는 가장 큰 자리에만 1이고 나머지는 0인 벡터인데, LSE는 그 1을 점수에 비례해 여러 자리로 나눠 줍니다. softmax가 「매끄러운 argmax」로 불리는 까닭이 여기서 나옵니다. 매끄러운 최댓값을 미분한 것이 매끄러운 argmax입니다.
이 계산이 교차엔트로피의 기울기 를 곧바로 줍니다. 손실이 이니 앞 항을 미분하면 , 뒤 항을 미분하면 정답 자리에만 1인 입니다. 지난 몇 편에서 야코비안으로 돌아서 얻은 결과가 한 줄로 나옵니다.
온도의 두 극단
점수를 온도 로 나눈 뒤 LSE를 구하고 다시 를 곱해 봅니다. 온도는 softmax에 넣기 전에 점수를 나누는 양수이고, 분포를 뾰족하게도 평평하게도 만드는 손잡이입니다. 두 줄 부등식의 자리에 를 넣고 를 곱하면 폭이 배가 됩니다.
를 0으로 보내면 폭 이 사라져 값이 최댓값으로 끼입니다. 에서 이면 800.318, 이면 800.0000045입니다. 이때 softmax는 가장 큰 자리 하나에 1을 몰아주는 argmax가 됩니다.
반대로 를 키우면 폭이 함께 자라 에서 809.21, 에서 907.88로 최댓값을 멀리 떠납니다. 이때 값을 끌고 가는 것은 이고, 그것을 빼고 나면 남는 것이 평균으로 갑니다. 에서 이 세 점수의 평균 798과 같습니다. softmax는 그 끝에서 모든 자리에 씩 나눠 주는 균등분포가 됩니다.
−inf 마스킹
더하기와 곱하기
인과 마스킹은 「 번째 토큰이 자기보다 뒤를 못 보게」 하는 일입니다. 순서대로 생각하면 softmax를 구한 다음 가릴 자리에 0을 곱하는 것이 자연스러워 보입니다. 실제 구현은 그렇게 하지 않고 softmax에 넣기 전에 를 더합니다.
이므로 가린 자리는 분자에서도 분모에서도 정확히 0입니다. 남은 자리들만으로 합이 1이 되므로 따로 정규화할 것이 없습니다.
0을 곱하는 쪽은 왜 안 될까요. 수학적으로는 곱한 뒤 다시 정규화하면 같은 답이 나옵니다. 문제는 그 사이에 계산이 이미 망가진다는 것입니다.
점수가 이고 셋째·넷째를 가려야 한다고 합시다.
0을 곱하는 쪽은 가릴 자리가 아직 살아 있으므로 최댓값이 900이고, 그것을 빼면 살릴 자리들은 이 되어 으로 전부 잘립니다. softmax 결과가 이고, 마스크를 곱하면 이며, 정규화하려는 순간 입니다.
를 더하는 쪽은 가릴 자리가 이미 라 최댓값이 살아 있는 것들 중에서 정해집니다. 이므로 이고, 결과는 입니다.
차이는 최댓값 자리를 누가 가져가는가입니다. 가려야 할 점수가 가장 크면 그것이 이 되고, 그러면 정작 살려야 할 자리들이 전부 0으로 잘립니다. 학습이 진행되며 점수가 커지는 것을 막을 방법이 없으니 이것은 언제든 일어날 수 있는 일입니다.
전부 가려진 행
를 더하는 데도 함정이 하나 있습니다. 한 행이 전부 이면 NaN이 나옵니다. 그 행의 최댓값이 이고
이기 때문입니다. 무한대끼리의 뺄셈은 정의되지 않습니다. 실제로 일어나는 경우가 둘 있습니다 — 패딩만 있는 행, 그리고 인과 마스크와 다른 마스크를 겹쳐서 볼 수 있는 자리가 하나도 안 남은 행입니다. 그래서 구현은 대신 아주 큰 음수를 쓰거나, 전부 가려진 행을 미리 찾아 따로 처리합니다.
큰 음수를 쓰면 한 행이 전부 가려져도 NaN 대신 균등분포가 나옵니다. float32에서 를 더하면 점수 몇 개쯤의 차이는 반올림에 묻혀 네 자리가 모두 같은 이 되고, 최댓값을 빼면 전부 이라 입니다. 틀린 분포지만 NaN처럼 번지지는 않고, 그 행은 어차피 패딩이라 뒤에서 버려집니다. 한 자리라도 살아 있으면 이 0으로 잘려 결과는 를 더한 것과 같습니다.
float16의 큰 음수
「아주 큰 음수」가 얼마나 커야 하는지는 자료형이 정합니다. 요즘 학습과 추론은 16비트 부동소수점인 float16을 자주 쓰는데, 이 형식이 적을 수 있는 가장 큰 유한값이 65,504입니다. 그래서 float32에서 멀쩡하던 를 float16으로 바꾸면 그 순간 가 됩니다. 한 자리라도 살아 있는 행에서는 그래도 괜찮지만, 전부 가려진 행에서는 앞 소절의 NaN이 그대로 돌아옵니다. 직접 돌려 보면 float16에서 를 더한 행이 을 냅니다.
그렇다고 를 더하는 것도 안전하지 않습니다. 점수가 인 자리에 더하면 합이 표현 범위를 넘어 다시 가 됩니다. 이 근처에서 float16의 눈금 간격은 32라, 조금만 넘쳐도 반올림이 무한대 쪽으로 떨어집니다. 그래서 두 가지 방식이 쓰입니다. 하나는 점수에 더하지 않고 가릴 자리를 그 값으로 바꿔 넣는 방식이고, 이러면 넘칠 덧셈 자체가 없습니다. 다른 하나는 여유를 두고 쯤을 더하는 방식으로, 도 이미 0으로 잘리므로 가리는 효과는 같습니다.
마스크의 기울기
순전파만 보면 「0을 곱한 뒤 다시 정규화」와 「 더하기」는 넘치지 않는 한 같은 답을 냅니다. 역전파에서도 같습니다. 다시 정규화하면 남은 확률들의 비가 볼 수 있는 점수들의 차이만으로 정해지므로, 가린 점수를 아무리 흔들어도 출력이 안 변하고 그 자리의 기울기는 0입니다.
재정규화를 빼먹으면 이야기가 달라집니다. 0을 곱하기만 하면 가린 점수가 여전히 softmax의 분모에 들어 있어서, 그 점수를 키우면 볼 수 있는 자리의 가중치가 줄어듭니다. 점수 에 셋째·넷째를 가리고 값 를 섞는 예로 재 봅니다. 입력을 아주 조금 앞뒤로 움직여 출력이 변한 양을 움직인 양으로 나누는 기울기 어림을 유한차분이라 하는데, 그렇게 재면 가린 두 자리의 기울기가 과 로 0이 아닙니다. 모델이 미래 토큰의 점수를 움직여 현재의 출력을 바꾸는 법을 배우게 되는 셈이고, 순전파에서 가린 것이 역전파로 새어 들어옵니다. 를 더한 쪽은 그 자리의 확률이 0이라 softmax 야코비안의 그 줄이 통째로 0이 되어 같은 누출이 구조적으로 막힙니다.
온라인 softmax
블록 점화식
지금까지는 한 행의 점수를 전부 손에 들고 있다고 가정했습니다. 최댓값을 알아야 빼고, 그러려면 다 봐야 하니까요.
그런데 행이 길면 그럴 수 없습니다. 문맥이 128,000 토큰이면 한 행이 128,000개이고, 그것을 통째로 빠른 메모리에 올릴 수 없습니다. 그러면 앞쪽 일부만 보고 계산을 시작해서, 뒤쪽을 볼 때마다 고쳐 나갈 수 있을까요. 이렇게 입력을 한 번 흘려 보내며 고쳐 가는 계산을 온라인 softmax라고 부릅니다.
할 수 있습니다. 블록 하나가 들고 가야 할 것은 셋뿐입니다.
은 그 블록의 최댓값, 은 정규화되지 않은 합, 는 아직 나누지 않은 출력입니다. 이 셋만 있으면 두 블록을 합칠 수 있습니다.
블록 가 각각 , 를 들고 있다고 합시다. 합친 블록의 최댓값은 당연히
입니다. 합은 어떻게 될까요. 정의대로 적고 각 블록의 기준을 으로 옮깁니다.
첫 합에서 로 쪼개고 을 밖으로 뽑으면 안쪽이 그대로 입니다. 둘째 합도 같습니다.
LSE 항등식을 두 번 쓴 것 전부입니다. 기준점을 옮길 때 곱해 주는 것이 이고, 새 기준이 더 크므로 그 인자는 언제나 1 이하입니다. 여기서도 넘칠 수 없습니다.
를 앞의 둘과 마지막 하나로 쪼개 확인합니다.
| 블록 1: | 800 | |
| 블록 2: | 795 | |
| 합친 것 |
한 번에 계산한 값과 정확히 같습니다. 근사가 아니라 같은 수를 다른 순서로 더한 것뿐이라 그렇습니다.
출력까지 붙이면 어텐션 한 행이 스트리밍으로 계산됩니다. 값 벡터가 있는 예로 확인해 봅니다 — 점수 를 두 블록으로 쪼개면
이고 합치면 , , 마지막에 로 나누면 한 번에 계산한 출력과 소수점 열여섯째 자리까지 같습니다.
결합법칙
이 점화식에는 좋은 성질이 둘 있습니다.
- 결합법칙이 성립합니다. 블록을 몇 개로 쪼개든, 어떤 순서로 합치든 결과가 같습니다. 그래서 병렬로 계산한 조각들을 아무 순서로 모아도 됩니다.
- 메모리가 블록 크기에만 달려 있습니다. 행 전체 길이만 한 배열을 만들 필요가 없고, 들고 가는 것은 언제나 셋입니다.
FlashAttention이 하는 일이 정확히 이것입니다 — 점수 행렬을 통째로 만들지 않고, 와 를 블록 단위로 읽어 오면서 위 점화식으로 출력을 갱신합니다. 그 절차가 옳다는 근거가 지금 유도한 세 줄이고, 어떤 블록 크기가 왜 빠른지는 IO 쪽 이야기라 어텐션 비용 줄이기가 맡습니다.
오차와 블록 크기
「정확히 같다」는 실수의 대수에서 하는 말입니다. 부동소수점에서는 더하는 순서가 바뀌면 마지막 자릿수가 달라질 수 있고, 온라인 softmax는 순서를 바꾸는 것이 본업입니다. 합칠 때마다 에 곱셈 한 번과 덧셈 한 번이 더해지고, 각각이 상대오차를 최대 단위 반올림만큼 남깁니다. 단위 반올림은 한 번의 연산이 남길 수 있는 가장 큰 상대오차로, float32에서 약 입니다.
그러니 최악의 경우 오차는 합치는 횟수에 비례해 자랍니다. 행 128,000개를 64개씩 나누면 합치기가 2,000번이라, 어림한 위쪽 한계가 언저리까지 올라갑니다. 실제로 재 보면 훨씬 작습니다. 표준편차 3인 점수 128,000개를 float32로 블록 16·64·128·1,024개씩 합쳐 보면 LSE의 오차가 에서 사이였고, 블록 크기와 나란히 줄어들지도 않았습니다. 반올림 오차가 한쪽으로만 쌓이지 않고 서로 상쇄되기 때문이고, 그래서 구현들은 과 만은 입력이 16비트여도 float32로 들고 갑니다.
블록 크기를 키우면 합치는 횟수가 줄어 최악의 한계가 내려가고, 대신 블록 하나를 올려 둘 빠른 메모리가 그만큼 더 듭니다. 이 맞바꿈에서 실제로 블록 크기를 정하는 쪽은 오차가 아니라 메모리와 속도입니다. 오차는 위에서 본 것처럼 어느 크기에서나 이미 충분히 작습니다.
역전파의 LSE
행별 LSE
역전파는 순전파에서 만든 확률 를 다시 씁니다. softmax 야코비안이 였으니 없이는 한 줄도 못 갑니다. 보통의 구현은 그래서 를 통째로 저장해 둡니다.
그런데 순전파 끝에 행마다 LSE 하나를 남겨 두면 는 필요 없습니다. softmax의 정의에서 분모가 이므로
점수 만 다시 구하면 뺄셈 하나와 지수 하나로 확률이 돌아옵니다. 로 해 보면 이고, , , 로 앞에서 구한 softmax가 그대로 나옵니다. 지수의 인자가 이므로 여기서도 넘칠 일이 없습니다. 두 줄 부등식의 아래쪽이 이 안전을 보증합니다.
온라인 softmax가 끝난 자리에서 로 이 값이 공짜로 나온다는 점도 중요합니다. 새로 계산할 것이 없고, 들고 있던 을 한 수로 접어 저장하기만 하면 됩니다.
저장과 재계산
저장량을 세어 봅니다. 길이 4,096짜리 시퀀스라면 한 헤드의 는 개, float32로 64MiB입니다. LSE는 행마다 하나라 4,096개, 16KiB입니다. 질의가 개이고 키가 개이면 행렬 대신 길이 짜리 벡터 하나를 남기므로 저장이 배 줄어듭니다. 블록의 최댓값에 쓰던 과 헷갈리지 않게 키의 수는 대문자로 적었습니다. 헤드와 층의 수만큼 곱해 보면 이것이 긴 문맥 학습에서 메모리를 가장 크게 잡아먹는 항목 하나를 지우는 일이라는 것이 보입니다.
대가는 계산입니다. 를 되살리려면 를 역전파에서 한 번 더 구해야 하고, 그것은 순전파가 이미 한 행렬곱을 되풀이하는 일입니다. 저장해 두었다 꺼내 쓰는 대신 필요할 때 다시 계산하는 이 방식을 재계산이라고 부르고, 메모리를 계산으로 되사는 거래입니다. FlashAttention의 역전파가 이 거래를 택하는데, 빠른 메모리와 느린 메모리 사이를 오가는 시간이 행렬곱 한 번보다 비싼 GPU에서는 되사는 쪽이 오히려 빠를 때가 많습니다. 그 속도 계산은 IO 쪽 이야기이고, 여기서 확인한 것은 되사는 것이 가능하다는 대수입니다. 행마다 수 하나면 충분합니다.
코드로 확인하기
항등식과 마스킹
import math
def lse(x):
m = max(x)
return m + math.log(sum(math.exp(v - m) for v in x))
def softmax(x):
m = max(x)
e = [math.exp(v - m) for v in x]
s = sum(e)
return [v / s for v in e]
x = [800.0, 799.0, 795.0]
# ① 그냥 하면 터진다
try:
sum(math.exp(v) for v in x)
except OverflowError as err:
print("naive:", err) # naive: math range error
print("exp(-800):", math.exp(-800)) # exp(-800): 0.0
# ② 최댓값을 빼면 산다
print(repr(lse(x))) # 800.3181754292475
print([round(v, 6) for v in softmax(x)]) # [0.727475, 0.267623, 0.004902]
# 항등식이므로 m 을 아무거나 잡아도 값은 같다 — 계산할 수만 있다면
for m in [800.0, 1234.0, 795.0, 0.0]:
try:
print(m, round(m + math.log(sum(math.exp(v - m) for v in x)), 10))
except (OverflowError, ValueError) as err:
print(m, "계산 불가:", err)
# 800.0 800.3181754292
# 1234.0 800.3181754292
# 795.0 800.3181754292
# 0.0 계산 불가: math range error
# ③ 마스킹 — 0 을 곱하는 쪽은 NaN 이 된다
s, mask = [2.0, 1.0, 900.0, 0.5], [1, 1, 0, 0]
good = softmax([s[i] if mask[i] else -math.inf for i in range(4)])
print([round(v, 6) for v in good]) # [0.731059, 0.268941, 0.0, 0.0]
p = softmax(s)
print(p) # [0.0, 0.0, 1.0, 0.0] ← 900 이 최댓값을 가져갔다
masked = [p[i] * mask[i] for i in range(4)]
print(masked, sum(masked)) # [0.0, 0.0, 0.0, 0.0] 0.0
try:
print([v / sum(masked) for v in masked])
except ZeroDivisionError as err:
print("재정규화:", err) # 재정규화: float division by zero
# 한 행이 전부 -inf 이면 -inf 쪽도 NaN 이 된다
row = [-math.inf] * 4
print(max(row), row[0] - max(row)) # -inf nan
# 그래서 구현은 아주 큰 음수를 쓴다
print([round(v, 6) for v in softmax([2.0, 1.0, -1e9, -1e9])])
# [0.731059, 0.268941, 0.0, 0.0]
③의 마지막 줄은 float64라 가 멀쩡합니다. 같은 값을 float16으로 바꾸면 가 된다는 앞 절의 이야기는 NumPy의 np.float16(-1e9)가 -inf를 내는 것으로 확인할 수 있습니다.
온라인 합과 기울기
# ④ 온라인 softmax — 블록을 이어 붙여도 같은 값이 나온다
s = [2.0, 1.0, 4.0, 0.5]
V = [[1.0, -0.5], [0.0, 2.0], [-1.5, 0.5], [2.0, 1.0]]
def block(idx):
m = max(s[j] for j in idx)
l = sum(math.exp(s[j] - m) for j in idx)
o = [sum(math.exp(s[j] - m) * V[j][a] for j in idx) for a in range(2)]
return m, l, o
def merge(A, B):
(m1, l1, o1), (m2, l2, o2) = A, B
m = max(m1, m2)
a1, a2 = math.exp(m1 - m), math.exp(m2 - m)
return m, a1 * l1 + a2 * l2, [a1 * o1[a] + a2 * o2[a] for a in range(2)]
m, l, o = merge(block([0, 1]), block([2, 3]))
print(m, round(l, 6), [round(v / l, 6) for v in o])
# 4.0 1.21532 [-1.073191, 0.462515]
p = softmax(s) # 한 번에 계산한 것
one_shot = [sum(p[j] * V[j][a] for j in range(4)) for a in range(2)]
print([round(v, 6) for v in one_shot]) # [-1.073191, 0.462515]
print(max(abs(o[a] / l - one_shot[a]) for a in range(2))) # 2.220446049250313e-16
# 쪼개는 방식을 바꿔도 같다
alt = merge(merge(block([0]), block([2])), merge(block([1]), block([3])))
print(round(alt[1], 6), [round(v / alt[1], 6) for v in alt[2]])
# 1.21532 [-1.073191, 0.462515]
④의 마지막 두 줄이 결합법칙의 증거입니다. 순서대로 이어 붙인 것과 짝을 바꿔 이어 붙인 것이 같은 값을 냅니다 — 병렬로 계산해도 되는 근거가 이것입니다.
import random
# ⑤ LSE 의 기울기는 softmax 다 — 유한차분으로 잰다
x, h = [800.0, 799.0, 795.0], 1e-6
fd = []
for i in range(3):
up, dn = list(x), list(x)
up[i] += h
dn[i] -= h
fd.append((lse(up) - lse(dn)) / (2 * h))
print([round(v, 6) for v in fd]) # [0.727475, 0.267623, 0.004902]
print([round(v, 6) for v in softmax(x)]) # [0.727475, 0.267623, 0.004902]
# 두 줄 부등식: max <= LSE <= max + log n
print(max(x), round(lse(x), 6), round(max(x) + math.log(len(x)), 6))
# 800.0 800.318175 801.098612
print(round(math.log(50257), 3)) # 10.825
# 역전파가 쓰는 되살리기: LSE 하나로 확률 전부
L = lse(x)
print([round(math.exp(v - L), 6) for v in x]) # [0.727475, 0.267623, 0.004902]
# ⑥ 블록 1,000 개로 쪼갠 온라인 합 vs 한 번에
random.seed(0)
z = [random.gauss(0, 3) for _ in range(100_000)]
M, S = -math.inf, 0.0
for i in range(0, len(z), 100):
blk = z[i:i + 100]
mb = max(blk)
lb = sum(math.exp(v - mb) for v in blk)
mn = max(M, mb)
S = math.exp(M - mn) * S + math.exp(mb - mn) * lb
M = mn
print(repr(M + math.log(S))) # 15.864301228245147
print(repr(lse(z))) # 15.864301228245145
⑤의 유한차분이 여섯째 자리까지 softmax와 겹칩니다. ⑥은 점 십만 개를 백 개씩, 블록 천 개로 나눠 이어 붙인 값과 한 번에 계산한 값이 마지막 자리 하나만 다릅니다. float64에서 합치기 천 번이 남긴 차이가 이니, 앞 절에서 말한 「순서가 바뀌면 마지막 자릿수만 흔들린다」가 이 모양입니다.
정리
- log-sum-exp 항등식 은 을 합 밖으로 뽑은 것뿐이라 모든 에서 정확히 성립한다. 근사가 아니다.
- 을 최댓값으로 잡으면 모든 지수의 인자가 0 이하가 되어 넘칠 수 없고, 가장 큰 항이 정확히 1이라 분모가 0이 될 수도 없다. 평균도 점수가 좁게 모이면 되지만, 최댓값만이 점수의 폭과 무관하게 안전하다.
- 두 항이면 이고, 차이가 크면
log1p가 작은 항을 살린다. - 확률을 만들지 않고 로 로그확률을 바로 얻는다.
log_softmax가 하는 일이다. - 이라 LSE는 매끄러운 최댓값이고, 그 기울기가 정확히 softmax다. 온도를 0으로 보내면 최댓값으로, 키우면 을 뺀 나머지가 평균으로 간다.
- 마스킹은 softmax 앞에서 를 더해서 한다. 뒤에 0을 곱하면 가려야 할 점수가 최댓값을 가져가 살릴 자리들이 전부 0으로 잘리고, 재정규화를 빼먹으면 가린 자리로 기울기가 샌다.
- 한 행이 전부 면 NaN이다. 큰 음수로 대신하되, float16에서는 가 로 넘치므로 값을 바꿔 넣거나 쯤을 더한다.
- 블록 하나가 셋만 들고 가면 둘을 합칠 수 있고, 이 점화식은 결합법칙을 만족하며 정확하다. 부동소수점 오차는 마지막 자릿수에 머문다.
- 순전파가 행마다 LSE 하나를 남기면 역전파는 로 확률 행렬을 되살린다. 저장은 키의 수 배만큼 줄고, 대가는 를 한 번 더 구하는 재계산이다.
처음 장면으로 돌아가면, 50,257개의 로짓이 수백까지 자라도 모델이 멈추지 않는 이유가 이 몇 줄입니다. 손실은 최댓값을 뺀 LSE로 계산되어 넘치지 않고, 그 값은 최댓값에서 10.825 안쪽에 머물며, 그 기울기가 곧 확률에서 정답을 뺀 입니다. 여기까지가 7단원의 마지막에서 두 번째 자리입니다. 어텐션 식의 기호를 하나씩 뜯어 계산과 유도로 닫았고, 이제 그것을 실제로 돌릴 때 필요한 대수까지 갖췄습니다.
남은 기호가 하나 있습니다. 지금까지 다룬 어떤 식에도 토큰의 순서가 들어 있지 않았습니다 — 는 토큰을 섞어도 같은 값을 냅니다. 다음 글에서 위치를 수로 적는 두 방법과, RoPE가 왜 상대 위치만 남기는지를 회전행렬로 증명합니다.
읽어주셔서 감사합니다. 😊

