Adăugarea metodelor la clasa MultiHeadAttention
În acest exercițiu, vei construi restul clasei MultiHeadAttention de la zero, definind patru metode:
.split_heads(): împarte și transformă încorporările de intrare între capetele de atenție.compute_attention(): calculează ponderile de atenție scalate prin produs scalar, înmulțite cu matricea valorilor.combine_heads(): transformă ponderile de atenție înapoi în aceeași formă ca încorporările de intrare,x.forward(): apelează celelalte metode pentru a trece încorporările de intrare prin fiecare etapă
torch.nn a fost importat ca nn, torch.nn.functional este disponibil ca F, iar torch este de asemenea disponibil.
Acest exercițiu face parte din cursul
Modele Transformer cu PyTorch
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
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)