Masked Autoencoder(MAE)の理論と実装 — 画像の75%を隠して学ぶ自己教師あり表現学習

100ピースのジグソーパズルを想像してください。75ピースを取り除き、残った25ピースだけで元の絵を復元できるでしょうか? 人間なら、木の幹が見えれば枝や葉の位置を予測でき、空の色調からは雲の広がりを想像できます。この「少ない手がかりから全体を推測する能力」こそ、視覚世界の構造を深く理解している証拠です。

Masked Autoencoder(MAE)は、まさにこのジグソーパズルの発想をニューラルネットワークに持ち込んだ手法です。画像をパッチに分割し、その75%をランダムに隠した上で、残り25%の可視パッチだけから元の画像を復元するようモデルを訓練します。ラベルなしの画像だけで豊かな視覚表現を獲得できるため、自己教師あり学習(Self-Supervised Learning)の強力なアプローチとして注目されています。

ジグソーパズルの発想

上の図がMAEの発想です。画像の75%を隠し、残り25%だけから全体を復元させます。隣の色を引き延ばすだけでは埋まらないため、モデルは「木の幹が見えれば枝葉の位置を推測する」ようなシーン全体の構造理解を強いられます。これがラベルなしで豊かな視覚表現を学ぶ仕掛けです。

自然言語処理の世界では、BERTが文章の単語を一部マスクして残りから予測する「マスク言語モデリング」で大成功を収めました。MAEはその成功を画像領域に持ち込んだものですが、単純にBERTの方法をコピーしただけではありません。画像と言語の本質的な違い — 特に画像の空間的冗長性の高さ — を巧みに利用し、75%という高いマスキング率と非対称なエンコーダ・デコーダ設計という独自の工夫を導入しています。

MAEの理論を理解することは、以下のような場面で直接的に役立ちます。

  • 大規模視覚モデルの事前学習: ImageNetのようなラベル付きデータに頼らず、ウェブ上の膨大なラベルなし画像から汎用的な視覚表現を学習できます
  • 衛星画像・医療画像の解析: ラベル付けコストが高い専門分野では、少数のラベルで高精度を達成するためのファインチューニング基盤として威力を発揮します
  • マルチモーダル学習への拡張: MAEの設計思想は時系列データ、音声信号、動画へと広がっており、ドメインを超えた自己教師あり学習の基盤アーキテクチャになっています

本記事の内容

  • 画像の自己教師あり学習がなぜ難しかったかの歴史的背景
  • MAEの核心的アイデア — 画像の冗長性と高マスキング率の動機
  • 非対称エンコーダ・デコーダアーキテクチャの詳細設計
  • 損失関数と学習アルゴリズムの数式的定式化
  • 計算効率の定量的分析
  • PyTorchでのスクラッチ実装と学習実験
  • 下流タスクへの転移学習とMAEの発展

前提知識

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

画像なし
Vision Transformer(ViT)の理論と実装
画像をパッチに分割しTransformerで処理するViTの仕組みを解説します。MAEのエンコーダはViTそのものです。
画像なし
Transformerアーキテクチャの全体像
Self-Attention、Multi-Head Attention、Feed-Forward Networkの構造を解説します。
画像なし
マスク言語モデリング — BERTの事前学習手法
BERTのマスク予測タスクの仕組みを解説します。MAEの直接的なインスピレーション源です。
画像なし
オートエンコーダの種類と理論
通常のオートエンコーダ、変分オートエンコーダ、デノイジングオートエンコーダの比較を解説します。

なぜ画像の自己教師あり学習は難しかったのか

MAEの設計を理解するためには、まず「なぜ画像の自己教師あり学習は言語ほどうまくいかなかったのか」という歴史的背景を押さえる必要があります。BERTが2018年に自然言語処理に革命をもたらした後、同じアイデアを画像に適用しようとする試みが数多くなされましたが、長らく決定打に欠けていました。

言語と画像の本質的な違い

自然言語処理では、テキストは離散的なトークン列です。単語は語彙辞書の中の有限個の候補から選ばれるため、マスクされた単語の予測はクラス分類問題として自然に定式化できます。一方、画像は連続的な画素値の集合であり、マスクされたピクセルの予測は回帰問題になります。この違いが予測タスクの設計を複雑にしていました。

さらに根本的な違いは、言語トークンと画像パッチの情報密度の差です。自然言語の各単語は豊富な意味情報を持っています。「猫」という一語には、その動物の外見、行動パターン、サイズなど膨大な概念が凝縮されています。BERTでマスクされた単語を予測するには、文脈の深い意味的理解が不可欠です。

これに対して、画像の個々のピクセルや小さなパッチは、それ自体では意味情報をほとんど持ちません。隣接するピクセルの値は高い相関を持ち、ある画素の値は周囲の画素から容易に補間できてしまいます。つまり、画像には空間的冗長性(spatial redundancy)が大きく存在するのです。

言語と画像の情報密度の違い

図のように、言語は離散トークンで1語に多くの意味が詰まっているため、15%マスクでも文脈理解が要ります。一方、画像は連続値で隣接画素が強く相関するため、少しの欠損なら補間で埋まってしまいます。だからこそ画像では、意味理解を促すために高いマスキング率が必要になります。

従来手法の課題

この問題に対して、MAE以前の自己教師あり学習手法は主に2つのアプローチをとっていました。

1つ目は対照学習(Contrastive Learning)です。対照学習(Contrastive Learning)は、同じ画像を異なるデータ拡張で2つのビューに変換し、それらの表現が近くなるように(他の画像とは遠くなるように)学習します。SimCLR、MoCo、BYOLなどが代表的で、ラベルなしで優れた表現を獲得できましたが、大量のネガティブサンプルや運動量エンコーダなどのテクニックが必要で、設計が複雑でした。

2つ目はマスク画像モデリングの初期的な試みです。BEiT(2021年)はdVAE(離散変分オートエンコーダ)で画像をトークン化してからBERT式のマスク予測を行いましたが、dVAEの事前学習が必要という二段階のパイプラインが煩雑でした。iGPTは画素レベルの自己回帰モデルでしたが、計算コストが莫大で大規模画像への適用が困難でした。

MAEが解決したのは、まさにこの「言語と画像のギャップ」です。画像の高い空間的冗長性を弱点ではなく強みとして活かし、非常に高いマスキング率(75%)を採用することで、単純な画素値の補間では解けない真に意味的な予測タスクを生み出しました。

では、MAEはどのような設計でこの課題を解決したのでしょうか。次節で、その核心的なアイデアを詳しく見ていきます。

MAEの核心的アイデア — 画像の冗長性を逆手に取る

MAEの基本的な発想は、驚くほどシンプルです。ジグソーパズルのアナロジーに戻ると、パズルのピースの半分が残っていれば、周囲のピースの色やパターンの連続性から欠けたピースの見た目はかなり正確に推測できます。しかし、ピースの75%が欠けていたら? もはや「隣の色を引き延ばす」だけでは復元できません。木の幹のピースから葉の茂り方を想像し、空の色から雲の形を予測するような、シーン全体の構造理解が求められます。

高マスキング率の動機

画像の空間的冗長性が高いということは、低いマスキング率では予測タスクが「簡単すぎる」ことを意味します。たとえば画像パッチの15%(BERTと同じ比率)をマスクした場合、マスクされたパッチの周囲には大量の可視パッチが残っており、テクスチャの連続性だけで補間が可能です。モデルは画像の意味的な構造を理解することなく、単なるローパスフィルタのような働きで高い復元精度を達成できてしまいます。

He et al.の実験により、MAEは75%という極めて高いマスキング率で最高の性能を発揮することが示されました。この比率では、残り25%のパッチだけが手がかりとなります。猫の耳のパッチだけが見えている状態から猫の体全体のテクスチャを復元するには、「猫の耳の下には顔がある」「猫の体は特定のテクスチャパターンを持つ」といった高レベルの意味理解が不可欠です。これが、MAEが単なる画素補間ではなく意味的な表現学習を達成できる理由です。

マスキング率と難易度のトレードオフ

