MHA in LLama

a9umon·2025년 12월 31일

보충자료

목록 보기
3/9

참고한 블로그는 여기

SHA (Single-Head Attention)

  • SHA는 다음과 같이 attention score를 계산한다.
  • 원래는 Q를 모든 token을 각각 Q로 만들어 모든 token에 대한 importance를 계산해야 하지만, matrix operation을 통해 중첩되는 반복문 구조를 없앤다.

MHA (Multi-Head Attention)

병렬 Head의 역할

  • 독립적으로 Attention 연산

  • 각각의 Head 에서 Q1, K1, V1을 이용해서 attention 연산, Q2, K2, V2를 이용해서 attention 연산 → 이렇게 계산된 attention들은 마지막에 다시 병합됨

  • 각 Head에서 다른 관점으로 attention을 계산하므로 여러 관점을 학습하게 됨

  • 이 구조는 SHA가 여러개 붙어있는 구조와 동일함

    • 따라서 matrix operation을 사용하지 않으면, 반복문 세 개가 중첩된 구조를 갖게 됨.
    • 이를 방지하기 위해 matrix 연산을 사용함
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)
profile
이것저것 다 합니다.

0개의 댓글