수학

MATH / 중급 10번

행렬곱 네 가지 시선과 배치·헤드 shape 산수

같은 행렬곱을 원소·행·열·외적의 네 관점으로 손으로 계산하고, 어느 시선이 언제 편한지 정리합니다. 그 눈으로 (B, H, T, d_h) 텐서의 어느 축이 축약되는지 따라가며 멀티헤드의 reshape·transpose·matmul 순서를 shape만으로 재구성하고, shape 오류를 실행 전에 예측합니다.

PALDYN Team33 MIN READ

어텐션을 처음 직접 짜면 거의 반드시 이 줄을 만납니다.

RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x768 and 64x768)

숫자 네 개만 적혀 있고 어디가 잘못됐는지는 안 알려 줍니다. 그런데 이 오류는 머릿속에서 미리 잡을 수 있는 종류입니다. 곱하기 전에 두 텐서의 축을 나란히 적어 놓고 어느 축이 맞닿는지만 보면 되기 때문입니다.

이 글은 두 가지를 합니다. 먼저 지난 글에서 합성사상으로 정의한 행렬곱을 실제로 계산하는 네 가지 시선으로 갈아 끼우고, 그 다음 그 눈으로 (B,H,T,dh)(B, H, T, d_h) 짜리 텐서의 축 산수를 손으로 따라갑니다.

행렬곱의 네 시선

계산할 것은 하나입니다.

A=(120013),B=(210110),C=AB=(2331)A = \begin{pmatrix} 1 & 2 & 0 \\ 0 & 1 & 3\end{pmatrix}, \qquad B = \begin{pmatrix} 2 & 1 \\ 0 & 1 \\ 1 & 0\end{pmatrix}, \qquad C = AB = \begin{pmatrix} 2 & 3 \\ 3 & 1\end{pmatrix}

AA 는 2×32\times3, BB 는 3×23\times2 이므로 결과는 2×22\times2 입니다. 이 CC 를 네 번 다르게 구합니다. 값은 언제나 같고, 달라지는 것은 무엇이 보이느냐입니다.

원소·행·열·외적 네 시선에서 A와 B의 어느 부분이 쓰이는지

원소 시선

가장 익숙한 것입니다.

Cij=∑kAikBkjC_{ij} = \sum_k A_{ik}B_{kj}

C11C_{11} 은 AA 의 1행 (1,2,0)(1,2,0) 과 BB 의 1열 (2,0,1)(2,0,1) 의 내적이라 1⋅2+2⋅0+0⋅1=21\cdot2 + 2\cdot0 + 0\cdot1 = 2 입니다. 나머지도 같은 식으로

C12=1⋅1+2⋅1+0⋅0=3,C21=0⋅2+1⋅0+3⋅1=3,C22=0⋅1+1⋅1+3⋅0=1C_{12} = 1\cdot1 + 2\cdot1 + 0\cdot0 = 3, \quad C_{21} = 0\cdot2+1\cdot0+3\cdot1 = 3, \quad C_{22} = 0\cdot1+1\cdot1+3\cdot0 = 1

첨자가 하는 말이 전부 이 시선에 담겨 있습니다. ii 와 jj 는 결과에 남는 자유 첨자이고 kk 는 합해서 사라지는 더미 첨자입니다 — 3번 글에서 세운 구분이 그대로입니다. 합해서 사라지는 축을 축약되는 축이라고 부르겠습니다. 이 글에서 끝까지 쓸 말이라 여기서 이름을 붙여 둡니다.

행 시선과 열 시선

CC 를 통째로 보지 않고 한 줄씩 떼어 보면 다른 것이 보입니다. 먼저 행입니다.

(C의 i행)=∑kAik (B의 k행)(C\text{의 } i\text{행}) = \sum_k A_{ik}\,(B\text{의 } k\text{행})

CC 의 1행은 BB 의 행 세 개를 AA 의 1행을 계수로 삼아 섞은 것입니다.

1⋅(2,1)+2⋅(0,1)+0⋅(1,0)=(2,1)+(0,2)=(2,3)1\cdot(2,1) + 2\cdot(0,1) + 0\cdot(1,0) = (2,1) + (0,2) = (2,3)