同じ画像を15%・75%・90%でマスクした様子です。15%では可視部分が多すぎてテクスチャの補間で解けてしまい、90%では手がかりが少なすぎて復元が曖昧になります。75%が「補間では解けないが、意味理解で解ける」スイートスポットです。

マスキングから復元までの流れ

MAEの処理の全体像を俯瞰しましょう。

[入力画像] (H × W × C)
    ↓
[パッチ分割] → N個のパッチ
    ↓
[ランダムマスキング] → 可視パッチ(25%) + マスクパッチ(75%)
    ↓
[エンコーダ] → 可視パッチのみ処理 → 潜在表現
    ↓
[マスクトークン挿入 + 位置埋め込み復元]
    ↓
[デコーダ] → 全パッチ位置の表現を復元
    ↓
[再構成ヘッド] → マスクパッチのピクセル値を予測
    ↓
[損失計算] → マスクパッチのみでMSE

MAEの処理フロー

この処理フローの中で、最も重要な設計判断はエンコーダがマスクされたパッチを一切見ないという点です。この非対称な設計がMAEの計算効率と表現学習の質を同時に高めています。

次節では、このアーキテクチャの各コンポーネントを詳細に解説していきます。

アーキテクチャの詳細 — 非対称エンコーダ・デコーダ設計

MAEの最も独創的な設計上の選択は、エンコーダとデコーダを非対称に設計したことです。通常のオートエンコーダではエンコーダとデコーダは対称的な構造を持ちますが、MAEでは意図的にエンコーダを大きく、デコーダを軽量にしています。この非対称性には明確な理由があります。

非対称エンコーダ・デコーダ設計

図の通り、大型のエンコーダは可視25%のパッチだけを処理し、軽量なデコーダがマスクトークンを含む全パッチを受け取って復元します。重い計算をエンコーダ側の少数パッチに集中させ、デコーダは事前学習後に捨てる——この割り切りが効率と表現力を両立させます。

エンコーダ:可視パッチだけを処理する大型ViT

MAEのエンコーダは、標準的なVision Transformer(ViT)です。ただし、通常のViTが全パッチを入力として受け取るのに対し、MAEのエンコーダはマスクされていない可視パッチのみを入力とします。

処理の流れを具体的に追いましょう。まず、$H \times W$ ピクセルの画像を $P \times P$ のパッチに分割します。パッチの総数は $N = HW / P^2$ 個です。例えば $224 \times 224$ の画像を $16 \times 16$ のパッチに分割すると $N = 196$ 個になります。

次に、$N$ 個のパッチからランダムに $\lfloor \rho N \rfloor$ 個($\rho$ はマスキング率、通常0.75)を選んでマスクし、残りの $(1 – \rho)N$ 個の可視パッチだけをエンコーダに入力します。75%マスキングの場合、196個中49個のパッチのみがエンコーダに渡されます。

各可視パッチ $\bm{x}_i^{\text{patch}} \in \mathbb{R}^{P^2 C}$ は、線形射影によって $D$ 次元の埋め込みベクトルに変換されます。

$$ \bm{e}_i = \bm{W}_E \bm{x}_i^{\text{patch}} + \bm{b}_E, \quad \bm{W}_E \in \mathbb{R}^{D \times P^2C} $$

さらに、各パッチの元の位置情報を保持するため、位置エンコーディングを加算します。

$$ \bm{z}_i^{(0)} = \bm{e}_i + \bm{p}_{\sigma(i)} $$

ここで $\bm{p}_{\sigma(i)} \in \mathbb{R}^D$ は、そのパッチの元画像における位置 $\sigma(i)$ に対応する位置埋め込みベクトルです。マスキングによって可視パッチの並びは元の空間的順序とは異なりますが、位置埋め込みが元の位置情報を伝えるため、エンコーダはパッチの空間的配置を把握できます。

この入力がViTの $L$ 層のTransformerブロックを通過します。各ブロックはSelf-AttentionとFeed-Forward Networkから構成されます。

$$ \bm{z}’^{(\ell)} = \text{MSA}\!\left(\text{LN}\!\left(\bm{z}^{(\ell-1)}\right)\right) + \bm{z}^{(\ell-1)} $$

$$ \bm{z}^{(\ell)} = \text{FFN}\!\left(\text{LN}\!\left(\bm{z}’^{(\ell)}\right)\right) + \bm{z}’^{(\ell)} $$

ここで、$\text{MSA}$ はMulti-Head Self-Attention、$\text{LN}$ はLayer Normalization、$\text{FFN}$ はPosition-wise Feed-Forward Networkを表します。$\ell = 1, 2, \dots, L$ がブロックのインデックスです。

ここが効率性の鍵です。 エンコーダは $N$ 個のパッチではなく、$(1-\rho)N$ 個のパッチのみを処理します。Self-Attentionの計算量はトークン数の二乗に比例するため、75%マスキングの場合、Attentionの計算量は $(0.25)^2 = 0.0625$、つまり元の約 $1/16$ に削減されます。さらに、全結合層の計算量もトークン数に比例するため $1/4$ になります。全体として、エンコーダの計算コストは通常のViTの約 $1/4$ 以下です。

デコーダ:マスクトークンを含む軽量Transformer

エンコーダの出力は可視パッチに対応する潜在表現のみです。デコーダの役割は、この部分的な情報からマスクされたパッチのピクセル値を復元することです。

デコーダへの入力を構成する手順は以下の通りです。

ステップ1: エンコーダ出力の射影

