LLaMAアーキテクチャの設計思想 — RMSNorm・SwiGLU・GQA・RoPEを完全解説

ChatGPT の登場以降、大規模言語モデル(LLM)は爆発的に普及しました。しかし GPT-3 や GPT-4 はクローズドモデルであり、モデルの重みもアーキテクチャの詳細も公開されていません。研究者やエンジニアがモデルの中身を自由に調べたり改良したりすることは困難でした。

この状況を変えたのが、Meta AI が 2023 年に公開した LLaMA(Large Language Model Meta AI) です。LLaMA は GPT と同じ Decoder-only Transformer をベースとしながらも、4 つの重要な設計変更を加えることで、より少ないパラメータでGPT-3 と同等以上の性能 を実現しました。その設計変更とは、RMSNormSwiGLURoPE、そして GQA です。

「なぜ GPT と同じアーキテクチャをそのまま使わなかったのか?」「それぞれの変更はどんな問題を解決するのか?」 — 本記事ではこれらの疑問に答えながら、LLaMA のアーキテクチャを数式レベルで完全に理解します。

LLaMA のアーキテクチャを理解することは、以下のような場面で直接役立ちます。

  • ローカル LLM の運用: llama.cpp や vLLM でモデルを動かす際に、各コンポーネントの役割を理解していればデバッグやチューニングが効率的になります
  • ファインチューニング: LoRA や QLoRA でモデルを特化させる際に、どの層にアダプタを挿入すべきかの判断にアーキテクチャの理解が不可欠です
  • 後続モデルの理解: Mistral、Mixtral、Gemma、Qwen など、現在の主要なオープンソース LLM はほぼ全て LLaMA のアーキテクチャを踏襲しています。LLaMA を理解すれば、これらのモデルの差分だけを追えばよくなります

本記事の内容

  • LLaMA の位置づけと Chinchilla スケーリング則
  • GPT からの変更点の全体像
  • Pre-Norm(Pre-Layer Normalization)による学習安定化
  • RMSNorm の数理と計算効率
  • SwiGLU 活性化関数の直感と数式
  • RoPE(Rotary Position Embedding)の復習
  • GQA(Grouped Query Attention)のメモリ効率化
  • LLaMA のモデルサイズ一覧
  • PyTorch によるスクラッチ実装とパラメータ数の確認

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

画像なし
GPTのアーキテクチャと自己回帰生成を解説
GPTのアーキテクチャをMasked Self-AttentionによるTransformer Decoderの構造から解説します
画像なし
Transformer Decoderの構造を解説
Transformer Decoderの構造を数式とコードで理解します
画像なし
RoPE(Rotary Position Embedding)を解説
回転行列を用いた位置エンコーディングRoPEの数理を詳しく解説します
画像なし
Layer Normalizationを解説
Layer Normalizationの数理と実装を解説します

それでは、まず LLaMA がどのような背景で生まれたモデルなのかを確認しましょう。

LLaMA の位置づけ

オープンソース LLM の転機

2023 年 2 月、Meta AI は LLaMA(Large Language Model Meta AI)を研究目的で公開しました。7B、13B、33B、65B の 4 つのサイズが用意され、特に LLaMA-13B は GPT-3(175B)を多くのベンチマークで上回る という衝撃的な結果を示しました。パラメータ数が 13 分の 1 以下であるにもかかわらず、同等以上の性能を達成したのです。

この成功の背景には、データ量を重視する という明確な設計方針がありました。

Chinchilla スケーリング則の採用

2022 年に DeepMind が発表した Chinchilla の研究は、LLM の学習において「モデルサイズとデータ量のバランス」が極めて重要であることを示しました。それ以前の主流であった GPT-3 の方針は「モデルを大きくすればするほど性能が上がる」というものでしたが、Chinchilla の知見は異なっていました。

従来のスケーリング則(Kaplan et al., 2020)では、計算予算が増えたときにモデルサイズを優先的に大きくすべきとされていました。しかし Chinchilla のスケーリング則(Hoffmann et al., 2022)では、モデルサイズとトークン数を等しい比率で増やすべき であることが示されました。具体的には、パラメータ数 $N$ のモデルに対して、学習トークン数 $D$ はおよそ $D \approx 20N$ が最適とされます。

LLaMA はこの知見を積極的に取り入れました。LLaMA-7B は 1 兆トークン、LLaMA-65B は 1.4 兆トークンで学習されています。これは GPT-3(175B パラメータ、3000 億トークン)と比較すると、パラメータあたりのデータ量が圧倒的に多いことがわかります。

モデル パラメータ数 学習トークン数 トークン/パラメータ
GPT-3 175B 300B 1.7
Chinchilla 70B 1.4T 20.0
LLaMA-7B 7B 1.0T 142.9
LLaMA-65B 65B 1.4T 21.5

LLaMA-7B に至っては、Chinchilla の推奨比率の約 7 倍ものデータで学習されています。これは「推論コストを下げたい(モデルを小さくしたい)場合は、学習データを増やして補う」という実用的な戦略です。

つまり LLaMA の成功は、アーキテクチャの改良だけでなく、データ戦略の転換 によるところも大きいのです。では次に、アーキテクチャ面で GPT からどのような変更が加えられたかを全体像から見ていきましょう。

GPT からの変更点の全体像

LLaMA は GPT と同じ Decoder-only Transformer ですが、5 つの重要な変更が加えられています。以下の表で全体像を把握しましょう。

コンポーネント GPT LLaMA 変更の目的
正規化の位置 Post-Norm(各サブレイヤーの後) Pre-Norm(各サブレイヤーの前) 学習の安定化
正規化の手法 Layer Normalization RMSNorm 計算効率の向上
FFN の活性化関数 ReLU(GPT-1) / GELU(GPT-2以降) SwiGLU 性能向上
位置エンコーディング 学習可能な絶対位置埋め込み RoPE(回転位置埋め込み) 長文への汎化
アテンション Multi-Head Attention Grouped Query Attention(LLaMA 2以降) 推論時のメモリ効率化

これらの変更はそれぞれ独立した研究に基づいており、LLaMA の論文はこれらの「ベストプラクティスの集大成」と言えます。一つひとつは小さな改良に見えるかもしれませんが、組み合わせることで大きな効果を発揮します。

それぞれの変更を、GPT との差分を意識しながら詳しく見ていきましょう。最初は、正規化の位置に関する変更 — Pre-Norm です。

Pre-Norm(Pre-Layer Normalization)

Post-Norm と Pre-Norm の違い

