Transformer Multi-Head Attention MHA

Transformer Multi-Head Attention(MHA) 완벽 이해

Transformer의 핵심 구조 중 하나는 Multi-Head Attention(다중 헤드 어텐션)입니다.

앞에서 Scaled Dot-Product Attention을 이해했다면, 이제 자연스럽게 다음 질문이 생깁니다.

Attention을 왜 하나만 계산하지 않고 여러 개의 Head로 나눠서 계산할까?

Multi-Head Attention의 핵심 아이디어는 하나의 Attention만 사용하는 대신, Q(Query), K(Key), V(Value)를 여러 조의 Head로 나누어 서로 다른 관점에서 Attention을 계산한 뒤 다시 하나로 합치는 것입니다.

link:https://proceedings.neurips.cc/paper_files/paper/2017/file/3f5ee243547dee91fbd053c1c4a845aa-Paper.pdf

이번 글에서는 다음 PyTorch 코드를 기준으로 Multi-Head Attention의 내부 동작을 하나씩 살펴보겠습니다.


1. 전체 코드

import copy
import torch
import torch.nn as nn


def clones(module, N):
    return nn.ModuleList([
        copy.deepcopy(module)
        for _ in range(N)
    ])


class MultiHeadedAttention(nn.Module):

    def __init__(self, head, embedding_dim, dropout=0.1):
        super().__init__()

        assert embedding_dim % head == 0

        self.d_k = embedding_dim // head number, h
        self.head = head

        self.linears = clones(
            nn.Linear(embedding_dim, embedding_dim),
            4
        )

        self.attn = None

        self.dropout = nn.Dropout(p=dropout)


    def forward(self, query, key, value, mask=None):

        if mask is not None:
            mask = mask.unsqueeze(1)

        batch_size = query.size(0)

        query, key, value = [
            model(x)
            .view(batch_size, -1, self.head, self.d_k)
            .transpose(1, 2)

            for model, x in zip(
                self.linears,
                (query, key, value)
            )
        ]

        x, self.attn = attention(
            query,
            key,
            value,
            mask=mask,
            dropout=self.dropout
        )

        x = (
            x.transpose(1, 2)
            .contiguous()
            .view(
                batch_size,
                -1,
                self.head * self.d_k
            )
        )

        return self.linears[-1](x)

전체 코드가 처음에는 복잡해 보이지만 실제 흐름은 다음과 같습니다.

Input
  ↓
Q, K, V Linear Projection
  ↓
Head 여러 개로 분할
  ↓
각 Head에서 Attention 계산
  ↓
Head들을 다시 결합
  ↓
Output Linear
  ↓
Multi-Head Attention Output

이제 하나씩 살펴보겠습니다.


2. clones() 함수는 무엇인가?

먼저 다음 함수가 등장합니다.

def clones(module, N):
    return nn.ModuleList([
        copy.deepcopy(module)
        for _ in range(N)
    ])

clones()는 동일한 구조를 가진 PyTorch Layer를 N개 생성하는 함수입니다.

예를 들어,

clones(
    nn.Linear(512, 512),
    4
)

를 실행하면 개념적으로 다음과 같은 결과가 만들어집니다.

nn.ModuleList([
    nn.Linear(512, 512),
    nn.Linear(512, 512),
    nn.Linear(512, 512),
    nn.Linear(512, 512)
])

여기서 중요한 점이 있습니다.

구조는 동일하지만 각각 독립적인 Layer입니다.

즉,

self.linears[0].weight
self.linears[1].weight
self.linears[2].weight
self.linears[3].weight

는 서로 다른 학습 가능한 Parameter입니다.

copy.deepcopy()를 사용하는 이유도 바로 이것입니다.


3. 왜 Linear Layer가 4개 필요할까?

Multi-Head Attention의 생성자에는 다음 코드가 있습니다.

self.linears = clones(
    nn.Linear(embedding_dim, embedding_dim),
    4
)

여기서 4를 보고

Head가 4개라는 뜻인가?

라고 생각하기 쉽지만 그렇지 않습니다.

이 4개는 각각 다음 역할을 합니다.

