참고한 블로그는 여기

독립적으로 Attention 연산
각각의 Head 에서 Q1, K1, V1을 이용해서 attention 연산, Q2, K2, V2를 이용해서 attention 연산 → 이렇게 계산된 attention들은 마지막에 다시 병합됨
각 Head에서 다른 관점으로 attention을 계산하므로 여러 관점을 학습하게 됨

이 구조는 SHA가 여러개 붙어있는 구조와 동일함
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads, seq_length):
super(MultiHeadAttention), self.__init__()
self.num_heads = num_heads
self.d_model = d_model
self.seq_length = seq_length
self.head_dim = d_model // num_heads
self.query = nn.Linear(d_model, d_model, bias=False)
self.key = nn.Linear(d_model, d_model, bias=False)
self.value = nn.Linear(d_model, d_model, bias=False)
self.fc_out = nn.Linaer(d_model, d_model, bias=False)
def forward(self, x, mask=None):
''' matrix 연산을 사용한 부분 '''
batch_size, seq_length, _ = x.size()
Q = self.query(x).view(batch_size, seq_length, self.num_heads, self.head_dim)
K = self.key(x).view(batch_size, seq_length, self.num_heads, self.head_dim)
V = self.value(x).view(batch_size, seq_length, self.num_heads, self.head_dim)
attention_scores = Q @ K.transpose(-2, -1) / (self.head_dim ** 0.5) # (B, L, H, D) @ (B, L, D, H) -> (B, L, H, H)
attention_probs = F.softmax(attention_scores, dim=-1) # (B, L, H, H)
attention_output = attention_probs @ V # (B, L, H, H) @ (B, L, H, D) -> (B, L, H, D)
attention_output = attention_output.transpose(1,2).contiguous().view(batch_size, seq_length, self.d_model)