Transformer には「正規化をどこに置くか」という設計上の選択があります。オリジナルの Transformer(Vaswani et al., 2017)と GPT-1 では、サブレイヤー(Self-Attention や FFN)の に正規化を行う Post-Norm が使われていました。

Post-Norm の残差接続は次のように書けます。

$$ \bm{H}^{(l)} = \text{Norm}(\bm{H}^{(l-1)} + \text{SubLayer}(\bm{H}^{(l-1)})) $$

一方、Pre-Norm では正規化をサブレイヤーの に適用します。

$$ \bm{H}^{(l)} = \bm{H}^{(l-1)} + \text{SubLayer}(\text{Norm}(\bm{H}^{(l-1)})) $$

数式の違いは Norm の位置が変わっただけに見えますが、勾配の流れに大きな影響を与えます。

なぜ Pre-Norm が学習安定性に優れるのか

Pre-Norm の最大の利点は、残差接続を通じて勾配が直接的に流れるパス が確保されることです。

Post-Norm の場合、残差接続の出力が正規化を通過するため、勾配は正規化層のヤコビアンを経由しなければなりません。これは層が深くなるほど勾配が不安定になる原因となります。実際、Post-Norm のTransformer では、学習率のウォームアップなしでは学習が発散しやすいことが知られています。

Pre-Norm の場合、残差接続は正規化を通らず、そのまま次の層へ加算されます。損失 $\mathcal{L}$ から入力 $\bm{H}^{(0)}$ までの勾配を考えると、Pre-Norm では各層の残差接続を通じた「ショートカットパス」が常に存在するため、勾配消失が起きにくくなります。

イメージとしては、Post-Norm は「階段を一段上がるたびに検問(正規化)を通る」のに対し、Pre-Norm は「各階にはエレベーター(残差接続の直接パス)があり、検問は作業部屋の入口にだけある」という違いです。

GPT-2 以降、多くの大規模モデルで Pre-Norm が採用されるようになりました。LLaMA もこの流れに従っています。さらに LLaMA では、Pre-Norm に加えて最終的な出力にも追加の RMSNorm を適用しています。

Pre-Norm の採用により学習が安定することはわかりました。では次に、その正規化手法自体を Layer Normalization から RMSNorm に置き換えた理由を見ていきましょう。

RMSNorm(Root Mean Square Normalization)

Layer Normalization の振り返り

まず、標準的な Layer Normalization を復習しましょう。Layer Norm は入力ベクトル $\bm{x} \in \mathbb{R}^d$ に対して、平均を引いてから標準偏差で割る という操作を行います。

$$ \text{LayerNorm}(\bm{x}) = \frac{\bm{x} – \mu}{\sigma} \odot \bm{\gamma} + \bm{\beta} $$

ここで $\mu$ と $\sigma$ はベクトル $\bm{x}$ の要素の平均と標準偏差です。

$$ \mu = \frac{1}{d} \sum_{i=1}^{d} x_i, \quad \sigma = \sqrt{\frac{1}{d} \sum_{i=1}^{d} (x_i – \mu)^2 + \epsilon} $$

$\bm{\gamma}, \bm{\beta} \in \mathbb{R}^d$ は学習可能なスケール・シフトパラメータ、$\epsilon$ はゼロ除算防止の小さな定数、$\odot$ は要素ごとの積です。

RMSNorm: 平均を引かない正規化

Layer Normalization は「平均を引く(re-centering)」と「分散で割る(re-scaling)」の 2 つの操作から成ります。RMSNorm の核心的なアイデアは、re-centering は実はそれほど重要ではなく、re-scaling だけで十分 だという仮説に基づいています。

日常的なアナロジーで考えてみましょう。身長のデータを正規化するとき、「全員の身長から平均を引いて、ばらつきで割る」のが Layer Norm です。一方 RMSNorm は「平均を引かず、二乗平均平方根で割るだけ」です。平均を引く操作を省略しても、各次元のスケールを揃えるという本質的な目的は達成できます。

RMSNorm は次のように定義されます。

$$ \text{RMSNorm}(\bm{x}) = \frac{\bm{x}}{\text{RMS}(\bm{x})} \odot \bm{\gamma} $$

ここで $\text{RMS}(\bm{x})$ は二乗平均平方根(Root Mean Square)です。

$$ \text{RMS}(\bm{x}) = \sqrt{\frac{1}{d} \sum_{i=1}^{d} x_i^2 + \epsilon} $$

Layer Norm との違いを整理すると次のようになります。

Layer Norm RMSNorm
平均の計算 必要($\mu$) 不要
分散/RMSの計算 $(x_i – \mu)^2$ の平均 $x_i^2$ の平均
シフトパラメータ $\bm{\beta}$ あり なし
計算量 $O(d)$(2パス) $O(d)$(1パス)

計算効率の向上

RMSNorm が Layer Norm より効率的な理由を具体的に見てみましょう。

Layer Norm では、まず平均 $\mu$ を計算し(第 1 パス)、次に $\mu$ を使って分散を計算する(第 2 パス)必要があります。つまり、データを 2 回走査しなければなりません。

一方 RMSNorm は、$x_i^2$ の平均を 1 回のパスで計算するだけです。平均 $\mu$ の計算もシフトパラメータ $\bm{\beta}$ の適用も不要なため、メモリアクセスと演算量の両方で有利です。

Zhang & Sennrich(2019)の実験では、RMSNorm は Layer Norm と比較して 学習速度が 7〜64% 向上 し、性能は同等以上であったと報告されています。大規模モデルでは正規化が頻繁に呼ばれるため(各 Transformer ブロック内で 2 回)、この効率化の効果は累積的に大きくなります。

$i$ 番目の要素に対する RMSNorm の出力を明示的に書くと、次のようになります。

$$ \text{RMSNorm}(\bm{x})_i = \frac{x_i}{\sqrt{\frac{1}{d}\sum_{j=1}^{d}x_j^2 + \epsilon}} \cdot \gamma_i $$

この式を見ると、各要素 $x_i$ がベクトル全体の「エネルギー」(二乗和の平均)で正規化されていることがわかります。これは信号処理における RMS(実効値)の概念と全く同じです。

RMSNorm によって正規化の計算が効率化されました。次は、FFN(Feed-Forward Network)における活性化関数の変更 — ReLU/GELU から SwiGLU への切り替えを見ていきましょう。

SwiGLU 活性化関数

GPT の FFN を振り返る

Transformer の各ブロックには、Self-Attention に続いて Feed-Forward Network(FFN)があります。GPT における標準的な FFN は次の形をしています。