2행도 마찬가지로 0⋅(2,1)+1⋅(0,1)+3⋅(1,0)=(3,1)0\cdot(2,1) + 1\cdot(0,1) + 3\cdot(1,0) = (3,1) 입니다. 이 시선은 행이 곧 하나의 데이터 표본일 때 유용합니다. X @ W 에서 XX 의 한 행이 토큰 하나라면, 출력의 그 행은 그 토큰 하나만으로 결정된다는 것이 이 식에서 바로 보입니다 — 배치의 다른 행은 끼어들지 않습니다.

이제 열입니다. 같은 곱을 열로 떼면 섞이는 쪽이 반대가 됩니다.

(C의 j열)=A (B의 j열)=∑kBkj (A의 k열)(C\text{의 } j\text{열}) = A\,(B\text{의 } j\text{열}) = \sum_k B_{kj}\,(A\text{의 } k\text{열})

BB 의 1열이 (2,0,1)(2,0,1) 이므로

2⋅(1,0)+0⋅(2,1)+1⋅(0,3)=(2,3)2\cdot(1,0) + 0\cdot(2,1) + 1\cdot(0,3) = (2,3)

지난 글의 "AxA\mathbf{x} 는 AA 열들의 선형결합"이 그대로 반복된 것입니다. 그래서 이 시선이 기하와 가장 가깝습니다 — CC 의 열은 전부 AA 의 열이 만드는 span 안에 있고, 곱을 아무리 해도 그 밖으로는 못 나갑니다. 행 시선이 "표본 하나가 어떻게 변했나"를 답한다면, 열 시선은 "결과가 어느 공간에 갇히나"를 답합니다.

외적 시선

마지막 시선은 더미 첨자 kk 를 바깥으로 꺼냅니다.

C=∑k(A의 k열)(B의 k행)C = \sum_k (A\text{의 } k\text{열})(B\text{의 } k\text{행})

열벡터 하나와 행벡터 하나를 곱하면 행렬 한 장이 나오고, 이것을 외적이라고 합니다. kk 마다 한 장씩 나온 판을 전부 더한 것이 CC 입니다.

A의 열과 B의 행이 만드는 판 세 장을 겹치면 C가 된다

(2100)+(0201)+(0030)=(2331)\begin{pmatrix} 2 & 1 \\ 0 & 0\end{pmatrix} + \begin{pmatrix} 0 & 2 \\ 0 & 1\end{pmatrix} + \begin{pmatrix} 0 & 0 \\ 3 & 0\end{pmatrix} = \begin{pmatrix} 2 & 3 \\ 3 & 1\end{pmatrix}

판 한 장은 열 하나와 행 하나만으로 완전히 결정됩니다. 이 사실이 나중에 "행렬을 얇은 조각 몇 장으로 근사한다"는 이야기의 출발점이 되고, LoRA의 ΔW=BA\Delta W = BA 가 정확히 판 rr 장의 합입니다. 그 계산은 「중급 16번 · 저계수 근사와 LoRA」의 몫입니다.

판이 따로따로 만들어진다는 것에는 당장 쓸모가 하나 더 있습니다. kk 를 조각내어 따로 더해도 된다는 것입니다. kk 가 1부터 3까지일 때 판 세 장을 한꺼번에 더하든, 앞의 두 장을 먼저 더해 두고 나중에 세 번째를 얹든 결과가 같습니다. 축약되는 축이 아주 길어서 한 번에 못 들고 있을 때 그 축을 토막 내 부분합으로 쌓아 가는 방식이 여기서 나옵니다 — 어텐션의 긴 문맥을 조각으로 나눠 처리하는 구현들이 딛고 있는 것이 이 한 줄입니다.

합의 순서와 곱셈 횟수

네 시선이 같은 값을 주는 이유는 따로 증명할 것이 없습니다. 원소 시선의 식을 전부 펼치면 i,j,ki, j, k 세 첨자에 대한 삼중합이고, 네 시선은 그 세 합을 어느 것부터 도느냐만 다릅니다. 원소 시선은 i,ji, j 를 바깥에 두고 kk 를 안에서 돌고, 외적 시선은 kk 를 바깥에 두고 i,ji, j 를 안에서 돕니다. 유한한 항의 덧셈은 순서를 바꿔도 값이 같으므로 네 값이 같을 수밖에 없습니다.

