DeBERTaのDisentangled Attention — コンテンツと位置を分離した注意機構

「The bank is near the river」と「I went to the bank to deposit money」という2つの文を考えてみてください。どちらにも “bank” という単語が含まれていますが、意味はまったく異なります。私たち人間は、bankの「意味」を考える際に、周囲の単語の内容(river、deposit money)と、bankが文のどの位置にあるか(主語なのか目的語なのか)を自然に使い分けています。

ところが、BERTやRoBERTaでは、単語の意味ベクトル(コンテンツ)と位置情報を最初に足し算で合成し、その混合ベクトルをAttention計算に使います。一度足し合わせてしまうと、Attentionが「この単語の意味が重要だから注目した」のか「この位置が重要だから注目した」のかを区別できません。これは、絵の具を混ぜてしまうと元の色に分離できないのと同じです。

DeBERTa(Decoding-enhanced BERT with disentangled attention)は、He et al.(2021, Microsoft)が提案したモデルで、この「コンテンツと位置の混合問題」を正面から解決します。コンテンツと位置を別々のベクトルとして保持し、Attention計算では3つの独立した注意行列(Content-to-Content、Content-to-Position、Position-to-Content)を個別に計算します。さらに、Enhanced Mask Decoder(EMD)という仕組みで、MLM(Masked Language Model)の最終予測層にだけ絶対位置情報を注入します。

この設計の結果、DeBERTaはSuperGLUEベンチマークで人間のベースラインを初めて超えたモデルとなり、NLP分野に大きなインパクトを与えました。

DeBERTaの考え方を理解すると、以下の場面で応用が利きます。

  • 高精度な自然言語理解(NLU): テキスト分類、質問応答、自然言語推論(NLI)など、位置と内容の関係を正確に捉える必要があるタスクで、BERT/RoBERTaを超える精度を実現します
  • SuperGLUEベンチマークの攻略: 人間のベースラインを超える性能を持つDeBERTaの設計原理を理解することで、最先端のNLPモデルの設計思想を把握できます
  • Attention機構の改良研究: コンテンツと位置を分離するという発想は、DeBERTaに限らず、新しいAttention機構を設計する際の重要な設計指針となります

本記事の内容

  • BERTにおける位置情報の扱いとその限界
  • Disentangled Attentionの直感的理解と数学的定式化
  • Content-to-Content、Content-to-Position、Position-to-Contentの3つの注意行列
  • 相対位置バイアスの計算と距離クリッピング
  • Enhanced Mask Decoder(EMD)による絶対位置の注入
  • PyTorchでのDisentangled AttentionとEMDの実装
  • DeBERTa V2/V3の改良点
  • BERT、RoBERTa、DeBERTaの性能比較

前提知識

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

画像なし
BERTのアーキテクチャを解説
BERTのTransformer Encoderベースの構造、入力表現、事前学習タスクを解説します
Positional Encodingの理論
sin/cosによる位置符号化の数式と、なぜこの設計が有効なのかを解説します
画像なし
Multi-Head Attentionの仕組み
複数のAttentionヘッドで異なる部分空間の情報を捉える仕組みを解説します
画像なし
RoBERTaの改良点と性能向上
BERTの学習設定を最適化して性能を向上させたRoBERTaの設計判断を解説します

従来の位置情報の扱いとその限界

Transformerベースのモデルにおいて、位置情報をどう扱うかは長年の課題です。Self-Attentionは集合演算であり、入力トークンの順序を本質的には区別しません。したがって、位置情報を何らかの形で明示的に注入しなければ、「猫が犬を追いかける」と「犬が猫を追いかける」を区別できなくなります。この問題に対して、これまでにいくつかのアプローチが提案されてきました。

BERTの方法: 絶対位置の加算

BERTは、入力の各トークンに対して3つの埋め込みベクトルを足し合わせます。

$$ \bm{h}_i = \bm{e}_{\text{token}}(x_i) + \bm{e}_{\text{segment}}(s_i) + \bm{e}_{\text{position}}(i) $$

ここで $\bm{e}_{\text{token}}(x_i)$ はトークン $x_i$ の意味を表すトークン埋め込み、$\bm{e}_{\text{segment}}(s_i)$ は文Aか文Bかを示すセグメント埋め込み、$\bm{e}_{\text{position}}(i)$ は位置 $i$ に対応する学習可能な位置埋め込みです。

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

問題1: 情報の混合(Entanglement)。加算によってコンテンツと位置が1つのベクトルに混ざると、後段のAttention計算で両者を分離して扱うことが困難になります。たとえば、Attention重みが大きい理由が「単語の意味が関連しているから」なのか「位置が近いから」なのかを分解できません。

問題2: 絶対位置の汎化性。BERTは最大系列長512で学習されるため、$\bm{e}_{\text{position}}(0), \bm{e}_{\text{position}}(1), \dots, \bm{e}_{\text{position}}(511)$ の512個の位置ベクトルを学習します。しかし学習データに含まれない長さの入力に対する外挿(extrapolation)が苦手です。また、「2つのトークンの距離が3」という相対的な関係は、位置 $(1, 4)$ でも $(100, 103)$ でも同じはずですが、絶対位置埋め込みではこの等価性を直接表現できません。

Transformer-XLの方法: 相対位置バイアス

Shaw et al.(2018)やDai et al.(2019, Transformer-XL / XLNet)は、Attention計算に相対位置のバイアス項を導入しました。Transformer-XLでは、標準的なAttentionスコアを以下のように分解します。

$$ A_{ij} = \underbrace{\bm{x}_i \bm{W}_q^T \bm{W}_k \bm{x}_j^T}_{(a)} + \underbrace{\bm{x}_i \bm{W}_q^T \bm{W}_k \bm{R}_{i-j}^T}_{(b)} + \underbrace{\bm{u}^T \bm{W}_k \bm{x}_j^T}_{(c)} + \underbrace{\bm{v}^T \bm{W}_k \bm{R}_{i-j}^T}_{(d)} $$

項(a)はコンテンツ同士、項(b)はコンテンツから相対位置、項(c)は位置に依存しないグローバルなコンテンツバイアス、項(d)はグローバルな位置バイアスを表します。

この方法は相対位置を扱えるため外挿性能が向上しますが、コンテンツと位置を完全には分離していません。コンテンツの射影($\bm{W}_k$)が位置に関する項(b)(d)にも共有されているため、コンテンツ用のパラメータと位置用のパラメータが絡み合っています。

共通の限界

BERT方式もTransformer-XL方式も、コンテンツと位置の情報を同じパラメータ空間で処理するという点で共通しています。コンテンツベクトルと位置ベクトルが同じQuery/Key行列を通ることで、モデルは「コンテンツに注目すべきか、位置に注目すべきか」のトレードオフを暗黙的に解決しなければなりません。

ここで自然に浮かぶ疑問があります。コンテンツと位置を最初から別々のベクトルとして保持し、Attention計算でも独立に扱えばどうなるでしょうか。これがDeBERTaのDisentangled Attentionの核心的なアイデアです。

Disentangled Attentionの直感

DeBERTaのDisentangled Attention(分離注意機構)の発想を、日常的なアナロジーで理解しましょう。

図書館での本探しのアナロジー

あなたが図書館で参考になる本を探しているとします。本を選ぶ基準は大きく2つあります。

  1. 内容の関連性: 自分が探しているトピックと本の内容が合致しているか
  2. 棚の位置: 自分が今いる場所から近い棚にあるか、同じセクションにあるか

BERTのやり方は、「この本は内容スコア7点+位置スコア3点=合計10点」のように、内容と位置の評価を最初から足し合わせてしまいます。合計点だけを見て本を選ぶので、「内容が素晴らしいから選んだのか、たまたま近くにあったから選んだのか」がわからなくなります。

DeBERTaのやり方は、内容の評価と位置の評価を別々のチャネルで管理します。「この本は内容として非常に関連がある」「しかも位置的にも近い」「だから合わせて注目度が高い」と、それぞれの理由を分離したまま最終的なスコアを出します。