$$ \text{FFN}(\bm{x}) = \text{GELU}(\bm{x}\bm{W}_1 + \bm{b}_1)\bm{W}_2 + \bm{b}_2 $$

ここで $\bm{W}_1 \in \mathbb{R}^{d \times d_{\text{ff}}}$、$\bm{W}_2 \in \mathbb{R}^{d_{\text{ff}} \times d}$ です。$d_{\text{ff}}$ は通常 $4d$ に設定されます。入力を一度高次元空間に射影し、活性化関数を適用してから元の次元に戻すという構造です。

GLU(Gated Linear Unit)の直感

SwiGLU を理解するために、まずその源流である GLU(Gated Linear Unit) のアイデアを押さえましょう。

通常の FFN では、入力を 1 つの重み行列で変換してから活性化関数を通します。これに対して GLU は、2 つの重み行列を使い、一方の出力を「ゲート」として使う というアイデアです。

日常的なアナロジーで言えば、通常の FFN は「すべての情報を同じフィルターに通す」のに対し、GLU は「情報を 2 つのルートに分け、一方のルートが『どの情報を通すか』を制御する蛇口の役割をする」イメージです。水道の蛇口は水の流量を調節しますが、GLU のゲートも同様に情報の流量を制御します。

GLU の数式は次の通りです。

$$ \text{GLU}(\bm{x}) = (\bm{x}\bm{W}_1) \odot \sigma(\bm{x}\bm{W}_{\text{gate}}) $$

ここで $\sigma$ はシグモイド関数、$\odot$ は要素ごとの積です。$\bm{x}\bm{W}_1$ が「情報の値」、$\sigma(\bm{x}\bm{W}_{\text{gate}})$ が「0〜1 のゲート値」を計算し、両者の要素ごとの積を取ることで、ゲートが開いている次元の情報だけが通過します。

SwiGLU の定義

SwiGLU は、Shazeer(2020)が提案した GLU の変種で、ゲートの活性化関数にシグモイドではなく SiLU(Swish) を使います。

$$ \text{SwiGLU}(\bm{x}) = (\bm{x}\bm{W}_1) \odot \text{SiLU}(\bm{x}\bm{W}_{\text{gate}}) $$

SiLU(Sigmoid Linear Unit)は次のように定義されます。

$$ \text{SiLU}(z) = z \cdot \sigma(z) = \frac{z}{1 + e^{-z}} $$

SiLU は ReLU と似た形をしていますが、2 つの重要な違いがあります。

  1. 滑らかさ: ReLU は $z = 0$ で微分不連続ですが、SiLU はどこでも微分可能です
  2. 負の領域の振る舞い: ReLU は負の入力を完全にゼロにしますが、SiLU は小さな負の値を許容します($z \approx -1.28$ で最小値 $\approx -0.28$ を取ります)

この「少しだけ負の値を通す」性質により、SiLU は勾配がゼロになる問題(dying ReLU 問題)を回避しつつ、ReLU に近い非線形性を保持します。

LLaMA の FFN 全体を書き下すと次のようになります。

$$ \text{FFN}_{\text{SwiGLU}}(\bm{x}) = \left[(\bm{x}\bm{W}_1) \odot \text{SiLU}(\bm{x}\bm{W}_{\text{gate}})\right] \bm{W}_2 $$

ここで $\bm{W}_1, \bm{W}_{\text{gate}} \in \mathbb{R}^{d \times d_{\text{ff}}}$、$\bm{W}_2 \in \mathbb{R}^{d_{\text{ff}} \times d}$ です。

隠れ次元の調整

SwiGLU では重み行列が 2 つ($\bm{W}_1$ と $\bm{W}_{\text{gate}}$)に増えるため、パラメータ数が増加します。GPT の FFN が $\bm{W}_1, \bm{W}_2$ の 2 つの行列で $2 \times d \times d_{\text{ff}}$ パラメータを持つのに対し、SwiGLU は 3 つの行列で $3 \times d \times d_{\text{ff}}$ パラメータを持ちます。

パラメータ数を同等に保つため、LLaMA では隠れ次元 $d_{\text{ff}}$ を調整しています。GPT では $d_{\text{ff}} = 4d$ でしたが、LLaMA では次のように設定されています。

$$ d_{\text{ff}} = \frac{2}{3} \times 4d = \frac{8}{3}d $$

この調整を確認しましょう。GPT の FFN のパラメータ数は $2 \times d \times 4d = 8d^2$ です。SwiGLU の場合、$d_{\text{ff}} = \frac{8}{3}d$ とすると、パラメータ数は次のようになります。

$$ 3 \times d \times \frac{8}{3}d = 8d^2 $$

このようにして、SwiGLU を導入しつつもパラメータ数を増やさない設計になっています。実際の LLaMA の実装では、$d_{\text{ff}}$ をさらに 256 の倍数に丸めるなどの調整が加えられています。

Shazeer の実験では、SwiGLU は ReLU や GELU と比較して、同じパラメータ数でより低いパープレキシティ(perplexity) を達成することが示されています。パラメータの使い方が効率的なのです。

SwiGLU により FFN の表現力が向上しました。次は、位置情報の表現方法の変更 — 学習可能な絶対位置埋め込みから RoPE(回転位置埋め込み)への切り替えを見ていきましょう。

RoPE(Rotary Position Embedding)

なぜ位置エンコーディングを変更するのか

GPT では、各トークンの位置を表現するために 学習可能な絶対位置埋め込み を使っていました。これは位置 $t$ に対して学習可能なベクトル $\bm{p}_t \in \mathbb{R}^d$ を用意し、トークン埋め込みに加算するものです。

$$ \bm{h}_t = \bm{e}_t + \bm{p}_t $$

この方法はシンプルですが、2 つの問題があります。

  1. 学習時の最大系列長を超えられない: 位置 $1, 2, \dots, T_{\text{max}}$ の埋め込みしか学習されないため、$T_{\text{max}}$ を超える位置の情報は表現できません
  2. 相対的な位置関係の表現が間接的: 「トークン $i$ とトークン $j$ の距離が $k$」という情報は、$\bm{p}_i$ と $\bm{p}_j$ の差として間接的に表現されるだけです

RoPE の基本的なアイデア

RoPE(Rotary Position Embedding)は、Su et al.(2021)が提案した位置エンコーディングで、2 次元平面での回転 を使って位置情報を表現します。

核心的なアイデアは、位置 $m$ にあるトークンの Query/Key ベクトルを、位置 $m$ に応じた角度だけ回転させるというものです。2 つのトークンの内積を取ると、回転角の差だけが残るため、自然に 相対位置 がエンコードされます。

