Methoden zur Klasse MultiHeadAttention hinzufügen
In dieser Übung baust du den restlichen Teil der Klasse MultiHeadAttention von Grund auf, indem du vier Methoden definierst:
.split_heads(): Eingabe-Embeddings zwischen den Attention-Heads aufteilen und transformieren.compute_attention(): die skalierten Skalarprodukt-Attention-Gewichte berechnen und mit der Values-Matrix multiplizieren.combine_heads(): die Attention-Gewichte zurück in dieselbe Form wie die Eingabe-Embeddingsxtransformieren.forward(): die anderen Methoden aufrufen, um die Eingabe-Embeddings durch jeden Verarbeitungsschritt zu leiten
torch.nn wurde als nn importiert, torch.nn.functional ist als F verfügbar, und torch ist ebenfalls verfügbar.
Diese Übung ist Teil des Kurses
<Kurs>Transformer-Modelle mit PyTorch</Kurs>Interaktive praktische Übung
Versuche dich an dieser Übung, indem du diesen Beispielcode vervollständigst.
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)