그래서 곱셈 횟수도 넷이 전부 같습니다. m×pm \times p 와 p×np \times n 을 곱하면 어느 순서로 돌든 AikBkjA_{ik}B_{kj} 꼴의 곱이 m⋅n⋅pm\cdot n\cdot p 번 일어나고, 그것을 더하는 덧셈이 m⋅n⋅(p−1)m\cdot n\cdot(p-1) 번입니다. 둘을 합쳐 흔히 2mnp2mnp 번의 부동소수점 연산으로 셉니다. 시선을 바꿔서 빨라지는 것이 아니라 메모리를 읽는 순서가 달라질 뿐이고, 실제 구현이 시선을 고르는 기준도 계산량이 아니라 그쪽입니다.

알고 싶은 것 편한 시선
특정 한 칸의 값 원소
표본 하나가 어떻게 변했는가 행
결과가 어느 공간에 갇히는가 열
곱을 몇 장으로 줄일 수 있는가 외적

shape 산수의 세 규칙

이제 축이 넷 이상인 텐서로 갑니다.

축이 맞닿는 자리

torch.matmul 이나 @ 는 다음 세 줄로 움직입니다.

  1. 뒤의 두 축만 행렬로 본다. 앞의 축은 전부 배치 취급이다.
  2. 맞닿은 축은 값이 같아야 하고, 곱한 뒤 사라진다. 왼쪽의 마지막 축과 오른쪽의 끝에서 둘째 축이다.
  3. 앞 축은 그대로 따라 나온다. 한쪽이 1이면 브로드캐스트로 늘어난다.

Q와 K를 뒤집은 것을 곱할 때 어느 축이 따라오고 어느 축이 사라지는지

어텐션 점수에 그대로 적용하면

(B,H,T,dh)×(B,H,dh,T) ⟶ (B,H,T,T)(B, H, T, d_h) \times (B, H, d_h, T) \ \longrightarrow\ (B, H, T, T)

BB 와 HH 는 곱셈에 끼지 않고 따라 나오고, dhd_h 는 양쪽이 같아서 사라지고, 바깥의 TT 둘이 결과의 행과 열이 됩니다. 질의 TT 개마다 키 TT 개에 대한 점수가 서는 표가 이 shape의 뜻입니다.

마스크의 브로드캐스트

규칙 3에 나온 브로드캐스트는 크기가 1인 축을 필요한 만큼 늘려 상대와 짝을 맞추는 규칙입니다. 값을 복사해 진짜로 늘리는 것이 아니라 같은 자리를 여러 번 읽을 뿐이라 메모리가 늘지 않습니다.

가장 자주 만나는 자리가 어텐션 마스크입니다. 인과 마스크는 "몇 번째 토큰이 몇 번째 토큰을 볼 수 있는가"만 담으므로 헤드마다 다를 이유가 없고, 그래서 (1,1,T,T)(1, 1, T, T) 나 (B,1,T,T)(B, 1, T, T) 로 만들어 둡니다. 이것을 점수 (B,H,T,T)(B, H, T, T) 에 더하면 크기 1인 헤드 축이 HH 로 늘어나 같은 마스크가 열두 헤드에 그대로 걸립니다. 헤드 수만큼 복사본을 만들어 두는 구현을 가끔 보는데, 규칙 3을 알면 그 복사가 통째로 필요 없다는 것이 보입니다.

축이 하나인 벡터

축이 하나뿐인 텐서에는 규칙 1이 그대로 안 걸립니다. 뒤의 두 축을 떼어 낼 축이 하나밖에 없기 때문입니다. 이때는 앞이나 뒤에 크기 1인 축을 임시로 붙였다가 결과에서 떼는 처리가 들어갑니다.