3つの注意の種類

DeBERTaは、トークン $i$ がトークン $j$ にどれだけ注目するかを計算する際、以下の3つの観点を独立に評価します。

Content-to-Content(C2C): トークン $i$ の意味と、トークン $j$ の意味がどれだけ関連するか。これは標準的なSelf-Attentionと同じ計算です。「bank」と「river」は意味的に関連が深いので、C2Cスコアが高くなります。

Content-to-Position(C2P): トークン $i$ の意味が、トークン $j$ の相対位置にどれだけ注目すべきか。たとえば、動詞が「すぐ後ろ(相対位置+1)にある単語」に強く注目するパターンを学習できます。動詞 “ate” は、直後に来やすい目的語の位置に注目します。

Position-to-Content(P2C): トークン $i$ の相対位置が、トークン $j$ の意味にどれだけ注目すべきか。たとえば、「文頭から3番目の位置」が名詞に注目しやすいといった、構文的なパターンを学習できます。

注意すべき点として、DeBERTaでは Position-to-Position(P2P) の項は含めません。He et al. は、P2Pは相対位置同士の関連であり、全ての入力で同じ値になるため有用な情報を提供しないと判断しました。実験でもP2Pを加えることによる性能向上は確認されていません。

この3つの注意行列を独立に計算し、最後に足し合わせることで、モデルは「なぜこのトークンに注目したのか」を内容と位置のレベルで分解できるようになります。では、この直感をどのように数学的に定式化するのかを次に見ていきましょう。

Disentangled Attentionの数学的定式化

2つのベクトル系列の定義

DeBERTaでは、各トークンに対して2つの独立したベクトルを保持します。

位置 $i$ のトークンに対して、コンテンツベクトル $\bm{H}_i$ はトークン埋め込みから得られる意味表現です。

$$ \bm{H}_i = \text{TokenEmbedding}(x_i) $$

一方、位置ベクトル $\bm{P}_{i|j}$ は、トークン $i$ と $j$ の相対位置 $\delta(i, j) = i – j$ を符号化した表現です。BERTのように絶対位置ベクトルを使うのではなく、相対位置ベクトルを使います。

$$ \bm{P}_{i|j} = \bm{R}_{\delta(i, j)}, \quad \delta(i, j) = \text{clip}(i – j, -k, k) $$

ここで $\bm{R}_{\delta}$ は相対位置 $\delta$ に対応する学習可能な埋め込みベクトル、$k$ は最大相対距離のクリッピングパラメータです。クリッピングの意味は後述します。

重要なのは、$\bm{H}_i$ と $\bm{P}_{i|j}$ が加算されない点です。BERTではこれらを足し合わせていましたが、DeBERTaでは最後まで別々のベクトルとして扱います。

Query・Key・Valueの分離射影

標準的なAttentionでは、1つの混合ベクトル $\bm{h}_i$ からQuery、Key、Valueを生成します。DeBERTaでは、コンテンツと位置のそれぞれに対して専用の射影行列を用意します。

コンテンツベクトル $\bm{H}_i$ に対するQuery、Key、Valueの射影は以下のとおりです。

$$ \bm{q}_i^c = \bm{H}_i \bm{W}_q^c, \quad \bm{k}_j^c = \bm{H}_j \bm{W}_k^c, \quad \bm{v}_j^c = \bm{H}_j \bm{W}_v^c $$

位置ベクトル $\bm{P}_{i|j}$ に対するQueryとKeyの射影は以下のとおりです。

$$ \bm{q}_i^p = \bm{P}_{i|j} \bm{W}_q^p, \quad \bm{k}_j^p = \bm{P}_{i|j} \bm{W}_k^p $$

ここで $\bm{W}_q^c, \bm{W}_k^c, \bm{W}_v^c$ はコンテンツ用の射影行列、$\bm{W}_q^p, \bm{W}_k^p$ は位置用の射影行列です。位置にはValueの射影がない点に注意してください。Attentionの出力値はコンテンツのValueのみから計算されます。位置は「どこに注目するか」の重みの計算にのみ関与します。

3つの注意行列の計算

Disentangled Attentionスコア $\tilde{A}_{ij}$ は、3つの項の和として定義されます。

$$ \tilde{A}_{ij} = \underbrace{(\bm{H}_i \bm{W}_q^c)(\bm{H}_j \bm{W}_k^c)^T}_{\text{Content-to-Content}} + \underbrace{(\bm{H}_i \bm{W}_q^c)(\bm{P}_{i|j} \bm{W}_k^p)^T}_{\text{Content-to-Position}} + \underbrace{(\bm{P}_{j|i} \bm{W}_q^p)(\bm{H}_j \bm{W}_k^c)^T}_{\text{Position-to-Content}} $$

各項が何を計算しているかを整理しましょう。

第1項 C2C(Content-to-Content)は、トークン $i$ のコンテンツQueryとトークン $j$ のコンテンツKeyの内積です。標準的なSelf-Attentionと同じ計算であり、2つのトークンの意味的な関連度を測ります。

$$ A_{ij}^{c2c} = (\bm{H}_i \bm{W}_q^c)(\bm{H}_j \bm{W}_k^c)^T $$

第2項 C2P(Content-to-Position)は、トークン $i$ のコンテンツQueryと、相対位置の位置Keyの内積です。トークン $i$ の意味が、トークン $j$ との相対的な距離にどれだけ注目するかを測ります。

$$ A_{ij}^{c2p} = (\bm{H}_i \bm{W}_q^c)(\bm{P}_{i|j} \bm{W}_k^p)^T $$

たとえば、前置詞 “in” は直後(相対位置+1)にくる名詞に強く注目する傾向があり、この項がそのパターンを捉えます。

第3項 P2C(Position-to-Content)は、相対位置の位置Queryとトークン $j$ のコンテンツKeyの内積です。トークン $i$ からの相対的な距離が、トークン $j$ の意味にどれだけ注目すべきかを測ります。

$$ A_{ij}^{p2c} = (\bm{P}_{j|i} \bm{W}_q^p)(\bm{H}_j \bm{W}_k^c)^T $$

P2Cは、たとえば「文頭から2つ目の位置にある単語が主語として名詞に注目しやすい」という構文的なパターンを学習します。

3つのスコアを足し合わせてスケーリングを適用し、softmaxで正規化すると最終的なAttention重みが得られます。

$$ A_{ij} = \frac{\tilde{A}_{ij}}{\sqrt{3d_h}} $$

$$ \alpha_{ij} = \text{softmax}_j(A_{ij}) $$

分母の $\sqrt{3d_h}$ は、3つの項を足し合わせることでスコアの分散が約3倍になることを補正するためのスケーリングファクタです。$d_h$ はヘッドあたりの次元数です。標準Attentionでは $\sqrt{d_h}$ で割りますが、DeBERTaでは3つの項を足すので $\sqrt{3d_h}$ で割ります。

Attentionの出力

Attention重み $\alpha_{ij}$ はコンテンツのValueに対してのみ適用されます。

$$ \bm{o}_i = \sum_j \alpha_{ij} \bm{H}_j \bm{W}_v^c $$

位置情報は「どこに注目するか」の重み計算にのみ関与し、「何を出力するか」にはコンテンツのValueだけが使われます。これは直感にも合っています。Attentionの出力は次の層への入力となる意味表現であり、そこに位置の値を混ぜるべきではないからです。

Disentangled Attentionの数学的な構造がわかったところで、相対位置のバイアスをどのように計算するかをもう少し詳しく見てみましょう。

相対位置バイアスの計算

相対位置 $\delta(i, j)$ の定義

DeBERTaでは、2つのトークンの位置関係を相対位置 $\delta(i, j) = i – j$ で表します。たとえば、位置3のトークンから見た位置5のトークンの相対位置は $\delta(3, 5) = 3 – 5 = -2$ です。負の値は「自分より後ろにある」ことを意味し、正の値は「自分より前にある」ことを意味します。