linears[0] → Query 변환
linears[1] → Key 변환
linears[2] → Value 변환
linears[3] → 최종 Output 변환

즉,

수식으로 표현하면 앞의 세 Linear Layer는 다음과 같습니다.

Q=XWQQ = XW^Q

K=XWKK = XW^K

V=XWVV = XW^V

그리고 여러 Head의 결과를 결합한 뒤 마지막 Linear Layer를 적용합니다.

Concat(head1,,headh)WOConcat(head_1,\dots,head_h)W^O

따라서 총 4개의 Linear Projection이 필요합니다.


4. headembedding_dim

생성자에서 다음 두 값을 전달받습니다.

def __init__(
    self,
    head,
    embedding_dim,
    dropout=0.1
):

예를 들어 Transformer 원 논문의 대표적인 설정을 생각해보겠습니다.

embedding_dim = 512
head = 8

전체 Embedding Dimension은 512이고 Attention Head는 8개입니다.

그러면 하나의 Head가 담당하는 차원은 다음과 같습니다.

dk=5128=64d_k = \frac{512}{8}=64

코드에서는 다음과 같이 계산합니다.

self.d_k = embedding_dim // head

따라서

self.d_k = 64

가 됩니다.


5. 왜 embedding_dim % head == 0이어야 할까?

다음 코드도 중요합니다.

assert embedding_dim % head == 0

Embedding Vector를 Head 개수만큼 동일하게 나누어야 하기 때문입니다.

예를 들어,

embedding_dim = 512
head = 8

512 / 8 = 64

이므로 문제없이 나눌 수 있습니다.

512차원

Head 1 → 64
Head 2 → 64
Head 3 → 64
Head 4 → 64
Head 5 → 64
Head 6 → 64
Head 7 → 64
Head 8 → 64

그리고

64 × 8 = 512

가 됩니다.

반대로 다음과 같은 경우라면,

embedding_dim = 512
head = 7

512가 7로 나누어떨어지지 않기 때문에 동일한 크기의 Head로 분리할 수 없습니다.

따라서 assert를 통해 잘못된 설정을 미리 방지합니다.


6. Mask에 왜 unsqueeze(1)을 사용할까?

forward()를 보면 다음 코드가 있습니다.

if mask is not None:
    mask = mask.unsqueeze(1)

Multi-Head Attention에서는 Attention을 여러 Head에서 동시에 계산합니다.

예를 들어 원래 Mask의 Shape이

(batch_size, seq_len, seq_len)

이라면,

mask.unsqueeze(1)

을 통해

(batch_size, 1, seq_len, seq_len)

으로 변경할 수 있습니다.

여기서 추가된 1 차원은 Head 차원에 대응하기 위한 차원입니다.

이렇게 만들어두면 PyTorch의 Broadcasting을 이용하여 동일한 Mask를 모든 Head에 적용할 수 있습니다.

Mask

(batch, 1, seq, seq)

          ↓ Broadcasting

Head 1 ─ Mask
Head 2 ─ Mask
Head 3 ─ Mask
...
Head 8 ─ Mask

즉, Head마다 Mask를 따로 8개 만들 필요가 없습니다.

예제 코드

아래는 unsqueeze(1)이 실제로 어떻게 동작하는지 보여주는 간단한 PyTorch 예제입니다.

import torch

batch_size = 2 
seq_len = 4 
num_heads = 8

# (batch, seq, seq) 형태의 mask
mask = torch.randint(0, 2, (batch_size, seq_len, seq_len))
#print("mask: \n", mask )
print("원래 mask shape:", mask.shape)

# head 차원 추가
mask = mask.unsqueeze(1)
print("unsqueeze 후 shape:", mask.shape)

# attention score (예시)
attn_score = torch.randn(batch_size, num_heads, seq_len, seq_len)
print("attn_score:", attn_score.shape)
#print("attn_score: \n", attn_score )

# broadcasting으로 mask 적용 가능
masked_score = attn_score.masked_fill(mask == 0, float("-inf"))
print("masked score shape:", masked_score.shape)

실행 결과:

이처럼 unsqueeze(1)을 사용하면 모든 head에 동일한 mask를 자동으로 공유할 수 있어 구현이 훨씬 간단해집니다.


7. Batch Size 가져오기

다음으로 입력 Tensor의 Batch Size를 가져옵니다.

batch_size = query.size(0)

예를 들어 Query의 Shape이

(32, 10, 512)

라면,

32  → batch_size
10  → sequence length
512 → embedding dimension

이므로

batch_size = 32

가 됩니다.


8. Multi-Head Attention에서 가장 중요한 코드

이제 이 코드에서 가장 중요한 부분입니다.

query, key, value = [ 
model(x).view( batch_size, -1, self.head, self.d_k ).transpose(1,2)
              for model, x in zip(self.linears, (query, key, value))
]


처음 보면 상당히 복잡해 보입니다.

하지만 이 코드는 사실 3단계 과정(Q, K, V 생성 → Head 분할 → 차원 재배치)을 한 번에 압축해 놓은 것입니다.



1) zip() : Q, K, V와 Linear Layer를 1:1로 매칭

python
zip(self.linears, (query, key, value))

이 부분은 다음처럼 매칭됩니다.

  • Linear(0) → Query
  • Linear(1) → Key
  • Linear(2) → Value

즉, 각각의 입력을 서로 다른 Linear Layer에 통과시킵니다.


2) 그런데 왜 4번째 Linear는 실행되지 않을까?

핵심 이유는 zip()이 3개까지만 묶기 때문입니다.

self.linears = [Linear_Q, Linear_K, Linear_V, Linear_Out]
(query, key, value)

여기서 중요한 규칙이 있습니다.

> `zip()`은 가장 짧은 iterable 기준으로 동작한다.

즉,

- self.linears → 4개
- (query, key, value) → 3개

이 동작을 코드로 직접 확인해보면 이해가 훨씬 쉬워집니다.

예를 들어 다음과 같은 간단한 예제를 생각해보겠습니다.

python
linears = ["Linear_Q", "Linear_K", "Linear_V", "Linear_Out"]
inputs = ["query", "key", "value"]

for model, x in zip(linears, inputs):
    print(model, "↔", x)

실행 결과는 다음과 같습니다.

Linear_Q ↔ query
Linear_K ↔ key
Linear_V ↔ value

여기서 중요한 포인트는 마지막 Linear_Out은 아예 사용되지 않는다는 점입니다.

즉, zip은 다음처럼 동작합니다.

(linears[0], inputs[0]) → OK
(linears[1], inputs[1]) → OK
(linears[2], inputs[2]) → OK
(linears[3], inputs[3]) → ❌ inputs가 없음 → 중단

그래서 Multi-Head Attention 코드에서도

  • Q → Linear_Q
  • K → Linear_K
  • V → Linear_V

까지만 처리되고,

self.linears[-1]

인 Output Linear은 Attention 계산이 끝난 뒤 따로 사용되는 구조입니다.

이 경우 결과는 다음처럼 잘립니다.

Linear_Q   ↔ Query
Linear_K   ↔ Key
Linear_V   ↔ Value

그리고 Linear_Out은 매칭 대상이 없기 때문에 아예 루프에 들어가지 못합니다.


3) 4번째 Linear는 언제 사용되나?

4번째 Linear는 여기서 사용되지 않고, 완전히 다른 단계에서 사용됩니다.

코드 마지막 부분을 보면:

return self.linears[-1](x)

즉, 4번째 Linear는 다음 역할을 합니다.

Multi-Head Attention 결과를 다시 원래 embedding space로 projection


4) 정리

  • zip() 때문에 Q, K, V까지만 처리됨
  • Linear 4개 중 마지막 1개는 zip 범위 밖이라 실행되지 않음
  • 4번째 Linear는 Attention 이후 “출력 변환용”으로 따로 사용됨

한 줄 요약

👉 4번째 Linear는 zip()이 Q, K, V까지만 묶기 때문에 실행되지 않고, Attention 계산이 끝난 뒤 최종 출력 변환 단계에서 따로 사용된다.


2) model(x) : Q, K, V 각각 Linear 변환

model(x)

