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

nanoMoEの深層解剖:Mixture of Expertsにおける条件分岐を排除した並列演算実装

Sponsored · Your Podcast. Everywhere. Effortlessly. Share. Educate. Inspire. Entertain. You do you. We'll handle the rest.

nanoMoEの深層解剖:Mixture of Expertsにおける条件分岐を排除した並列演算実装

大規模言語モデル(LLM)の効率的なスケールアップを実現するアーキテクチャ「Mixture of Experts (MoE)」を、PyTorchを用いてゼロから実装・解説した資料です。

■ 概要
MoEの肝となる「Router」の意思決定プロセスから、複数の「Expert(MLP)」へのデータ分配、そして並列演算による出力の統合までを、行列の形状推移(Shape Tracking)に焦点を当てて可視化しました。

■ 本スライドの核となる技術トピック
・RouterのGPU最適化実装: トークンごとの条件分岐(if文/forループ)を避け、topk、scatter、one_hot、cumsum を駆使したGPUフレンドリーなマスキング手法をステップバイステップで解説。
・Expert Capacityの導入: 特定のExpertへの負荷集中を防ぐための容量制限(Capacity Factor)の実装と、あふれたトークンの切り捨て処理の論理構造。
・行列演算による並列処理(bmmの活用): view や permute を用いて、バラバラなExpertへの入力を一つのテンソルとしてまとめ、一括で処理する「分配と統合」のメカニズム。

■ 実装した主なコンポーネント
・Routerクラス: Gating Logitの計算とTop-K選択。
・MLPExpertsクラス: 複数のExpertを一つのパラメータ群として保持し、バッチ行列乗算(bmm)で並列処理。
・MOELayerクラス: RouterとExpertsを統合し、エンドツーエンドの推論フローを構築。

■ 開発環境
Python
PyTorch (bmm, einsum, torch.amp 等の活用)

Avatar for Chigen SEN

Chigen SEN

March 11, 2026

More Decks by Chigen SEN

Other Decks in Programming