系列長 $n$ のとき、相対位置の取り得る範囲は $-(n-1)$ から $+(n-1)$ です。しかし、$n = 512$ のとき相対位置は $-511$ から $+511$ までの1023通りになり、これだけの位置ベクトルを学習するのはパラメータ効率の面で無駄です。実際には、500トークン離れた位置関係の精密な区別はほとんど必要ありません。

距離クリッピング

そこでDeBERTaは、相対位置をある最大距離 $k$ でクリッピングします。

$$ \delta'(i, j) = \text{clip}(\delta(i, j), -k, k) = \max(-k, \min(k, i – j)) $$

デフォルトでは $k = 512$ が使われます。クリッピングにより、相対位置ベクトルの種類は $2k + 1$ 個に限定されます。$k = 512$ の場合は $-512, -511, \dots, 0, \dots, 511, 512$ の1025個です。

クリッピングの直感的な意味は、「ある程度以上離れたトークン同士は、正確な距離よりも『遠い』という情報だけで十分」ということです。自然言語では、直近の数単語の位置関係は重要ですが(主語と動詞の距離、冠詞と名詞の距離など)、100単語以上離れたトークン同士の正確な距離は構文解析にほとんど影響しません。

相対位置埋め込みの参照

クリッピングされた相対位置 $\delta’$ から位置ベクトルを取得するには、ルックアップテーブル(埋め込みテーブル)を使います。

$$ \bm{R}_{\delta’} = \text{Embedding}(\delta’ + k) $$

$\delta’ + k$ はインデックスを非負にするためのシフトです。$\delta’ = -k$ のとき $\delta’ + k = 0$、$\delta’ = 0$ のとき $\delta’ + k = k$、$\delta’ = k$ のとき $\delta’ + k = 2k$ となり、$0$ から $2k$ までのインデックスで $2k + 1$ 個のベクトルを参照できます。

この仕組みにより、位置情報は相対距離に基づいたコンパクトなテーブルで表現されます。絶対位置を使うBERTと比べて、同じ相対距離にある全てのトークン対が同じ位置バイアスを共有するため、位置パターンの汎化性能が向上します。

なぜ相対位置がコンテンツ分離と相性が良いのか

ここで重要な点を強調しておきます。相対位置とDisentangled Attentionの組み合わせが特に効果的である理由です。

絶対位置を使うBERTでは、位置情報がトークン埋め込みに加算されるため、Attentionスコアの中に「絶対位置同士の内積」という項が暗黙的に含まれます。これは位置5と位置10のペアと、位置100と位置105のペアで異なるバイアスを生みます(同じ相対距離5なのに)。

一方、DeBERTaでは相対位置ベクトルを使い、かつそれをコンテンツとは完全に分離しているため、「相対距離が同じなら同じ位置バイアス」という自然な性質が保証されます。コンテンツと位置を分離したからこそ、相対位置の利点を最大限に活かせるのです。

しかし、相対位置だけでは捉えきれない情報もあります。それが「絶対位置」です。次のセクションでは、DeBERTaがこの問題にどう対処するかを見ていきます。

Enhanced Mask Decoder (EMD)

なぜ絶対位置が必要なのか

Disentangled Attentionは相対位置のみを扱いますが、自然言語には絶対位置が本質的に重要な場面があります。その最たる例がMLM(Masked Language Model)の予測です。

次の例を考えてみてください。

  • 入力: “[MASK] went to the store”
  • 入力: “The store [MASK] near my house”

1つ目の文では [MASK] は文頭にあるので主語(人名や代名詞)を予測すべきです。2つ目では [MASK] は「The store」の直後にあるのでbe動詞(”is” や “was”)を予測すべきです。

相対位置だけでは、[MASK] が「文頭にある」のか「文中にある」のかを区別できません。相対位置は2つのトークン間の距離を表すものであり、文全体の中での絶対的な位置を表すものではないからです。He et al. はこの問題を「絶対位置の欠落」と呼び、MLMの予測精度に影響することを実験的に確認しました。

EMDの設計思想

ここでDeBERTaが採る戦略は非常に巧妙です。Transformerの全層で絶対位置を使うのではなく、MLMの予測層(デコーダ層)でのみ絶対位置を注入します。

この設計には明確な理由があります。

  • Encoder層(1層目〜L-1層目): 文脈表現を構築する段階では、相対位置だけで十分です。「この単語とあの単語の距離は3」という情報でAttentionパターンを学習できます。ここで絶対位置を混ぜると、Disentangled Attentionの利点(コンテンツと位置の分離)が失われてしまいます。
  • Decoder層(最終層、MLM予測用): マスクされたトークンを具体的に予測する段階では、「文のどの位置にあるか」が手がかりになります。文頭なら主語、文末なら句読点や助動詞が来やすいなど、絶対位置に依存した統計的パターンを活用すべきです。

つまり、EMDは「文脈の構築には相対位置を、最終予測には絶対位置を」という二段階の戦略を実現しています。

EMDの計算

EMDの具体的な計算は以下のとおりです。Transformerの最終層の出力を $\bm{H}^L$ とし、絶対位置埋め込みを $\bm{I}$ とします。

まず、コンテンツ表現に絶対位置を統合した新しい入力を作ります。

$$ \bm{H}^0_{\text{emd}} = \bm{H}^L + \bm{I} $$

ここで $\bm{I} = [\bm{e}_{\text{pos}}(0), \bm{e}_{\text{pos}}(1), \dots, \bm{e}_{\text{pos}}(n-1)]$ は、BERTと同様の学習可能な絶対位置埋め込みです。$\bm{H}^L$ は $L$ 層のTransformerで十分にリッチな文脈表現が構築された後のベクトルであり、ここに絶対位置を加算します。

次に、この統合された表現をさらに1層(または少数の層)のTransformerブロックに通して、最終的なMLM予測を行います。

$$ \bm{H}^1_{\text{emd}} = \text{TransformerLayer}(\bm{H}^0_{\text{emd}}) $$

$$ P(x_i \mid \bm{x}) = \text{softmax}(\bm{H}^1_{\text{emd}, i} \bm{W}_{\text{vocab}} + \bm{b}_{\text{vocab}}) $$

このTransformerLayerには、通常のSelf-Attention(絶対位置が統合済みなので標準的なAttentionで構いません)を使用することも、Disentangled Attentionを使用することもできます。論文ではEMD層でもDisentangled Attentionを使い、そこにさらに絶対位置の効果を加えることで最良の結果を得ています。

EMDのアーキテクチャ上の位置づけ

DeBERTa全体のアーキテクチャをまとめると、以下のようになります。

  1. 入力層: トークン埋め込みのみ(絶対位置は加算しない)。相対位置ベクトルはAttention計算時に動的に参照
  2. Encoder($L$ 層): Disentangled Attention + FFN の繰り返し。コンテンツと相対位置を分離して処理
  3. Enhanced Mask Decoder(1〜2層): Encoderの出力に絶対位置を加算し、追加のTransformer層でMLM予測を生成

このアーキテクチャは、ファインチューニング時にも有効です。分類タスクでは [CLS] トークンの最終表現を使いますが、EMDで絶対位置が注入された後の表現を使うことで、文中での位置関係もエンコードされた、より豊かな表現を利用できます。

理論の全体像が見えたところで、PyTorchでDisentangled AttentionとEMDを実装してみましょう。

PyTorch実装: Disentangled Attention

まず、Disentangled Attentionの核心部分を実装します。コンテンツベクトルと位置ベクトルを分離して3つの注意行列を計算するロジックを、ステップごとに構築していきます。

相対位置インデックスの生成

最初に、相対位置のインデックス行列を生成する関数を実装します。

import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import matplotlib.pyplot as plt