エンコーダの出力 $\bm{z}_i^{(L)} \in \mathbb{R}^D$ を、デコーダの次元 $D’$ に線形射影します。一般にデコーダはエンコーダより小さく設計されるため $D’ < D$ です(論文のデフォルトでは $D = 1024$, $D' = 512$)。

$$ \hat{\bm{z}}_i = \bm{W}_{\text{dec}} \bm{z}_i^{(L)} + \bm{b}_{\text{dec}}, \quad \bm{W}_{\text{dec}} \in \mathbb{R}^{D’ \times D} $$

ステップ2: マスクトークンの挿入

マスクされた位置には、学習可能な共有ベクトル $\bm{m} \in \mathbb{R}^{D’}$(マスクトークン)を挿入します。マスクトークンは全てのマスク位置で同一のベクトルですが、次のステップで位置埋め込みが加算されるため、各位置の区別は可能です。

ステップ3: 位置埋め込みの加算

全 $N$ 個のトークン(可視パッチの射影 + マスクトークン)に対して、元画像での位置に対応する位置埋め込み $\bm{p}’_j \in \mathbb{R}^{D’}$ を加算します。

$$ \tilde{\bm{z}}_j = \begin{cases} \hat{\bm{z}}_j + \bm{p}’_j & (j \in \text{可視パッチ}) \\ \bm{m} + \bm{p}’_j & (j \in \text{マスクパッチ}) \end{cases} $$

この操作により、マスクトークンは「自分がどの位置に対応しているか」を知ることができます。デコーダのSelf-Attentionで可視パッチの情報とマスクトークンの位置情報が組み合わさることで、各マスク位置に適した画素値が復元されます。

ステップ4: デコーダTransformerブロック

全 $N$ トークンをデコーダの $L’$ 層のTransformerブロックに通します(論文のデフォルトは $L’ = 8$ 層、エンコーダの $L = 24$ 層に比べて軽量)。

ステップ5: 再構成ヘッド

デコーダの最終層出力から、マスク位置に対応するトークンを取り出し、線形射影でパッチのピクセル値 $P^2 C$ 次元に変換します。

$$ \hat{\bm{x}}_j = \bm{W}_{\text{rec}} \tilde{\bm{z}}_j^{(L’)} + \bm{b}_{\text{rec}}, \quad \bm{W}_{\text{rec}} \in \mathbb{R}^{P^2 C \times D’} $$

なぜデコーダを軽量にするのか

MAEの目的は「良いデコーダ」を作ることではなく、「良いエンコーダの表現」を学習することです。デコーダは事前学習のときだけ使われ、ファインチューニングでは捨てられます。したがって、デコーダに多くのパラメータと計算を割く必要はありません。

むしろ、デコーダを軽量にすることで、復元タスクの「重労働」がエンコーダ側に集中します。エンコーダは25%のパッチだけから画像全体を表現する豊かな潜在表現を構築しなければならず、これがエンコーダの表現力を高めるのです。He et al.の実験では、デコーダの深さを1層にしても下流タスクの性能低下は小さく、8層程度で十分であることが確認されています。

この非対称な設計は、「事前学習中の計算効率」と「学習される表現の質」の両方に寄与するエレガントな解です。次に、具体的なマスキング戦略について見ていきましょう。

マスキング戦略 — なぜランダムマスキングが最適なのか

マスキングの方法は、自己教師あり学習タスクの難易度と学習される表現の質を直接的に左右します。直感的には、マスクするパッチの選び方によって「何を復元させるか」というタスクの性質が根本的に変わるからです。たとえば、画像の左半分を全てマスクするのと、ランダムに散らばったパッチをマスクするのでは、求められる推論能力が異なります。

ランダムマスキング

MAEが採用するのは一様ランダムマスキングです。$N$ 個のパッチから、マスクする $\lfloor \rho N \rfloor$ 個を一様ランダムに選択します。アルゴリズムとしては、$N$ 個のインデックスをランダムにシャッフルし、先頭 $(1-\rho)N$ 個を可視パッチ、残りをマスクパッチとするだけです。

ランダムマスキングが効果的な理由は、画像全体に均等にマスクが分布するため、モデルが特定の空間パターンに依存した「ショートカット」を学習しにくい点にあります。画像の中央だけにマスクが集中する、あるいは規則的なグリッドパターンでマスクされるといった偏りがないため、モデルはあらゆる空間位置の関係性を学ぶ必要があります。

他のマスキング戦略との比較

He et al.はランダムマスキングの他に、ブロックマスキング(連続する矩形領域をマスク)とグリッドマスキング(規則的な格子パターンでマスク)も実験しています。

ブロックマスキングは、大きな連続領域が欠落するため、一見するとタスクが難しくなりそうですが、実際には可視パッチが画像の一部に偏って集中するため、その領域の局所的な情報だけで多くのマスクパッチを復元できてしまいます。また、マスクされた大きな矩形領域の内部では境界付近のパッチが特に推測しやすくなり、学習信号に偏りが生じます。

グリッドマスキングでは、可視パッチが等間隔に配置されるため、どのマスクパッチにも近くに可視パッチが存在します。これにより、隣接パッチの情報による単純な補間でも高い復元精度が得られてしまい、高レベルな意味理解を促す学習信号として不十分です。

実験結果は明確で、ランダムマスキングが最も高い下流タスク性能をもたらしました。

マスキング戦略の比較

図の3戦略を比べると、ブロック(中央)は可視パッチが一箇所に偏って局所補間で解けてしまい、グリッド(右)は各マスクの近くに必ず可視パッチがあり補間が容易です。ランダム(左)は全体に均等に散らばるため、特定位置に依存したショートカットを学習しにくく、最も良い表現が得られます。

75%の最適性

マスキング率 $\rho$ は、「タスクの難しさ」と「利用可能な情報量」のトレードオフを決定するハイパーパラメータです。He et al.は、マスキング率を40%から90%まで変化させた系統的な実験を行い、75%付近が最適であることを示しました。

この値が最適である直感的な理由は次のように説明できます。マスキング率が低い(40%以下)場合、残りの可視パッチが多すぎて、テクスチャの補間だけで復元でき、意味的理解を必要としません。マスキング率が高すぎる(90%以上)場合、手がかりが少なすぎて復元タスク自体が曖昧になり、安定した学習が難しくなります。75%はこのバランスが最もよいスイートスポットなのです。

注目すべきは、この最適なマスキング率がBERTの15%よりはるかに高い点です。これは先に述べた画像と言語の情報密度の違いを反映しています。各パッチの情報密度が低い画像では、意味的な予測タスクを生み出すために多くのパッチを隠す必要があるのです。

マスキング戦略が決まったところで、次は「何を復元させるか」という予測ターゲットと損失関数の設計を見ていきます。

損失関数 — マスクパッチのピクセルを復元する

復元すべきターゲットの選択は、自己教師あり学習において極めて重要な設計判断です。BEiTのようにトークン化された離散表現を予測するアプローチもありますが、MAEはよりシンプルに生のピクセル値を直接復元します。

再構成ターゲット

ここで言う「ピクセル値」とは何かを正確に定義しましょう。各パッチは $P \times P \times C$ ピクセルで構成されます($C = 3$ はRGBチャンネル)。これを1次元に平坦化した $P^2 C$ 次元のベクトル $\bm{x}_j^{\text{patch}} \in \mathbb{R}^{P^2 C}$ が再構成ターゲットです。

実際には、He et al.はパッチごとに正規化(パッチ内のピクセル値を平均0・分散1に標準化)した値をターゲットとすることで性能が向上することを報告しています。正規化により、各パッチの「どんなテクスチャパターンを持っているか」という相対的な構造が強調され、画像全体の明るさやコントラストのような大域的な情報の影響が軽減されます。

パッチ正規化による再構成ターゲット

図のように、生のピクセル値(左)には明るさやコントラストの情報が含まれますが、パッチごとに平均0・分散1へ正規化(右)すると、相対的なテクスチャ構造が際立ちます。He et al.はこの正規化ターゲットを使うと性能が向上することを報告しています。

MSE損失

MAEの損失関数は、マスクパッチのみに対する平均二乗誤差(Mean Squared Error, MSE)です。直感的には、「隠したピースだけを答え合わせする」仕組みです。可視パッチはモデルがそのまま見ているため、それらの復元精度を評価しても意味がありません。

マスクパッチの集合を $\mathcal{M}$ とすると、損失関数は次のように定義されます。

$$ \mathcal{L} = \frac{1}{|\mathcal{M}|} \sum_{j \in \mathcal{M}} \left\| \hat{\bm{x}}_j – \bm{x}_j^{\text{patch}} \right\|_2^2 $$

ここで $\hat{\bm{x}}_j \in \mathbb{R}^{P^2 C}$ はデコーダが予測したパッチ $j$ のピクセル値、$\bm{x}_j^{\text{patch}} \in \mathbb{R}^{P^2 C}$ は対応する正解のピクセル値です。$|\mathcal{M}|$ はマスクパッチの数です。

パッチ正規化を適用する場合は、各マスクパッチの正解ピクセル値を事前に正規化しておきます。

パッチ $j$ 内のピクセル値の平均を $\mu_j$、標準偏差を $\sigma_j$ として、正規化されたターゲットは次のようになります。

$$ \tilde{\bm{x}}_j = \frac{\bm{x}_j^{\text{patch}} – \mu_j}{\sigma_j + \epsilon} $$

$\epsilon$ はゼロ除算を防ぐ小さな定数です。この場合、損失関数のターゲットが $\tilde{\bm{x}}_j$ に置き換わります。

$$ \mathcal{L}_{\text{norm}} = \frac{1}{|\mathcal{M}|} \sum_{j \in \mathcal{M}} \left\| \hat{\bm{x}}_j – \tilde{\bm{x}}_j \right\|_2^2 $$

なぜピクセル復元なのか

BEiTが離散トークン予測を採用したのに対し、MAEが生のピクセル値復元というシンプルなアプローチで成功した理由は何でしょうか。鍵は高マスキング率との組み合わせにあります。

75%という高いマスキング率のもとでは、残りのパッチから元のピクセル値を正確に復元すること自体が十分に難しいタスクです。低いマスキング率ではピクセル復元が容易すぎて表面的なパターン学習に陥りますが、高いマスキング率ではピクセルレベルの復元であっても意味的な理解が必要になります。したがって、追加のトークナイザ(dVAEなど)を導入する必要がなく、パイプラインが簡素化されます。

さらに、ピクセル復元は連続値の回帰問題であるため、離散トークン予測に必要な語彙サイズの設計やトークナイザの品質への依存がありません。この「シンプルだが高マスキング率との組み合わせで強力」という設計哲学が、MAEの美しさの本質です。

損失関数の全体像がわかったところで、次はこれらの要素を数式で統一的に定式化しましょう。

数式による定式化

ここまで個別に解説してきたMAEの各要素を、統一的な数式の枠組みで整理します。この定式化を通じて、MAEの全体像を一つの数学的フレームワークとして把握できるようになります。

パッチ分割とマスキング

入力画像 $\bm{X} \in \mathbb{R}^{H \times W \times C}$ を $P \times P$ パッチに分割し、$N = HW/P^2$ 個のパッチ列 $\{\bm{x}_1, \bm{x}_2, \dots, \bm{x}_N\}$ を得ます。各パッチ $\bm{x}_i \in \mathbb{R}^{P^2 C}$ は平坦化されたピクセル値ベクトルです。

インデックス集合 $\{1, 2, \dots, N\}$ のランダムな順列 $\pi$ を生成し、先頭 $N_v = \lfloor (1-\rho)N \rfloor$ 個のインデックスを可視集合 $\mathcal{V} = \{\pi(1), \dots, \pi(N_v)\}$、残りをマスク集合 $\mathcal{M} = \{\pi(N_v+1), \dots, \pi(N)\}$ とします。

エンコーダの処理

可視パッチのみを線形射影し、位置埋め込みを加算します。

$$ \bm{z}_i^{(0)} = \bm{W}_E \bm{x}_{\pi(i)} + \bm{b}_E + \bm{p}_{\pi(i)}, \quad i = 1, \dots, N_v $$

ここで $\bm{W}_E \in \mathbb{R}^{D \times P^2 C}$ はパッチ埋め込み行列、$\bm{p}_k \in \mathbb{R}^D$ は位置 $k$ に対応する位置埋め込みです。

$L$ 層のTransformerブロックを適用します。$\ell = 1, \dots, L$ について、

まず、Layer Normalization後にMulti-Head Self-Attentionを計算し、残差接続を加えます。

$$ \bm{z}’^{(\ell)} = \text{MSA}\!\left(\text{LN}\!\left(\bm{z}^{(\ell-1)}\right)\right) + \bm{z}^{(\ell-1)} $$

次に、再びLayer Normalization後にFFNを適用し、残差接続を加えます。

$$ \bm{z}^{(\ell)} = \text{FFN}\!\left(\text{LN}\!\left(\bm{z}’^{(\ell)}\right)\right) + \bm{z}’^{(\ell)} $$

エンコーダ出力: $\{\bm{z}_1^{(L)}, \dots, \bm{z}_{N_v}^{(L)}\}$(可視パッチのみ、$N_v$ 個)

デコーダの処理

エンコーダ出力を線形射影し、マスクトークンを挿入して全 $N$ 位置のトークン列を構成します。

$$ \tilde{\bm{z}}_j^{(0)} = \begin{cases} \bm{W}_{\text{dec}} \bm{z}_j^{(L)} + \bm{b}_{\text{dec}} + \bm{p}’_j & (j \in \mathcal{V}) \\ \bm{m} + \bm{p}’_j & (j \in \mathcal{M}) \end{cases} $$

ここで $\bm{m} \in \mathbb{R}^{D’}$ は学習可能なマスクトークン、$\bm{p}’_j \in \mathbb{R}^{D’}$ はデコーダの位置埋め込みです。

$L’$ 層のTransformerブロックを適用します($\ell = 1, \dots, L’$)。エンコーダと同じ構造ですが、全 $N$ トークンに対して適用します。

再構成と損失

デコーダ最終層からマスク位置のトークンを取り出し、ピクセル値に射影します。

$$ \hat{\bm{x}}_j = \bm{W}_{\text{rec}} \tilde{\bm{z}}_j^{(L’)} + \bm{b}_{\text{rec}}, \quad j \in \mathcal{M} $$

損失関数はマスクパッチのみの平均二乗誤差です。

$$ \mathcal{L} = \frac{1}{|\mathcal{M}|} \sum_{j \in \mathcal{M}} \left\| \hat{\bm{x}}_j – \bm{x}_j \right\|_2^2 $$

この定式化全体を一つのパイプラインとして見ると、MAEは「ランダムに部分観測された画像パッチから、未観測パッチのピクセル値を条件付き回帰で予測する」フレームワークであることがわかります。

定式化を踏まえて、次はこのアーキテクチャがなぜ計算効率に優れるのかを定量的に分析しましょう。

なぜ効率的か — 計算量の分析

MAEの非対称エンコーダ・デコーダ設計は、事前学習の計算効率を劇的に向上させます。ここでは、その効率性を定量的に分析します。実際にどれだけの計算が節約されるのかを、具体的な数値で実感しましょう。

Transformerブロックの計算量

Transformerブロックの計算量は、主にSelf-Attentionと Feed-Forward Networkの2つの成分から構成されます。入力トークン数を $n$、埋め込み次元を $d$ とすると、次のようになります。

まず、Multi-Head Self-Attentionについて考えます。Query・Key・Valueの線形射影にそれぞれ $O(nd^2)$ の計算量がかかります。

次に、$\bm{Q}\bm{K}^\top$ の計算に $O(n^2 d)$ がかかります。

出力の射影に $O(nd^2)$ がかかります。

合計すると、MSA全体で $O(4nd^2 + 2n^2 d)$ の計算量になります。

FFNは2つの線形層($d \to 4d \to d$)で構成されるため、計算量は $O(8nd^2)$ です。

したがって、1ブロックあたりの計算量は以下のようになります。

$$ \text{Cost}_{\text{block}}(n) = O(12nd^2 + 2n^2d) $$

通常、$d \gg n$(例: $d = 1024$, $n = 196$)の場合、$nd^2$ 項が支配的です。つまり、計算量はトークン数 $n$ にほぼ線形に比例します。

MAEの計算量削減

標準ViTがすべてのパッチ $N$ 個を処理するのに対し、MAEのエンコーダは可視パッチ $N_v = (1-\rho)N$ 個だけを処理します。エンコーダの計算コスト比は次のようになります。

$$ \frac{\text{Cost}_{\text{MAE-enc}}}{\text{Cost}_{\text{ViT}}} \approx \frac{N_v}{N} = 1 – \rho $$

これは $nd^2$ 項が支配的な場合の近似です。$\rho = 0.75$ の場合、エンコーダの計算コストは標準ViTの約 25% になります。

ただし、$n^2 d$ 項(Self-Attentionのスコア計算)を考慮すると、削減効果はさらに大きくなります。

$$ \frac{N_v^2}{N^2} = (1 – \rho)^2 = 0.0625 $$

Self-Attentionの二次項は元の約 6.25% にまで削減されます。

75%マスクによる計算量削減

図の通り、可視25%しか処理しないため、エンコーダの線形項は約25%、Self-Attentionの二次項は約6.25%にまで減ります。デコーダは軽量なので、事前学習全体では標準ViTの約1/3のコストで済み、これが大規模モデルへのスケーラビリティを支えています。

デコーダの計算コスト

デコーダは全 $N$ トークンを処理しますが、(1)次元が小さい($D’ < D$)、(2)層数が少ない($L' < L$)ため、計算コストは抑えられています。

論文のデフォルト設定(ViT-Large)を具体的な数値で見てみましょう。

エンコーダ デコーダ
入力トークン数 49(可視パッチ) 196(全パッチ)
埋め込み次元 $d$ 1024 512
層数 $L$ 24 8
ヘッド数 16 16

エンコーダ1層あたりの $nd^2$ 項は $49 \times 1024^2 \approx 5.1 \times 10^7$ です。一方、デコーダ1層あたりは $196 \times 512^2 \approx 5.1 \times 10^7$ と、実はほぼ同じオーダーです。しかし、エンコーダが24層でデコーダが8層なので、全体ではデコーダはエンコーダの約 $1/3$ のコストで済みます。

全体の計算コストを通常のViT(24層、196トークン、$d=1024$)と比較すると、MAEの事前学習はおおよそ3倍以上高速です。He et al.は、ViT-Largeの事前学習がA100 GPU 8台で約31時間(1600エポック)で完了したと報告しており、これは同等性能の対照学習手法の数分の1の計算コストです。

この計算効率の高さは、大規模モデル(ViT-Huge, ViT-Giant)への事前学習のスケーラビリティを大きく向上させます。実際、MAEの論文タイトルが “Masked Autoencoders Are Scalable Vision Learners” であることが、この点を強調しています。

ここまでの理論的な理解を踏まえて、いよいよPyTorchで実際にMAEを実装していきましょう。

PyTorchでの実装

ここからは、MAEの各コンポーネントをPyTorchでスクラッチ実装します。まずは必要なライブラリのインポートと基本的なモジュールから始め、段階的にMAEの全体を組み上げていきます。

パッチ埋め込みとマスキング

最初に、画像をパッチに分割して埋め込みベクトルに変換し、ランダムマスキングを行うモジュールを実装します。

import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.font_manager as _fm
for _c in ['Hiragino Sans','Yu Gothic','Noto Sans CJK JP','IPAexGothic','Meiryo']:
    if any(_c==_f.name for _f in _fm.fontManager.ttflist):
        plt.rcParams['font.family']=_c; break
plt.rcParams['axes.unicode_minus']=False
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

class PatchEmbed(nn.Module):
    """画像をパッチに分割し、線形射影で埋め込みベクトルに変換"""
    def __init__(self, img_size=32, patch_size=4, in_channels=3, embed_dim=192):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        self.num_patches = (img_size // patch_size) ** 2
        # 畳み込みでパッチ分割と線形射影を同時に実行
        self.proj = nn.Conv2d(
            in_channels, embed_dim,
            kernel_size=patch_size, stride=patch_size
        )

    def forward(self, x):
        # x: (B, C, H, W) -> (B, num_patches, embed_dim)
        x = self.proj(x)           # (B, embed_dim, H/P, W/P)
        x = x.flatten(2).transpose(1, 2)  # (B, num_patches, embed_dim)
        return x

PatchEmbed は畳み込み層を使ってパッチ分割と線形射影を1ステップで実行しています。カーネルサイズとストライドをパッチサイズに設定することで、重複のないパッチ分割が実現できます。CIFAR-10($32 \times 32$)を $4 \times 4$ パッチで分割すると $64$ 個のパッチが生成されます。

Transformerブロック

次に、エンコーダとデコーダで共通に使うTransformerブロックを実装します。

class TransformerBlock(nn.Module):
    """Pre-Norm Transformer ブロック(MSA + FFN + 残差接続)"""
    def __init__(self, dim, num_heads, mlp_ratio=4.0, dropout=0.0):
        super().__init__()
        self.norm1 = nn.LayerNorm(dim)
        self.attn = nn.MultiheadAttention(
            dim, num_heads, dropout=dropout, batch_first=True
        )
        self.norm2 = nn.LayerNorm(dim)
        mlp_hidden = int(dim * mlp_ratio)
        self.mlp = nn.Sequential(
            nn.Linear(dim, mlp_hidden),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(mlp_hidden, dim),
            nn.Dropout(dropout),
        )

    def forward(self, x):
        # Self-Attention + 残差接続
        x_norm = self.norm1(x)
        attn_out, _ = self.attn(x_norm, x_norm, x_norm)
        x = x + attn_out
        # FFN + 残差接続
        x = x + self.mlp(self.norm2(x))
        return x

Pre-Norm構成(Layer Normalizationを Attention/FFN の前に適用)を採用しています。これはViTやMAEの論文で標準的に使われている構成です。nn.MultiheadAttentionbatch_first=True により、テンソル形状が (batch, seq_len, dim) の直感的な並びになります。

MAEエンコーダ

エンコーダは可視パッチのみを処理するViTです。ランダムマスキングのロジックも含みます。

class MAEEncoder(nn.Module):
    """MAEエンコーダ — 可視パッチのみを処理するViT"""
    def __init__(self, img_size=32, patch_size=4, in_channels=3,
                 embed_dim=192, depth=6, num_heads=6, mlp_ratio=4.0):
        super().__init__()
        self.patch_embed = PatchEmbed(img_size, patch_size, in_channels, embed_dim)
        num_patches = self.patch_embed.num_patches
        # 学習可能な位置埋め込み
        self.pos_embed = nn.Parameter(
            torch.zeros(1, num_patches, embed_dim)
        )
        # Transformerブロック
        self.blocks = nn.ModuleList([
            TransformerBlock(embed_dim, num_heads, mlp_ratio)
            for _ in range(depth)
        ])
        self.norm = nn.LayerNorm(embed_dim)
        # 初期化
        nn.init.trunc_normal_(self.pos_embed, std=0.02)

    def random_masking(self, x, mask_ratio):
        """ランダムマスキング: 可視パッチとマスクパッチに分割"""
        B, N, D = x.shape
        num_keep = int(N * (1 - mask_ratio))

        # 各サンプルに対してランダムノイズを生成し、argsortでシャッフル
        noise = torch.rand(B, N, device=x.device)
        ids_shuffle = torch.argsort(noise, dim=1)
        ids_restore = torch.argsort(ids_shuffle, dim=1)

        # 可視パッチのインデックスを取り出す
        ids_keep = ids_shuffle[:, :num_keep]

        # 可視パッチの埋め込みを取得
        x_visible = torch.gather(
            x, dim=1,
            index=ids_keep.unsqueeze(-1).expand(-1, -1, D)
        )
        # マスク(0=可視, 1=マスク)を生成
        mask = torch.ones(B, N, device=x.device)
        mask[:, :num_keep] = 0
        mask = torch.gather(mask, dim=1, index=ids_restore)

        return x_visible, mask, ids_restore

    def forward(self, x, mask_ratio=0.75):
        # パッチ埋め込み + 位置埋め込み
        x = self.patch_embed(x)       # (B, N, D)
        x = x + self.pos_embed        # 位置情報を加算
        # ランダムマスキング
        x, mask, ids_restore = self.random_masking(x, mask_ratio)
        # Transformerブロックを通す(可視パッチのみ)
        for block in self.blocks:
            x = block(x)
        x = self.norm(x)
        return x, mask, ids_restore

random_masking メソッドがMAEの核心部分です。torch.rand でランダムノイズを生成し、argsort でシャッフル順を得ることで、バッチ全体で効率的にランダムマスキングを実行しています。ids_restore はデコーダでマスクトークンを正しい位置に挿入するために保持します。torch.gather によって可視パッチだけを抽出する操作は、GPU上でバッチ並列に実行されるため高速です。

MAEデコーダ

デコーダはエンコーダ出力にマスクトークンを挿入し、全パッチの再構成を行います。

class MAEDecoder(nn.Module):
    """MAEデコーダ — マスクトークンを含む軽量Transformer"""
    def __init__(self, num_patches=64, encoder_dim=192,
                 decoder_dim=96, depth=2, num_heads=3,
                 mlp_ratio=4.0, patch_size=4, in_channels=3):
        super().__init__()
        self.decoder_embed = nn.Linear(encoder_dim, decoder_dim)
        # 学習可能なマスクトークン
        self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
        # デコーダの位置埋め込み
        self.decoder_pos_embed = nn.Parameter(
            torch.zeros(1, num_patches, decoder_dim)
        )
        # Transformerブロック
        self.blocks = nn.ModuleList([
            TransformerBlock(decoder_dim, num_heads, mlp_ratio)
            for _ in range(depth)
        ])
        self.norm = nn.LayerNorm(decoder_dim)
        # 再構成ヘッド: デコーダ次元 → パッチのピクセル値
        self.head = nn.Linear(decoder_dim, patch_size ** 2 * in_channels)
        # 初期化
        nn.init.trunc_normal_(self.mask_token, std=0.02)
        nn.init.trunc_normal_(self.decoder_pos_embed, std=0.02)

    def forward(self, x, ids_restore):
        # エンコーダ出力をデコーダ次元に射影
        x = self.decoder_embed(x)     # (B, N_visible, decoder_dim)
        B, N_vis, D = x.shape
        N = ids_restore.shape[1]  # 全パッチ数

        # マスクトークンをマスク位置の数だけ複製
        mask_tokens = self.mask_token.expand(B, N - N_vis, -1)

        # 可視トークンとマスクトークンを結合し、元の順序に復元
        x_full = torch.cat([x, mask_tokens], dim=1)
        x_full = torch.gather(
            x_full, dim=1,
            index=ids_restore.unsqueeze(-1).expand(-1, -1, D)
        )
        # 位置埋め込みを加算
        x_full = x_full + self.decoder_pos_embed

        # Transformerブロック
        for block in self.blocks:
            x_full = block(x_full)
        x_full = self.norm(x_full)
        # 再構成ヘッド
        x_rec = self.head(x_full)     # (B, N, patch_size^2 * C)
        return x_rec

デコーダでは、ids_restore を使って可視トークンとマスクトークンを元の空間的順序に並べ替えています。この並べ替えの後に位置埋め込みを加算することで、各トークンが正しい空間位置の情報を持つようになります。再構成ヘッドの出力次元は $P^2 \times C$(パッチあたりのピクセル数)です。

MAEモデル全体

エンコーダとデコーダを統合し、損失計算まで含むMAEモデルの全体を組み上げます。

class MAE(nn.Module):
    """Masked Autoencoder (MAE) — 完全なモデル"""
    def __init__(self, img_size=32, patch_size=4, in_channels=3,
                 encoder_dim=192, encoder_depth=6, encoder_heads=6,
                 decoder_dim=96, decoder_depth=2, decoder_heads=3,
                 mlp_ratio=4.0, mask_ratio=0.75, norm_pix_loss=True):
        super().__init__()
        self.patch_size = patch_size
        self.mask_ratio = mask_ratio
        self.norm_pix_loss = norm_pix_loss
        num_patches = (img_size // patch_size) ** 2

        self.encoder = MAEEncoder(
            img_size, patch_size, in_channels,
            encoder_dim, encoder_depth, encoder_heads, mlp_ratio
        )
        self.decoder = MAEDecoder(
            num_patches, encoder_dim, decoder_dim,
            decoder_depth, decoder_heads, mlp_ratio,
            patch_size, in_channels
        )

    def patchify(self, imgs):
        """画像テンソルをパッチ列に変換"""
        P = self.patch_size
        B, C, H, W = imgs.shape
        h, w = H // P, W // P
        x = imgs.reshape(B, C, h, P, w, P)
        x = x.permute(0, 2, 4, 3, 5, 1)  # (B, h, w, P, P, C)
        x = x.reshape(B, h * w, P * P * C)  # (B, N, P^2*C)
        return x

    def unpatchify(self, x):
        """パッチ列を画像テンソルに再構成"""
        P = self.patch_size
        B, N, _ = x.shape
        h = w = int(N ** 0.5)
        C = x.shape[-1] // (P * P)
        x = x.reshape(B, h, w, P, P, C)
        x = x.permute(0, 5, 1, 3, 2, 4)  # (B, C, h, P, w, P)
        imgs = x.reshape(B, C, h * P, w * P)
        return imgs

    def forward_loss(self, imgs, pred, mask):
        """マスクパッチのみでMSE損失を計算"""
        target = self.patchify(imgs)  # (B, N, P^2*C)
        if self.norm_pix_loss:
            # パッチごとに正規化
            mean = target.mean(dim=-1, keepdim=True)
            var = target.var(dim=-1, keepdim=True)
            target = (target - mean) / (var + 1e-6).sqrt()
        loss = (pred - target) ** 2
        loss = loss.mean(dim=-1)  # (B, N) — パッチごとの平均MSE
        # マスクパッチのみで平均
        loss = (loss * mask).sum() / mask.sum()
        return loss

    def forward(self, imgs):
        # エンコーダ: 可視パッチのみ処理
        latent, mask, ids_restore = self.encoder(imgs, self.mask_ratio)
        # デコーダ: 全パッチの再構成
        pred = self.decoder(latent, ids_restore)
        # 損失計算: マスクパッチのみ
        loss = self.forward_loss(imgs, pred, mask)
        return loss, pred, mask

patchifyunpatchify は、画像テンソルとパッチ列の間の変換を行うユーティリティメソッドです。forward_loss では、mask テンソル(マスクパッチ=1、可視パッチ=0)を使って、マスクされた位置のみの損失を計算しています。norm_pix_loss=True の場合、各パッチの平均と分散で正規化されたターゲットに対して損失を計算します。

これでMAEの全コンポーネントが揃いました。次に、このモデルを実際にデータで学習させてみましょう。

学習実験と復元結果の可視化

実装したMAEの動作を確認するため、CIFAR-10データセットを使った学習実験を行います。CIFAR-10は $32 \times 32$ ピクセルの小さな画像ですが、MAEの基本的な振る舞いを検証するには十分です。

学習ループの実装

def train_mae(model, dataloader, epochs=50, lr=1.5e-4, device='cpu'):
    """MAEの事前学習ループ"""
    model = model.to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.05)
    # コサインアニーリングスケジューラ
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=epochs, eta_min=1e-6
    )

    loss_history = []
    model.train()

    for epoch in range(epochs):
        epoch_loss = 0.0
        num_batches = 0
        for imgs, _ in dataloader:  # ラベルは使わない(自己教師あり)
            imgs = imgs.to(device)
            loss, pred, mask = model(imgs)

            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

            epoch_loss += loss.item()
            num_batches += 1

        scheduler.step()
        avg_loss = epoch_loss / num_batches
        loss_history.append(avg_loss)

        if (epoch + 1) % 10 == 0:
            print(f"Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}")

    return loss_history

このコードでは、ラベル _ を一切使わずに学習を行っています。これが自己教師あり学習の特徴です。AdamWオプティマイザとコサインアニーリングスケジューラの組み合わせは、MAEの論文で採用されている設定に準じています。

実験の実行

# データの準備
transform = transforms.Compose([
    transforms.RandomHorizontalFlip(),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.4914, 0.4822, 0.4465],
                         std=[0.2470, 0.2435, 0.2616]),
])
dataset = datasets.CIFAR10(root='./data', train=True,
                           download=True, transform=transform)
