어텐션을 처음 직접 짜면 거의 반드시 이 줄을 만납니다.
RuntimeError: mat1 and mat2 shapes cannot be multiplied (32x768 and 64x768)
숫자 네 개만 적혀 있고 어디가 잘못됐는지는 안 알려 줍니다. 그런데 이 오류는 머릿속에서 미리 잡을 수 있는 종류입니다. 곱하기 전에 두 텐서의 축을 나란히 적어 놓고 어느 축이 맞닿는지만 보면 되기 때문입니다.
이 글은 두 가지를 합니다. 먼저 지난 글에서 합성사상으로 정의한 행렬곱을 실제로 계산하는 네 가지 시선으로 갈아 끼우고, 그 다음 그 눈으로 짜리 텐서의 축 산수를 손으로 따라갑니다.
행렬곱의 네 시선
계산할 것은 하나입니다.
는 , 는 이므로 결과는 입니다. 이 를 네 번 다르게 구합니다. 값은 언제나 같고, 달라지는 것은 무엇이 보이느냐입니다.
원소 시선
가장 익숙한 것입니다.
은 의 1행 과 의 1열 의 내적이라 입니다. 나머지도 같은 식으로
첨자가 하는 말이 전부 이 시선에 담겨 있습니다. 와 는 결과에 남는 자유 첨자이고 는 합해서 사라지는 더미 첨자입니다 — 3번 글에서 세운 구분이 그대로입니다. 합해서 사라지는 축을 축약되는 축이라고 부르겠습니다. 이 글에서 끝까지 쓸 말이라 여기서 이름을 붙여 둡니다.
행 시선과 열 시선
를 통째로 보지 않고 한 줄씩 떼어 보면 다른 것이 보입니다. 먼저 행입니다.
의 1행은 의 행 세 개를 의 1행을 계수로 삼아 섞은 것입니다.
2행도 마찬가지로 입니다. 이 시선은 행이 곧 하나의 데이터 표본일 때 유용합니다. X @ W 에서 의 한 행이 토큰 하나라면, 출력의 그 행은 그 토큰 하나만으로 결정된다는 것이 이 식에서 바로 보입니다 — 배치의 다른 행은 끼어들지 않습니다.
이제 열입니다. 같은 곱을 열로 떼면 섞이는 쪽이 반대가 됩니다.
의 1열이 이므로
지난 글의 " 는 열들의 선형결합"이 그대로 반복된 것입니다. 그래서 이 시선이 기하와 가장 가깝습니다 — 의 열은 전부 의 열이 만드는 span 안에 있고, 곱을 아무리 해도 그 밖으로는 못 나갑니다. 행 시선이 "표본 하나가 어떻게 변했나"를 답한다면, 열 시선은 "결과가 어느 공간에 갇히나"를 답합니다.
외적 시선
마지막 시선은 더미 첨자 를 바깥으로 꺼냅니다.
열벡터 하나와 행벡터 하나를 곱하면 행렬 한 장이 나오고, 이것을 외적이라고 합니다. 마다 한 장씩 나온 판을 전부 더한 것이 입니다.
판 한 장은 열 하나와 행 하나만으로 완전히 결정됩니다. 이 사실이 나중에 "행렬을 얇은 조각 몇 장으로 근사한다"는 이야기의 출발점이 되고, LoRA의 가 정확히 판 장의 합입니다. 그 계산은 「중급 16번 · 저계수 근사와 LoRA」의 몫입니다.
판이 따로따로 만들어진다는 것에는 당장 쓸모가 하나 더 있습니다. 를 조각내어 따로 더해도 된다는 것입니다. 가 1부터 3까지일 때 판 세 장을 한꺼번에 더하든, 앞의 두 장을 먼저 더해 두고 나중에 세 번째를 얹든 결과가 같습니다. 축약되는 축이 아주 길어서 한 번에 못 들고 있을 때 그 축을 토막 내 부분합으로 쌓아 가는 방식이 여기서 나옵니다 — 어텐션의 긴 문맥을 조각으로 나눠 처리하는 구현들이 딛고 있는 것이 이 한 줄입니다.
합의 순서와 곱셈 횟수
네 시선이 같은 값을 주는 이유는 따로 증명할 것이 없습니다. 원소 시선의 식을 전부 펼치면 세 첨자에 대한 삼중합이고, 네 시선은 그 세 합을 어느 것부터 도느냐만 다릅니다. 원소 시선은 를 바깥에 두고 를 안에서 돌고, 외적 시선은 를 바깥에 두고 를 안에서 돕니다. 유한한 항의 덧셈은 순서를 바꿔도 값이 같으므로 네 값이 같을 수밖에 없습니다.
그래서 곱셈 횟수도 넷이 전부 같습니다. 와 을 곱하면 어느 순서로 돌든 꼴의 곱이 번 일어나고, 그것을 더하는 덧셈이 번입니다. 둘을 합쳐 흔히 번의 부동소수점 연산으로 셉니다. 시선을 바꿔서 빨라지는 것이 아니라 메모리를 읽는 순서가 달라질 뿐이고, 실제 구현이 시선을 고르는 기준도 계산량이 아니라 그쪽입니다.
| 알고 싶은 것 | 편한 시선 |
|---|---|
| 특정 한 칸의 값 | 원소 |
| 표본 하나가 어떻게 변했는가 | 행 |
| 결과가 어느 공간에 갇히는가 | 열 |
| 곱을 몇 장으로 줄일 수 있는가 | 외적 |
shape 산수의 세 규칙
이제 축이 넷 이상인 텐서로 갑니다.
축이 맞닿는 자리
torch.matmul 이나 @ 는 다음 세 줄로 움직입니다.
- 뒤의 두 축만 행렬로 본다. 앞의 축은 전부 배치 취급이다.
- 맞닿은 축은 값이 같아야 하고, 곱한 뒤 사라진다. 왼쪽의 마지막 축과 오른쪽의 끝에서 둘째 축이다.
- 앞 축은 그대로 따라 나온다. 한쪽이 1이면 브로드캐스트로 늘어난다.
어텐션 점수에 그대로 적용하면
와 는 곱셈에 끼지 않고 따라 나오고, 는 양쪽이 같아서 사라지고, 바깥의 둘이 결과의 행과 열이 됩니다. 질의 개마다 키 개에 대한 점수가 서는 표가 이 shape의 뜻입니다.
마스크의 브로드캐스트
규칙 3에 나온 브로드캐스트는 크기가 1인 축을 필요한 만큼 늘려 상대와 짝을 맞추는 규칙입니다. 값을 복사해 진짜로 늘리는 것이 아니라 같은 자리를 여러 번 읽을 뿐이라 메모리가 늘지 않습니다.
가장 자주 만나는 자리가 어텐션 마스크입니다. 인과 마스크는 "몇 번째 토큰이 몇 번째 토큰을 볼 수 있는가"만 담으므로 헤드마다 다를 이유가 없고, 그래서 나 로 만들어 둡니다. 이것을 점수 에 더하면 크기 1인 헤드 축이 로 늘어나 같은 마스크가 열두 헤드에 그대로 걸립니다. 헤드 수만큼 복사본을 만들어 두는 구현을 가끔 보는데, 규칙 3을 알면 그 복사가 통째로 필요 없다는 것이 보입니다.
축이 하나인 벡터
축이 하나뿐인 텐서에는 규칙 1이 그대로 안 걸립니다. 뒤의 두 축을 떼어 낼 축이 하나밖에 없기 때문입니다. 이때는 앞이나 뒤에 크기 1인 축을 임시로 붙였다가 결과에서 떼는 처리가 들어갑니다.
가 이고 가 이면 W @ x 는 를 로 보고 곱한 뒤 그 1을 떼어 을 냅니다. 반대로 가 일 때 x @ W 는 으로 보고 곱합니다. 같은 벡터가 어느 쪽에 서느냐에 따라 행처럼도 열처럼도 쓰이는 것이 이 임시 축의 효과이고, 그래서 둘 다 오류 없이 통합니다. 대신 결과의 축 개수가 입력보다 하나 적어지므로, 배치 차원을 기대하고 이어 붙이던 다음 줄에서 엉뚱한 자리가 어긋납니다.
오류 없이 통과하는 브로드캐스트
브로드캐스트는 편한 만큼 조용합니다. 짜리와 짜리를 더하면 양쪽의 1이 각각 로 늘어나 가 오류 없이 나옵니다. 길이 인 두 벡터를 원소별로 더하려던 것이었다면 결과가 개가 아니라 개가 된 것인데, 어디에서도 경고가 나지 않습니다.
이면 이 한 줄이 400만 개짜리 텐서를 만듭니다. 메모리가 갑자기 뛰거나 손실이 이상한 값에서 멈추는 자리를 거슬러 올라가면 이런 줄이 앉아 있는 경우가 많습니다. 크기 1인 축은 언제나 늘어날 수 있다는 것을 규칙으로 외워 두고, 1이 섞인 shape을 볼 때마다 한 번 멈추는 편이 낫습니다.
멀티헤드의 축 흐름
세 규칙만 있으면 멀티헤드 어텐션의 view 와 transpose 순서를 외우지 않고 재구성할 수 있습니다. , , 로 둡니다.
view와 transpose의 순서
각 단계가 왜 그 자리에 있어야 하는지가 규칙에서 나옵니다.
view(B, T, 12, 64)— 768을 12×64로 쪼갭니다. 숫자는 하나도 안 움직이고 칸을 다시 나눈 것뿐입니다.transpose(1, 2)— 규칙 1이 "뒤 두 축만 행렬"이라고 했으므로, 곱하고 싶은 를 뒤로 보내야 합니다. 헤드 축이 앞으로 나가 배치 취급을 받게 되는 것이 이 한 줄의 목적입니다.- — 가 맞닿아 사라지고 가 됩니다.
- 가중치 @ V — 이번에는 가 맞닿아 사라지고 로 돌아옵니다. 같은 곱셈 규칙인데 사라지는 축이 다른 것이 어텐션의 두 단계입니다.
- 되돌리기 —
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를 정하는 제약
이 자리에서 흔한 오해가 를 손잡이로 보는 것입니다. 헤드 수 와 머리 폭 를 따로 정할 수 있을 것 같지만, view 가 요구하는 것은 한 줄입니다. 768을 쪼개는 것이라 는 로 저절로 정해지고, 고를 수 있는 것은 뿐입니다.
그리고 도 아무 값이나 되지 않습니다. 768을 나누어떨어뜨려야 하므로 12는 64, 16은 48, 8은 96이 되지만 10은 76.8이라 안 됩니다. "헤드를 열 개로 늘려 보자"가 안 되는 이유가 설계 철학이 아니라 나눗셈이라는 것이 이 줄의 전부입니다.
보폭과 contiguous
마지막 줄에서 view 가 아니라 reshape 을 쓴 이유가 있습니다. 텐서는 숫자가 한 줄로 늘어선 데이터 한 벌과, 축마다 "이 축에서 한 칸 가려면 몇 칸 건너뛰어야 하는가"를 적어 둔 수 한 벌로 되어 있습니다. 뒤쪽을 보폭이라고 부릅니다.
transpose 는 숫자를 하나도 안 옮기고 보폭의 순서만 바꿔 둡니다. 그래서 공짜인데, 그 결과는 "읽는 순서"와 "메모리에 놓인 순서"가 어긋난 상태입니다. view 는 메모리에 놓인 순서 그대로 칸을 다시 나누는 함수라 그 상태를 거부합니다. reshape 은 필요하면 복사까지 해 주므로 안전하고, contiguous() 는 그 복사를 명시적으로 먼저 해 두는 것입니다. 둘 다 공짜가 아니라 데이터를 한 벌 더 쓰는 값을 치릅니다 — transpose 가 공짜인 대신 그다음이 비싸지는 셈입니다.
transpose 뒤에 view 를 쓰면 나는 오류는 shape 산수가 아니라 메모리 배치의 문제라서, 이 글의 세 줄 규칙으로는 예측되지 않는 유일한 자리입니다.
GQA와 MQA의 헤드 수
요즘 모델은 와 의 헤드 수를 다르게 둡니다. 질의는 32 헤드인데 키와 값은 8 헤드만 두고 넷씩 나눠 쓰는 식이고, 8을 1까지 내린 것이 MQA입니다. 캐시로 들고 있어야 할 가 4분의 1이나 32분의 1로 줄어드는 것이 목적입니다.
shape으로 보면 이것은 새 연산이 아니라 규칙 3을 한 번 더 쓴 것입니다. 를 로 보고 질의를 로 보면 크기 1인 축이 4로 늘어나 짝이 맞습니다. 값을 복사하지 않고도 한 벌의 키가 네 질의 헤드에 걸리는 것이라, 절약되는 것이 계산이 아니라 메모리라는 점도 같은 그림에서 읽힙니다.
점수 표의 메모리와 연산량
세 규칙은 shape이 맞는지만 알려 주지만, 같은 shape에서 비용도 함께 읽힙니다.
원소 수로 세기
점수 표 의 원소 수는 그냥 네 수의 곱입니다. , , 이면
로 4억 개가 조금 넘습니다. fp16이면 한 개가 2바이트이므로 805MB입니다. 가중치가 아니라 중간 결과 한 장이 그만큼입니다. 모델 파라미터를 세어 메모리를 가늠하다가 실제로 넘치는 자리가 대개 여기입니다.
T가 두 배면 네 배
이 수에서 만 두 번 곱해진다는 것이 핵심입니다. 배치와 헤드는 한 번씩 들어가는데 는 행에도 열에도 서므로, 문맥을 두 배로 늘리면 표가 네 배가 됩니다.
| 원소 수 | 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 세기
연산량도 같은 규칙으로 셉니다. 앞에서 와 의 곱이 라고 했으니, 배치 축 와 는 그냥 앞에 곱해집니다.
두 단계가 정확히 같은 값입니다 — 앞은 가 축약되고 뒤는 가 축약되는데, 세 수의 곱이 로 같기 때문입니다. 위의 설정에 넣으면 한 단계가 51.5 GFLOP이라 둘이 103 GFLOP입니다. 메모리와 연산량이 둘 다 을 따라간다는 것이 이 절의 결론이고, 그래서 둘 중 하나만 줄이는 해법이 없습니다.
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인 축을 안 뗀 채 더하기. 과 가 만나 가 되는 그 자리입니다.
세 경우의 공통점은 어긋난 쪽이 1이거나, 곱이 같다는 것입니다. 규칙이 허용하도록 만들어진 자리라 검사로는 못 막고, 중간 텐서의 shape을 한 번 찍어 보는 것 말고는 방법이 없습니다.
확인하는 순서
오류 메시지는 뒤 두 축만 적어 주므로 앞 축은 안 보입니다. 순서를 정해 두면 빠릅니다.
- 뒤 두 축 — 맞닿는 자리가 같은 값인지. 메시지에 적힌 것이 이것입니다.
- 앞 축 — 배치와 헤드가 서로 같거나 한쪽이 1인지. 메시지에 안 적히는 자리입니다.
- 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이 값을 하나로 정해 주지만 비용은 정해 주지 않습니다. 계산 순서에 따라 중간 결과의 크기가 달라지기 때문입니다.
가 , 가 , 가 일 때 를 생각해 봅니다. 로 가면 가 이라 중간에 100만 개짜리가 서고 연산이 4천만 번입니다. 로 가면 가 이라 중간 결과가 100개이고 연산은 40만 번입니다. 같은 식, 같은 답, 100배 차이입니다.
torch.einsum 은 인자 셋 이상이면 순서를 스스로 고르지만, 축 길이가 비슷할 때는 그 선택이 최선이 아닐 수 있습니다. 중간 결과가 얼마나 커지는지를 손으로 한 번 세어 보는 것이 그래서 여전히 필요합니다.
축 이름을 먼저 정하기
einsum이 규칙을 없애 주는 것은 아닙니다. 문자열을 잘못 쓰면 오류 없이 다른 값이 나오므로, 어느 문자가 축약되는지를 스스로 읽을 수 있어야 하는 것은 그대로입니다.
대신 습관 하나가 바뀝니다. 코드를 쓰기 전에 축마다 문자를 정해 두면 — 배치는 b, 헤드는 h, 토큰 자리는 i와 j, 머리 폭은 d — 그다음부터는 원하는 결과의 문자열을 적는 일만 남습니다. 축의 물리적 순서를 기억할 필요가 없어지므로 transpose 를 넣었다 빼는 왕복 자체가 사라집니다. 문자를 먼저 정하는 이 한 걸음이 view·transpose 실수를 줄이는 가장 값싼 방법입니다.
정리
- 같은 곱을 네 시선으로 계산할 수 있습니다. 원소는 행·열의 내적, 행 시선은 오른쪽 행들의 혼합, 열 시선은 왼쪽 열들의 혼합, 외적 시선은 판의 합입니다. 넷은 삼중합을 도는 순서만 다르므로 값도 곱셈 횟수 도 같습니다.
- 외적 시선은 축약되는 축을 조각낼 수 있게 해 줍니다. 판 몇 장을 먼저 더해 두었다가 나중에 얹어도 되고, 긴 문맥을 조각으로 도는 구현이 여기에 섭니다.
- shape 산수는 세 줄입니다. 뒤 두 축만 곱하고, 맞닿은 축은 같아야 하며 사라지고, 앞 축은 따라 나옵니다. 크기 1인 축은 브로드캐스트로 늘어나는데, 이것이 마스크를 한 벌만 두게 해 주는 동시에 과 를 말없이 로 만드는 자리이기도 합니다.
- 멀티헤드의
view·transpose순서는 외울 것이 아닙니다. 곱할 두 축을 뒤로 보낸다는 목적에서 매번 재구성되고, 는 손잡이가 아니라 입니다. transpose뒤의view실패만 예외입니다. shape이 아니라 보폭의 문제라reshape으로 풀고, 그 대신 복사 비용을 뭅니다.- 점수 표는 을 따라갑니다. , , 에서 4.03억 개·805MB이고, 를 두 배로 하면 메모리도 FLOPs도 네 배입니다.
- 오류가 안 나는 어긋남이 더 비쌉니다. 배치 축 1, 곱이 같은
view, 크기 1인 축의 덧셈 — 셋 다 규칙이 허락하는 자리라 shape을 찍어 보는 것 말고는 방법이 없습니다.
다음 글은 시선을 다시 기하로 돌립니다 — 행렬 하나로 회전·전단·스케일을 전부 만들고, 회전행렬 를 세워 를 곱셈으로 직접 확인합니다.
읽어주셔서 감사합니다. 😊