def build_relative_position(seq_len, max_relative_distance):
    """相対位置インデックス行列を生成する

    Args:
        seq_len: 系列長
        max_relative_distance: 最大相対距離 k

    Returns:
        relative_pos: (seq_len, seq_len) の相対位置インデックス行列
                      値は [0, 2*k] の範囲(-k を 0 にシフト済み)
    """
    # 位置インデックス [0, 1, ..., seq_len-1]
    positions = torch.arange(seq_len, dtype=torch.long)

    # 相対位置行列: positions[i] - positions[j]
    # (seq_len, 1) - (1, seq_len) -> (seq_len, seq_len)
    relative_pos = positions.unsqueeze(1) - positions.unsqueeze(0)

    # クリッピング: [-k, k] に制限
    relative_pos = torch.clamp(relative_pos, -max_relative_distance, max_relative_distance)

    # 非負にシフト: [-k, k] -> [0, 2k]
    relative_pos = relative_pos + max_relative_distance

    return relative_pos

この関数を使って、相対位置行列がどのような構造になるか確認してみましょう。

# 相対位置行列の可視化
seq_len = 8
k = 4  # 小さい k で構造を見やすくする
rel_pos = build_relative_position(seq_len, k)

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

# 生のインデックス(シフト前)
raw_rel = rel_pos - k  # シフトを戻す
im1 = axes[0].imshow(raw_rel.numpy(), cmap='RdBu_r', vmin=-k, vmax=k)
axes[0].set_title('Relative Position (raw)', fontsize=13, fontweight='bold')
axes[0].set_xlabel('Key position j')
axes[0].set_ylabel('Query position i')
for i in range(seq_len):
    for j in range(seq_len):
        axes[0].text(j, i, f'{raw_rel[i, j].item():+d}', ha='center', va='center', fontsize=9)
plt.colorbar(im1, ax=axes[0])

# クリッピング後のインデックス
im2 = axes[1].imshow(rel_pos.numpy(), cmap='viridis', vmin=0, vmax=2*k)
axes[1].set_title(f'Relative Position Index (clipped, k={k})', fontsize=13, fontweight='bold')
axes[1].set_xlabel('Key position j')
axes[1].set_ylabel('Query position i')
for i in range(seq_len):
    for j in range(seq_len):
        axes[1].text(j, i, f'{rel_pos[i, j].item()}', ha='center', va='center', fontsize=9,
                     color='white' if rel_pos[i, j].item() < k else 'black')
plt.colorbar(im2, ax=axes[1])

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

左のヒートマップは生の相対位置 $\delta(i, j) = i – j$ を示しており、対角線上が0(自分自身との距離)、右上が負(jがiより後ろ)、左下が正(jがiより前)になっています。右のヒートマップはクリッピングとシフト後のインデックスを示しており、$k = 4$ で値が0〜8の範囲に収まっていることが確認できます。対角線から離れた位置では、クリッピングにより同じインデックス値が並んでいます。これが「遠く離れたトークン同士は同じ位置バイアスを共有する」というクリッピングの効果です。

Disentangled Attentionモジュールの実装

次に、3つの注意行列を計算するDisentangled Attentionモジュールを実装します。

class DisentangledAttention(nn.Module):
    """DeBERTaのDisentangled Attention

    コンテンツと位置を分離した3つの注意行列(C2C, C2P, P2C)を計算する。
    """

    def __init__(self, d_model, n_heads, max_relative_distance=512, dropout=0.1):
        super().__init__()
        assert d_model % n_heads == 0

        self.d_model = d_model
        self.n_heads = n_heads
        self.d_head = d_model // n_heads
        self.max_relative_distance = max_relative_distance

        # コンテンツ用の射影行列
        self.W_q_c = nn.Linear(d_model, d_model)  # Content Query
        self.W_k_c = nn.Linear(d_model, d_model)  # Content Key
        self.W_v_c = nn.Linear(d_model, d_model)  # Content Value

        # 位置用の射影行列(Valueは不要)
        self.W_q_p = nn.Linear(d_model, d_model)  # Position Query
        self.W_k_p = nn.Linear(d_model, d_model)  # Position Key

        # 相対位置埋め込みテーブル: 2k+1 個のベクトル
        self.rel_embeddings = nn.Embedding(
            2 * max_relative_distance + 1, d_model
        )

        # 出力射影
        self.W_o = nn.Linear(d_model, d_model)
        self.dropout = nn.Dropout(dropout)

        # スケーリングファクタ: 3つの項を足すので sqrt(3 * d_head)
        self.scale = (3 * self.d_head) ** 0.5

    def _split_heads(self, x):
        """(batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_head)"""
        batch, seq_len, _ = x.shape
        x = x.view(batch, seq_len, self.n_heads, self.d_head)
        return x.transpose(1, 2)

    def _merge_heads(self, x):
        """(batch, n_heads, seq_len, d_head) -> (batch, seq_len, d_model)"""
        batch, _, seq_len, _ = x.shape
        x = x.transpose(1, 2).contiguous()
        return x.view(batch, seq_len, self.d_model)

    def forward(self, H, attention_mask=None):
        """
        Args:
            H: コンテンツ表現 (batch, seq_len, d_model)
            attention_mask: (batch, 1, 1, seq_len) or None

        Returns:
            output: (batch, seq_len, d_model)
            attention_weights: (batch, n_heads, seq_len, seq_len)
        """
        batch, seq_len, _ = H.shape

        # --- コンテンツの Query, Key, Value を計算 ---
        q_c = self._split_heads(self.W_q_c(H))  # (B, H, S, D)
        k_c = self._split_heads(self.W_k_c(H))  # (B, H, S, D)
        v_c = self._split_heads(self.W_v_c(H))  # (B, H, S, D)

        # --- 相対位置埋め込みの取得 ---
        rel_pos_idx = build_relative_position(seq_len, self.max_relative_distance)
        rel_pos_idx = rel_pos_idx.to(H.device)  # (S, S)
        rel_emb = self.rel_embeddings(rel_pos_idx)  # (S, S, d_model)

        # --- 1. Content-to-Content (C2C) ---
        # 標準的な Attention: q_c @ k_c^T
        A_c2c = torch.matmul(q_c, k_c.transpose(-1, -2))  # (B, H, S, S)

        # --- 2. Content-to-Position (C2P) ---
        # q_c[i] @ (rel_emb[i,j] @ W_k_p)^T
        k_p = self.W_k_p(rel_emb)  # (S, S, d_model)
        # ヘッド分割: (S, S, n_heads, d_head) -> (n_heads, S, S, d_head)
        k_p = k_p.view(seq_len, seq_len, self.n_heads, self.d_head).permute(2, 0, 1, 3)
        # q_c: (B, H, S, D), k_p: (H, S, S, D)
        # A_c2p[b, h, i, j] = q_c[b, h, i, :] @ k_p[h, i, j, :]
        A_c2p = torch.einsum('bhid,hijd->bhij', q_c, k_p)  # (B, H, S, S)

        # --- 3. Position-to-Content (P2C) ---
        # (rel_emb[j,i] @ W_q_p) @ k_c[j]^T
        q_p = self.W_q_p(rel_emb)  # (S, S, d_model)
        q_p = q_p.view(seq_len, seq_len, self.n_heads, self.d_head).permute(2, 0, 1, 3)
        # q_p[h, j, i, :] は位置jからiへの相対位置Query
        # P2Cでは、位置 j|i の相対位置Queryと k_c[j] の内積を取る
        # A_p2c[b, h, i, j] = q_p[h, j, i, :] @ k_c[b, h, j, :]
        A_p2c = torch.einsum('hjid,bhjd->bhij', q_p, k_c)  # (B, H, S, S)

        # --- 3つの注意スコアを統合 ---
        A = (A_c2c + A_c2p + A_p2c) / self.scale

        # マスクの適用
        if attention_mask is not None:
            A = A + attention_mask  # mask は 0 or -inf

        # softmax で正規化
        attn_weights = F.softmax(A, dim=-1)
        attn_weights = self.dropout(attn_weights)

        # --- Attention 出力(Value はコンテンツのみ)---
        output = torch.matmul(attn_weights, v_c)  # (B, H, S, D)
        output = self._merge_heads(output)          # (B, S, d_model)
        output = self.W_o(output)                   # (B, S, d_model)

        return output, attn_weights

