始める無料で始める

MultiHeadAttention クラスにメソッドを追加する

この演習では、MultiHeadAttention クラスの残りを一から実装し、次の4つのメソッドを定義します。

  • .split_heads(): 入力埋め込みをアテンションヘッド間に分割・変換します
  • .compute_attention(): スケールド・ドット積により計算したアテンション重みを values 行列に掛け合わせます
  • .combine_heads(): アテンション重みを入力埋め込み x と同じ形状に戻します
  • .forward(): これらのメソッドを呼び出し、入力埋め込みを各処理に通します

torch.nnnn としてインポート済みで、torch.nn.functionalF として、torch も利用可能です。

この演習はコースの一部です

PyTorchで学ぶTransformerモデル

コースを見る

実践的なインタラクティブ演習

このサンプルコードを完成させて、この演習に挑戦してみましょう。

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)
コードを編集して実行