여기서 하는 일은 단순합니다.

  • 입력 (embedding)을
  • 학습 가능한 가중치로 변환해서
  • 새로운 Q, K, V 표현을 만드는 과정입니다.

즉:

Q = XW^Q
K = XW^K
V = XW^V

3) view() : embedding을 head 개수만큼 분할

.view(batch_size, -1, self.head, self.d_k)

이 단계가 Multi-Head Attention의 핵심입니다.

예를 들어:

  • embedding_dim = 512
  • head = 8
  • d_k = 64

이라면

(32, 10, 512)
        ↓
(32, 10, 8, 64)

즉, 하나의 512차원 벡터를 8개의 64차원 head로 분리합니다.


4) transpose(1, 2) : head를 앞으로 이동

.transpose(1, 2)

이 작업은 계산을 쉽게 하기 위한 차원 재배치입니다.

변화:

(32, 10, 8, 64)
        ↓
(32, 8, 10, 64)

이제 구조는 다음과 같습니다:

  • batch
  • head
  • sequence
  • feature

즉, 각 head가 독립적으로 attention을 계산할 수 있는 형태가 됩니다.


5) 최종 결과

이 전체 과정을 거치면 Q, K, V는 모두 다음 형태가 됩니다:

(batch_size, head, seq_len, d_k)

예:

(32, 8, 10, 64)

핵심 요약

이 코드는 사실 아래 3단계를 한 줄로 압축한 것입니다:

  1. Q, K, V 각각 Linear 변환
  2. embedding을 head 개수만큼 분할
  3. attention 계산을 위한 차원 재배치

즉, Multi-Head Attention의 핵심은:

“하나의 벡터를 여러 개의 작은 관점(head)으로 나눠서 각각 attention을 수행하는 구조”

입니다.

하지만 세 단계로 나누면 이해하기 쉽습니다.

1. Linear Projection
2. Head 분리
3. 차원 순서 변경

각 단계의 의미를 조금 더 풀어서 보면 다음과 같습니다.

1. Linear Projection
   → Q, K, V 각각을 학습 가능한 Linear Layer로 변환합니다.
   → 입력을 새로운 표현 공간으로 매핑하는 과정입니다.

   예를 들어 Query 텐서가 (batch, seq, embedding_dim)이라면,
   Linear Layer를 통과한 뒤에도 shape은 동일하지만
   각 토큰의 의미 표현이 "학습된 방식으로 변환"됩니다.

   즉, 단순한 복사가 아니라
   "Attention을 잘 수행하기 위한 특징 공간으로 재구성"하는 단계입니다.

2. Head 분리
   → embedding_dim을 head 개수만큼 나누어
     (head, d_k) 구조로 쪼갭니다.
   → 하나의 큰 벡터를 여러 개의 작은 벡터로 분해하는 단계입니다.

   예를 들어 embedding_dim = 512, head = 8이라면
   하나의 512차원 벡터를 64차원 벡터 8개로 나누는 것입니다.

   이 과정을 통해 각 Head는 서로 다른 부분 공간(subspace)에서
   독립적으로 Attention을 학습할 수 있게 됩니다.

3. 차원 순서 변경
   → (batch, seq, head, d_k) 형태를
     (batch, head, seq, d_k)로 바꿉니다.
   → 각 Head가 독립적으로 Attention을 계산할 수 있도록 만드는 과정입니다.

   이 변환이 중요한 이유는
   Attention 연산이 "Head 단위로 병렬 처리"되기 때문입니다.

   즉, Head를 앞쪽 차원으로 올려야
   각 Head별로 QKᵀ 연산을 한 번에 수행할 수 있습니다.

9. zip()은 무엇을 연결하는가?

먼저 다음 부분을 살펴보겠습니다.

zip(
    self.linears,
    (query, key, value)
)

self.linears에는 4개의 Linear Layer가 있습니다.

Linear_Q
Linear_K
Linear_V
Linear_Output

반면 입력은 3개입니다.

query
key
value

zip()은 짧은 쪽을 기준으로 묶기 때문에 다음 세 쌍만 만들어집니다.

Linear 0 ↔ Query
Linear 1 ↔ Key
Linear 2 ↔ Value