実装のポイントを整理します。

  • コンテンツと位置で別々の射影行列を持ちます。$\bm{W}_q^c, \bm{W}_k^c, \bm{W}_v^c$ がコンテンツ用、$\bm{W}_q^p, \bm{W}_k^p$ が位置用です
  • 相対位置埋め込みテーブルは $2k + 1$ 個のベクトルを保持します。これは全ヘッドで共有されます
  • C2Pの計算では、torch.einsum を使ってQuery位置 $i$ ごとに異なる位置Key $k_p[i, j]$ との内積を効率的に計算しています
  • P2Cの計算も同様に torch.einsum で処理しています
  • スケーリングは $\sqrt{3d_h}$ で行います(3つの項の和であるため)

動作確認と注意重みの可視化

実装が正しく動作するか確認し、3つの注意行列がどのような重みパターンを示すかを可視化してみましょう。

# パラメータ設定
d_model = 64
n_heads = 4
seq_len = 10
batch_size = 1
max_k = 8

# モデルの初期化
torch.manual_seed(42)
attn = DisentangledAttention(d_model, n_heads, max_relative_distance=max_k)

# ダミー入力
H = torch.randn(batch_size, seq_len, d_model)

# フォワードパス
output, weights = attn(H)

print(f"入力形状:   {H.shape}")
print(f"出力形状:   {output.shape}")
print(f"重み形状:   {weights.shape}")
print(f"重みの和:   {weights[0, 0].sum(dim=-1)}")  # 各行の和が1になることを確認

入力と出力の形状が一致し、Attention重みの各行の和が1.0になっていることから、softmaxの正規化が正しく機能していることが確認できます。これはDisentangled Attentionが標準的なAttentionと同じインターフェースを持ちながら、内部で3つの独立した注意行列を計算していることを示しています。

# 注意重みの可視化(ヘッドごと)
fig, axes = plt.subplots(1, 4, figsize=(18, 4))

for h in range(n_heads):
    w = weights[0, h].detach().numpy()
    im = axes[h].imshow(w, cmap='Blues', vmin=0)
    axes[h].set_title(f'Head {h}', fontsize=12, fontweight='bold')
    axes[h].set_xlabel('Key position')
    axes[h].set_ylabel('Query position')
    plt.colorbar(im, ax=axes[h], fraction=0.046)

plt.suptitle('Disentangled Attention Weights by Head', fontsize=14, fontweight='bold', y=1.02)
plt.tight_layout()
plt.savefig('disentangled_attention_heads.png', dpi=150, bbox_inches='tight')
plt.show()

4つのAttentionヘッドがそれぞれ異なるパターンを示していることが確認できます。ランダム初期化の段階では明確なパターンは見えませんが、学習後は各ヘッドが異なる言語的パターン(近接単語への注目、構文的な依存関係、長距離の意味的関連など)を学習することが期待されます。特にDisentangled Attentionでは、あるヘッドがC2P(コンテンツから位置への注目)を強く学習し、別のヘッドがP2C(位置からコンテンツへの注目)を重視するような分業が起こり得ます。

次に、3つの注意行列の寄与を個別に確認してみましょう。

# 3つの注意行列の個別計算と可視化
torch.manual_seed(42)
attn2 = DisentangledAttention(d_model, n_heads, max_relative_distance=max_k)