dataloader = DataLoader(dataset, batch_size=256, shuffle=True,
                        num_workers=2, drop_last=True)

# モデルの構築(CIFAR-10用の小型設定)
model = MAE(
    img_size=32,
    patch_size=4,
    in_channels=3,
    encoder_dim=192,
    encoder_depth=6,
    encoder_heads=6,
    decoder_dim=96,
    decoder_depth=2,
    decoder_heads=3,
    mask_ratio=0.75,
    norm_pix_loss=True,
)

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
loss_history = train_mae(model, dataloader, epochs=50, lr=1.5e-4, device=device)

CIFAR-10の小型設定ではエンコーダが6層・192次元、デコーダが2層・96次元です。ImageNetでの本格的な実験ではViT-Large(24層・1024次元)が使われますが、ここでは動作確認のために規模を縮小しています。

学習曲線の可視化

plt.figure(figsize=(8, 4))
plt.plot(loss_history, color='#00bcd4', linewidth=2)
plt.xlabel('エポック')
plt.ylabel('再構成損失(MSE)')
plt.title('MAE事前学習の損失')
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()

学習曲線では、損失が最初の数エポックで急激に低下し、その後は緩やかに収束していく様子が観察されるはずです。最初の急激な低下は、モデルが画像の大域的な色調やテクスチャ分布といった基本的なパターンを素早く学習していることに対応します。その後の緩やかな収束は、より細かなテクスチャ構造や物体の形状といった高レベルの特徴を徐々に獲得している段階です。

