Context Window拡張技術 — 学習済みモデルのコンテキスト長を後から伸ばす手法を徹底解説

LLaMAは4,096トークンで学習されました。しかし実務では、1本の論文(約8,000トークン)を丸ごと要約したい、数十ファイルにまたがるコードベースを一度に解析したい、といった場面が頻繁に発生します。学習時のコンテキスト長を超える入力をそのまま与えると、モデルの出力品質は劇的に劣化します。Perplexityが数百から数千に跳ね上がり、もはや意味のあるテキストを生成できなくなるのです。

「では、もっと長いコンテキストで最初から学習し直せばよいのでは?」と思うかもしれません。しかし、LLaMAクラスのモデルを4,096トークンから32,768トークンに再学習するには、膨大なGPU時間とデータが必要です。学習済みの知識を活かしつつ、コンテキスト長だけを拡張できれば、コストを大幅に削減できます。

2023年以降、RoPE(Rotary Position Embedding)の数学的性質を巧みに利用して、学習済みモデルのコンテキスト長を後から拡張する手法が次々と提案されました。これがContext Window拡張技術です。

Context Window拡張技術を理解すると、以下のことが可能になります。

  • 長文処理の設計: 学習済みモデルの能力を活かしたまま、処理可能な系列長を数倍〜数十倍に伸ばす方法の理解
  • 最新LLMの仕組み: LLaMA 2 Long、Code Llama、Mistralなどが採用するコンテキスト拡張の原理を把握
  • ファインチューニング戦略: 少量の追加学習でコンテキスト長を効率的に拡張する設計判断の習得
  • 位置エンコーディングの深い理解: RoPEの周波数構造がなぜ外挿に失敗し、補間なら成功するのかの数学的直感

本記事の内容

  • コンテキスト長の限界 — なぜ学習時の長さを超えると破綻するのか
  • Position Interpolation(PI)— 位置を線形に圧縮する最初のアイデア
  • NTK-Aware Scaling — 高周波成分を保護するスケーリング
  • YaRN — 周波数帯ごとに最適な戦略を選ぶ統合手法
  • Dynamic NTK Scaling — 入力長に応じて動的にbaseを調整
  • Pythonでの位置補間の実装と可視化
  • Perplexityへの影響の実験

前提知識

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

画像なし
RoPE(回転位置埋め込み)の数学的導出と実装
RoPEの回転行列による位置情報埋め込み、相対位置の表現、長系列への外挿を解説します。
画像なし
位置エンコーディングの各種手法
sin/cos、学習可能PE、RoPE、ALiBiの全体像を俯瞰します。
画像なし
LLaMAアーキテクチャの全体像
LLaMAの構造とRoPEの組み込み方を解説します。
画像なし
KVキャッシュの仕組みと実装
推論時のKey/Valueキャッシュの仕組みを解説します。

コンテキスト長の限界 — なぜ外挿は失敗するのか

RoPEの復習

Context Window拡張技術を理解するには、まずRoPEの仕組みを正確に把握する必要があります。ここでは本記事で使う最低限の数学を復習します。

RoPEは、位置 $m$ のQuery/Keyベクトルに、位置に応じた回転を適用します。$d$ 次元のベクトルを $d/2$ 個の2次元ペアに分割し、各ペア $(q_{2i}, q_{2i+1})$ に対して角度 $m\theta_i$ の回転を行います。

$$ \bm{R}(m, \theta_i) = \begin{pmatrix} \cos(m\theta_i) & -\sin(m\theta_i) \\ \sin(m\theta_i) & \cos(m\theta_i) \end{pmatrix} $$

ここで、各次元ペアに割り当てられる基本角周波数 $\theta_i$ は

$$ \theta_i = b^{-2i/d}, \quad i = 0, 1, \ldots, d/2 – 1 $$

と定義されます。$b$ はbase値で、元のRoPEでは $b = 10000$ が使われます。

この設計の意味を直感的に理解しましょう。$i = 0$(最初の次元ペア)では $\theta_0 = 1$ で、位置が1つ進むごとに1ラジアン回転します。波長は $2\pi \approx 6.28$ トークンです。一方、$i = d/2 – 1$(最後の次元ペア)では $\theta_{d/2-1} = 10000^{-1} = 0.0001$ で、波長は $2\pi \times 10000 \approx 62{,}832$ トークンです。つまり、低次元のペアは近い位置の違いに敏感で、高次元のペアは遠い位置の違いに敏感です。

アテンションスコアと相対位置

RoPEの最も重要な性質は、位置 $m$ のQueryと位置 $n$ のKeyの内積が相対位置 $m – n$ のみに依存することです。

$$ \langle \bm{R}(m)\bm{q},\, \bm{R}(n)\bm{k} \rangle = \sum_{i=0}^{d/2-1} \langle \bm{R}(m\theta_i)\bm{q}_i,\, \bm{R}(n\theta_i)\bm{k}_i \rangle $$

回転行列の性質 $\bm{R}(\alpha)^\top \bm{R}(\beta) = \bm{R}(\beta – \alpha)$ を使うと、各ペアの内積は

$$ \langle \bm{R}(m\theta_i)\bm{q}_i,\, \bm{R}(n\theta_i)\bm{k}_i \rangle = \langle \bm{q}_i,\, \bm{R}((n – m)\theta_i)\bm{k}_i \rangle $$

のように相対位置 $n – m$ のみで決まります。この性質がRoPEの強みであり、同時にコンテキスト拡張の鍵でもあります。

外挿が失敗する理由

モデルが最大位置 $L_{\text{train}}$(例: 4096)までで学習された場合、各次元ペア $i$ が学習中に経験する回転角の範囲は

$$ m\theta_i \in [0,\, L_{\text{train}} \cdot \theta_i] $$

です。推論時に位置 $m > L_{\text{train}}$ のトークンが入力されると、一部の次元ペア(特に $\theta_i$ が大きい低次元ペア)で、回転角 $m\theta_i$ が学習時に経験した範囲を超えます。

具体的に数値で確認しましょう。$d = 128$、$b = 10000$、$L_{\text{train}} = 4096$ のとき

  • $i = 0$: $\theta_0 = 1$、学習時の最大回転角 = $4096$ ラジアン(約651回転)。$2\pi$ の周期性により、回転角自体は問題になりません
  • $i = 32$: $\theta_{32} = 10000^{-0.5} = 0.01$、学習時の最大回転角 = $40.96$ ラジアン(約6.5回転)。位置8192では $81.92$ ラジアン(約13回転)で、未経験の角度領域に入ります
  • $i = 63$: $\theta_{63} = 10000^{-63/64} \approx 0.000107$、学習時の最大回転角 = $0.438$ ラジアン。位置8192では $0.876$ ラジアンで、これも未経験