WW 가 (dout,din)(d_{\text{out}}, d_{\text{in}}) 이고 xx 가 (din,)(d_{\text{in}},) 이면 W @ x 는 xx 를 (din,1)(d_{\text{in}}, 1) 로 보고 곱한 뒤 그 1을 떼어 (dout,)(d_{\text{out}},) 을 냅니다. 반대로 xx 가 (dout,)(d_{\text{out}},) 일 때 x @ W 는 (1,dout)(1, d_{\text{out}}) 으로 보고 곱합니다. 같은 벡터가 어느 쪽에 서느냐에 따라 행처럼도 열처럼도 쓰이는 것이 이 임시 축의 효과이고, 그래서 둘 다 오류 없이 통합니다. 대신 결과의 축 개수가 입력보다 하나 적어지므로, 배치 차원을 기대하고 이어 붙이던 다음 줄에서 엉뚱한 자리가 어긋납니다.

오류 없이 통과하는 브로드캐스트

브로드캐스트는 편한 만큼 조용합니다. (T,1)(T, 1) 짜리와 (1,T)(1, T) 짜리를 더하면 양쪽의 1이 각각 TT 로 늘어나 (T,T)(T, T) 가 오류 없이 나옵니다. 길이 TT 인 두 벡터를 원소별로 더하려던 것이었다면 결과가 TT 개가 아니라 T2T^2 개가 된 것인데, 어디에서도 경고가 나지 않습니다.

T=2048T = 2048 이면 이 한 줄이 400만 개짜리 텐서를 만듭니다. 메모리가 갑자기 뛰거나 손실이 이상한 값에서 멈추는 자리를 거슬러 올라가면 이런 줄이 앉아 있는 경우가 많습니다. 크기 1인 축은 언제나 늘어날 수 있다는 것을 규칙으로 외워 두고, 1이 섞인 shape을 볼 때마다 한 번 멈추는 편이 낫습니다.

멀티헤드의 축 흐름

세 규칙만 있으면 멀티헤드 어텐션의 view 와 transpose 순서를 외우지 않고 재구성할 수 있습니다. H=12H = 12, dh=64d_h = 64, dmodel=768d_{\text{model}} = 768 로 둡니다.

입력 (B,T,768)이 view·transpose·matmul을 거쳐 다시 (B,T,768)로 돌아오는 축의 흐름

view와 transpose의 순서

각 단계가 왜 그 자리에 있어야 하는지가 규칙에서 나옵니다.

  • view(B, T, 12, 64) — 768을 12×64로 쪼갭니다. 숫자는 하나도 안 움직이고 칸을 다시 나눈 것뿐입니다.
  • transpose(1, 2) — 규칙 1이 "뒤 두 축만 행렬"이라고 했으므로, 곱하고 싶은 (T,dh)(T, d_h) 를 뒤로 보내야 합니다. 헤드 축이 앞으로 나가 배치 취급을 받게 되는 것이 이 한 줄의 목적입니다.
  • QK⊤QK^\top — dhd_h 가 맞닿아 사라지고 (B,12,T,T)(B, 12, T, T) 가 됩니다.
  • 가중치 @ V — 이번에는 TT 가 맞닿아 사라지고 (B,12,T,64)(B, 12, T, 64) 로 돌아옵니다. 같은 곱셈 규칙인데 사라지는 축이 다른 것이 어텐션의 두 단계입니다.
  • 되돌리기 — transpose(1, 2) 로 헤드를 제자리에 놓고 reshape 로 12×64를 768로 붙입니다.
import torch

B, T, H, d_h = 2, 5, 12, 64
x = torch.randn(B, T, H * d_h)

q = x.view(B, T, H, d_h).transpose(1, 2)     # (2, 12, 5, 64)
k = x.view(B, T, H, d_h).transpose(1, 2)
v = x.view(B, T, H, d_h).transpose(1, 2)

scores = q @ k.transpose(-2, -1)             # (2, 12, 5, 5)
out = scores.softmax(-1) @ v                 # (2, 12, 5, 64)
out = out.transpose(1, 2).reshape(B, T, H * d_h)
print(scores.shape, out.shape)               # (2,12,5,5) (2,5,768)