復元結果の可視化

学習したMAEが実際にどの程度の品質で画像を復元できるかを可視化します。

@torch.no_grad()
def visualize_reconstruction(model, dataloader, device, num_images=8):
    """マスク・可視パッチ・復元結果を並べて表示"""
    model.eval()
    imgs, _ = next(iter(dataloader))
    imgs = imgs[:num_images].to(device)

    loss, pred, mask = model(imgs)
    pred = model.unpatchify(pred)  # (B, C, H, W)

    # 正規化を元に戻す
    mean = torch.tensor([0.4914, 0.4822, 0.4465]).to(device).view(1, 3, 1, 1)
    std = torch.tensor([0.2470, 0.2435, 0.2616]).to(device).view(1, 3, 1, 1)
    imgs_denorm = imgs * std + mean
    pred_denorm = pred * std + mean
    pred_denorm = pred_denorm.clamp(0, 1)

    # マスク可視化用:マスク位置をグレーで表示
    P = model.patch_size
    mask_img = mask.unsqueeze(-1).repeat(1, 1, P * P * 3)
    mask_img = model.unpatchify(mask_img)  # (B, C, H, W)
    visible = imgs_denorm * (1 - mask_img) + 0.5 * mask_img

    # 復元画像:可視部分は元画像、マスク部分は復元結果
    reconstructed = imgs_denorm * (1 - mask_img) + pred_denorm * mask_img

    fig, axes = plt.subplots(4, num_images, figsize=(2 * num_images, 8))
    titles = ['元画像', '75%マスク', '復元', '可視+復元']

    for i in range(num_images):
        for row, img_tensor in enumerate([
            imgs_denorm, visible, pred_denorm, reconstructed
        ]):
            img = img_tensor[i].cpu().permute(1, 2, 0).numpy()
            img = np.clip(img, 0, 1)
            axes[row, i].imshow(img)
            axes[row, i].axis('off')
            if i == 0:
                axes[row, i].set_ylabel(titles[row], fontsize=10)

    plt.suptitle('MAEの復元結果(75%マスク)', fontsize=14)
    plt.tight_layout()
    plt.show()

