Dodawanie metod do klasy MultiHeadAttention
W tym ćwiczeniu zbudujesz resztę klasy MultiHeadAttention od podstaw, definiując cztery metody:
.split_heads(): dzieli i przekształca wejściowe osadzenia (embeddings) między głowice atencji.compute_attention(): oblicza skalowane iloczyny skalarnych wag atencji pomnożonych przez macierz wartości.combine_heads(): przekształca wagi atencji z powrotem do kształtu zgodnego z wejściowymi osadzeniami,x.forward(): wywołuje pozostałe metody, przepuszczając wejściowe osadzenia przez każdy etap przetwarzania
torch.nn jest zaimportowany jako nn, torch.nn.functional jest dostępny jako F, a torch jest również dostępny.
To ćwiczenie jest częścią kursu
Modele Transformer w PyTorch
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
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)