첫 줄의 view(B, T, 12, 64) 를 view(B, T, 64, 12) 로 잘못 적으면 어떻게 될까요. 768을 64×12로 쪼개도 곱은 똑같이 768이라 shape 검사는 그대로 통과합니다. 그다음 transpose(1, 2) 도 통과하고, 점수 표도 나오고, 학습도 돌아갑니다. 다만 한 헤드가 들고 있는 64개 숫자가 원래 묶여 있던 64개가 아니라 12칸씩 건너뛴 다른 묶음이 됩니다. 오류가 아니라 조용히 다른 값이라, 손실이 안 떨어지는 것으로만 드러납니다.

d_h를 정하는 제약

이 자리에서 흔한 오해가 dhd_h 를 손잡이로 보는 것입니다. 헤드 수 HH 와 머리 폭 dhd_h 를 따로 정할 수 있을 것 같지만, view 가 요구하는 것은 H⋅dh=dmodelH \cdot d_h = d_{\text{model}} 한 줄입니다. 768을 쪼개는 것이라 dhd_h 는 dmodel/Hd_{\text{model}}/H 로 저절로 정해지고, 고를 수 있는 것은 HH 뿐입니다.

그리고 HH 도 아무 값이나 되지 않습니다. 768을 나누어떨어뜨려야 하므로 12는 64, 16은 48, 8은 96이 되지만 10은 76.8이라 안 됩니다. "헤드를 열 개로 늘려 보자"가 안 되는 이유가 설계 철학이 아니라 나눗셈이라는 것이 이 줄의 전부입니다.

보폭과 contiguous

마지막 줄에서 view 가 아니라 reshape 을 쓴 이유가 있습니다. 텐서는 숫자가 한 줄로 늘어선 데이터 한 벌과, 축마다 "이 축에서 한 칸 가려면 몇 칸 건너뛰어야 하는가"를 적어 둔 수 한 벌로 되어 있습니다. 뒤쪽을 보폭이라고 부릅니다.

transpose 는 숫자를 하나도 안 옮기고 보폭의 순서만 바꿔 둡니다. 그래서 공짜인데, 그 결과는 "읽는 순서"와 "메모리에 놓인 순서"가 어긋난 상태입니다. view 는 메모리에 놓인 순서 그대로 칸을 다시 나누는 함수라 그 상태를 거부합니다. reshape 은 필요하면 복사까지 해 주므로 안전하고, contiguous() 는 그 복사를 명시적으로 먼저 해 두는 것입니다. 둘 다 공짜가 아니라 데이터를 한 벌 더 쓰는 값을 치릅니다 — transpose 가 공짜인 대신 그다음이 비싸지는 셈입니다.

transpose 뒤에 view 를 쓰면 나는 오류는 shape 산수가 아니라 메모리 배치의 문제라서, 이 글의 세 줄 규칙으로는 예측되지 않는 유일한 자리입니다.

GQA와 MQA의 헤드 수

요즘 모델은 QQ 와 K,VK, V 의 헤드 수를 다르게 둡니다. 질의는 32 헤드인데 키와 값은 8 헤드만 두고 넷씩 나눠 쓰는 식이고, 8을 1까지 내린 것이 MQA입니다. 캐시로 들고 있어야 할 K,VK, V 가 4분의 1이나 32분의 1로 줄어드는 것이 목적입니다.

shape으로 보면 이것은 새 연산이 아니라 규칙 3을 한 번 더 쓴 것입니다. KK 를 (B,8,1,T,dh)(B, 8, 1, T, d_h) 로 보고 질의를 (B,8,4,T,dh)(B, 8, 4, T, d_h) 로 보면 크기 1인 축이 4로 늘어나 짝이 맞습니다. 값을 복사하지 않고도 한 벌의 키가 네 질의 헤드에 걸리는 것이라, 절약되는 것이 계산이 아니라 메모리라는 점도 같은 그림에서 읽힙니다.

점수 표의 메모리와 연산량

세 규칙은 shape이 맞는지만 알려 주지만, 같은 shape에서 비용도 함께 읽힙니다.

원소 수로 세기