visualize_reconstruction(model, dataloader, device)

可視化結果では4行の画像が表示されます。1行目は元の画像、2行目はマスキング後の画像(75%がグレーで隠されている)、3行目はデコーダの復元出力、4行目は可視部分の元画像とマスク部分の復元結果を組み合わせたものです。

75%がマスクされた状態から復元された画像を見ると、大まかな色調と形状は概ね再現されているものの、細部のテクスチャは元画像と異なることがわかります。これは、MAEが「ピクセル単位の完璧な復元」ではなく「意味的に妥当な復元」を学習していることを示しています。CIFAR-10の $32 \times 32$ という小さな解像度では微細なテクスチャの復元が難しいですが、ImageNet($224 \times 224$)など高解像度の画像ではより鮮明な復元結果が得られます。

マスキング率の影響を観察する

マスキング率を変えると復元の難易度がどう変わるかも視覚的に確認してみましょう。

@torch.no_grad()
def compare_mask_ratios(model, dataloader, device, ratios=[0.25, 0.50, 0.75, 0.90]):
    """マスキング率ごとの復元結果を比較"""
    model.eval()
    imgs, _ = next(iter(dataloader))
    imgs = imgs[:4].to(device)

    mean = torch.tensor([0.4914, 0.4822, 0.4465]).to(device).view(1, 3, 1, 1)
    std = torch.tensor([0.2470, 0.2435, 0.2616]).to(device).view(1, 3, 1, 1)

    fig, axes = plt.subplots(len(ratios), 4, figsize=(8, 2 * len(ratios)))

    for row, ratio in enumerate(ratios):
        # マスキング率を一時的に変更
        original_ratio = model.mask_ratio
        model.mask_ratio = ratio
        model.encoder.eval()
        model.decoder.eval()

        loss, pred, mask = model(imgs)
        pred = model.unpatchify(pred)
        imgs_d = imgs * std + mean
        pred_d = (pred * std + mean).clamp(0, 1)

        P = model.patch_size
        mask_img = mask.unsqueeze(-1).repeat(1, 1, P * P * 3)
        mask_img = model.unpatchify(mask_img)
        visible = imgs_d * (1 - mask_img) + 0.5 * mask_img

        for col in range(4):
            vis = visible[col].cpu().permute(1, 2, 0).numpy().clip(0, 1)
            axes[row, col].imshow(vis)
            axes[row, col].axis('off')
            if col == 0:
                axes[row, col].set_ylabel(f'{int(ratio*100)}%', fontsize=12)

        model.mask_ratio = original_ratio

    plt.suptitle('マスキング率の効果', fontsize=14)
    plt.tight_layout()
    plt.show()