따라서 마지막

Linear 3

은 여기서 사용되지 않습니다.

마지막 Linear는 Attention 계산이 모두 끝난 뒤 사용됩니다.

return self.linears[-1](x)

10. model(x) : Q, K, V Linear Projection

먼저

model(x)

가 실행됩니다.

예를 들어 Query가

(batch_size, seq_len, embedding_dim)

(32, 10, 512)

라면,

nn.Linear(512, 512)

를 통과하더라도 Shape은 그대로입니다.

(32, 10, 512)
        ↓
Linear(512, 512)
        ↓
(32, 10, 512)

하지만 값은 달라집니다.

즉, 단순히 입력을 복사하는 것이 아니라 학습 가능한 Weight Matrix를 통해 새로운 표현으로 변환합니다.


11. view() : 하나의 Embedding을 여러 Head로 나누기

다음으로,

.view(
    batch_size,
    -1,
    self.head,
    self.d_k
)

를 수행합니다.

예를 들어,

batch_size = 32
seq_len = 10
embedding_dim = 512
head = 8
d_k = 64

라면 원래 Shape은

(32, 10, 512)

입니다.

이를

(32, 10, 8, 64)

로 변경합니다.

즉,

512

↓

8 × 64

로 분리한 것입니다.

중요한 점은 512차원의 벡터를 8개의 64차원 벡터로 나눈 것입니다.


12. transpose(1, 2)는 왜 필요할까?

현재 Shape은

(batch, seq_len, head, d_k)

즉,

(32, 10, 8, 64)

입니다.

하지만 Attention을 Head별로 계산하려면 Head 차원을 앞으로 가져오는 것이 편리합니다.

그래서

.transpose(1, 2)

를 사용합니다.

결과는

(32, 8, 10, 64)

입니다.

즉,

Before

(batch, seq, head, d_k)

(32, 10, 8, 64)


After

(batch, head, seq, d_k)

(32, 8, 10, 64)

가 됩니다.

이제 각 Head가 독립적으로 Attention을 계산할 수 있는 형태가 되었습니다.


13. Q, K, V의 최종 Shape

결과적으로 Q, K, V는 모두 다음과 같은 Shape을 갖습니다.

(batch_size, head, seq_len, d_k)

예를 들어,

Query → (32, 8, 10, 64)
Key   → (32, 8, 10, 64)
Value → (32, 8, 10, 64)

가 됩니다.

이를 그림으로 생각하면 다음과 같습니다.

Query
(32, 10, 512)

      ↓ Linear

(32, 10, 512)

      ↓ view

(32, 10, 8, 64)

      ↓ transpose

(32, 8, 10, 64)
      ↑
    8 Heads

Key와 Value도 동일한 과정을 거칩니다.


14. 각 Head에서 Attention 계산

이제 준비된 Q, K, V를 기존의 attention() 함수에 전달합니다.

x, self.attn = attention(
    query,
    key,
    value,
    mask=mask,
    dropout=self.dropout
)

Scaled Dot-Product Attention의 핵심 계산은 다음과 같습니다.

softmax(QKTdk)Vsoftmax \left( \frac{QK^T}{\sqrt{d_k}} \right)V

Mask가 존재한다면 Softmax 이전에 적용됩니다.

전체 흐름은 다음과 같습니다.

Q × Kᵀ
   ↓
Scale by √d_k
   ↓
Mask
   ↓
Softmax
   ↓
Dropout
   ↓
× V
   ↓
Attention Output

중요한 점은 이 계산이 8개의 Head에서 동시에 각각 수행된다는 것입니다.


15. 왜 여러 Head를 사용할까?

하나의 Attention만 사용하면 하나의 표현 공간에서 단어 간 관계를 학습합니다.

반면 여러 Head를 사용하면 서로 다른 Head가 서로 다른 관계를 학습할 가능성이 생깁니다.

예를 들어 다음 문장이 있다고 해보겠습니다.

The student studies Transformer carefully.

각 Head가 학습 과정에서 다음과 같은 관계에 집중할 수 있습니다.