問題は、Attentionの重みを計算するsoftmaxが、学習時に見たことのない角度パターンに対して予測不能な振る舞いをすることです。学習時の角度範囲内では、モデルはAttentionの重みを適切に配分する方法を学んでいます。しかし範囲外では、数式としては値が算出されるものの、その値がモデルの期待する分布から逸脱し、結果としてPerplexityが急上昇します。

これは「訓練データの分布外に出た」という、機械学習における典型的な汎化失敗です。RoPEの外挿問題は、位置エンコーディングという特定のコンポーネントにおける分布シフトと理解できます。

では、この問題をどう解決するか。最もシンプルなアイデアは「外挿ではなく補間する」ことです。

Position Interpolation(PI)— 外挿から補間へ

基本アイデア

Position Interpolation(PI)は、Chen et al. (2023) が提案した、驚くほどシンプルかつ強力な手法です。アイデアは一言で説明できます。「位置インデックスを縮小して、拡張後の全位置が学習時の範囲に収まるようにする」。

元のRoPEでは、位置 $m$ の回転角は $m\theta_i$ でした。拡張後の最大コンテキスト長を $L’$(例: 16384)、元の最大長を $L$(例: 4096)としたとき、PIは位置インデックスを $s = L’/L$(スケール比)で除算します。

$$ m\theta_i \quad \longrightarrow \quad \frac{m}{s}\theta_i $$

つまり、位置 $m$ を $m/s$ に圧縮してからRoPEを適用します。$s = 4$ のとき、位置16384は $16384/4 = 4096$ に圧縮され、ちょうど元の最大位置と一致します。

アナロジーで理解する

PIを直感的に理解するために、「目盛り」のアナロジーを使いましょう。

元のRoPEは、長さ4096の物差しに1刻みで目盛りを振ったようなものです。位置1は目盛り1、位置4096は目盛り4096。この物差しで5000の位置を測ろうとすると、物差しの外にはみ出してしまいます(外挿)。

PIは、同じ物差しの目盛りの間隔を $1/4$ に縮めます。位置1は目盛り0.25、位置4は目盛り1、位置16384は目盛り4096。全ての位置が物差しの範囲に収まります(補間)。

代償は、隣接トークン間の角度差が $1/s$ に縮小することです。元々1ラジアンの差があった隣接トークンが0.25ラジアンの差になり、位置の「分解能」が下がります。ただし、実験的にはこの分解能の低下は、短いファインチューニング(1000ステップ程度)で補償できることが示されています。

数学的定式化

PIにおけるRoPEの回転角を正式に定義します。スケール比 $s = L’/L$ を導入すると

$$ \theta_i^{\text{PI}} = \frac{\theta_i}{s} = \frac{b^{-2i/d}}{s} $$

これは等価的に、baseの値を $b’ = b \cdot s^{d/(d-2)}$ に変更する操作… ではありません。PIは単純に周波数を一律に $1/s$ 倍します。全ての次元ペアが同じ割合で圧縮される点が、次に紹介するNTK-Aware Scalingとの決定的な違いです。

位置 $m$ と位置 $n$ の相対アテンションスコアへの寄与は

$$ \langle \bm{q}_i,\, \bm{R}\left(\frac{(n-m)\theta_i}{s}\right)\bm{k}_i \rangle $$

となります。相対位置 $(n-m)$ が同じでも、スケール比 $s$ によってアテンションスコアが変化します。これが「短いファインチューニングが必要な理由」です。モデルは新しいスケーリングでのAttention分布を微調整する必要があります。

PIの問題点

PIの弱点は、全次元を一律にスケーリングする点にあります。先ほど見たように、RoPEの各次元ペアは異なる周波数を持ちます。

  • 高周波成分($i$ が小さい、$\theta_i$ が大きい): 波長が短く、学習時に何百回も回転を経験している。1つ分の位置差で大きな角度変化が生じるため、隣接トークンの区別に重要
  • 低周波成分($i$ が大きい、$\theta_i$ が小さい): 波長が長く、学習時にもほとんど回転していない。外挿問題が最も深刻

PIは全次元を等しく圧縮するため、元々問題のない高周波成分まで不必要に歪めてしまいます。高周波成分は学習時に十分な回転を経験しており、回転角が $2\pi$ の周期を何度も超えているため、外挿しても問題が起きにくいのです。にもかかわらず、高周波成分の分解能を下げてしまうのはもったいない。

この問題を解決するのが、次に紹介するNTK-Aware Scalingです。高周波成分をなるべく保存しつつ、低周波成分だけをスケーリングする、より賢い戦略です。

NTK-Aware Scaling — 高周波を保護する

Neural Tangent Kernelからの着想

NTK-Aware Scaling(bloc97, 2023)は、Neural Tangent Kernel(NTK)理論からの着想に基づいています。NTK理論では、ニューラルネットワークの学習ダイナミクスにおいて高周波成分の学習が困難であることが知られています。この知見をRoPEのコンテキスト拡張に応用したのがこの手法です。

核心的なアイデアは、PIのように位置インデックスをスケーリングするのではなく、RoPEのbase値 $b$ を大きくすることです。

base変更の効果

元のRoPEの角周波数は $\theta_i = b^{-2i/d}$ でした。baseを $b$ から $b’$ に変更すると

$$ \theta_i’ = (b’)^{-2i/d} $$

ここで $b’ > b$ とすると、全ての $\theta_i’ < \theta_i$(角周波数が小さくなる = 波長が長くなる)ですが、次元によって変化率が異なります

具体的に、角周波数の変化率を計算しましょう。

$$ \frac{\theta_i’}{\theta_i} = \frac{(b’)^{-2i/d}}{b^{-2i/d}} = \left(\frac{b}{b’}\right)^{2i/d} $$