compare_mask_ratios(model, dataloader, device)

マスキング率の違いによる可視パッチ量

この比較を見ると、25%マスキングでは元画像のほとんどが見えており復元は容易であること、75%では大部分が隠されておりモデルが画像の構造を理解している必要があること、90%ではごく少数のパッチしか見えず復元が非常に困難であることが視覚的に確認できます。MAEが75%で最適化されている理由が体感的に理解できるのではないでしょうか。

学習と復元の実験が完了しました。次に、事前学習されたMAEエンコーダを下流タスクにファインチューニングする方法を見ていきます。

ファインチューニングと下流タスクへの転移

MAEの事前学習の目的は、デコーダによる画像復元そのものではなく、エンコーダの豊かな視覚表現の獲得です。事前学習が完了したら、デコーダを取り外し、エンコーダにタスク固有のヘッドを取り付けてファインチューニングを行います。

事前学習から下流タスクへの転移

流れは図の通りです。事前学習が済んだらデコーダは破棄し、エンコーダだけを残します。全体を更新するファインチューニングと、エンコーダを凍結して線形層だけ学ぶ線形プローブの2通りがあり、特にラベルが少ない状況でMAE事前学習の効果が大きく出ます。

ファインチューニングの手順

ファインチューニングの一般的な手順は以下の通りです。

  1. 事前学習済みMAEからエンコーダの重みを抽出する
  2. エンコーダを通常のViTとして使用する(全パッチを入力、マスキングなし)
  3. 分類ヘッドを追加する(CLSトークンまたはGlobal Average Pooling + 線形層)
  4. ラベル付きデータで全パラメータまたは一部パラメータを更新する
class ViTClassifier(nn.Module):
    """MAEで事前学習したエンコーダを使う分類モデル"""
    def __init__(self, pretrained_encoder, num_classes=10, embed_dim=192):
        super().__init__()
        # 事前学習済みエンコーダの重みを流用
        self.patch_embed = pretrained_encoder.patch_embed
        self.pos_embed = pretrained_encoder.pos_embed
        self.blocks = pretrained_encoder.blocks
        self.norm = pretrained_encoder.norm
        # 分類ヘッド(Global Average Pooling + 線形層)
        self.head = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        # 全パッチを入力(マスキングなし)
        x = self.patch_embed(x)
        x = x + self.pos_embed
        for block in self.blocks:
            x = block(x)
        x = self.norm(x)
        # Global Average Pooling
        x = x.mean(dim=1)  # (B, D)
        x = self.head(x)   # (B, num_classes)
        return x

このコードでは、事前学習済みエンコーダのパッチ埋め込み、位置埋め込み、Transformerブロック、Layer Normalizationの全重みを新しい分類モデルに移植しています。マスキングは行わず、全パッチをエンコーダに入力する点がMAE事前学習との違いです。分類にはGlobal Average Pooling(全パッチの表現の平均)を使用していますが、CLSトークンを追加する方式でも同様に動作します。

ファインチューニングの実行

def finetune(model, train_loader, test_loader, epochs=30, lr=1e-3, device='cpu'):
    """分類タスクでのファインチューニング"""
    model = model.to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.05)
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=epochs
    )
    criterion = nn.CrossEntropyLoss()

    train_accs, test_accs = [], []

    for epoch in range(epochs):
        # 学習
        model.train()
        correct, total = 0, 0
        for imgs, labels in train_loader:
            imgs, labels = imgs.to(device), labels.to(device)
            logits = model(imgs)
            loss = criterion(logits, labels)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            correct += (logits.argmax(1) == labels).sum().item()
            total += labels.size(0)
        scheduler.step()
        train_accs.append(correct / total)

        # テスト
        model.eval()
        correct, total = 0, 0
        with torch.no_grad():
            for imgs, labels in test_loader:
                imgs, labels = imgs.to(device), labels.to(device)
                logits = model(imgs)
                correct += (logits.argmax(1) == labels).sum().item()
                total += labels.size(0)
        test_accs.append(correct / total)

        if (epoch + 1) % 10 == 0:
            print(f"Epoch {epoch+1}: Train Acc={train_accs[-1]:.4f}, "
                  f"Test Acc={test_accs[-1]:.4f}")

    return train_accs, test_accs