Head 1 → student ↔ studies
Head 2 → studies ↔ Transformer
Head 3 → studies ↔ carefully
Head 4 → The ↔ student
...

단, 특정 Head가 반드시 특정 문법 관계를 담당하도록 사람이 지정하는 것은 아닙니다.

각 Head가 어떤 관계에 집중할지는 학습 과정에서 결정됩니다.

이것이 Multi-Head Attention의 중요한 장점입니다.


16. Attention 결과의 Shape

Attention 계산이 끝난 후 x의 Shape은 여전히 다음과 같은 형태입니다.

(batch_size, head, seq_len, d_k)

예를 들어,

(32, 8, 10, 64)

입니다.

하지만 다음 Layer에 전달하려면 다시 원래의 embedding_dim = 512 형태로 합쳐야 합니다.


17. Head들을 다시 하나로 합치기

다음 코드가 그 역할을 합니다.

x = (
    x.transpose(1, 2)
    .contiguous()
    .view(
        batch_size,
        -1,
        self.head * self.d_k
    )
)

먼저,

x.transpose(1, 2)

를 수행합니다.

(32, 8, 10, 64)

↓

(32, 10, 8, 64)

앞에서 Head를 분리하기 위해 했던 transpose()를 다시 되돌리는 것입니다.


18. contiguous()는 왜 사용할까?

여기서 중요한 코드가 하나 등장합니다.

.contiguous()

PyTorch에서 transpose()는 실제 메모리 데이터를 새롭게 정렬하는 것이 아니라 Tensor를 바라보는 차원 순서를 변경합니다.

따라서 메모리상 데이터가 연속적으로 배치되어 있지 않을 수 있습니다.

그 상태에서 바로

view()

를 사용하면 문제가 발생할 수 있습니다.

그래서

x.transpose(1, 2).contiguous()

를 통해 데이터를 연속적인 메모리 구조로 만들어준 후 view()를 사용합니다.

쉽게 말하면,

transpose
    ↓
차원 순서는 변경됨
    ↓
메모리 배치는 아직 복잡할 수 있음
    ↓
contiguous()
    ↓
연속적인 메모리 형태로 정리
    ↓
view()

라고 이해하면 됩니다.


19. view()를 이용해 Head 결합하기

마지막으로,

.view(
    batch_size,
    -1,
    self.head * self.d_k
)

를 수행합니다.

현재

head = 8
d_k = 64

이므로

8×64=5128\times64=512

입니다.

따라서

(32, 10, 8, 64)

↓

(32, 10, 512)

가 됩니다.

즉, 여러 Head의 결과를 다시 하나의 Embedding Vector로 합칩니다.

이 과정이 바로 논문에서 말하는 Concat에 해당합니다.

Head 1 → 64
Head 2 → 64
Head 3 → 64
...
Head 8 → 64

        ↓ Concat

64 × 8 = 512

        ↓

512-dimensional Vector

20. 마지막 Linear Layer

이제 마지막 코드입니다.

return self.linears[-1](x)

앞에서 self.linears에는 총 4개의 Linear Layer가 있었습니다.

linears[0] → Q
linears[1] → K
linears[2] → V
linears[3] → Output

따라서

self.linears[-1]

은 마지막 Output Linear Layer를 의미합니다.

현재 x

(batch_size, seq_len, embedding_dim)

형태입니다.

예를 들어,

(32, 10, 512)

입니다.

마지막으로

Linear(512, 512)

를 통과합니다.

Concat 결과
(32, 10, 512)

      ↓

Linear(512, 512)

      ↓

최종 Output
(32, 10, 512)

이것이 Multi-Head Attention의 최종 출력입니다.


21. 전체 Shape 변화 한 번에 보기

예를 들어 다음과 같이 설정했다고 가정하겠습니다.

batch_size = 32
seq_len = 10
embedding_dim = 512
head = 8
d_k = 64

그러면 전체 Shape 변화는 다음과 같습니다.

Input Q/K/V
(32, 10, 512)

        ↓ Linear Projection

Q/K/V
(32, 10, 512)

        ↓ view()

(32, 10, 8, 64)

        ↓ transpose(1, 2)