점수 표 (B,H,T,T)(B, H, T, T) 의 원소 수는 그냥 네 수의 곱입니다. B=8B = 8, H=12H = 12, T=2048T = 2048 이면

8×12×2048×2048=402,653,1848 \times 12 \times 2048 \times 2048 = 402{,}653{,}184

로 4억 개가 조금 넘습니다. fp16이면 한 개가 2바이트이므로 805MB입니다. 가중치가 아니라 중간 결과 한 장이 그만큼입니다. 모델 파라미터를 세어 메모리를 가늠하다가 실제로 넘치는 자리가 대개 여기입니다.

T가 두 배면 네 배

이 수에서 TT 만 두 번 곱해진다는 것이 핵심입니다. 배치와 헤드는 한 번씩 들어가는데 TT 는 행에도 열에도 서므로, 문맥을 두 배로 늘리면 표가 네 배가 됩니다.

TT 원소 수 fp16 메모리
512 2,517만 50MB
1,024 1.01억 201MB
2,048 4.03억 805MB
4,096 16.1억 3.2GB

512에서 4096으로 여덟 배 늘렸더니 64배가 됐습니다. 문맥 길이의 벽이라고 부르는 것이 이 표입니다 — 파라미터는 한 개도 안 늘었는데 중간 결과만으로 카드가 찹니다. 점수 표를 통째로 만들지 않고 조각으로 나눠 도는 구현들이 겨냥하는 것이 정확히 이 열입니다.

FLOPs 세기

연산량도 같은 규칙으로 셉니다. 앞에서 m×pm \times p 와 p×np \times n 의 곱이 2mnp2mnp 라고 했으니, 배치 축 BB 와 HH 는 그냥 앞에 곱해집니다.

QK⊤: 2⋅B⋅H⋅T2⋅dh,(가중치)V: 2⋅B⋅H⋅T2⋅dhQK^\top:\ 2 \cdot B \cdot H \cdot T^2 \cdot d_h, \qquad (\text{가중치})V:\ 2 \cdot B \cdot H \cdot T^2 \cdot d_h

두 단계가 정확히 같은 값입니다 — 앞은 dhd_h 가 축약되고 뒤는 TT 가 축약되는데, 세 수의 곱이 T⋅T⋅dhT \cdot T \cdot d_h 로 같기 때문입니다. 위의 설정에 넣으면 한 단계가 51.5 GFLOP이라 둘이 103 GFLOP입니다. 메모리와 연산량이 둘 다 T2T^2 을 따라간다는 것이 이 절의 결론이고, 그래서 둘 중 하나만 줄이는 해법이 없습니다.

T를 512에서 4096으로 올릴 때 점수 표의 메모리와 FLOPs가 각각 네 배씩 뛰는 두 벌 막대그래프

shape 오류 읽기

첫머리 오류의 해독

RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x768 and 64x768)

세 줄 규칙으로 읽습니다. 왼쪽의 마지막 축은 768, 오른쪽의 끝에서 둘째 축은 64입니다. 맞닿아야 할 두 값이 768과 64라 어긋났습니다.

원인은 지난 글에서 본 저장 순서입니다. nn.Linear(768, 64).weight 의 모양이 (64, 768) 이라 x @ W 는 768과 64를 맞대게 되고, 맞는 것은 x @ W.T 입니다. 오류 메시지의 두 숫자 중 어느 쪽을 뒤집어야 하는지가 이 한 줄에서 결정됩니다.

자주 나오는 어긋남은 몇 가지로 정리됩니다.

증상 어긋난 곳 고치는 법
(a×d) 와 (e×d) 오른쪽을 전치하지 않았다 @ W.T
(a×d) 와 (d×b) 인데도 실패 앞의 배치 축이 서로 다르다 브로드캐스트되게 맞추거나 expand
결과 축이 하나 더 많다 벡터를 (n,) 이 아니라 (n,1) 로 만들었다 squeeze 또는 인덱싱으로 축 제거
view 만 실패 transpose 뒤라 메모리가 이어져 있지 않다 reshape 또는 contiguous().view(...)

오류 없이 지나가는 어긋남