このファインチューニングの結果、MAEで事前学習したエンコーダは、ランダム初期化から学習した同じアーキテクチャのViTと比べて、特にラベル数が限られた設定で高い精度を達成します。MAEの事前学習により、エンコーダは画像の構造的な特徴(エッジ、テクスチャ、形状、空間配置)を既に学習しているため、少量のラベルでタスク固有の判別能力を素早く獲得できるのです。

線形プロービング

ファインチューニングとは別に、線形プロービング(Linear Probing)という評価手法もあります。これはエンコーダの重みを凍結し、分類ヘッドの線形層のみを学習する方法です。エンコーダの表現の質を直接的に評価できるため、事前学習手法の比較によく使われます。

def linear_probing(model, train_loader, test_loader, epochs=30,
                   lr=1e-2, device='cpu'):
    """エンコーダを凍結し線形層のみを学習"""
    model = model.to(device)
    # エンコーダの重みを凍結
    for name, param in model.named_parameters():
        if 'head' not in name:
            param.requires_grad = False
    # 分類ヘッドのみ最適化
    optimizer = torch.optim.SGD(model.head.parameters(), lr=lr,
                                momentum=0.9, weight_decay=0.0)
    criterion = nn.CrossEntropyLoss()

    for epoch in range(epochs):
        model.train()
        for imgs, labels in train_loader:
            imgs, labels = imgs.to(device), labels.to(device)
            logits = model(imgs)
            loss = criterion(logits, labels)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

    # テスト精度
    model.eval()
    correct, total = 0, 0
    with torch.no_grad():
        for imgs, labels in test_loader:
            imgs, labels = imgs.to(device), labels.to(device)
            logits = model(imgs)
            correct += (logits.argmax(1) == labels).sum().item()
            total += labels.size(0)
    print(f"Linear Probing Test Accuracy: {correct/total:.4f}")

線形プロービングは、エンコーダの凍結された表現の上に単一の線形層しか学習しないため、エンコーダが出力する特徴ベクトルがどれだけ「直接的に分類に使える形」になっているかを測定します。MAEは生成的な事前学習(ピクセル復元)を行っているため、対照学習と比べると線形プロービングのスコアはやや低くなる傾向がありますが、ファインチューニングでは同等かそれ以上の性能を発揮します。これは、MAEが学習する表現が「表面的な分類特徴」ではなく「画像の深い構造理解」に基づいていることを示唆しています。

ここまでで、MAEの事前学習からファインチューニングまでの一連のパイプラインを実装しました。最後に、MAEのアイデアが画像以外のドメインにどのように拡張されているかを見ていきましょう。

MAEの発展 — 画像を超えて

MAEの「高マスキング率 + 非対称エンコーダ・デコーダ」という設計原理は、画像に限らずさまざまなデータモダリティに応用されています。ここでは主要な発展を概観します。

ビデオMAE(VideoMAE)

動画は時空間的な冗長性が画像よりさらに高いデータです。隣接するフレーム間はほぼ同一の内容であり、静止画以上にマスキングが有効に機能します。VideoMAE(2022年)は、動画を時空間パッチ(例: $2 \times 16 \times 16$ の時空間キューブ)に分割し、90%以上という極端に高いマスキング率を適用します。

動画のフレーム間冗長性により、75%程度のマスキングでは時間的に隣接するフレームの情報からマスクパッチを容易に補間できてしまうため、ビデオではさらに高いマスキング率が必要です。この知見は、MAEの中核原理 — 冗長性が高いほど高いマスキング率が最適 — が一般的に成り立つことを示しています。

音声MAE(Audio-MAE)

音声信号もスペクトログラム(時間-周波数の2次元表現)として扱えば、画像と類似した構造を持ちます。Audio-MAE(2022年)はスペクトログラムをパッチに分割してMAEを適用し、音声の自己教師あり表現学習を行います。

音声では時間方向と周波数方向で冗長性の度合いが異なるため、マスキング戦略にはドメイン固有の工夫が加えられています。時間方向では連続的な変化が多いのに対し、周波数方向では倍音構造のような離散的なパターンが存在するため、非一様なマスキングが効果的な場合があります。

時系列MAE

時系列データへのMAEの応用も広がっています。多変量センサーデータなどの時系列に対して、時間方向にマスキングを行い、欠損部分を復元する自己教師あり学習が行われています。

時系列MAEが特に有効なのは、正常データは大量に得られる一方で異常データのラベルが極めて少ない異常検知の場面です。自己教師あり事前学習で正常パターンの表現を獲得し、少数の異常ラベルでファインチューニングすることで、ラベル不足を補えます。これは製造設備の監視や各種インフラの予兆検知など、幅広い分野に応用できます。

マルチモーダルMAE

CLIPやImageBindなどのマルチモーダル学習への拡張も進んでいます。MultiMAE(2022年)は画像・深度マップ・セマンティックセグメンテーションなど複数のモダリティを同時に扱い、モダリティ間の相関を活かした復元を行います。

これらの発展を俯瞰すると、MAEの設計原理は以下のように一般化できます。

  1. データの冗長性を利用して高いマスキング率を設定する
  2. エンコーダは観測された部分のみを効率的に処理する
  3. 軽量デコーダで欠損部分を復元するタスクにより、エンコーダに意味的理解を強制する
  4. 事前学習後はエンコーダのみを下流タスクに転用する

この原理は、データの冗長性が存在するあらゆるドメインに適用可能です。画像・動画・音声・時系列・点群・分子構造など、自己教師あり学習の汎用的なフレームワークとしてMAEの影響力は今後も広がっていくでしょう。

まとめ

本記事では、Masked Autoencoder(MAE)の理論と実装を解説しました。

  • MAEの動機: 画像は言語と比べて空間的冗長性が高いため、15%程度の低いマスキング率では意味的な表現学習に至りません。75%という高マスキング率を採用することで、テクスチャの補間では解けない真に意味的な予測タスクを生み出しています

  • 非対称エンコーダ・デコーダ設計: エンコーダ(大型ViT)は可視パッチ(25%)のみを処理するため、計算コストが通常のViTの約1/4に削減されます。デコーダ(軽量Transformer)はマスクトークンを含む全パッチを処理して復元を行いますが、事前学習後は不要です

  • ランダムマスキングとMSE損失: シンプルな一様ランダムマスキングとピクセル値のMSE損失の組み合わせが、ブロックマスキングや離散トークン予測よりも効果的です。この「シンプルさ × 高マスキング率 = 強力な表現学習」という等式がMAEの設計哲学です

  • スケーラビリティ: エンコーダが処理するトークン数を大幅に削減できるため、大規模モデル(ViT-Huge, ViT-Giant)への事前学習が現実的な計算コストで実現でき、スケールアップに伴う性能向上が持続します

  • 広範な応用: MAEの設計原理は動画、音声、時系列、マルチモーダルデータに拡張されており、自己教師あり学習の汎用的なフレームワークとして定着しつつあります

MAEは、「データの冗長性を逆手に取る」というシンプルな着眼点から、効率性・表現の質・スケーラビリティの全てを両立させた優れた手法です。自然言語処理のBERTが生み出した「マスク予測による自己教師あり学習」のパラダイムを、画像の特性に合わせて再設計することで、視覚表現学習に新たな道を切り拓きました。

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

画像なし
Vision Transformer(ViT)の理論と実装
MAEのエンコーダの基盤であるViTの詳細なアーキテクチャと実装を解説します。
画像なし
対照学習(Contrastive Learning)の理論
MAEとは異なるアプローチの自己教師あり学習手法を比較して理解を深めましょう。
画像なし
CLIPの理論と実装
画像とテキストを結びつけるマルチモーダル表現学習を解説します。