# 内部で各項を個別に取得するための簡易計算
with torch.no_grad():
    q_c = attn2._split_heads(attn2.W_q_c(H))
    k_c = attn2._split_heads(attn2.W_k_c(H))

    rel_pos_idx = build_relative_position(seq_len, max_k).to(H.device)
    rel_emb = attn2.rel_embeddings(rel_pos_idx)

    # C2C
    A_c2c = torch.matmul(q_c, k_c.transpose(-1, -2))

    # C2P
    k_p = attn2.W_k_p(rel_emb).view(seq_len, seq_len, n_heads, d_model // n_heads).permute(2, 0, 1, 3)
    A_c2p = torch.einsum('bhid,hijd->bhij', q_c, k_p)

    # P2C
    q_p = attn2.W_q_p(rel_emb).view(seq_len, seq_len, n_heads, d_model // n_heads).permute(2, 0, 1, 3)
    A_p2c = torch.einsum('hjid,bhjd->bhij', q_p, k_c)

# Head 0 の各項を可視化
fig, axes = plt.subplots(1, 3, figsize=(16, 4.5))
titles = ['Content-to-Content (C2C)', 'Content-to-Position (C2P)', 'Position-to-Content (P2C)']
matrices = [A_c2c, A_c2p, A_p2c]

for idx, (title, mat) in enumerate(zip(titles, matrices)):
    data = mat[0, 0].detach().numpy()
    im = axes[idx].imshow(data, cmap='RdBu_r', vmin=-np.abs(data).max(), vmax=np.abs(data).max())
    axes[idx].set_title(title, fontsize=12, fontweight='bold')
    axes[idx].set_xlabel('Key position j')
    axes[idx].set_ylabel('Query position i')
    plt.colorbar(im, ax=axes[idx], fraction=0.046)

plt.suptitle('Three Attention Components (Head 0, before softmax)', fontsize=14, fontweight='bold', y=1.02)
plt.tight_layout()
plt.savefig('three_attention_components.png', dpi=150, bbox_inches='tight')
plt.show()

3つの注意行列のヒートマップを見ると、それぞれが明確に異なる構造を持っていることがわかります。C2C(左)は入力トークンの意味的な関連を反映するため、ランダム初期化でも各トークン対ごとに異なる値を取ります。C2P(中央)は対角線に沿った帯状のパターンを示す傾向があり、これは相対位置が近いトークンほど高い注目を受けやすいことを反映しています。P2C(右)も同様に対角線に沿った構造を持ちますが、C2Pとは異なるパターンです。学習が進むと、C2Pは「動詞が後続する目的語の位置に注目する」ようなパターンを、P2Cは「主語の位置にある名詞に注目する」ようなパターンを獲得していきます。

PyTorch実装: Enhanced Mask Decoder

続いて、Enhanced Mask Decoderの概念実装を行います。EMDは、Encoder の出力に絶対位置を注入してMLM予測を行うモジュールです。

class EnhancedMaskDecoder(nn.Module):
    """DeBERTaのEnhanced Mask Decoder

    Encoder出力に絶対位置を注入し、追加のTransformer層で
    MLM予測を生成する。
    """

    def __init__(self, d_model, n_heads, max_seq_len=512,
                 max_relative_distance=512, vocab_size=30522, dropout=0.1):
        super().__init__()

        # 絶対位置埋め込み
        self.absolute_pos_embedding = nn.Embedding(max_seq_len, d_model)

        # EMD用のTransformer層(Disentangled Attention + FFN)
        self.emd_attention = DisentangledAttention(
            d_model, n_heads, max_relative_distance, dropout
        )
        self.emd_norm1 = nn.LayerNorm(d_model)
        self.emd_ffn = nn.Sequential(
            nn.Linear(d_model, d_model * 4),
            nn.GELU(),
            nn.Linear(d_model * 4, d_model),
            nn.Dropout(dropout)
        )
        self.emd_norm2 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)

        # MLM予測ヘッド
        self.mlm_head = nn.Sequential(
            nn.Linear(d_model, d_model),
            nn.GELU(),
            nn.LayerNorm(d_model),
            nn.Linear(d_model, vocab_size)
        )

    def forward(self, encoder_output, masked_positions=None):
        """
        Args:
            encoder_output: Encoder の最終出力 (batch, seq_len, d_model)
            masked_positions: マスクされた位置 (batch, n_masked) or None

        Returns:
            logits: MLM予測のlogits
        """
        batch, seq_len, d_model = encoder_output.shape

        # --- 絶対位置の注入 ---
        positions = torch.arange(seq_len, device=encoder_output.device)
        abs_pos_emb = self.absolute_pos_embedding(positions)  # (S, D)

        # Encoder出力 + 絶対位置埋め込み
        H_emd = encoder_output + abs_pos_emb.unsqueeze(0)  # (B, S, D)

        # --- EMD Transformer Layer ---
        # Self-Attention + Residual + LayerNorm
        attn_out, _ = self.emd_attention(H_emd)
        H_emd = self.emd_norm1(H_emd + self.dropout(attn_out))

        # FFN + Residual + LayerNorm
        ffn_out = self.emd_ffn(H_emd)
        H_emd = self.emd_norm2(H_emd + ffn_out)

        # --- MLM予測 ---
        if masked_positions is not None:
            # マスク位置のみの表現を取得
            batch_idx = torch.arange(batch, device=H_emd.device).unsqueeze(1)
            masked_repr = H_emd[batch_idx, masked_positions]  # (B, n_masked, D)
            logits = self.mlm_head(masked_repr)
        else:
            # 全位置のMLM予測
            logits = self.mlm_head(H_emd)

        return logits

EMDの動作を確認します。

# EMDの動作確認
torch.manual_seed(42)

d_model = 64
n_heads = 4
max_seq_len = 32
vocab_size = 1000
batch_size = 2
seq_len = 10

emd = EnhancedMaskDecoder(
    d_model=d_model,
    n_heads=n_heads,
    max_seq_len=max_seq_len,
    max_relative_distance=8,
    vocab_size=vocab_size
)

# ダミーのEncoder出力
encoder_output = torch.randn(batch_size, seq_len, d_model)

# マスク位置(各サンプルで2箇所をマスク)
masked_positions = torch.tensor([[2, 5], [1, 7]])

# フォワードパス
logits = emd(encoder_output, masked_positions)

print(f"Encoder出力形状:  {encoder_output.shape}")
print(f"マスク位置:       {masked_positions}")
print(f"MLM logits形状:   {logits.shape}")
print(f"予測結果(Top-3):")
top3 = torch.topk(logits[0, 0], k=3)
for rank, (val, idx) in enumerate(zip(top3.values, top3.indices), 1):
    print(f"  第{rank}位: token_id={idx.item()}, logit={val.item():.3f}")

EMDモジュールの出力形状が (batch_size, n_masked, vocab_size) = (2, 2, 1000) であり、マスクされた各位置に対して語彙サイズ分のlogitsが生成されていることが確認できます。ランダム初期化の段階ではTop-3の予測に意味はありませんが、学習後にはこのlogitsからsoftmaxを取って最も確率の高いトークンを予測します。EMDが絶対位置を注入しているため、同じコンテンツ表現でも文頭と文中で異なる予測を行えるようになります。

DeBERTa全体のアーキテクチャ

最後に、Encoder層とEMDを組み合わせたDeBERTaの全体像を実装します。

class DeBERTaEncoder(nn.Module):
    """DeBERTaのEncoder(Disentangled Attention を使う Transformer Encoder)"""

    def __init__(self, d_model, n_heads, n_layers, max_relative_distance=512, dropout=0.1):
        super().__init__()
        self.layers = nn.ModuleList()
        for _ in range(n_layers):
            layer = nn.ModuleDict({
                'attention': DisentangledAttention(d_model, n_heads, max_relative_distance, dropout),
                'norm1': nn.LayerNorm(d_model),
                'ffn': nn.Sequential(
                    nn.Linear(d_model, d_model * 4),
                    nn.GELU(),
                    nn.Linear(d_model * 4, d_model),
                    nn.Dropout(dropout)
                ),
                'norm2': nn.LayerNorm(d_model),
            })
            self.layers.append(layer)
        self.dropout = nn.Dropout(dropout)

    def forward(self, H, attention_mask=None):
        for layer in self.layers:
            # Self-Attention + Residual + LayerNorm
            attn_out, _ = layer['attention'](H, attention_mask)
            H = layer['norm1'](H + self.dropout(attn_out))
            # FFN + Residual + LayerNorm
            ffn_out = layer['ffn'](H)
            H = layer['norm2'](H + ffn_out)
        return H


class DeBERTaModel(nn.Module):
    """DeBERTaの全体モデル(事前学習用)"""

    def __init__(self, vocab_size=30522, d_model=768, n_heads=12,
                 n_layers=12, max_seq_len=512, max_relative_distance=512, dropout=0.1):
        super().__init__()

        # トークン埋め込み(位置埋め込みは加算しない)
        self.token_embedding = nn.Embedding(vocab_size, d_model)
        self.embedding_norm = nn.LayerNorm(d_model)
        self.embedding_dropout = nn.Dropout(dropout)

        # Encoder(Disentangled Attention ベース)
        self.encoder = DeBERTaEncoder(
            d_model, n_heads, n_layers, max_relative_distance, dropout
        )

        # Enhanced Mask Decoder
        self.emd = EnhancedMaskDecoder(
            d_model, n_heads, max_seq_len, max_relative_distance, vocab_size, dropout
        )

    def forward(self, input_ids, attention_mask=None, masked_positions=None):
        # トークン埋め込み(絶対位置は加算しない)
        H = self.token_embedding(input_ids)
        H = self.embedding_norm(H)
        H = self.embedding_dropout(H)

        # Encoder: Disentangled Attention で文脈を構築
        encoder_output = self.encoder(H, attention_mask)

        # EMD: 絶対位置を注入してMLM予測
        logits = self.emd(encoder_output, masked_positions)

        return logits

    def count_parameters(self):
        return sum(p.numel() for p in self.parameters() if p.requires_grad)


# モデル全体の確認
model = DeBERTaModel(
    vocab_size=1000,
    d_model=64,
    n_heads=4,
    n_layers=2,
    max_seq_len=32,
    max_relative_distance=8,
    dropout=0.1
)

input_ids = torch.randint(0, 1000, (2, 10))
masked_positions = torch.tensor([[2, 5], [1, 7]])

logits = model(input_ids, masked_positions=masked_positions)

print(f"入力形状:       {input_ids.shape}")
print(f"出力形状:       {logits.shape}")
print(f"パラメータ数:   {model.count_parameters():,}")
print(f"\nモデル構造:")
print(f"  トークン埋め込み: vocab_size=1000, d_model=64")
print(f"  Encoder: 2層 x DisentangledAttention(4 heads)")
print(f"  EMD: 1層 + MLM head")

このコードの出力から、DeBERTaの全体アーキテクチャが正しく動作していることが確認できます。入力のトークンIDから、マスク位置に対する語彙分布の予測までの一連の処理が、Disentangled Attention(コンテンツと位置の分離)とEMD(絶対位置の注入)を経て行われています。ここでは小規模な設定(d_model=64, 2層, 語彙1000)ですが、DeBERTa-Baseの実際の設定(d_model=768, 12層, 語彙50265)ではパラメータ数は約1億4000万となり、BERTの1億1000万より多くなります。増分の大部分は、位置用の射影行列($\bm{W}_q^p, \bm{W}_k^p$)と相対位置埋め込みテーブルに由来します。

それでは、DeBERTa V2/V3での改良点を見ていきましょう。

DeBERTa V2/V3の改良

初代DeBERTa(V1)の成功を受けて、Microsoftは V2 と V3 で更なる改良を加えました。これらの改良は、Disentangled Attentionの基本思想はそのまま活かしつつ、学習効率と性能をさらに向上させるものです。

DeBERTa V2の改良

1. 語彙サイズの拡大

DeBERTa V1 はGPT-2と同じBPEトークナイザ(語彙サイズ50,265)を使用していましたが、V2 では SentencePiece に基づく大規模語彙(128,000トークン)に拡大しました。語彙が大きくなると、未知語(UNK)の発生率が下がり、特に中国語や日本語のような文字数の多い言語でのサブワード分割が改善されます。

ただし、語彙サイズの拡大は埋め込み行列のパラメータ数を増加させます。語彙サイズ $V$ で埋め込み次元 $d$ のとき、埋め込み行列のパラメータ数は $V \times d$ です。$V = 128,000$、$d = 1536$(DeBERTa-XLarge)のとき約2億パラメータとなり、モデル全体のパラメータ数に占める割合が大きくなります。

2. 埋め込み行列の分解

語彙サイズの拡大に伴うパラメータ増加を抑えるため、V2 ではALBERTと同様の埋め込み行列分解を導入しました。$V \times d$ の巨大な行列を、$V \times d_e$ と $d_e \times d$ の2つの行列に分解します($d_e \ll d$)。

$$ \bm{E}(x) = \bm{E}_{\text{small}}(x) \cdot \bm{W}_{\text{proj}} $$

ここで $\bm{E}_{\text{small}} \in \mathbb{R}^{V \times d_e}$ は低次元の埋め込み行列、$\bm{W}_{\text{proj}} \in \mathbb{R}^{d_e \times d}$ は射影行列です。$d_e = 256$, $d = 1536$ の場合、パラメータ数は $128,000 \times 256 + 256 \times 1536 \approx 3,300$万となり、分解しない場合の約2億から大幅に削減されます。

3. nGiE(n-Gram Induced Embeddings)

V2 では、トークンの埋め込みを計算する際に、そのトークンの前後 $n$ 個のトークン(n-gram)の情報を畳み込みで取り込む仕組みを追加しました。これにより、Transformer の Self-Attention に入る前の段階で、局所的な文脈情報をある程度反映した埋め込みが得られます。

DeBERTa V3の改良: ELECTRA方式の事前学習

V3 の最大の革新は、事前学習タスクをMLMからReplaced Token Detection(RTD)に変更したことです。これはELECTRAで提案された事前学習方式であり、DeBERTaのアーキテクチャと組み合わせることで強力な相乗効果を生みます。

RTDでは、小さなGeneratorがマスク位置にもっともらしい偽のトークンを生成し、大きなDiscriminator(DeBERTa)が全てのトークンに対して「本物か偽物か」を2値分類します。MLMでは入力の15%のマスク位置でしか学習信号が得られませんが、RTDでは入力の100%のトークンで学習信号が得られるため、学習効率が大幅に向上します。

$$ \mathcal{L}_{\text{RTD}} = -\sum_{i=1}^{n} \left[ y_i \log D(\tilde{x}_i) + (1 – y_i) \log(1 – D(\tilde{x}_i)) \right] $$

ここで $\tilde{x}_i$ はGeneratorが生成した(可能性のある)偽トークンを含む入力系列、$y_i$ はトークン $i$ が本物(元のトークンと同一)なら1、偽物なら0のラベル、$D(\cdot)$ はDiscriminatorの出力確率です。

V3 ではさらに、GeneratorとDiscriminator間の埋め込み共有を改良しました。ELECTRA のオリジナル実装では2つのモデルが埋め込み行列を完全に共有していましたが、V3 では新しい勾配分離(Gradient Disentangled Embedding Sharing, GDES)を導入し、Generator の学習がDiscriminator の埋め込みに悪影響を与えることを防いでいます。具体的には、Discriminator の埋め込みは Discriminator の損失のみで更新され、Generator の損失からの勾配は stop-gradient で遮断されます。

V1/V2/V3の設計上の位置づけ

特徴 V1 V2 V3
Disentangled Attention あり あり あり
EMD あり あり あり
語彙サイズ 50K 128K 128K
埋め込み分解 なし あり あり
事前学習タスク MLM MLM RTD(ELECTRA方式)
埋め込み共有 GDES

V3 は、V1/V2のアーキテクチャ上の利点(Disentangled Attention + EMD)を保ちながら、ELECTRAの効率的な事前学習を取り込んだ、DeBERTaシリーズの集大成といえます。

実際の性能差はどの程度なのでしょうか。次のセクションで、ベンチマークスコアを確認しましょう。

性能比較: DeBERTa vs BERT vs RoBERTa

DeBERTaの設計判断がベンチマークスコアにどう反映されるかを、主要なNLUベンチマークで比較します。

SuperGLUEベンチマーク

SuperGLUEは、GLUEの後継として設計された、より高難度な自然言語理解ベンチマークです。8つのタスクで構成され、人間のベースラインは89.8点です。DeBERTaが歴史的に重要なのは、このベンチマークで人間のベースラインを初めて超えたモデルだという点です。

以下のコードで、各モデルの性能を可視化してみましょう。

import matplotlib.pyplot as plt
import numpy as np

# SuperGLUEベンチマークのスコア(論文・リーダーボードから)
models = ['BERT-Large', 'RoBERTa-Large', 'DeBERTa-Large', 'DeBERTa V3-Large', 'Human']
superglue_avg = [69.0, 84.6, 88.3, 91.4, 89.8]

colors = ['#42A5F5', '#66BB6A', '#FF7043', '#AB47BC', '#BDBDBD']

fig, ax = plt.subplots(figsize=(10, 6))
bars = ax.barh(models, superglue_avg, color=colors, edgecolor='white', linewidth=1.5, height=0.6)

# 人間のベースラインに赤い破線を追加
ax.axvline(x=89.8, color='red', linestyle='--', alpha=0.7, label='Human Baseline (89.8)')

# スコアをバーの横に表示
for bar, score in zip(bars, superglue_avg):
    ax.text(bar.get_width() + 0.5, bar.get_y() + bar.get_height()/2,
            f'{score:.1f}', ha='left', va='center', fontsize=12, fontweight='bold')

ax.set_xlabel('SuperGLUE Score', fontsize=13)
ax.set_title('SuperGLUE Benchmark Comparison', fontsize=15, fontweight='bold')
ax.set_xlim(60, 97)
ax.legend(fontsize=11)
ax.grid(axis='x', alpha=0.3)

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

このグラフから、DeBERTaシリーズの性能向上の軌跡が明確に読み取れます。BERT-Large(69.0)からRoBERTa-Large(84.6)への15.6ポイントの改善は、学習レシピの最適化によるものです。RoBERTa-LargeからDeBERTa-Large(88.3)への3.7ポイントの改善は、Disentangled AttentionとEMDというアーキテクチャの革新によるものであり、アーキテクチャの変更がいかに効果的かを示しています。さらにDeBERTa V3-Large(91.4)は、RTD事前学習の導入により人間のベースライン(89.8, 赤い破線)を1.6ポイント上回っています。

GLUEベンチマークでの詳細比較

GLUEベンチマークの個別タスクでも比較してみましょう。

import matplotlib.pyplot as plt
import numpy as np

# GLUEベンチマークの個別タスクスコア
tasks = ['MNLI', 'QQP', 'QNLI', 'SST-2', 'CoLA', 'STS-B', 'MRPC', 'RTE']

bert_scores =     [86.6, 91.3, 92.3, 93.2, 60.6, 90.0, 88.0, 70.4]
roberta_scores =  [90.2, 92.2, 94.7, 96.4, 68.0, 92.4, 90.9, 86.6]
deberta_scores =  [91.1, 92.3, 95.3, 96.8, 70.5, 92.8, 92.0, 88.3]
deberta_v3 =      [91.8, 92.7, 96.0, 97.1, 72.4, 93.0, 92.5, 92.4]

x = np.arange(len(tasks))
width = 0.2

fig, ax = plt.subplots(figsize=(16, 7))

b1 = ax.bar(x - 1.5*width, bert_scores, width, label='BERT-Large', color='#42A5F5', alpha=0.85)
b2 = ax.bar(x - 0.5*width, roberta_scores, width, label='RoBERTa-Large', color='#66BB6A', alpha=0.85)
b3 = ax.bar(x + 0.5*width, deberta_scores, width, label='DeBERTa-Large', color='#FF7043', alpha=0.85)
b4 = ax.bar(x + 1.5*width, deberta_v3, width, label='DeBERTa V3-Large', color='#AB47BC', alpha=0.85)

ax.set_xlabel('GLUE Tasks', fontsize=13)
ax.set_ylabel('Score', fontsize=13)
ax.set_title('GLUE Benchmark: Detailed Task Comparison', fontsize=15, fontweight='bold')
ax.set_xticks(x)
ax.set_xticklabels(tasks, fontsize=11)
ax.legend(fontsize=11, loc='lower right')
ax.set_ylim(55, 100)
ax.grid(axis='y', alpha=0.3)

# 各バーの上にスコアを表示
for bars in [b1, b2, b3, b4]:
    for bar in bars:
        height = bar.get_height()
        ax.annotate(f'{height:.1f}',
                    xy=(bar.get_x() + bar.get_width() / 2, height),
                    xytext=(0, 3), textcoords="offset points",
                    ha='center', va='bottom', fontsize=7)

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

# 改善幅の表示
print("=== DeBERTa V3 vs BERT-Large の改善幅 ===")
for task, bert, dv3 in zip(tasks, bert_scores, deberta_v3):
    diff = dv3 - bert
    print(f"  {task:8s}: {bert:.1f} -> {dv3:.1f} (+{diff:.1f})")
avg_improvement = np.mean(np.array(deberta_v3) - np.array(bert_scores))
print(f"  平均改善: +{avg_improvement:.1f}")

GLUEの個別タスク比較から、DeBERTaの改善が全タスクにわたって一貫していることが読み取れます。特にCoLA(+11.8ポイント)とRTE(+22.0ポイント)での改善が顕著です。CoLAは文法的な正しさの判定タスクであり、RTE は自然言語推論タスクです。これらはいずれもデータ量が少ないタスクであり、Disentangled Attentionによる効率的な位置・内容の学習と、RTDによる効率的な事前学習が、限られたファインチューニングデータでも高い転移学習性能を発揮することを示しています。QQPやSTS-Bのようにデータが豊富なタスクでは改善幅は控えめですが、それでもDeBERTa V3が全タスクでBERTを上回っている点は注目に値します。

Disentangled Attentionの寄与の分析

DeBERTaの性能向上がどの要素に起因するかを、消去実験(Ablation Study)の結果から見てみましょう。

import matplotlib.pyplot as plt
import numpy as np

# He et al. (2021) の消去実験データ(MNLIの精度)
configs = [
    'BERT-Large\n(Baseline)',
    '+ Relative\nPosition',
    '+ Disentangled\nAttention',
    '+ EMD',
    'DeBERTa\n(Full)'
]
mnli_scores = [86.6, 88.1, 89.5, 90.1, 91.1]

fig, ax = plt.subplots(figsize=(11, 6))

colors = ['#90CAF9', '#42A5F5', '#1E88E5', '#1565C0', '#0D47A1']
bars = ax.bar(range(len(configs)), mnli_scores, color=colors, edgecolor='white', linewidth=1.5, width=0.65)

# 改善幅の矢印
for i in range(1, len(configs)):
    diff = mnli_scores[i] - mnli_scores[i-1]
    ax.annotate(f'+{diff:.1f}',
                xy=(i, mnli_scores[i]),
                xytext=(i, mnli_scores[i] + 0.4),
                ha='center', fontsize=11, color='red', fontweight='bold')

# スコア表示
for i, (bar, score) in enumerate(zip(bars, mnli_scores)):
    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() - 0.5,
            f'{score:.1f}', ha='center', va='top', fontsize=12, fontweight='bold', color='white')

