LoslegenKostenlos starten

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-Embeddings x transformieren
  • .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>
Kurs ansehen

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)
Code bearbeiten und ausführen