Thêm các phương thức vào lớp MultiHeadAttention
Trong bài tập này, bạn sẽ tự xây dựng phần còn lại của lớp MultiHeadAttention từ đầu bằng cách định nghĩa bốn phương thức:
.split_heads(): tách và biến đổi các embedding đầu vào giữa các head attention.compute_attention(): tính trọng số attention scaled dot-product và nhân với ma trận values.combine_heads(): biến đổi trọng số attention về lại cùng hình dạng với embedding đầu vào,x.forward(): gọi các phương thức khác để đưa embedding đầu vào đi qua từng bước xử lý
torch.nn đã được nhập là nn, torch.nn.functional có sẵn là F, và torch cũng có sẵn.
Bài tập này là một phần của khóa học
Các mô hình Transformer với PyTorch
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
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)