(32, 8, 10, 64)

        ↓
Scaled Dot-Product Attention
        ↓

(32, 8, 10, 64)

        ↓ transpose(1, 2)

(32, 10, 8, 64)

        ↓ contiguous()
        ↓ view()

(32, 10, 512)

        ↓ Output Linear

(32, 10, 512)

Shape만 정확하게 이해해도 Multi-Head Attention 코드의 절반 이상을 이해했다고 볼 수 있습니다.


22. 코드 한 줄의 의미 다시 보기

이제 처음에는 복잡해 보였던 코드를 다시 보겠습니다.

query, key, value = [
    model(x)
    .view(
        batch_size,
        -1,
        self.head,
        self.d_k
    )
    .transpose(1, 2)

    for model, x in zip(
        self.linears,
        (query, key, value)
    )
]

이제 이 코드는 사실 다음 세 문장의 압축이라는 것을 알 수 있습니다.

① Q, K, V를 각각 Linear Layer에 통과시킨다.

② embedding_dim을 head × d_k로 나눈다.

③ Head별 Attention 계산을 위해
   (batch, head, seq, d_k) 형태로 변경한다.

그리고 Attention 계산 이후의

x = (
    x.transpose(1, 2)
    .contiguous()
    .view(
        batch_size,
        -1,
        self.head * self.d_k
    )
)

는 반대로,

① Head 차원을 다시 뒤로 보낸다.

② 메모리를 연속적으로 정리한다.

③ 여러 Head를 하나의 embedding_dim으로 합친다.

라는 의미입니다.


23. Multi-Head Attention 전체 구조

최종적으로 전체 과정을 연결하면 다음과 같습니다.

Transformer 논문의 Multi-Head Attention 수식과 정확히 연결하면 다음과 같습니다.

headi=Attention(QWiQ,KWiK,VWiV)head_i = Attention ( QW_i^Q, KW_i^K, VW_i^V )

그리고

MultiHead(Q,K,V)=Concat(head1,,headh)WOMultiHead(Q, K, V ) = Concat(head_1,\dots,head_h)W^O

즉, 코드의

self.linears[0]
self.linears[1]
self.linears[2]

가 Q, K, V를 만드는 Projection에 해당하고,

attention(...)

이 각 Head의 Attention 계산에 해당하며,

.view(
    batch_size,
    -1,
    self.head * self.d_k
)

Concat에 해당하고,

self.linears[-1](x)

가 마지막 (W^O) Projection에 해당합니다.


24. 핵심 정리

Multi-Head Attention 코드에서 반드시 기억해야 할 핵심은 다음과 같습니다.

  • clones(..., 4)4는 Head 개수가 아니다.
  • 4개의 Linear Layer는 각각 Q, K, V, Output Projection에 사용된다.
  • 실제 Head 개수는 head 변수로 결정된다.
  • embedding_dimhead로 나누어떨어져야 한다.
  • 각 Head의 차원은 d_k = embedding_dim // head이다.
  • view()를 통해 Embedding Dimension을 여러 Head로 분리한다.
  • transpose()를 통해 Head별 Attention 계산이 편리한 형태로 변경한다.
  • 각 Head에서 Scaled Dot-Product Attention을 독립적으로 계산한다.
  • Attention 계산 후 여러 Head의 결과를 다시 Concat한다.
  • 마지막 Linear Layer를 통과시켜 최종 Multi-Head Attention 출력을 만든다.

결국 Multi-Head Attention의 핵심 흐름은 아주 간단하게 정리할 수 있습니다.

Q, K, V 생성
      ↓
여러 Head로 분할
      ↓
각 Head에서 Attention 계산
      ↓
Head 결과 Concat
      ↓
Output Linear

즉, Multi-Head Attention은 하나의 Attention을 단순히 여러 번 복사하는 것이 아니라, Embedding 공간을 여러 Head로 나누어 다양한 관점에서 단어 간 관계를 학습하고 그 결과를 다시 하나로 결합하는 구조라고 이해하면 됩니다.

댓글 남기기

이메일 주소는 공개되지 않습니다. 필수 필드는 *로 표시됩니다