시작하기무료로 시작하기

MultiHeadAttention 클래스에 메서드 추가하기

이 연습 문제에서는 네 가지 메서드를 정의해 MultiHeadAttention 클래스를 처음부터 완성해 볼 거예요.

  • .split_heads(): 입력 임베딩을 어텐션 헤드 수만큼 분할하고 변환해요.
  • .compute_attention(): 스케일드 도트 프로덕트 어텐션 가중치를 계산하고 value 행렬과 곱해요.
  • .combine_heads(): 어텐션 가중치를 입력 임베딩 x와 동일한 형태로 되돌려요.
  • .forward(): 위 메서드들을 호출해 입력 임베딩을 각 단계로 전달해요.

torch.nnnn으로 임포트되어 있고, torch.nn.functionalF로, torch도 사용 가능해요.

이 연습은 강의의 일부입니다

PyTorch로 배우는 Transformer 모델

강의 보기

실습형 인터랙티브 연습

이 예제를 이 샘플 코드를 완성하여 풀어보세요.

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super().__init__()
        self.num_heads = num_heads
        self.d_model = d_model
        self.head_dim = d_model // num_heads
        self.query_linear = nn.Linear(d_model, d_model, bias=False)
        self.key_linear = nn.Linear(d_model, d_model, bias=False)
        self.value_linear = nn.Linear(d_model, d_model, bias=False)
        self.output_linear = nn.Linear(d_model, d_model)

    def split_heads(self, x, batch_size):
        seq_length = x.size(1)
        # Split the input embeddings and permute
        x = x.____
        return x.permute(0, 2, 1, 3)
코드 편집 및 실행