ax.set_ylabel('MNLI-m Accuracy (%)', fontsize=13)
ax.set_title('DeBERTa Ablation Study (MNLI-m)', fontsize=15, fontweight='bold')
ax.set_xticks(range(len(configs)))
ax.set_xticklabels(configs, fontsize=10)
ax.set_ylim(85, 92.5)
ax.grid(axis='y', alpha=0.3)

# 累積改善の表示
total = mnli_scores[-1] - mnli_scores[0]
ax.text(len(configs)-1, 92, f'Total: +{total:.1f}', ha='center', fontsize=12,
        fontweight='bold', color='darkred',
        bbox=dict(boxstyle='round,pad=0.3', facecolor='lightyellow', edgecolor='darkred'))

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

# 各改良の寄与率
print("=== 各改良の寄与率(MNLI-m)===")
total = mnli_scores[-1] - mnli_scores[0]
for i in range(1, len(configs)):
    diff = mnli_scores[i] - mnli_scores[i-1]
    pct = diff / total * 100
    label = configs[i].replace('\n', ' ')
    print(f"  {label}: +{diff:.1f} ({pct:.0f}%)")
print(f"  合計: +{total:.1f}")

消去実験の結果は、DeBERTaの各コンポーネントがそれぞれ独立して性能に貢献していることを示しています。相対位置の導入(+1.5ポイント)は、BERTの絶対位置加算よりも位置情報をうまく活用できていることを意味します。Disentangled Attention(+1.4ポイント)は、相対位置を導入した上でさらにコンテンツと位置を分離することの効果を示しており、これが本論文の中核的な貢献です。EMD(+0.6ポイント)は、相対位置だけでは捉えきれない絶対位置の情報を最終予測層に注入する効果を示しています。Disentangled AttentionとEMDを合わせた改善幅は約2.0ポイントであり、相対位置の導入(+1.5ポイント)と同等以上の貢献をしています。つまり、「位置の扱い方を変える」だけでなく「コンテンツと位置を分離する」ことが、性能向上の本質であることがわかります。