고치기 어려운 것은 오류가 나는 쪽이 아닙니다. 디버깅에서 오래 잡아먹는 자리는 셋이고, 셋 다 메시지가 없습니다.

  • 배치 축이 1이라 브로드캐스트로 맞아 버리는 자리. 배치 8짜리와 배치 1짜리를 곱하면 규칙 3이 1을 8로 늘려 통과시킵니다. 한 표본의 값이 여덟 표본 전부에 걸린 것인데 shape은 정상입니다.
  • 앞에서 본 view(B, T, 64, 12). 곱이 같으면 어떤 쪼개기도 통과합니다.
  • 크기 1인 축을 안 뗀 채 더하기. (T,1)(T,1) 과 (1,T)(1,T) 가 만나 (T,T)(T,T) 가 되는 그 자리입니다.

세 경우의 공통점은 어긋난 쪽이 1이거나, 곱이 같다는 것입니다. 규칙이 허용하도록 만들어진 자리라 검사로는 못 막고, 중간 텐서의 shape을 한 번 찍어 보는 것 말고는 방법이 없습니다.

확인하는 순서

오류 메시지는 뒤 두 축만 적어 주므로 앞 축은 안 보입니다. 순서를 정해 두면 빠릅니다.

  1. 뒤 두 축 — 맞닿는 자리가 같은 값인지. 메시지에 적힌 것이 이것입니다.
  2. 앞 축 — 배치와 헤드가 서로 같거나 한쪽이 1인지. 메시지에 안 적히는 자리입니다.
  3. dtype — fp16과 fp32가 섞였는지. shape은 멀쩡한데 곱이 거부되는 나머지 경우가 대개 이쪽입니다.

그리고 실행 전에 미리 훑는 방법이 있습니다. 실제 값을 만들지 않고 shape만 굴려 보는 것입니다.

import torch

with torch.device("meta"):                       # 값을 할당하지 않는 장치
    x = torch.empty(8, 2048, 768)
    w = torch.empty(768, 64)
    print((x @ w).shape)                         # (8, 2048, 64)

meta 장치 위의 텐서는 메모리를 한 바이트도 안 잡고 shape 계산만 통과시킵니다. 카드에 안 올라가는 크기라도 축이 맞는지는 이 줄로 먼저 확인할 수 있습니다.

einsum과 축 이름

축 순서를 맞추려고 transpose 를 넣는 일 자체를 없애는 방법이 있습니다. 3번 글에서 만든 einsum 문자열입니다.

문자열 한 줄이 적는 것

import torch

Q = torch.randn(2, 12, 5, 64)                    # (B, H, T, d_h)
K = torch.randn(2, 12, 5, 64)

s1 = Q @ K.transpose(-2, -1)
s2 = torch.einsum('bhid,bhjd->bhij', Q, K)
print(torch.allclose(s1, s2))                    # True

'bhid,bhjd->bhij' 한 줄에 세 줄 규칙이 전부 적혀 있습니다. 양쪽에 나오지만 결과에 없는 d 가 축약되는 축이고, 양쪽에 똑같이 있고 결과에도 있는 b, h 가 따라 나오는 배치 축이며, 한쪽에만 있는 i, j 가 바깥 축입니다. 축의 자리를 맞출 필요가 없으니 transpose 도 필요 없습니다.

행렬곱만 적히는 것이 아닙니다. 입력이 하나뿐이어도 문자열이 성립합니다.

문자열 하는 일
'ij->ji' 전치
'ii->i' 대각 성분 뽑기
'ij->i' 행마다 합
'ij->' 전부 합

읽는 법은 하나뿐입니다 — 화살표 오른쪽에 없는 문자가 합해져 사라집니다. 'ii->i' 에서 같은 문자를 두 번 쓴 것은 "두 축의 번호가 같은 자리만 본다"는 뜻이고, 그래서 대각선이 나옵니다.

텐서가 셋 이상일 때

곱할 것이 셋이 되면 einsum이 값을 하나로 정해 주지만 비용은 정해 주지 않습니다. 계산 순서에 따라 중간 결과의 크기가 달라지기 때문입니다.