$d$ 次元のベクトルを $d/2$ 個の 2 次元ペアに分割し、$k$ 番目のペア $(x_{2k-1}, x_{2k})$ に対して角度 $m\theta_k$ の回転を適用します。

$$ \begin{pmatrix} x’_{2k-1} \\ x’_{2k} \end{pmatrix} = \begin{pmatrix} \cos m\theta_k & -\sin m\theta_k \\ \sin m\theta_k & \cos m\theta_k \end{pmatrix} \begin{pmatrix} x_{2k-1} \\ x_{2k} \end{pmatrix} $$

ここで $\theta_k$ は次のように定義されます。

$$ \theta_k = 10000^{-2k/d}, \quad k = 1, 2, \dots, d/2 $$

低次元のペアは高い周波数(大きな $\theta_k$)で回転し、高次元のペアは低い周波数で回転します。これにより、近い位置のトークンは「急速に回転する次元」で区別され、遠い位置のトークンは「ゆっくり回転する次元」で区別されます。

RoPE の重要な性質

RoPE が LLM に適している理由は、主に 3 つの性質にあります。

第一に、相対位置の自然な表現 です。位置 $m$ の Query ベクトル $\bm{q}_m$ と位置 $n$ の Key ベクトル $\bm{k}_n$ の内積を取ると、結果は $m – n$(相対位置)のみに依存します。これは回転行列の性質 $\bm{R}_m^T \bm{R}_n = \bm{R}_{n-m}$ に由来します。

第二に、長文への外挿可能性 です。回転角は任意の整数 $m$ に対して定義できるため、学習時に見なかった長さの入力にも(ある程度)対応できます。実際、NTK-aware Scaling や YaRN などの拡張手法を組み合わせることで、学習時の数倍の系列長に外挿できることが示されています。

第三に、計算効率 です。RoPE は追加のパラメータを必要としません。回転角は位置と次元から決定論的に計算されるため、学習すべきパラメータが増えないのです。

RoPE の詳細な数理については、以下の記事で詳しく解説しています。

画像なし
RoPE(Rotary Position Embedding)を解説
回転行列を用いた位置エンコーディングRoPEの数理を詳しく解説します

ここまでで、LLaMA の 4 つの変更のうち 3 つ(Pre-Norm + RMSNorm、SwiGLU、RoPE)を見てきました。これらは全て LLaMA 1 から採用されている技術です。最後の変更点 — GQA は LLaMA 2 で導入された、推論効率に関わる重要な改良です。

GQA(Grouped Query Attention)

Multi-Head Attention の推論時のボトルネック

GPT で使われる標準的な Multi-Head Attention(MHA)では、$h$ 個のヘッドがそれぞれ独立した Query、Key、Value を持ちます。

$$ \text{head}_i = \text{Attention}(\bm{Q}_i, \bm{K}_i, \bm{V}_i), \quad i = 1, \dots, h $$

学習時にはこれで問題ありませんが、推論時(テキスト生成時) にはボトルネックになります。自己回帰生成では、新しいトークンを 1 つ生成するたびに、過去の全トークンの Key と Value を参照する必要があります。この過去の Key/Value を保持するメモリ領域を KV キャッシュ と呼びます。

KV キャッシュのサイズは次のように見積もれます。

$$ \text{KV cache size} = 2 \times L \times h \times T \times d_{\text{head}} \times \text{bytes} $$

ここで $L$ はレイヤー数、$h$ はヘッド数、$T$ は系列長、$d_{\text{head}} = d / h$ はヘッドあたりの次元数です。先頭の $2$ は Key と Value の 2 つ分です。

例えば LLaMA-70B($L = 80, h = 64, d_{\text{head}} = 128$)で系列長 $T = 4096$ の場合、FP16(2 bytes)で KV キャッシュのサイズを計算してみましょう。

$$ 2 \times 80 \times 64 \times 4096 \times 128 \times 2 = 107{,}374{,}182{,}400 \approx 100 \text{ GB} $$

これはモデルの重み自体のサイズ(約 140 GB in FP16)に匹敵する巨大なメモリ消費です。バッチサイズを増やしたり、長い文脈を扱ったりすると、KV キャッシュがメモリの支配的なボトルネックとなります。

MQA と GQA: Key/Value ヘッドの共有

この問題に対する解決策が、Key/Value のヘッドを共有するというアイデアです。

Multi-Query Attention(MQA) は、Shazeer(2019)が提案した手法で、全ての Query ヘッドで 1 つの Key/Value ペアを共有 します。

$$ \text{MQA:} \quad \text{head}_i = \text{Attention}(\bm{Q}_i, \bm{K}_{\text{shared}}, \bm{V}_{\text{shared}}) $$

これにより KV キャッシュのサイズは $1/h$ に削減されます。しかし、全ヘッドが同じ Key/Value を見るため、表現力が低下する という問題がありました。

Grouped Query Attention(GQA) は、Ainslie et al.(2023)が提案した MHA と MQA の中間的なアプローチです。$h$ 個の Query ヘッドを $g$ 個のグループに分け、各グループが 1 つの Key/Value ペアを共有 します。

$$ \text{GQA:} \quad \text{head}_i = \text{Attention}(\bm{Q}_i, \bm{K}_{g(i)}, \bm{V}_{g(i)}) $$

ここで $g(i) = \lfloor i \cdot n_{\text{kv}} / h \rfloor$ は Query ヘッド $i$ が属するグループのインデックスで、$n_{\text{kv}}$ は KV ヘッドの数です。

3 つの方式を整理しましょう。

方式 Query ヘッド数 KV ヘッド数 KV キャッシュ比率 特徴
MHA $h$ $h$ $1.0$ 最高の表現力
GQA $h$ $n_{\text{kv}}$($1 < n_{\text{kv}} < h$) $n_{\text{kv}} / h$ 表現力と効率のバランス
MQA $h$ $1$ $1 / h$ 最高の効率

LLaMA 2 の 70B モデルでは、$h = 64$(Query ヘッド数)に対して $n_{\text{kv}} = 8$(KV ヘッド数)の GQA が採用されています。つまり、8 個の Query ヘッドが 1 つの KV ペアを共有するため、KV キャッシュのサイズは MHA の $8/64 = 1/8$ に削減されます。

先ほどの LLaMA-70B の例では、KV キャッシュが約 100 GB だったものが約 12.5 GB に削減されるということです。この差は、実用的な推論環境では極めて大きな違いとなります。

GQA の数式

GQA の計算を数式で整理します。$i$ 番目の Query ヘッドの Attention は次のように計算されます。