まとめ

本記事では、DeBERTaのDisentangled AttentionとEnhanced Mask Decoderの設計思想と数学的定式化を解説しました。

  • 従来の限界: BERTはコンテンツと位置を足し算で混合し、Transformer-XLも射影行列を共有していたため、Attentionが「内容で注目しているのか、位置で注目しているのか」を分離できなかった
  • Disentangled Attention: コンテンツと位置を独立したベクトルとして保持し、C2C(意味的関連)、C2P(意味→位置)、P2C(位置→意味)の3つの注意行列を個別に計算する。P2Pは有用な情報を持たないため省略する
  • 相対位置バイアス: 絶対位置ではなく相対位置を使い、クリッピングによって $2k+1$ 個の位置ベクトルにコンパクトに表現する。コンテンツとの分離により、相対位置の利点を最大限に活かせる
  • Enhanced Mask Decoder: Encoder全体では相対位置だけを使い、MLM予測の最終層でのみ絶対位置を注入する「二段階戦略」。文脈構築には相対位置、最終予測には絶対位置という使い分けが性能向上の鍵
  • DeBERTa V3: ELECTRA方式のRTD事前学習を採用し、入力の100%から学習信号を得る効率化。GDESにより埋め込み共有の課題も解決
  • 性能: SuperGLUEで人間のベースラインを初めて超え、GLUEの全タスクでBERT/RoBERTaを上回る

DeBERTaの「コンテンツと位置を分離する」という設計原理は、Attention機構の改良に留まらず、モデルの表現力と解釈性を同時に向上させる重要な知見です。今後、新しいAttention機構を設計する際の基本的な指針として活用できるでしょう。

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

画像なし
ELECTRAの判別的事前学習
DeBERTa V3が採用したReplaced Token Detectionの仕組みをGenerator/Discriminatorの構成から解説します
画像なし
ALBERT・DistilBERT — BERTの軽量化・効率化手法
パラメータ共有やモデル蒸留によるBERTの軽量化手法を解説します
画像なし
RoBERTaの改良点と性能向上
BERTの学習設定を最適化して性能を向上させたRoBERTaの設計判断を解説します