Transcript

  1. 1.1 MoEとは? 4 Normalization Masked Self-Attention + Normalization FFN +

    Normalization Masked Self-Attention + Normalization MoE Layer + Decoder Block MoE Block
  2. 1. なぜMoEを使うのか? 5 Normalization Masked Self-Attention + Normalization FFN +

    Normalization Masked Self-Attention + Normalization MoE Layer + Decoder Block MoE Block 典型的なFFN: 全結合、膨大なパラメータ数 https://www.programming-ocean.com/ai-architecutres/FFN-achitecture.php
  3. 1. なぜMoEを使うのか? 6 Normalization Masked Self-Attention + Normalization FFN +

    Normalization Masked Self-Attention + Normalization MoE Layer + Decoder Block MoE Block FFNと比べて、MoEは 1. 強み: a. 訓練・推論時の計算コスト (FLOPs)を大幅削減 b. 計算量を増やさずにパラメー タ数を拡大可能 2. 課題: a. VRAM消費量は巨大(全 Expertをメモリに保持する必 要あり) b. 特定のExpertに処理が偏りや すく、過学習しやすい 典型的なFFN: 全結合、膨大なパラメータ数 https://www.programming-ocean.com/ai-architecutres/FFN-achitecture.php
  4. 10 ✖ 避けたいアプローチ: トークンごとに if 文で条件分岐を行い、 for ループで特定のExpertに処理を割り当てる。(条件分岐は計算 効率を著しく低下させる) ◦

    採用するアプローチ: 入力テンソルを Expert へ射影 → Softmaxで確率分布へ変換  → Top-K選択 → マスクを生成 → 入 力テンソル × マスク → 選ばれたExpertの出力のみが反映される つまり、CPU的思考からGPU的思考へ。誤差逆伝播(バックプロパゲーション)を通じて Router自体を学習さ せる。 Router:どのように実装するのか?
  5. class Router(nn.Module): def __init__(self, config): # 略 def forward(self, x):

    device_type = 'cuda' if torch.cuda.is_available() else 'cpu' ctx = nullcontext() if not self.router_use_full_prec else torch.amp.autocast(device_type=device_type, enabled=False) with ctx: B, T, _ = x.size() num_tokens = B * T logits = self.w_g(x) top_k_logits, top_k_indices = logits.topk(self.top_k, dim=-1) router_probs = torch.full_like(logits, float('-inf')) router_probs.scatter_(-1, top_k_indices, top_k_logits) router_probs = F.softmax(router_probs, dim=-1) Router 11 ステップ 1 (Gating Logit の計算) [B, T, d] × [d, E] → [B, T, E] これによって、各トークン(B × T個)に対し て、各 expert( E 個) への好みスコア (logit)を求める。 例:[B=2, T=4, E=3] -0.8, 2.1, -1.4 -2.5, 0.9, 1.8 1.2, -0.3, 0.4 0.5, 1.7, -2.1 Token 0 Token 1 Token 2 Token 3 2.4, -1.1, -1.6 0.7, 2.3, -3.0 -1.3, 0.2, 1.5 -0.4, -2.2, 2.8 Token 0 Token 1 Token 2 Token 3 Batch 1 Batch 2 E1 E2 E3 E1 E2 E3
  6. Router 12 class Router(nn.Module): def __init__(self, config): # 略 def

    forward(self, x): device_type = 'cuda' if torch.cuda.is_available() else 'cpu' ctx = nullcontext() if not self.router_use_full_prec else torch.amp.autocast(device_type=device_type, enabled=False) with ctx: B, T, _ = x.size() num_tokens = B * T logits = self.w_g(x) top_k_logits, top_k_indices = logits.topk(self.top_k, dim=-1) router_probs = torch.full_like(logits, float('-inf')) # [B, T, n_exp] router_probs.scatter_(-1, top_k_indices, top_k_logits) router_probs = F.softmax(router_probs, dim=-1) ステップ 2 (Top-K Expert の選択): [B, T, E] → [B, T, k] top_k_logits(Logit値): [B, T, k] top_k_indices(Index): [B, T, k] これによって、k 個の expert を選ぶ。 -0.8, 2.1 0.9, 1.8 1.2, 0.4 0.5, 1.7 Token 0 Token 1 Token 2 Token 3 2.4, -1.1 0.7, 2.3 0.2, 1.5 -0.4, 2.8 Token 0 Token 1 Token 2 Token 3 Batch 1 Batch 2 K1 K2 K1 K2
  7. Router 13 class Router(nn.Module): def __init__(self, config): # 略 def

    forward(self, x): device_type = 'cuda' if torch.cuda.is_available() else 'cpu' ctx = nullcontext() if not self.router_use_full_prec else torch.amp.autocast(device_type=device_type, enabled=False) with ctx: B, T, _ = x.size() num_tokens = B * T logits = self.w_g(x) top_k_logits, top_k_indices = logits.topk(self.top_k, dim=-1) router_probs = torch.full_like(logits, float('-inf')) # [B, T, n_exp] router_probs.scatter_(-1, top_k_indices, top_k_logits) router_probs = F.softmax(router_probs, dim=-1) ステップ 3 (Softmax の適用): 選ばれた k 個の expert への割り当て率の 合計が 1 になる。 0.052, 0.948 0.289, 0.711 0.690, 0.310 0.231, 0.769 Token 0 Token 1 Token 2 Token 3 0.971, 0.029 0.168, 0.832 0.214, 0.786 0.039, 0.961 Token 0 Token 1 Token 2 Token 3 Batch 1 Batch 2 K1 K2 K1 K2
  8. Router ステップ 4 (Expert Masking) 簡単な(?)例を挙げて説明します! 14 # 続き exp_capacity

    = self.get_capacity(num_tokens) exp_mask = F.one_hot(top_k_indices, num_classes=self.n_exp) # [B, T, k, n_exp] exp_mask = exp_mask.view(num_tokens, self.top_k, self.n_exp) # [B * T, k, n_exp] exp_mask = exp_mask.permute(1, 0, 2) # [k, B * T, n_exp] exp_rank = exp_mask.reshape(self.top_k * num_tokens, self.n_exp) # [k * B * T, n_exp] exp_rank = torch.cumsum(exp_rank, dim=0) - 1 # [k * B * T, n_exp] exp_rank = exp_rank.reshape(self.top_k, num_tokens, self.n_exp) # [k, B * T, n_exp] exp_mask *= torch.lt(exp_rank, exp_capacity) # [k, B*T, n_exp] used_capacity = torch.sum(exp_mask, dim=(0, 1)) # [n_exp] exp_rank = torch.sum(exp_mask * exp_rank, dim=-1) # [k, B * T] router_probs = router_probs.view(num_tokens, self.n_exp)[None, :] # [1, B * T, n_exp] exp_weights = exp_mask * router_probs # [k, B * T, n_exp] exp_rank_sc = F.one_hot(exp_rank, num_classes=exp_capacity) # [k, B * T, exp_capacity] cb_weight = torch.sum(exp_weights.unsqueeze(3) * exp_rank_sc.unsqueeze(2), dim=0) sec_mask = cb_weight.bool() return used_capacity, cb_weight, sec_mask
  9. Expert Masking(例) 15 例: B, T = 1, 4 #

    バッチサイズ1, トークン数4 (計4トークン) n_exp = 3 # エキスパートは3人 (0, 1, 2) top_k = 2 # 各トークンは2つのエキスパートを選べる capacity = 1 # 各エキスパートの定員 Step 1: Shape: [B, T, Top-K] Token 0 : [0, 1] Token 1 : [0, 2] Token 2 : [0, 1] Token 3 : [0, 2]
  10. 16 Step 2, 3: exp_mask = F.one_hot(top_k_indices, num_classes=self.n_exp) exp_mask =

    exp_mask.view(num_tokens, self.top_k, self.n_exp) exp_mask = exp_mask.permute(1, 0, 2) Shape: [B*T , Top-K, n_exp] → [Top-K, B*T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 1 0 0 1 0 0 1 0 0 1 0 0 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 0 1 0 0 0 1 0 1 0 0 0 1 E2 E3 Top-1 (K=0) Top-2 (K=1) exp_mask 説明: exp_mask(マスク) • 役割: 「各トークンが、どのExpertに割り当てら れたか」を 1 か 0 で表すマスクテンソルであ る。
  11. 17 Step 4: exp_rank = exp_mask.reshape(self.top_k * num_tokens, self.n_exp) Shape:

    [Top-K * B * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 1 0 0 1 0 0 1 0 0 1 0 0 E2 E3 Token 4 Token 5 Token 6 Token 7 0 1 0 0 0 1 0 1 0 0 0 1 exp_rank
  12. 18 Step 5: exp_rank = torch.cumsum(exp_rank, dim=0) - 1 Shape:

    [Top-K * B * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 0 -1 -1 1 -1 -1 2 -1 -1 3 -1 -1 E2 E3 Token 4 Token 5 Token 6 Token 7 3 0 -1 3 0 0 3 1 0 3 1 1 exp_rank Token 0 Token 1 Token 2 Token 3 E1 1 0 0 1 0 0 1 0 0 1 0 0 E2 E3 Token 4 Token 5 Token 6 Token 7 0 1 0 0 0 1 0 1 0 0 0 1 exp_rank cumsum(...)-1 説明: exp_rank(順位) • 役割: ある特定のExpertに 割り当てられたトークン群の 中で、「そのトークンが何番 目に割り当てられたか(整 理番号)」を示すテンソルで す。
  13. 19 Step 6: exp_rank = exp_rank.reshape(self.top_k, num_tokens, self.n_exp) Shape: [Top-K,

    B * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 0 -1 -1 1 -1 -1 2 -1 -1 3 -1 -1 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 3 0 -1 3 0 0 3 1 0 3 1 1 E2 E3 Top-1 (K=0) Top-2 (K=1) exp_rank 説明: • 再びTop-1とTop-2のブロックに分割する。 • 効果:次のStep 7 で、Top-1のトークン が優先的に処理されることが保証される。
  14. 20 Step 7: exp_mask *= torch.lt(exp_rank, exp_capacity=1) Shape: [Top-K, B

    * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 T T T F T T F T T F T T E2 E3 Token 0 Token 1 Token 2 Token 3 E1 F T T F T T F F T F F F E2 E3 Top-1 (K=0) Top-2 (K=1) torch.lt (...) Token 0 Token 1 Token 2 Token 3 E1 0 -1 -1 1 -1 -1 2 -1 -1 3 -1 -1 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 3 0 -1 3 0 0 3 1 0 3 1 1 E2 E3 Top-1 (K=0) Top-2 (K=1) exp_rank 説明: • 各Expertの処理上限( Expert Capacity)を適用するステップ。 Capacity=1(1個まで)のため、順番 (exp_rank)が容量を超えたトークンを False にして、切り捨てを行う。 • つまり、0番目(最初に来たトークン)だけ を True として残し、それ以降は False にして処理から除外する。 • ※ -1 の箇所が True になっているが、 exp_mask が 0 なので、掛け算すれば 最終的に 0 になるため問題ない。
  15. 21 Step 7: exp_mask *= torch.lt(exp_rank, exp_capacity=1) Shape: [Top-K, B

    * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 1 0 0 1 0 0 1 0 0 1 0 0 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 0 1 0 0 0 1 0 1 0 0 0 1 E2 E3 Top-1 (K=0) Top-2 (K=1) exp_mask Token 0 Token 1 Token 2 Token 3 E1 T F Token 0 Token 1 Token 2 Token 3 T F F T F E2 E3 Top-1 (K=0) Top-2 (K=1) 更新された exp_mask * Token 0 Token 1 Token 2 Token 3 E1 T T T F T T F T T F T T E2 E3 Token 0 Token 1 Token 2 Token 3 E1 F T T F T T F F T F F F E2 E3 Top-1 (K=0) Top-2 (K=1) torch.lt (...)
  16. 22 Step 8: exp_rank = torch.sum(exp_mask * exp_rank, dim=-1) Shape:

    [Top-K, B * T , n_exp] Token 0 Token 1 Token 2 Token 3 E1 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) 更新された exp_mask T F F F F F F F F F F F F T F F F T F F F F F F E2 E3 E2 E3 E1 * Token 0 Token 1 Token 2 Token 3 E1 0 -1 -1 1 -1 -1 2 -1 -1 3 -1 -1 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 3 0 -1 3 0 0 3 1 0 3 1 1 E2 E3 Top-1 (K=0) Top-2 (K=1) exp_rank Token 0 Token 1 Token 2 Token 3 E1 1 0 0 0 0 0 0 0 0 0 0 0 E2 E3 Token 0 Token 1 Token 2 Token 3 E1 0 1 0 0 0 1 0 0 0 0 0 0 E2 E3 Top-1 (K=0) Top-2 (K=1) 更新されたexp_rank
  17. 23 Step 8: exp_rank = torch.sum(exp_mask * exp_rank, dim=-1) Shape:

    [Top-K, B * T] Token 0 Token 1 Token 2 Token 3 1 0 0 0 Token 0 Token 1 Token 2 Token 3 1 1 0 0 Top-1 (K=0) Top-2 (K=1) 更新されたexp_rank 説明: • 横に広がっていた「E1, E2, E3...」expertの列が 潰され、 • 最終的に有効(Active)なトークンが 1、切り捨て られたトークンが 0
  18. 24 Step 9: router_probs = router_probs.view(num_tokens, n_exp)[None, :] Shape: [B,

    T, n_exp] → [1, B * T, n_exp] View Token 0 Token 1 Token 2 Token 3 E1 0.6 0.3 0.1 0.5 0.2 0.3 0.7 0.2 0.1 0.4 0.1 0.5 E2 E3 router_probs Token 0 Token 1 Token 2 Token 3 E1 0.6 0.3 0.1 0.5 0.2 0.3 0.7 0.2 0.1 0.4 0.1 0.5 E2 E3 更新されたrouter_probs 説明: • [Top-K…] テンソルとの掛け 算(ブロードキャスト)に向け て、確率テンソルの形 (Shape)を整える準備 • B = 1, T = 4 のため、shape の変化なし(1, 4, 3)
  19. 25 Step 10: exp_weights = exp_mask * router_probs Shape: [Top-K,

    B * T, n_exp] Token 0 Token 1 Token 2 Token 3 E1 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) 更新された exp_mask T F F F F F F F F F F F F T F F F T F F F F F F E2 E3 E2 E3 E1 * Token 0 Token 1 Token 2 Token 3 E1 0.6 0.3 0.1 0.5 0.2 0.3 0.7 0.2 0.1 0.4 0.1 0.5 E2 E3 更新されたrouter_probs Token 0 Token 1 Token 2 Token 3 E1 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) exp_weights 0.6 0 0 0 0 0 0 0 0 0 0 0 0 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E2 E3 E1
  20. 26 Step 11: exp_rank_sc = F.one_hot(exp_rank, num_classes=exp_capacity) Shape: [Top-K, B

    * T, n_capacity] Token 0 Token 1 Token 2 Token 3 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) exp_rank_sc 1 1 1 1 1 1 1 1
  21. 27 Step 12: cb_weight = torch.sum(exp_weights.unsqueeze(3) * exp_rank_sc.unsqueeze(2), dim=0) Shape:

    [Top-K, B * T, n_expert, n_capacity] Token 0 Token 1 Token 2 Token 3 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) exp_rank_sc 1 1 1 1 1 1 1 1 * Token 0 Token 1 Token 2 Token 3 E1 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) exp_weights 0.6 0 0 0 0 0 0 0 0 0 0 0 0 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E2 E3 E1 Token 0 Token 1 Token 2 Token 3 E1 Token 0 Token 1 Token 2 Token 3 Top-1 (K=0) Top-2 (K=1) 0.6 0 0 0 0 0 0 0 0 0 0 0 0 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E2 E3 E1
  22. 28 Step 12: cb_weight = torch.sum(exp_weights.unsqueeze(3) * exp_rank_sc.unsqueeze(2), dim=0) Shape:

    [B * T, n_expert, n_capacity] Expert 3 Expert 2 Router Expert 1 inputs + cb_weightによって、「どのトークンが、どの Expertの出力を、どれくらいの割合で受け取るか」 を決定す る Token 0 Token 1 Token 2 Token 3 0.6 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E1 Token 0 Token 0 Token 1
  23. 29 Step 13: sec_mask = cb_weight.bool() Shape: [B * T,

    n_expert, n_capacity] Boolean Token 0 Token 1 Token 2 Token 3 0.6 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E1 Token 0 Token 1 Token 2 Token 3 T T F F F T F F F F F F E2 E3 E1 sec_maskによって、各トークンをどの「 Expert」に送るかを決定する
  24. MLPExperts 32 拡張: [B, T, C] × [B, C, 4C]

    → [B, T, 4C] class MLPExperts(nn.Module): def __init__(self, config): super().__init__() self.bias = config.bias self.c_fc = nn.Parameter(torch.empty(config.n_exp, config.n_embd, 4 * config.n_embd)) self.c_proj = nn.Parameter(torch.empty(config.n_exp, 4 * config.n_embd, config.n_embd)) self.fc_bias = nn.Parameter(torch.empty(config.n_exp, 1, 4 * config.n_embd)) if self.bias else None self.proj_bias = nn.Parameter(torch.empty(config.n_exp, 1, config.n_embd)) if self.bias else None self.gelu = nn.GELU() self.dropout = nn.Dropout(config.dropout) def forward(self, x): x = torch.bmm(x, self.c_fc) if self.bias: x += self.fc_bias x = self.gelu(x) x = torch.bmm(x, self.c_proj) if self.bias: x += self.proj_bias x = self.dropout(x) return x 射影: [B, T, 4C] × [B, 4C, C] → [B, T, C]           MLP Expert 拡張 射影 Drop-out
  25. MOELayer 33 class MOELayer(nn.Module): def __init__(self, config): super().__init__() self.router =

    Router(config) self.experts = MLPExperts(config) def forward(self, x: torch.Tensor): B, T, n_embd = x.size() num_tokens = (B * T) used_capacity, exp_weight, exp_mask = self.router(x) x = x.view(num_tokens, n_embd) exp_batches = exp_mask.permute(1, 2, 0).type_as(x) @ x exp_out = self.experts(exp_batches) exp_weight = exp_weight.view(num_tokens, -1) exp_out = exp_out.view(-1, n_embd) output = exp_weight @ exp_out return output.view(B, T, n_embd) [B, T, n_embd]→ [B * T, n_embd] 残りは簡単な(?)例を挙げて説明します!
  26. 34 Step 1: 分配 exp_batches = exp_mask.permute(1, 2, 0).type_as(x) @

    x Shape: [B * T, n_expert, n_capacity] Token 0 Token 1 Token 2 Token 3 1 1 0 0 0 1 0 0 0 0 0 0 E2 E3 E1 Token 0 Token 1 Token 2 Token 3 T T F F F T F F F F F F E2 E3 E1
  27. 35 Step 1: 分配 exp_batches = exp_mask.permute(1, 2, 0).type_as(x) @

    x Shape: [n_expert, n_capacity, B * T] Token 0 Token 1 Token 2 Token 3 1 0 0 0 1 0 0 0 0 1 0 0 E2 E3 E1 Token 0 Token 1 Token 2 Token 3 1 1 0 0 0 1 0 0 0 0 0 0 E2 E3 E1
  28. 36 Step 1: 分配 exp_batches = exp_mask.permute(1, 2, 0).type_as(x) @

    x Shape: [n_expert, n_capacity, B * T] * [B * T, n_embd] -> [n_expert, n_capacity, n_embd] Token 0 Token 1 Token 2 Token 3 1 0 0 0 1 0 0 0 0 1 0 0 E2 E3 E1 * Token 0 Token 1 Token 2 Token 3 A0 A1 B0 B1 C0 C1 D0 D1 A0 A1 A0 A1 B0 B1 E2 E3 E1 Emb 1 Emb 2 Emb 1 Emb 2
  29. MoE Layer 37 Step 2: Expert処理 exp_out = self.experts(exp_batches) Shape:

    [n_expert, n_capacity, n_embd] A0 A1 A0 A1 B0 B1 E2 E3 E1 Emb 1 Emb 2           MLP Expert 拡 張 射 影 Drop-out           MLP Expert 拡 張 射 影 Drop-out           MLP Expert 拡 張 射 影 Drop-out a0 a1 a’0 a’1 b0 b1 E2 E3 E1 Emb 1 Emb 2 MoE Layer
  30. 38 Step 3: 行列計算ための変形 exp_weight = exp_weight.view(num_tokens, -1) Shape:[B *

    T, n_expert* n_capacity] a0 a1 a’0 a’1 b0 b1 E2 E3 E1 Emb 1 Emb 2 Token 0 Token 1 Token 2 Token 3 0.6 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E1 exp_out = exp_out.view(-1, n_embd) Shape:[n_expert * n_capacity, n_embd]
  31. 39 Step 4: 行列の計算 output = exp_weight @ exp_out Shape:[B

    * T, n_embd] a0 a1 a’0 a’1 b0 b1 E2 E3 E1 Emb 1 Emb 2 Token 0 Token 1 Token 2 Token 3 0.6 0.3 0 0 0 0.3 0 0 0 0 0 0 E2 E3 E1 * 0.6a0 + 0.3a’0 0.6a1 + 0.3a’1 0.3b0 0.3b1 0 0 0 0 Token 0 Token 1 Token 2 Token 3 Emb 1 Emb 2
  32. 40 Step 5: 形の復元 return output.view(B, T, n_embd) Shape:[B, T,

    n_embd] 0.6a0 + 0.3a’0 0.6a1 + 0.3a’1 0.3b0 0.3b1 0 0 0 0 Token 0 Token 1 Token 2 Token 3 Emb 1 Emb 2