まず、Query は各ヘッド固有の重み行列で計算されます。

$$ \bm{Q}_i = \bm{H}\bm{W}^Q_i, \quad \bm{W}^Q_i \in \mathbb{R}^{d \times d_{\text{head}}} $$

Key と Value は、グループごとに共有された重み行列で計算されます。

$$ \bm{K}_{g(i)} = \bm{H}\bm{W}^K_{g(i)}, \quad \bm{V}_{g(i)} = \bm{H}\bm{W}^V_{g(i)} $$

ここで $\bm{W}^K_{g(i)}, \bm{W}^V_{g(i)} \in \mathbb{R}^{d \times d_{\text{head}}}$ です。

Attention の計算自体は通常と同じです。

$$ \text{Attention}(\bm{Q}_i, \bm{K}_{g(i)}, \bm{V}_{g(i)}) = \text{softmax}\left(\frac{\bm{Q}_i \bm{K}_{g(i)}^T}{\sqrt{d_{\text{head}}}}\right)\bm{V}_{g(i)} $$

最後に、全ヘッドの出力を結合して出力射影を適用します。

$$ \text{GQA}(\bm{H}) = \text{Concat}(\text{head}_1, \dots, \text{head}_h)\bm{W}^O $$

ここで $\bm{W}^O \in \mathbb{R}^{d \times d}$ は出力射影行列です。

GQA のパラメータ数を MHA と比較しましょう。MHA では Query、Key、Value それぞれに $h$ 個のヘッド分の重み行列が必要で、合計 $3hd \cdot d_{\text{head}} = 3d^2$ パラメータ(出力射影を除く)です。GQA では Query が $h$ ヘッド、Key/Value が $n_{\text{kv}}$ ヘッドなので、$(h + 2n_{\text{kv}})d \cdot d_{\text{head}}$ パラメータとなります。$h = 64, n_{\text{kv}} = 8$ の場合、パラメータ数の比は次のようになります。

$$ \frac{h + 2n_{\text{kv}}}{3h} = \frac{64 + 16}{192} = \frac{80}{192} \approx 0.42 $$

つまり、Attention のパラメータ数自体も約 58% に削減されます。

Ainslie et al. の実験では、GQA は MHA とほぼ同等の品質を維持しつつ、MQA に近い推論速度を達成することが報告されています。LLaMA 2 では 70B モデルのみに GQA が採用されましたが、LLaMA 3 以降では全サイズで GQA が標準となっています。

ここまでで LLaMA のアーキテクチャの全コンポーネントを解説しました。次に、LLaMA ファミリーの各モデルのサイズと設定を整理しましょう。

LLaMA のモデルサイズ一覧

LLaMA 1(2023年2月)

LLaMA 1 は 4 つのサイズで公開されました。全モデルで Attention は標準的な MHA を使用しています。

パラメータ 7B 13B 33B 65B
隠れ次元 $d$ 4096 5120 6656 8192
レイヤー数 $L$ 32 40 52 64
ヘッド数 $h$ 32 40 52 64
FFN 隠れ次元 $d_{\text{ff}}$ 11008 13824 17920 22016
学習トークン数 1.0T 1.0T 1.4T 1.4T
コンテキスト長 2048 2048 2048 2048

$d_{\text{ff}}$ の値を確認しましょう。$\frac{8}{3} \times 4096 = 10922.67$ ですが、256 の倍数に丸めると $10922.67 \to 11008 = 43 \times 256$ となります。他のサイズでも同様の丸め処理が行われています。

LLaMA 2(2023年7月)

LLaMA 2 では 3 つのサイズが公開され、70B モデルで GQA が導入されました。また、コンテキスト長が 2048 から 4096 に倍増しています。

パラメータ 7B 13B 70B
隠れ次元 $d$ 4096 5120 8192
レイヤー数 $L$ 32 40 80
Query ヘッド数 $h$ 32 40 64
KV ヘッド数 $n_{\text{kv}}$ 32 (MHA) 40 (MHA) 8 (GQA)
FFN 隠れ次元 $d_{\text{ff}}$ 11008 13824 28672
学習トークン数 2.0T 2.0T 2.0T
コンテキスト長 4096 4096 4096

LLaMA 2 の 70B モデルでは、LLaMA 1 の 65B からいくつかの変更があります。レイヤー数が 64 から 80 に増加し、GQA により KV ヘッドが 8 に削減されています。学習トークン数も全サイズで 2 兆トークンに統一されました。

LLaMA 2 はさらに、Chat モデル(LLaMA 2-Chat)も合わせて公開されました。これは RLHF(Reinforcement Learning from Human Feedback)によるアラインメントが施されたモデルで、対話形式のタスクに特化しています。

LLaMA の各バージョンのモデル構成がわかったところで、いよいよこれらのコンポーネントを PyTorch で実装してみましょう。コードを通じて、各パーツがどのように組み合わさるかを具体的に確認します。

PyTorch での実装

ここからは、LLaMA のアーキテクチャを構成する主要コンポーネントを PyTorch で実装します。各クラスを個別に実装した後、それらを組み合わせて LLaMA の Transformer ブロックを構築します。

RMSNorm クラス

まず、RMSNorm を実装します。

import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class RMSNorm(nn.Module):
    """Root Mean Square Layer Normalization"""
    def __init__(self, dim: int, eps: float = 1e-6):
        super().__init__()
        self.eps = eps
        # 学習可能なスケールパラメータ γ
        self.weight = nn.Parameter(torch.ones(dim))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # RMS(x) = sqrt(mean(x^2) + eps)
        rms = torch.sqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
        # x / RMS(x) * γ
        return (x / rms) * self.weight

# --- 動作確認 ---
torch.manual_seed(42)
dim = 8
norm = RMSNorm(dim)

x = torch.randn(2, 4, dim)  # (batch=2, seq_len=4, dim=8)
out = norm(x)

print(f"入力の形状: {x.shape}")
print(f"出力の形状: {out.shape}")
print(f"入力の二乗平均平方根: {x.pow(2).mean(dim=-1).sqrt()}")
print(f"出力の二乗平均平方根: {out.pow(2).mean(dim=-1).sqrt()}")

出力の二乗平均平方根が入力に比べて 1 に近い値に正規化されていることが確認できます。RMSNorm は平均を引かないため、出力の平均は必ずしもゼロにはなりません。しかし、各次元のスケールが揃うことで、後続の層が安定して学習できるようになります。パラメータは $\bm{\gamma}$ の $d$ 個のみで、Layer Norm のように $\bm{\beta}$ を持たないため、パラメータ数も半分です。