AA 가 1000×101000\times10, BB 가 10×100010\times1000, CC 가 1000×101000\times10 일 때 ABCABC 를 생각해 봅니다. (AB)C(AB)C 로 가면 ABAB 가 1000×10001000\times1000 이라 중간에 100만 개짜리가 서고 연산이 4천만 번입니다. A(BC)A(BC) 로 가면 BCBC 가 10×1010\times10 이라 중간 결과가 100개이고 연산은 40만 번입니다. 같은 식, 같은 답, 100배 차이입니다.

torch.einsum 은 인자 셋 이상이면 순서를 스스로 고르지만, 축 길이가 비슷할 때는 그 선택이 최선이 아닐 수 있습니다. 중간 결과가 얼마나 커지는지를 손으로 한 번 세어 보는 것이 그래서 여전히 필요합니다.

축 이름을 먼저 정하기

einsum이 규칙을 없애 주는 것은 아닙니다. 문자열을 잘못 쓰면 오류 없이 다른 값이 나오므로, 어느 문자가 축약되는지를 스스로 읽을 수 있어야 하는 것은 그대로입니다.

대신 습관 하나가 바뀝니다. 코드를 쓰기 전에 축마다 문자를 정해 두면 — 배치는 b, 헤드는 h, 토큰 자리는 i와 j, 머리 폭은 d — 그다음부터는 원하는 결과의 문자열을 적는 일만 남습니다. 축의 물리적 순서를 기억할 필요가 없어지므로 transpose 를 넣었다 빼는 왕복 자체가 사라집니다. 문자를 먼저 정하는 이 한 걸음이 view·transpose 실수를 줄이는 가장 값싼 방법입니다.

정리

  • 같은 곱을 네 시선으로 계산할 수 있습니다. 원소는 행·열의 내적, 행 시선은 오른쪽 행들의 혼합, 열 시선은 왼쪽 열들의 혼합, 외적 시선은 판의 합입니다. 넷은 삼중합을 도는 순서만 다르므로 값도 곱셈 횟수 2mnp2mnp 도 같습니다.
  • 외적 시선은 축약되는 축을 조각낼 수 있게 해 줍니다. 판 몇 장을 먼저 더해 두었다가 나중에 얹어도 되고, 긴 문맥을 조각으로 도는 구현이 여기에 섭니다.
  • shape 산수는 세 줄입니다. 뒤 두 축만 곱하고, 맞닿은 축은 같아야 하며 사라지고, 앞 축은 따라 나옵니다. 크기 1인 축은 브로드캐스트로 늘어나는데, 이것이 마스크를 한 벌만 두게 해 주는 동시에 (T,1)(T,1) 과 (1,T)(1,T) 를 말없이 (T,T)(T,T) 로 만드는 자리이기도 합니다.
  • 멀티헤드의 view·transpose 순서는 외울 것이 아닙니다. 곱할 두 축을 뒤로 보낸다는 목적에서 매번 재구성되고, dhd_h 는 손잡이가 아니라 dmodel/Hd_{\text{model}}/H 입니다.
  • transpose 뒤의 view 실패만 예외입니다. shape이 아니라 보폭의 문제라 reshape 으로 풀고, 그 대신 복사 비용을 뭅니다.
  • 점수 표는 T2T^2 을 따라갑니다. B=8B=8, H=12H=12, T=2048T=2048 에서 4.03억 개·805MB이고, TT 를 두 배로 하면 메모리도 FLOPs도 네 배입니다.
  • 오류가 안 나는 어긋남이 더 비쌉니다. 배치 축 1, 곱이 같은 view, 크기 1인 축의 덧셈 — 셋 다 규칙이 허락하는 자리라 shape을 찍어 보는 것 말고는 방법이 없습니다.

다음 글은 시선을 다시 기하로 돌립니다 — 2×22\times2 행렬 하나로 회전·전단·스케일을 전부 만들고, 회전행렬 R(θ)R(\theta) 를 세워 R(α)R(β)=R(α+β)R(\alpha)R(\beta) = R(\alpha+\beta) 를 곱셈으로 직접 확인합니다.


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

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