$b’ > b$ のとき $b/b’ < 1$ なので、$i$ が大きいほど(低周波成分ほど)変化率が小さく(より強く圧縮され)、$i = 0$ のとき変化率は $(b/b')^0 = 1$(全く変化しない)です。

これはまさに望ましい性質です。

  • $i = 0$(最高周波数): 変化率 = 1。高周波成分は完全に保存される
  • $i$ が中程度: 中程度の圧縮
  • $i = d/2 – 1$(最低周波数): 最も強く圧縮される。外挿問題が深刻な低周波成分を効果的にスケーリング

スケーリング公式

コンテキスト長を $s$ 倍に拡張するとき、NTK-Aware Scalingではbase値を次のように設定します。

$$ b’ = b \cdot s^{d/(d-2)} $$

なぜこの形になるのでしょうか。PIと等価な「全体の圧縮率」を達成しつつ、圧縮を次元間で非一様に分配するためです。

PIでの全次元にわたる「総圧縮量」を考えます。PIは全次元で一律に $1/s$ のスケーリングを行います。NTK-Aware Scalingは低周波を多く、高周波を少なく圧縮しますが、全次元を合わせた「実効的な圧縮率」が同程度になるようにbaseを設定します。

$d/2$ 個の次元ペアにわたる圧縮率の幾何平均を考えると

$$ \prod_{i=0}^{d/2-1} \frac{\theta_i’}{\theta_i} = \prod_{i=0}^{d/2-1} \left(\frac{b}{b’}\right)^{2i/d} = \left(\frac{b}{b’}\right)^{\frac{2}{d}\sum_{i=0}^{d/2-1} i} $$

ここで指数の和は $\sum_{i=0}^{d/2-1} i = \frac{(d/2-1)(d/2)}{2}$ です。これを計算すると

$$ \left(\frac{b}{b’}\right)^{\frac{2}{d} \cdot \frac{(d/2-1)(d/2)}{2}} = \left(\frac{b}{b’}\right)^{\frac{d/2-1}{2}} $$

PIの幾何平均 $1/s$ と等しくおくと

$$ \left(\frac{b}{b’}\right)^{(d/2-1)/2} = \frac{1}{s} $$

両辺を $(d/2-1)/2$ 乗根とって整理すると

$$ \frac{b’}{b} = s^{2/(d/2-1)} = s^{d/(d(d/2-1)/2)} $$

ここで正確には $\frac{b’}{b} = s^{2/(d/2-1)}$ です。$d$ が十分大きいとき $d/2-1 \approx d/2$ なので

$$ b’ \approx b \cdot s^{d/(d-2)} $$

となります。

PIとNTK-Awareの比較

両手法の違いを、各次元ペアの有効波長で整理します。元のRoPEで次元ペア $i$ の波長は $\lambda_i = 2\pi / \theta_i = 2\pi \cdot b^{2i/d}$ です。

手法 位置 $m$ の回転角 高周波への影響 低周波への影響
元のRoPE $m \cdot b^{-2i/d}$ そのまま そのまま
PI $\frac{m}{s} \cdot b^{-2i/d}$ 一律 $1/s$ に圧縮 一律 $1/s$ に圧縮
NTK-Aware $m \cdot (b’)^{-2i/d}$ ほぼ保存 強く圧縮

NTK-Aware Scalingは、PIよりも高周波の情報を保持するため、ファインチューニングなしでも一定の品質を維持できます。これは「ゼロショット拡張」と呼ばれ、ファインチューニングのコストを削減できる実用的な利点です。

しかし、NTK-Aware Scalingにも改善の余地があります。base変更は次元間の圧縮配分を改善しましたが、その配分は指数的なカーブに固定されており、最適な配分とは限りません。次に紹介するYaRNは、各次元を「高周波」「中間」「低周波」の3つのグループに分類し、それぞれに最適な戦略を適用します。

YaRN — 周波数帯ごとに最適な戦略を選ぶ

YaRNの設計思想

YaRN(Yet another RoPE extensioN method、Peng et al., 2023)は、RoPEの各次元ペアを波長に基づいて3つのカテゴリに分類し、各カテゴリに異なる戦略を適用する統合的な手法です。

YaRNの核心的な洞察は、「コンテキスト拡張に必要な処理は、次元ペアの波長によって本質的に異なる」というものです。

  • 波長が非常に短い次元ペア(高周波): 学習時に何百回も回転を経験しており、$2\pi$ の周期性により外挿しても問題なし。触る必要がない
  • 波長が非常に長い次元ペア(低周波): 学習時に1回転すらしていない。未経験の角度領域に外挿される。PIのような線形補間が必要
  • 波長が中間の次元ペア(中間周波数): 部分的に外挿問題が生じる。高周波と低周波の間を滑らかに補間する必要がある

波長による分類

YaRNは波長のしきい値を2つ定義します。

$$ \lambda_{\text{low}} = \frac{2\pi \cdot L_{\text{train}}}{\alpha_{\text{low}}}, \quad \lambda_{\text{high}} = \frac{2\pi \cdot L_{\text{train}}}{\alpha_{\text{high}}} $$

ここで $\alpha_{\text{low}}$ と $\alpha_{\text{high}}$ はハイパーパラメータです。元の論文では $\alpha_{\text{low}} = 1$, $\alpha_{\text{high}} = 32$ が推奨されています。

次元ペア $i$ の波長 $\lambda_i = 2\pi \cdot b^{2i/d}$ に基づいて

$$ \gamma_i = \begin{cases} 0 & \text{if } \lambda_i < \lambda_{\text{high}} \quad (\text{高周波: そのまま保存}) \\ 1 & \text{if } \lambda_i > \lambda_{\text{low}} \quad (\text{低周波: 完全に補間}) \\ \frac{\lambda_i – \lambda_{\text{high}}}{\lambda_{\text{low}} – \lambda_{\text{high}}} & \text{otherwise} \quad (\text{中間: 線形ランプ}) \end{cases} $$

この $\gamma_i \in [0, 1]$ が、次元ペア $i$ における「PIの混合比率」を決定します。

YaRNの回転角

$\gamma_i$ を使って、各次元ペアの回転角を次のように定義します。

$$ \theta_i^{\text{YaRN}} = \frac{\theta_i}{(1 – \gamma_i) \cdot 1 + \gamma_i \cdot s} = \frac{\theta_i}{1 + \gamma_i(s – 1)} $$

この式の意味を確認しましょう。

  • $\gamma_i = 0$(高周波): $\theta_i^{\text{YaRN}} = \theta_i$。元のRoPEと同じ。スケーリングなし
  • $\gamma_i = 1$(低周波): $\theta_i^{\text{YaRN}} = \theta_i / s$。PIと同じ。完全にスケーリング
  • $0 < \gamma_i < 1$(中間): PIとスケーリングなしの線形補間

つまり、YaRNはPIとNTK-Aware Scalingの「いいとこ取り」です。高周波は完全に保存(NTK-Awareのように)、低周波は完全にスケーリング(PIのように)、中間はなめらかに接続します。

温度スケーリング

YaRNにはもう一つ重要な要素があります。Attentionのlogitsに対する温度スケーリングです。

コンテキスト長が拡張されると、Attentionの分布が変化します。元のモデルはトークン数 $L_{\text{train}}$ のAttention分布で学習していますが、拡張後は $L’$ 個のトークンにAttentionを分配する必要があります。トークン数が増えると、各トークンへのAttention重みの平均値が下がり、Attentionの「エントロピー」が上がります。

YaRNはこの問題を、Attentionスコアに温度係数 $\sqrt{t}$ を乗じることで補正します。

$$ \text{Attention}^{\text{YaRN}} = \text{softmax}\left(\frac{\sqrt{t} \cdot \bm{Q}\bm{K}^\top}{\sqrt{d_k}}\right)\bm{V} $$

推奨値は $t = 0.1 \ln(s) + 1$ です。$s = 4$ のとき $t \approx 1.139$、$\sqrt{t} \approx 1.067$ で、Attentionスコアがわずかにシャープになります。これにより、拡張前と同等のAttentionの「集中度」が保たれます。

YaRNは複数の工夫を組み合わせた手法ですが、各要素の動機が明確で、実装も比較的シンプルです。次に、YaRNとは異なるアプローチで動的にスケーリングを行う手法を紹介します。

Dynamic NTK Scaling — 入力長に応じた動的調整

静的スケーリングの問題

これまで紹介したPI、NTK-Aware Scaling、YaRNは全て、拡張後のコンテキスト長 $L’$ を事前に決めてスケーリングを適用する「静的」な手法です。しかし、実際のLLM推論では、入力の長さは推論時まで分かりません。

$L’ = 16384$ でスケーリングを設定したモデルに、実際には512トークンの短い入力が来た場合を考えましょう。不要なスケーリングが適用され、短い入力の処理品質が下がる可能性があります。逆に、$L’ = 8192$ で設定したのに32768トークンの入力が来た場合、スケーリングが不十分で品質が劣化します。

Dynamic NTKの仕組み

Dynamic NTK Scaling(Emozilla, 2023)は、入力系列の実際の長さに基づいてbase値を動的に調整します。

推論時の現在の系列長を $L_{\text{current}}$ とします。$L_{\text{current}} \leq L_{\text{train}}$ のときは元のRoPEをそのまま使い(スケーリング不要)、$L_{\text{current}} > L_{\text{train}}$ のときだけNTK-Awareスケーリングを適用します。

$$ b’_{\text{dynamic}} = \begin{cases} b & \text{if } L_{\text{current}} \leq L_{\text{train}} \\ b \cdot \left(\frac{\alpha \cdot L_{\text{current}}}{L_{\text{train}}}\right)^{d/(d-2)} & \text{if } L_{\text{current}} > L_{\text{train}} \end{cases} $$

ここで $\alpha$ は安全マージンの係数で、通常 $\alpha = 1$ が使われます。

ポイントは、推論の各ステップで系列長が1トークンずつ増えるたびに、base値が微小に変化する点です。自己回帰生成(トークンを1つずつ生成する処理)では、生成が進むにつれて系列が長くなり、base値もそれに追従して少しずつ大きくなります。

実装上の注意点

Dynamic NTK Scalingは実装上の重要な注意点があります。base値が変わると、過去の全トークンのRoPE回転角も変わるため、理論的にはKVキャッシュ内の全てのKey値を再計算する必要があります

KVキャッシュは、過去のトークンのKey/Valueベクトルを保存して再計算を避ける最適化手法です。しかし、Dynamic NTK ScalingではKey値がbase値に依存するため、base値が変わるとキャッシュが無効になります。

実用上は、以下の2つのアプローチが取られます。

  1. キャッシュ再計算: base変更のたびに全てのKey値を再計算する。正確だが計算コストが高い
  2. 近似: キャッシュは再計算せず、新しいトークンのみ新しいbase値で計算する。厳密には不正確だが、base値の変化が微小なので実用上問題ない

多くの実装では、近似アプローチ(方法2)が採用されています。base値の変化は1トークンごとに微小であるため、キャッシュの再計算を省略しても品質への影響はほとんどありません。

ここまでで4つのContext Window拡張手法を学びました。いよいよ、これらの手法をPythonで実装し、各手法の振る舞いの違いを可視化しましょう。

Pythonによる位置補間の実装と可視化

RoPEの周波数スペクトルの比較

まず、各手法がRoPEの周波数(角速度)をどのように変化させるかを可視化します。横軸に次元インデックス、縦軸に角周波数をプロットすることで、高周波の保存度合いと低周波の圧縮度合いが一目で分かります。

import numpy as np
import matplotlib.pyplot as plt

# パラメータ設定
d = 128            # 次元数
base = 10000       # 元のbase
L_train = 4096     # 学習時の最大コンテキスト長
L_target = 16384   # 拡張先のコンテキスト長
s = L_target / L_train  # スケール比 = 4

dim_indices = np.arange(d // 2)  # i = 0, 1, ..., 63

# --- 元のRoPE ---
theta_original = base ** (-2 * dim_indices / d)

# --- Position Interpolation (PI) ---
theta_pi = theta_original / s

# --- NTK-Aware Scaling ---
base_ntk = base * s ** (d / (d - 2))
theta_ntk = base_ntk ** (-2 * dim_indices / d)

# --- YaRN ---
alpha_low = 1
alpha_high = 32
lambda_low = 2 * np.pi * L_train / alpha_low
lambda_high = 2 * np.pi * L_train / alpha_high

wavelengths = 2 * np.pi / theta_original  # 各次元の波長
gamma = np.clip((wavelengths - lambda_high) / (lambda_low - lambda_high), 0, 1)
theta_yarn = theta_original / (1 + gamma * (s - 1))

# 可視化
fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# 左: 角周波数(対数スケール)
ax = axes[0]
ax.semilogy(dim_indices, theta_original, 'k-', linewidth=2, label='Original RoPE')
ax.semilogy(dim_indices, theta_pi, 'b--', linewidth=2, label='Position Interpolation')
ax.semilogy(dim_indices, theta_ntk, 'r-.', linewidth=2, label='NTK-Aware Scaling')
ax.semilogy(dim_indices, theta_yarn, 'g:', linewidth=2.5, label='YaRN')
ax.set_xlabel('Dimension pair index $i$', fontsize=12)
ax.set_ylabel(r'Angular frequency $\theta_i$', fontsize=12)
ax.set_title('Angular Frequency Spectrum', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# 右: 元のRoPEに対する比率
ax = axes[1]
ax.plot(dim_indices, theta_pi / theta_original, 'b--', linewidth=2, label='PI ratio')
ax.plot(dim_indices, theta_ntk / theta_original, 'r-.', linewidth=2, label='NTK-Aware ratio')
ax.plot(dim_indices, theta_yarn / theta_original, 'g:', linewidth=2.5, label='YaRN ratio')
ax.axhline(y=1.0, color='k', linewidth=1, linestyle='-', alpha=0.5, label='No change')
ax.axhline(y=1/s, color='gray', linewidth=1, linestyle=':', alpha=0.5, label=f'1/s = {1/s:.2f}')
ax.set_xlabel('Dimension pair index $i$', fontsize=12)
ax.set_ylabel(r'$\theta_i^{\prime} / \theta_i$', fontsize=12)
ax.set_title('Frequency Ratio (Scaled / Original)', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)
ax.set_ylim(-0.05, 1.15)

plt.tight_layout()
plt.savefig('context_extension_frequency_spectrum.png', dpi=150, bbox_inches='tight')
plt.show()

左のグラフから、以下のことが読み取れます。

  1. PIは全次元の周波数を一律に $1/4$ に下げている(青の破線が黒の実線から一定量だけ下方にシフト)。高周波も低周波も同じ割合で圧縮されます
  2. NTK-Aware Scalingは低次元(高周波)をほぼ保存し、高次元(低周波)を強く圧縮している(赤の一点鎖線が、左端では黒に近く、右端ではPIに近い)。これがbase変更の効果です
  3. YaRNは高周波側で元のRoPEと完全に一致し、低周波側でPIと一致する(緑の点線が、左端では黒と重なり、右端ではPIの青と重なる)。中間は滑らかに遷移しています

右のグラフはこの違いをさらに明確にしています。PIの比率は全次元で0.25(一定)、NTK-Awareは指数的なカーブ、YaRNは区分的な関数(高周波は1.0、低周波は0.25、中間は線形)です。

位置に対する回転角の可視化

次に、各手法で特定の次元ペアの回転角が位置に対してどう変化するかを可視化します。

import numpy as np
import matplotlib.pyplot as plt

# パラメータ
d = 128
base = 10000
L_train = 4096
L_target = 16384
s = L_target / L_train
positions = np.arange(0, L_target + 1, 1)

# 3つの次元ペアを選択(高周波・中間・低周波)
dim_pairs = [0, 16, 48]
dim_labels = ['$i=0$ (high freq)', '$i=16$ (mid freq)', '$i=48$ (low freq)']

fig, axes = plt.subplots(1, 3, figsize=(16, 5))

for ax_idx, (dim_i, label) in enumerate(zip(dim_pairs, dim_labels)):
    ax = axes[ax_idx]

    theta_orig = base ** (-2 * dim_i / d)

    # 元のRoPE
    angle_orig = positions * theta_orig
    # PI
    angle_pi = (positions / s) * theta_orig
    # NTK-Aware
    base_ntk = base * s ** (d / (d - 2))
    theta_ntk_i = base_ntk ** (-2 * dim_i / d)
    angle_ntk = positions * theta_ntk_i
    # YaRN
    wavelength_i = 2 * np.pi / theta_orig
    alpha_low, alpha_high = 1, 32
    lam_low = 2 * np.pi * L_train / alpha_low
    lam_high = 2 * np.pi * L_train / alpha_high
    gamma_i = np.clip((wavelength_i - lam_high) / (lam_low - lam_high), 0, 1)
    theta_yarn_i = theta_orig / (1 + gamma_i * (s - 1))
    angle_yarn = positions * theta_yarn_i

    # mod 2piで表示(回転角のラップアラウンド)
    ax.plot(positions, angle_orig % (2 * np.pi), 'k.', markersize=0.3, alpha=0.4, label='Original')
    ax.plot(positions, angle_pi % (2 * np.pi), 'b.', markersize=0.3, alpha=0.4, label='PI')
    ax.plot(positions, angle_ntk % (2 * np.pi), 'r.', markersize=0.3, alpha=0.4, label='NTK-Aware')
    ax.plot(positions, angle_yarn % (2 * np.pi), 'g.', markersize=0.3, alpha=0.4, label='YaRN')

    ax.axvline(x=L_train, color='orange', linewidth=1.5, linestyle='--', alpha=0.7, label=f'$L_{{train}}={L_train}$')
    ax.set_xlabel('Position $m$', fontsize=11)
    ax.set_ylabel('Rotation angle mod $2\\pi$', fontsize=11)
    ax.set_title(label, fontsize=12)
    ax.legend(fontsize=8, markerscale=10)
    ax.grid(True, alpha=0.3)

plt.suptitle('Rotation Angles for Different Scaling Methods', fontsize=14, y=1.02)
plt.tight_layout()
plt.savefig('context_extension_rotation_angles.png', dpi=150, bbox_inches='tight')
plt.show()

3つのグラフから、以下のことが読み取れます。

  1. 高周波($i=0$): 元のRoPEでは位置が進むにつれ急速に回転し、$[0, 2\pi)$ 全体をまんべんなく使っています。学習時の範囲(オレンジの破線の左側)で既に十分な回転を経験しているため、外挿しても角度のパターンは大きく変わりません。YaRNとNTK-Awareは元のRoPEとほぼ重なっており、高周波を保存していることが確認できます。一方、PIは角度変化の速度が $1/4$ に低下しているのが見えます
  2. 中間周波数($i=16$): 回転の速度は中程度です。各手法の差が最も見やすい領域です。NTK-Awareは元のRoPEと比べてわずかに圧縮されており、YaRNはNTK-AwareとPIの中間の圧縮率を示しています
  3. 低周波($i=48$): 回転が非常にゆっくりで、学習時の範囲内でもわずかしか回転しません。外挿すると未経験の角度に大きく踏み込みます。PI、NTK-Aware、YaRNのいずれも角度変化を抑制して、拡張後も学習時の角度範囲を大きく超えないようにしています

YaRNの混合係数 $\gamma$ の可視化

YaRNの挙動を直感的に理解するため、各次元ペアの波長と、それに対応する $\gamma$ 値の関係を可視化します。

import numpy as np
import matplotlib.pyplot as plt

# パラメータ
d = 128
base = 10000
L_train = 4096
alpha_low = 1
alpha_high = 32

dim_indices = np.arange(d // 2)
theta = base ** (-2 * dim_indices / d)
wavelengths = 2 * np.pi / theta

lambda_low = 2 * np.pi * L_train / alpha_low
lambda_high = 2 * np.pi * L_train / alpha_high

gamma = np.clip((wavelengths - lambda_high) / (lambda_low - lambda_high), 0, 1)

fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# 左: 波長と閾値
ax = axes[0]
ax.semilogy(dim_indices, wavelengths, 'k-', linewidth=2, label='Wavelength $\\lambda_i$')
ax.axhline(y=lambda_low, color='red', linewidth=1.5, linestyle='--',
           label=f'$\\lambda_{{low}}$ = {lambda_low:.0f}')
ax.axhline(y=lambda_high, color='blue', linewidth=1.5, linestyle='--',
           label=f'$\\lambda_{{high}}$ = {lambda_high:.0f}')
ax.fill_between(dim_indices, lambda_high, lambda_low, alpha=0.1, color='green', label='Transition zone')
ax.set_xlabel('Dimension pair index $i$', fontsize=12)
ax.set_ylabel('Wavelength (tokens)', fontsize=12)
ax.set_title('Wavelength Spectrum and YaRN Thresholds', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# 右: gamma値
ax = axes[1]
ax.plot(dim_indices, gamma, 'g-', linewidth=2.5)
ax.fill_between(dim_indices, 0, gamma, alpha=0.2, color='green')
ax.set_xlabel('Dimension pair index $i$', fontsize=12)
ax.set_ylabel('$\\gamma_i$ (interpolation ratio)', fontsize=12)
ax.set_title('YaRN Mixing Coefficient $\\gamma_i$', fontsize=13)
ax.set_ylim(-0.05, 1.1)

# 領域のラベル
ax.annotate('No scaling\n(high freq)', xy=(3, 0.05), fontsize=10,
            color='blue', fontweight='bold')
ax.annotate('Full PI scaling\n(low freq)', xy=(45, 0.85), fontsize=10,
            color='red', fontweight='bold')
ax.annotate('Smooth\ntransition', xy=(20, 0.45), fontsize=10,
            color='green', fontweight='bold')
ax.grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig('yarn_gamma_visualization.png', dpi=150, bbox_inches='tight')
plt.show()

左のグラフから、RoPEの波長が次元インデックスに対して指数的に増加する様子がわかります。赤と青の破線が $\lambda_{\text{low}}$ と $\lambda_{\text{high}}$ の閾値を示しており、緑の領域がYaRNの「遷移ゾーン」です。波長がこの領域にある次元ペアは、高周波側の「そのまま保存」から低周波側の「完全にPI」へと滑らかに遷移します。

右のグラフは、$\gamma_i$ の値が次元インデックスに対してどう変化するかを示しています。低次元(高周波)では $\gamma = 0$(スケーリングなし)、高次元(低周波)では $\gamma = 1$(完全にPIスケーリング)、中間は線形ランプで接続されています。この区分的な設計が、YaRNの強みです。

Dynamic NTK Scalingのbase変化の可視化

Dynamic NTK Scalingでは、系列長に応じてbase値が変化します。この変化を可視化しましょう。

import numpy as np
import matplotlib.pyplot as plt

# パラメータ
d = 128
base = 10000
L_train = 4096

seq_lengths = np.arange(1, 20001)

# Dynamic NTK: 系列長に応じてbaseを計算
dynamic_base = np.where(
    seq_lengths <= L_train,
    base,
    base * (seq_lengths / L_train) ** (d / (d - 2))
)

fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# 左: base値の変化
ax = axes[0]
ax.plot(seq_lengths, dynamic_base, 'purple', linewidth=2)
ax.axvline(x=L_train, color='orange', linewidth=1.5, linestyle='--',
           label=f'$L_{{train}}={L_train}$')
ax.axhline(y=base, color='gray', linewidth=1, linestyle=':', alpha=0.7,
           label=f'Original base={base}')
ax.set_xlabel('Current sequence length', fontsize=12)
ax.set_ylabel('Dynamic base value', fontsize=12)
ax.set_title('Dynamic NTK: Base Value vs Sequence Length', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# 右: いくつかの次元の有効周波数
ax = axes[1]
dim_examples = [0, 10, 30, 63]
colors = ['red', 'blue', 'green', 'purple']

for dim_i, color in zip(dim_examples, colors):
    theta_dynamic = dynamic_base ** (-2 * dim_i / d)
    theta_static = base ** (-2 * dim_i / d)
    ax.plot(seq_lengths, theta_dynamic / theta_static, color=color,
            linewidth=1.5, label=f'$i={dim_i}$')

ax.axvline(x=L_train, color='orange', linewidth=1.5, linestyle='--', alpha=0.7)
ax.set_xlabel('Current sequence length', fontsize=12)
ax.set_ylabel(r'$\theta_i^{dynamic} / \theta_i^{original}$', fontsize=12)
ax.set_title('Dynamic NTK: Frequency Ratio per Dimension', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)
ax.set_ylim(0, 1.1)

plt.tight_layout()
plt.savefig('dynamic_ntk_scaling.png', dpi=150, bbox_inches='tight')
plt.show()

左のグラフから、系列長が $L_{\text{train}} = 4096$ を超えた時点でbase値が急速に増加し始めることがわかります。それ以前はbase = 10000で固定されており、元のRoPEと全く同じ挙動です。この「必要なときだけスケーリングする」性質が、Dynamic NTK Scalingの最大の利点です。

右のグラフは、各次元ペアの周波数が動的にどう変化するかを示しています。$i = 0$(最高周波数)はbase値が変わってもほとんど影響を受けず、比率が1に近いままです。一方、$i = 63$(最低周波数)は系列長が増えるにつれて急速に圧縮されます。これはNTK-Aware Scalingの「高周波保護」の性質が動的にも維持されていることを示しています。

相対位置のアテンションスコアへの影響

各手法のアテンションカーネルの比較

各スケーリング手法が相対位置のアテンションスコアにどう影響するかを可視化します。ランダムなQuery/Keyベクトルに各手法を適用し、相対位置に対するアテンションスコアの変化を観察しましょう。

import numpy as np
import matplotlib.pyplot as plt

def rope_attention_score(q, k, relative_pos, thetas):
    """RoPEを適用した内積を計算"""
    d_half = len(thetas)
    score = 0.0
    for i in range(d_half):
        angle = relative_pos * thetas[i]
        cos_a, sin_a = np.cos(angle), np.sin(angle)
        # 2次元回転を適用した内積
        q_2i, q_2i1 = q[2*i], q[2*i+1]
        k_2i, k_2i1 = k[2*i], k[2*i+1]
        score += q_2i * (k_2i * cos_a - k_2i1 * sin_a)
        score += q_2i1 * (k_2i * sin_a + k_2i1 * cos_a)
    return score

# パラメータ
d = 64
base = 10000
L_train = 4096
s = 4  # 4倍に拡張

np.random.seed(42)
q = np.random.randn(d) / np.sqrt(d)
k = np.random.randn(d) / np.sqrt(d)

dim_indices = np.arange(d // 2)

# 各手法のtheta
theta_orig = base ** (-2 * dim_indices / d)
theta_pi = theta_orig / s
base_ntk = base * s ** (d / (d - 2))
theta_ntk = base_ntk ** (-2 * dim_indices / d)

# YaRN
alpha_low, alpha_high = 1, 32
lam_low = 2 * np.pi * L_train / alpha_low
lam_high = 2 * np.pi * L_train / alpha_high
wavelengths = 2 * np.pi / theta_orig
gamma = np.clip((wavelengths - lam_high) / (lam_low - lam_high), 0, 1)
theta_yarn = theta_orig / (1 + gamma * (s - 1))

# 相対位置に対するスコアを計算
rel_positions = np.arange(0, L_train * s + 1, 4)

scores_orig = [rope_attention_score(q, k, rp, theta_orig) for rp in rel_positions]
scores_pi = [rope_attention_score(q, k, rp, theta_pi) for rp in rel_positions]
scores_ntk = [rope_attention_score(q, k, rp, theta_ntk) for rp in rel_positions]
scores_yarn = [rope_attention_score(q, k, rp, theta_yarn) for rp in rel_positions]

fig, axes = plt.subplots(2, 1, figsize=(14, 10))

# 上: 全範囲
ax = axes[0]
ax.plot(rel_positions, scores_orig, 'k-', linewidth=0.8, alpha=0.7, label='Original RoPE')
ax.plot(rel_positions, scores_pi, 'b-', linewidth=0.8, alpha=0.7, label='PI')
ax.plot(rel_positions, scores_ntk, 'r-', linewidth=0.8, alpha=0.7, label='NTK-Aware')
ax.plot(rel_positions, scores_yarn, 'g-', linewidth=0.8, alpha=0.7, label='YaRN')
ax.axvline(x=L_train, color='orange', linewidth=2, linestyle='--', alpha=0.7,
           label=f'$L_{{train}}={L_train}$')
ax.set_xlabel('Relative position $|m - n|$', fontsize=12)
ax.set_ylabel('Attention score (dot product)', fontsize=12)
ax.set_title('Attention Score vs Relative Position (Full Range)', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# 下: 近距離のズーム
ax = axes[1]
zoom_range = rel_positions < 200
ax.plot(rel_positions[zoom_range], np.array(scores_orig)[zoom_range],
        'k-', linewidth=1.5, label='Original RoPE')
ax.plot(rel_positions[zoom_range], np.array(scores_pi)[zoom_range],
        'b--', linewidth=1.5, label='PI')
ax.plot(rel_positions[zoom_range], np.array(scores_ntk)[zoom_range],
        'r-.', linewidth=1.5, label='NTK-Aware')
ax.plot(rel_positions[zoom_range], np.array(scores_yarn)[zoom_range],
        'g:', linewidth=2, label='YaRN')
ax.set_xlabel('Relative position $|m - n|$', fontsize=12)
ax.set_ylabel('Attention score (dot product)', fontsize=12)
ax.set_title('Attention Score vs Relative Position (Near Range)', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig('attention_score_comparison.png', dpi=150, bbox_inches='tight')
plt.show()

上のグラフから、以下の重要な特徴が読み取れます。

  1. 元のRoPEは $L_{\text{train}}$ を超えた領域でもスコアの振動パターンが続くが、このパターンは学習時に見たことのないものであるため、Attentionの重み付けが不正確になります。振動のエンベロープや位相が予測不能な状態です
  2. PIはスコアの振動を $s$ 倍に引き伸ばしており、拡張後の全範囲で学習時と同じ角度範囲に収めています。ただし、近距離のスコア変化が元のRoPEと異なり、短距離の位置分解能が低下しています
  3. NTK-AwareとYaRNは、近距離では元のRoPEと類似のパターンを保ちつつ、遠距離でのスコアの振る舞いを調整しています

下のズームしたグラフからは、近距離でのスコア変化の違いがよく見えます。NTK-AwareとYaRNが元のRoPEに近い振る舞いを示しているのに対し、PIは明らかに異なるパターンを示しています。これが、NTK-AwareやYaRNがファインチューニングなしでも一定の品質を維持できる理由です。

Perplexity劣化のシミュレーション

簡易的なPerplexityの傾向予測

実際のPerplexity測定にはGPUを用いたモデル推論が必要ですが、ここでは各手法がRoPEの角度分布をどの程度「学習時の分布内」に保つかを定量化し、Perplexity劣化の傾向を予測します。

各次元ペアについて、位置 $m$ での回転角 $m\theta_i$ が学習時の最大回転角 $L_{\text{train}} \cdot \theta_i^{\text{orig}}$ をどれだけ超えるかを「外挿度」として測定します。

import numpy as np
import matplotlib.pyplot as plt

def compute_extrapolation_ratio(thetas_scaled, theta_original, L_train, positions):
    """各位置・次元での外挿度を計算"""
    # (positions, dims) の2D配列
    angles_scaled = np.outer(positions, thetas_scaled)     # 各手法の角度
    max_angles_train = L_train * theta_original             # 学習時の最大角度

    # 外挿度: 学習時の最大角度を超える割合
    ratios = angles_scaled / max_angles_train[np.newaxis, :]
    # 1を超える次元の割合(位置ごと)
    extrapolation_fraction = np.mean(ratios > 1.0, axis=1)
    # 超過量の平均
    excess = np.mean(np.maximum(ratios - 1.0, 0), axis=1)
    return extrapolation_fraction, excess

# パラメータ
d = 128
base = 10000
L_train = 4096
s_values = [2, 4, 8, 16]

dim_indices = np.arange(d // 2)
theta_orig = base ** (-2 * dim_indices / d)

fig, axes = plt.subplots(2, 2, figsize=(14, 10))

for ax_idx, s in enumerate(s_values):
    ax = axes[ax_idx // 2, ax_idx % 2]
    L_target = L_train * s
    positions = np.arange(1, L_target + 1, max(1, L_target // 2000))

    # PI
    theta_pi = theta_orig / s
    _, excess_pi = compute_extrapolation_ratio(theta_pi, theta_orig, L_train, positions)

    # NTK-Aware
    base_ntk = base * s ** (d / (d - 2))
    theta_ntk = base_ntk ** (-2 * dim_indices / d)
    _, excess_ntk = compute_extrapolation_ratio(theta_ntk, theta_orig, L_train, positions)

    # YaRN
    alpha_low, alpha_high = 1, 32
    lam_low = 2 * np.pi * L_train / alpha_low
    lam_high = 2 * np.pi * L_train / alpha_high
    wavelengths = 2 * np.pi / theta_orig
    gamma = np.clip((wavelengths - lam_high) / (lam_low - lam_high), 0, 1)
    theta_yarn = theta_orig / (1 + gamma * (s - 1))
    _, excess_yarn = compute_extrapolation_ratio(theta_yarn, theta_orig, L_train, positions)

    # 元のRoPE(外挿)
    _, excess_orig = compute_extrapolation_ratio(theta_orig, theta_orig, L_train, positions)

    ax.plot(positions, excess_orig, 'k-', linewidth=1.5, alpha=0.7, label='Original (extrapolation)')
    ax.plot(positions, excess_pi, 'b--', linewidth=1.5, label='PI')
    ax.plot(positions, excess_ntk, 'r-.', linewidth=1.5, label='NTK-Aware')
    ax.plot(positions, excess_yarn, 'g:', linewidth=2, label='YaRN')
    ax.axvline(x=L_train, color='orange', linewidth=1.5, linestyle='--', alpha=0.5)
    ax.set_xlabel('Position', fontsize=11)
    ax.set_ylabel('Mean angle excess ratio', fontsize=11)
    ax.set_title(f'Scale factor $s = {s}$ ($L\' = {L_target}$)', fontsize=12)
    ax.legend(fontsize=9)
    ax.grid(True, alpha=0.3)

plt.suptitle('Extrapolation Severity for Different Extension Methods', fontsize=14)
plt.tight_layout()
plt.savefig('extrapolation_severity.png', dpi=150, bbox_inches='tight')
plt.show()

4つのグラフから、スケール比 $s$ が大きくなるにつれて各手法の差が顕著になることがわかります。

  1. 元のRoPE(黒)は、$L_{\text{train}}$ を超えると外挿度が線形に増加します。$s = 16$ では外挿度が非常に大きくなり、モデルの出力が完全に崩壊することが予測されます
  2. PI(青)は、定義上、外挿度がゼロに保たれます。全ての位置が学習時の角度範囲内に収まるため、角度の外挿による劣化は発生しません。ただし、このメトリクスには位置分解能の低下は反映されていません
  3. NTK-Aware(赤)は、低周波成分で若干の外挿が残ります。高周波を保護した代償として、全体の圧縮率がPIほど均等でないためです
  4. YaRN(緑)は、PIと同様にほぼゼロの外挿度を維持しつつ、高周波を保護しています。低周波は完全にPIスケーリングされるため外挿が発生せず、高周波は元から外挿が問題にならないため、両方の利点を享受しています

これらの結果は、実際のPerplexity測定の傾向と一致します。Chen et al. (2023) の報告では、$s = 4$ のときPIを適用した LLaMA-7B は1000ステップのファインチューニング後にPerplexityが安定し、元の4096トークン内の品質もほぼ維持されました。YaRN (Peng et al., 2023) はさらに少ないファインチューニングステップで同等以上の性能を達成しています。

各手法の比較まとめ

ここまでの4手法を整理します。

手法の特性比較表

特性 PI NTK-Aware YaRN Dynamic NTK
スケーリング対象 位置 $m$ を $m/s$ base $b$ を $b’$ 次元ごとに $\gamma_i$ 系列長に応じてbase
高周波の保護 なし あり(指数的) あり(完全保護) あり(指数的)
ファインチューニング 必要(1000+ steps) 少量または不要 少量(400 steps) 不要
温度スケーリング なし なし あり なし
動的調整 不可 不可 不可
実装の簡潔さ 非常に簡単 簡単 やや複雑 中程度
報告されたPerplexity 良好(FT後) 良好(ゼロショット) 最良(FT後) 良好(ゼロショット)

実用的な選択指針

各手法はそれぞれ異なるユースケースに適しています。

Position Interpolation(PI)を選ぶとき

  • ファインチューニングのコストを許容できる場合
  • 拡張後のコンテキスト長が事前に決まっている場合
  • 実装のシンプルさを優先する場合

NTK-Aware Scalingを選ぶとき

  • ファインチューニングなし(ゼロショット)で拡張したい場合
  • 短い入力での品質劣化を最小限にしたい場合
  • 中程度の拡張比($s \leq 4$)で十分な場合

YaRNを選ぶとき

  • 最高品質のコンテキスト拡張が必要な場合
  • 少量のファインチューニングが可能な場合
  • 大きな拡張比($s = 8$ 以上)が必要な場合

Dynamic NTK Scalingを選ぶとき

  • 入力長が事前にわからない場合
  • ファインチューニングなしで多様な長さの入力を処理したい場合
  • 実装のシンプルさとゼロショット性能のバランスを取りたい場合

歴史的な発展

これらの手法は2023年半ばに急速に発展しました。時系列で整理すると

  1. Position Interpolation(Chen et al., 2023年6月): 「外挿ではなく補間」という根本的な発想転換
  2. NTK-Aware Scaling(bloc97, 2023年6月): Redditの投稿から始まった、baseスケーリングのアイデア
  3. Dynamic NTK Scaling(Emozilla, 2023年6月): NTK-Awareを動的に適用するアイデア
  4. YaRN(Peng et al., 2023年8月): PI、NTK-Aware、温度スケーリングを統合した包括的手法

わずか2か月の間に、コミュニティの集合知によってContext Window拡張技術が大きく進化したことは注目に値します。この急速な発展は、RoPEの数学的構造が明確であったこと、そしてオープンなLLM(LLaMAなど)の登場により誰もが実験できる環境が整っていたことが大きな要因です。

まとめ

本記事では、学習済みLLMのコンテキスト長を後から拡張するContext Window拡張技術について解説しました。

  • コンテキスト長の限界: RoPEの回転角が学習時の範囲を超えると、Attentionスコアが予測不能になりPerplexityが急上昇する。これは位置エンコーディングにおける分布外汎化の失敗
  • Position Interpolation(PI): 位置インデックスを $1/s$ に圧縮し、全位置を学習時の角度範囲に収める。シンプルだが高周波も不要に圧縮する
  • NTK-Aware Scaling: base値を大きくすることで、高周波を保護しつつ低周波を圧縮する。ファインチューニングなしでも動作する
  • YaRN: 波長に基づいて次元ペアを3グループに分類し、高周波は保存、低周波はPIスケーリング、中間は線形補間する統合手法。温度スケーリングも含む
  • Dynamic NTK Scaling: 入力長に応じてbase値を動的に調整し、必要なときだけスケーリングを適用する

これらの技術は全て、RoPEの「回転角 = 位置 $\times$ 周波数」という単純な構造を利用しています。周波数を変える(NTK-Aware)、位置を変える(PI)、あるいは両方を組み合わせる(YaRN)という、数学的に明快な操作でコンテキスト長の拡張を実現している点が美しいです。

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

画像なし
RoPE(回転位置埋め込み)の数学的導出と実装
RoPEの回転行列の導出と基本性質を詳しく解説します。
画像なし
ALiBi(Attention with Linear Biases)の理論と実装
RoPEとは異なるアプローチで外挿性を向上させるALiBiを解説します。
画像なし
Sparse Attention(Longformer・BigBird)の理論と実装
コンテキスト長の拡張とは別の方向から長系列処理を実現するSparse Attentionを解説します。
画像なし
KVキャッシュの仕組みと実装
コンテキスト拡張と密接に関わるKVキャッシュの最適化手法を解説します。