Upgrade to Pro — share decks privately, control downloads, hide ads and more …

PyTorchによるGPT-2モデルのフルスクラッチ実装と内部構造の解説

Sponsored · Ship Features Fearlessly Turn features on and off without deploys. Used by thousands of Ruby developers.

 PyTorchによるGPT-2モデルのフルスクラッチ実装と内部構造の解説

PyTorchを用いて、大規模言語モデルの基礎となるGPT-2(Transformer Decoder)をゼロから実装した際の解説資料です。

■ 概要
「Attention Is All You Need」の論文に基づき、単なるライブラリの呼び出しではなく、行列演算のレベルからアーキテクチャを理解することを目的として作成しました。

■ 本スライドで重点を置いたポイント
・行列の形状(Shape)の可視化: (batch_size, seq_len, embed_dim) が各層でどのように変化するかを図解し、デバッグや実装時の混乱を防ぐ工夫をしています 。
・Multi-head Attentionの並列化: なぜ transpose(1, 2) が必要なのか、PyTorchのメモリ配置の観点から解説しました 。
・実践的な実装Tips: view と reshape の挙動の違いや contiguous() の必要性など、効率的なコーディングに欠かせない「ハマりどころ」についてもまとめています 。

■ 実装した主なモジュール
・Single-head / Multi-head Attention (Causal Mask対応)
・GPT Block (LayerNorm, MLP, Residual Connection)
・GPT Model (Token & Positional Embedding, Output Head)

Avatar for Chigen SEN

Chigen SEN

March 11, 2026

More Decks by Chigen SEN

Other Decks in Programming

Transcript

  1. 0. GPTの構造 Vaswani, Ashish, et al. "Attention is all you

    need." Advances in Neural Information Processing Systems 30 (2017).
  2. Encoder Decoder 0. GPTの構造 Vaswani, Ashish, et al. "Attention is

    all you need." Advances in Neural Information Processing Systems 30 (2017).
  3. Encoder   ↓ GPT(Generative pre-trained transformer) Decoder 0. GPTの構造 Vaswani, Ashish,

    et al. "Attention is all you need." Advances in Neural Information Processing Systems 30 (2017).
  4. W W X X Q, K, V Q, K, V

    1-1. Single-head Attentionの実装 線形変換(Linear Transformation) したがって、Q, K, Vのshapeは W X Q, K, V = *
  5. W W X X Q, K, V Q, K, V

    したがって、Q, K, Vのshapeは W X Q, K, V = * 「1つの単語を、いくつの数字のセットで表すか」 を決める 「1つの文章の中に、単語がいくつ並んでいるか」 を表す
  6. queries queries attn_logits attn_logits 1-1. Single-head Attentionの実装 attention logitsのshape X

    X queries attn_logits = * keys_transpos ed したがって、Attention Logitsのshapeは
  7. 1-2. Multi-head Attentionの実装 MyMultiheadAttentionクラス 768 12 768 12 False False

    False 0.1 ここで、Multi-Head: 64次元(768 / 12)の計算を 12個の独立したヘッド として扱う。もし、 Single-Headの場合: 768次元の計算を 1回 行う。
  8. 1-2. Multi-head Attentionの実装 MyMultiheadAttentionクラス 768 12 768 12 False False

    False 0.1 1. torch.ones(cfg.max_len, cfg.max_len) 2. torch.tril(...) 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 0 1 1 1 0 0 1 1 0 0 0 1 Batch_size, num_heads = 1
  9. 1-2. Multi-head Attentionの実装 MyMultiheadAttentionクラス shape: (B, H, L, d) 1

    1 1 1 - 1 1 1 - - 1 1 - - - 1 Dropout関数 Broadcasting
  10. 2. GPT2のアーキテクチャー GPTブロック x LayerNorm MHA Drop-out FFN x LayerNorm

    Drop-out x 入力層・出力層以外に中間層が1つ以上あるニュー ラルネットワーク
  11. 2. GPT2のアーキテクチャー GPTモデル 段階 テンソルの形 内容 self.emb(idx ) (batch, seq_len,

    embed_dim) トークンID → 各単語の埋め 込みベクトル torch.arange (seq_len) (seq_len,) [0, 1, 2, ..., seq_len-1] self.pos_emb (...) (seq_len, embed_dim) 各位置に対応する埋め込み ベクトル broadcasting (batch, seq_len, embed_dim) + (1, seq_len, embed_dim) 位置情報がバッチ全体に広 がる 加算後 x = x + ... (batch, seq_len, embed_dim) 形は全く同じ Transformer!
  12. torch.einsum # attn_logits = torch.matmul(queries, keys.transpose(2, 3)) attn_logits = torch.einsum('b

    h i d, b h j d -> b h i j', queries, keys) 説明:(b,h)は変更せず。(i,d)*(d,j) => (i,j) keys.transpose(2,3)せずに自動に計算する。 # z = torch.matmul(attn_weights, values) z = torch.einsum('b h i j, b h j d -> b h i d', attn_weights, values) 説明:(b,h)は変更せず。(i,j)*(j,d) => (i,d)
  13. viewとreshape import torch x = torch.tensor([[1, 2, 3], [4, 5,

    6]]) # 形状: (2, 3) y = x.transpose(0, 1) # 転置により形状は (3, 2) になるが、メモリ配置は変わらない # ここでエラー! try: z = y.view(6) except RuntimeError as e: print("エラー :", e) z = y.contiguous().view(6) # これなら OK
  14. viewとreshape view は、新しいメモリ領域を確保しないため、処理が 速い。 しかし、transpose(転置)や permute(順番を置き換える) などの操作を行った直後のテンソルは、メモリ上 の配置が非連続になっていることがあり、この状態で view を使うとエラーが発生。

    x = torch.tensor([[1, 2, 3], [4, 5, 6]]) y = x.transpose(0, 1) z = y.reshape(6) # エラーにならず、自動でコピーして整形してくれる print(z) 一方、reshape は自動的に検知して • データが連続している場合 → view() と同じ • データが連続していない場合 → contiguous().view() と同じ(メモリコピー発生、非効率的)