MultiHeadAttention 클래스에 메서드 추가하기
이 연습 문제에서는 네 가지 메서드를 정의해 MultiHeadAttention 클래스를 처음부터 완성해 볼 거예요.
.split_heads(): 입력 임베딩을 어텐션 헤드 수만큼 분할하고 변환해요..compute_attention(): 스케일드 도트 프로덕트 어텐션 가중치를 계산하고value행렬과 곱해요..combine_heads(): 어텐션 가중치를 입력 임베딩x와 동일한 형태로 되돌려요..forward(): 위 메서드들을 호출해 입력 임베딩을 각 단계로 전달해요.
torch.nn은 nn으로 임포트되어 있고, torch.nn.functional은 F로, 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)