SwiGLU FFN クラス

次に、SwiGLU を用いた FFN を実装します。

class SwiGLUFFN(nn.Module):
    """SwiGLU Feed-Forward Network"""
    def __init__(self, dim: int, hidden_dim: int = None):
        super().__init__()
        if hidden_dim is None:
            # 8/3 * dim を256の倍数に丸める
            hidden_dim = int(2 * (4 * dim) / 3)
            hidden_dim = ((hidden_dim + 255) // 256) * 256

        # 値を計算する線形変換
        self.w1 = nn.Linear(dim, hidden_dim, bias=False)
        # ゲートを計算する線形変換
        self.w_gate = nn.Linear(dim, hidden_dim, bias=False)
        # 出力射影
        self.w2 = nn.Linear(hidden_dim, dim, bias=False)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # SwiGLU(x) = (xW1) ⊙ SiLU(xW_gate)
        # その後 W2 で元の次元に戻す
        return self.w2(self.w1(x) * F.silu(self.w_gate(x)))

# --- 動作確認 ---
dim = 512
ffn = SwiGLUFFN(dim)

x = torch.randn(2, 10, dim)  # (batch=2, seq_len=10, dim=512)
out = ffn(x)

print(f"入力の形状: {x.shape}")
print(f"出力の形状: {out.shape}")
print(f"隠れ次元: {ffn.w1.out_features}")

# パラメータ数の確認
total_params = sum(p.numel() for p in ffn.parameters())
print(f"SwiGLU FFN パラメータ数: {total_params:,}")
print(f"標準 FFN (4d) パラメータ数: {2 * dim * 4 * dim:,}")

出力を確認すると、隠れ次元が $\frac{8}{3} \times 512 \approx 1365$ を 256 の倍数に丸めた 1536 になっていることがわかります。SwiGLU FFN は 3 つの重み行列($\bm{W}_1, \bm{W}_{\text{gate}}, \bm{W}_2$)を持ちますが、隠れ次元を調整することで、標準的な FFN($d_{\text{ff}} = 4d$、2 つの行列)とほぼ同等のパラメータ数になっています。バイアス項を省略していることも LLaMA の特徴で、これにより僅かながらパラメータ数とメモリが削減されます。

GQA 付き Attention クラス

次に、GQA を実装します。これが LLaMA のアーキテクチャで最も実装の注意が必要なコンポーネントです。

class GroupedQueryAttention(nn.Module):
    """Grouped Query Attention (GQA)"""
    def __init__(
        self,
        dim: int,
        n_heads: int,
        n_kv_heads: int = None,
    ):
        super().__init__()
        self.n_heads = n_heads
        # KVヘッド数(Noneの場合はMHA)
        self.n_kv_heads = n_kv_heads if n_kv_heads is not None else n_heads
        self.head_dim = dim // n_heads
        # 1つのKVヘッドあたり何個のQueryヘッドが共有するか
        self.n_rep = self.n_heads // self.n_kv_heads

        # Query: 全ヘッド分
        self.wq = nn.Linear(dim, n_heads * self.head_dim, bias=False)
        # Key/Value: KVヘッド分のみ
        self.wk = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=False)
        self.wv = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=False)
        # 出力射影
        self.wo = nn.Linear(n_heads * self.head_dim, dim, bias=False)

    def _repeat_kv(self, x: torch.Tensor) -> torch.Tensor:
        """KVヘッドを繰り返してQueryヘッド数に合わせる"""
        if self.n_rep == 1:
            return x
        bs, seq_len, n_kv_heads, head_dim = x.shape
        # (bs, seq_len, n_kv_heads, 1, head_dim)
        #   -> (bs, seq_len, n_kv_heads, n_rep, head_dim)
        #   -> (bs, seq_len, n_heads, head_dim)
        x = x.unsqueeze(3).expand(bs, seq_len, n_kv_heads, self.n_rep, head_dim)
        return x.reshape(bs, seq_len, self.n_heads, head_dim)

    def forward(
        self,
        x: torch.Tensor,
        mask: torch.Tensor = None,
    ) -> torch.Tensor:
        bs, seq_len, _ = x.shape

        # Query, Key, Value を計算
        q = self.wq(x).view(bs, seq_len, self.n_heads, self.head_dim)
        k = self.wk(x).view(bs, seq_len, self.n_kv_heads, self.head_dim)
        v = self.wv(x).view(bs, seq_len, self.n_kv_heads, self.head_dim)

        # ※ 実際のLLaMAではここでRoPEを適用
        # q, k = apply_rope(q, k, freqs_cis)

        # KVヘッドを繰り返してQueryヘッド数に合わせる
        k = self._repeat_kv(k)  # (bs, seq_len, n_heads, head_dim)
        v = self._repeat_kv(v)  # (bs, seq_len, n_heads, head_dim)

        # (bs, n_heads, seq_len, head_dim) に転置
        q = q.transpose(1, 2)
        k = k.transpose(1, 2)
        v = v.transpose(1, 2)

        # スケーリングドット積アテンション
        scale = math.sqrt(self.head_dim)
        scores = torch.matmul(q, k.transpose(-2, -1)) / scale

        # 因果マスクの適用
        if mask is not None:
            scores = scores.masked_fill(mask == 0, float('-inf'))

        attn_weights = F.softmax(scores, dim=-1)
        output = torch.matmul(attn_weights, v)

        # ヘッドを結合して出力射影
        output = output.transpose(1, 2).contiguous().view(bs, seq_len, -1)
        return self.wo(output)

GQA の実装のポイントは _repeat_kv メソッドです。KV ヘッドが $n_{\text{kv}}$ 個しかないのに対して Query ヘッドは $h$ 個あるため、計算時に KV ヘッドを $h / n_{\text{kv}}$ 回繰り返して次元を揃えます。メモリ上では KV は $n_{\text{kv}}$ ヘッド分しか保持しないため、KV キャッシュのサイズは $n_{\text{kv}} / h$ に削減されます。

動作を確認しましょう。

# --- GQA の動作確認 ---
dim = 512
n_heads = 8
n_kv_heads = 2  # 4つのQueryヘッドが1つのKVヘッドを共有

gqa = GroupedQueryAttention(dim, n_heads, n_kv_heads)
x = torch.randn(2, 10, dim)

# 因果マスク
seq_len = 10
mask = torch.tril(torch.ones(seq_len, seq_len)).unsqueeze(0).unsqueeze(0)

out = gqa(x, mask)
print(f"入力の形状: {x.shape}")
print(f"出力の形状: {out.shape}")
print(f"Queryヘッド数: {n_heads}")
print(f"KVヘッド数: {n_kv_heads}")
print(f"Queryヘッド / KVヘッド = {n_heads // n_kv_heads} (繰り返し数)")

# MHAとのパラメータ数比較
gqa_params = sum(p.numel() for p in gqa.parameters())
mha = GroupedQueryAttention(dim, n_heads, n_heads)  # MHA
mha_params = sum(p.numel() for p in mha.parameters())
print(f"\nGQA パラメータ数: {gqa_params:,}")
print(f"MHA パラメータ数: {mha_params:,}")
print(f"パラメータ削減率: {1 - gqa_params / mha_params:.1%}")

GQA のパラメータ数が MHA より少ないことが確認できます。Query の重み行列は同じサイズですが、Key/Value の重み行列が $n_{\text{kv}} / h$ 倍に縮小されているためです。出力射影 $\bm{W}^O$ のサイズは変わらない($d \times d$)ため、削減率は Attention 全体で見ると緩やかですが、推論時の KV キャッシュ削減効果は劇的です。

LLaMA ブロックの組み立て

最後に、これまでのコンポーネントを組み合わせて LLaMA の Transformer ブロックを構築します。

class LLaMABlock(nn.Module):
    """LLaMA Transformer ブロック"""
    def __init__(
        self,
        dim: int,
        n_heads: int,
        n_kv_heads: int = None,
        ffn_hidden_dim: int = None,
        norm_eps: float = 1e-6,
    ):
        super().__init__()
        # Pre-Norm用のRMSNorm(Attention前)
        self.attention_norm = RMSNorm(dim, eps=norm_eps)
        # GQA付きAttention
        self.attention = GroupedQueryAttention(dim, n_heads, n_kv_heads)
        # Pre-Norm用のRMSNorm(FFN前)
        self.ffn_norm = RMSNorm(dim, eps=norm_eps)
        # SwiGLU FFN
        self.ffn = SwiGLUFFN(dim, ffn_hidden_dim)

    def forward(
        self,
        x: torch.Tensor,
        mask: torch.Tensor = None,
    ) -> torch.Tensor:
        # Pre-Norm + Attention + 残差接続
        h = x + self.attention(self.attention_norm(x), mask)
        # Pre-Norm + FFN + 残差接続
        out = h + self.ffn(self.ffn_norm(h))
        return out


class LLaMAModel(nn.Module):
    """LLaMA モデル全体"""
    def __init__(
        self,
        vocab_size: int,
        dim: int,
        n_layers: int,
        n_heads: int,
        n_kv_heads: int = None,
        ffn_hidden_dim: int = None,
        norm_eps: float = 1e-6,
        max_seq_len: int = 2048,
    ):
        super().__init__()
        # トークン埋め込み(位置埋め込みはRoPEで代替)
        self.tok_embeddings = nn.Embedding(vocab_size, dim)
        # Transformerブロックの積層
        self.layers = nn.ModuleList([
            LLaMABlock(dim, n_heads, n_kv_heads, ffn_hidden_dim, norm_eps)
            for _ in range(n_layers)
        ])
        # 最終正規化
        self.norm = RMSNorm(dim, eps=norm_eps)
        # 出力ヘッド(語彙サイズへの射影)
        self.output = nn.Linear(dim, vocab_size, bias=False)

    def forward(
        self,
        tokens: torch.Tensor,
        mask: torch.Tensor = None,
    ) -> torch.Tensor:
        h = self.tok_embeddings(tokens)

        # 因果マスクの生成
        if mask is None:
            seq_len = tokens.shape[1]
            mask = torch.tril(torch.ones(seq_len, seq_len, device=tokens.device))
            mask = mask.unsqueeze(0).unsqueeze(0)

        # 各Transformerブロックを順に適用
        for layer in self.layers:
            h = layer(h, mask)

        h = self.norm(h)
        logits = self.output(h)
        return logits

このコードで注目すべき点は 3 つあります。第一に、LLaMABlock 内で RMSNorm が Attention と FFN のに適用されている点です。これが Pre-Norm のパターンです。第二に、LLaMAModel でトークン埋め込みの後に位置埋め込みを加算していない点です。LLaMA では RoPE を使うため、位置情報は Attention の計算時に注入されます(本実装では RoPE の適用部分は省略しています)。第三に、バイアス項が全ての線形層で bias=False に設定されている点です。これも LLaMA の特徴で、わずかながらパラメータ数の削減に貢献しています。

パラメータ数の確認

LLaMA-7B 相当の設定でモデルを構築し、パラメータ数を確認しましょう。

import numpy as np

# LLaMA-7B 相当の設定
config = {
    "vocab_size": 32000,
    "dim": 4096,
    "n_layers": 32,
    "n_heads": 32,
    "n_kv_heads": 32,       # LLaMA 1はMHA
    "ffn_hidden_dim": 11008,
    "max_seq_len": 2048,
}

model = LLaMAModel(**config)

# パラメータ数の集計
def count_parameters(model):
    """各コンポーネントのパラメータ数を集計"""
    counts = {}
    total = 0
    for name, param in model.named_parameters():
        total += param.numel()
        # コンポーネント別に集計
        component = name.split('.')[0]
        if 'attention' in name and 'norm' not in name:
            component = 'attention'
        elif 'ffn' in name and 'norm' not in name:
            component = 'ffn'
        elif 'norm' in name:
            component = 'norm'
        counts[component] = counts.get(component, 0) + param.numel()
    return total, counts

total, counts = count_parameters(model)

print("=" * 50)
print(f"LLaMA-7B 相当のパラメータ数")
print("=" * 50)
print(f"総パラメータ数: {total:,} ({total / 1e9:.2f}B)")
print()
for component, count in sorted(counts.items(), key=lambda x: -x[1]):
    print(f"  {component:20s}: {count:>15,} ({count/total*100:.1f}%)")

# 各ブロックの内訳
print()
print("--- 1ブロックあたりの内訳 ---")
block = model.layers[0]
for name, param in block.named_parameters():
    print(f"  {name:30s}: {str(list(param.shape)):20s} = {param.numel():>12,}")

block_total = sum(p.numel() for p in block.parameters())
print(f"  {'合計':30s}: {'':20s} = {block_total:>12,}")

出力から、LLaMA-7B のパラメータ数が約 6.7B(67 億)であることが確認できます。「7B」という名称はおおよその目安で、実際のパラメータ数はそれより若干少ない値になります。コンポーネント別に見ると、FFN(SwiGLU)が最も多くのパラメータを占めていることがわかります。これは $3 \times d \times d_{\text{ff}}$ のパラメータを持つためです。次に Attention が続き、正規化(RMSNorm)のパラメータは全体のごく一部です。

最後に、LLaMA 2-70B(GQA あり)の設定でもパラメータ数を確認してみましょう。

# LLaMA 2-70B 相当の設定
config_70b = {
    "vocab_size": 32000,
    "dim": 8192,
    "n_layers": 80,
    "n_heads": 64,
    "n_kv_heads": 8,         # GQA: 8つのKVヘッド
    "ffn_hidden_dim": 28672,
    "max_seq_len": 4096,
}

# メモリ節約のため実際にはインスタンス化せず計算
def estimate_params(config):
    """パラメータ数を数式から見積もる"""
    d = config["dim"]
    L = config["n_layers"]
    h = config["n_heads"]
    n_kv = config["n_kv_heads"]
    d_ff = config["ffn_hidden_dim"]
    V = config["vocab_size"]
    d_head = d // h

    # 各コンポーネント(1ブロック分)
    attn_q = d * (h * d_head)          # Query射影
    attn_k = d * (n_kv * d_head)       # Key射影
    attn_v = d * (n_kv * d_head)       # Value射影
    attn_o = (h * d_head) * d          # 出力射影
    attn_total = attn_q + attn_k + attn_v + attn_o

    ffn_total = d * d_ff * 3           # W1, W_gate, W2

    norm_total = d * 2                 # 2つのRMSNorm(γのみ)

    block_total = attn_total + ffn_total + norm_total

    # モデル全体
    embedding = V * d
    final_norm = d
    output_head = d * V

    total = L * block_total + embedding + final_norm + output_head

    return {
        "attention (per block)": attn_total,
        "ffn (per block)": ffn_total,
        "norm (per block)": norm_total,
        "block total": block_total,
        "embedding": embedding,
        "output_head": output_head,
        "total": total,
    }

for name, cfg in [("LLaMA 1 - 7B", config), ("LLaMA 2 - 70B", config_70b)]:
    est = estimate_params(cfg)
    print(f"\n{'=' * 50}")
    print(f"{name}")
    print(f"{'=' * 50}")
    for k, v in est.items():
        print(f"  {k:30s}: {v:>15,} ({v/1e9:.3f}B)")

# GQAによるKVキャッシュ削減の見積もり
print(f"\n{'=' * 50}")
print(f"KVキャッシュ比較 (LLaMA 2-70B, seq_len=4096, FP16)")
print(f"{'=' * 50}")
seq_len = 4096
d_head = 128
L = 80
bytes_per_elem = 2  # FP16

mha_kv = 2 * L * 64 * seq_len * d_head * bytes_per_elem
gqa_kv = 2 * L * 8 * seq_len * d_head * bytes_per_elem

print(f"  MHA (64 KV heads): {mha_kv / 1e9:.1f} GB")
print(f"  GQA ( 8 KV heads): {gqa_kv / 1e9:.1f} GB")
print(f"  削減率: {1 - gqa_kv / mha_kv:.1%}")

この結果から、LLaMA 2-70B のパラメータ数が約 68.9B であること、そして GQA によって KV キャッシュが MHA の 1/8 に削減されることが数値で確認できます。70B クラスのモデルでは KV キャッシュだけで 100 GB 近くになり得るため、GQA による 87.5% の削減は実用上極めて重要です。この削減がなければ、70B モデルを単一の GPU で推論することは事実上不可能でしょう。

GPT アーキテクチャとのコード上の差分

ここで、GPT と LLaMA のコード構造の違いを整理しておきます。

# GPT のブロック(擬似コード)
class GPTBlock:
    def forward(self, x):
        # Post-Norm: サブレイヤーの後に正規化
        x = self.layer_norm_1(x + self.mha(x))
        x = self.layer_norm_2(x + self.ffn(x))
        return x

# LLaMA のブロック(擬似コード)
class LLaMABlock:
    def forward(self, x):
        # Pre-Norm: サブレイヤーの前に正規化
        x = x + self.gqa(self.rms_norm_1(x))
        x = x + self.swiglu_ffn(self.rms_norm_2(x))
        return x

この擬似コードの比較から、変更点が「正規化の位置」「正規化の種類」「Attention の種類」「FFN の活性化関数」という 4 つの差し替えに集約されていることがわかります。基本的な残差接続の構造は共通しており、Decoder-only Transformer の骨格はそのまま保たれています。

まとめ

本記事では、LLaMA のアーキテクチャを GPT との差分に焦点を当てて解説しました。

  • Pre-Norm: 正規化をサブレイヤーの前に配置することで、残差接続を通じた勾配の流れを確保し、学習を安定化させます
  • RMSNorm: Layer Normalization から平均の計算(re-centering)を省略し、二乗平均平方根による正規化のみを行うことで、同等の性能を保ちつつ計算効率を向上させます
  • SwiGLU: FFN にゲート機構(GLU)と SiLU 活性化関数を組み合わせた SwiGLU を導入し、隠れ次元を $\frac{8}{3}d$ に調整することでパラメータ数を維持しつつ性能を向上させます
  • RoPE: 学習可能な絶対位置埋め込みを回転位置埋め込みに置き換え、相対位置の自然な表現と長文への汎化能力を実現します
  • GQA(LLaMA 2 以降): Key/Value ヘッドをグループ化して共有することで、推論時の KV キャッシュを大幅に削減し、メモリ効率を改善します

これらの変更はそれぞれ独立した研究に基づいていますが、LLaMA はこれらを巧みに組み合わせ、Chinchilla スケーリング則に基づくデータ戦略と合わせることで、オープンソース LLM の新たな基準を打ち立てました。

LLaMA のアーキテクチャは、Mistral、Mixtral、Gemma、Qwen、Yi など、後続のほぼ全てのオープンソース LLM に継承されています。本記事で学んだ内容は、これらのモデルを理解する共通基盤となります。

次のステップとして、以下の記事も参考にしてください。

画像なし
Mistral / Mixtral のアーキテクチャを解説
LLaMAをベースにSliding Window AttentionやMixture of Expertsを導入したMistral/Mixtralの設計を解説します
画像なし
LoRA(Low-Rank Adaptation)の理論と実装
大規模言語モデルを効率的にファインチューニングするLoRAの数理と実装を解説します
画像なし
QLoRA による効率的なファインチューニング
4bit量子化とLoRAを組み合わせたQLoRAで、消費者向けGPUでLLMをファインチューニングする方法を解説します