深層学習

Vision Transformer

画像を固定サイズのパッチ列に変換し、クラストークンと位置埋め込みを加えてTransformerへ入力する画像認識モデルです。

  • A|中核
  • 深層学習

数式の記号で止まったら 記号の読み方 (∂・⊙・転置・上付き添字を、読み方から)

ひとことで言うと

Vision Transformer(ViT)は、画像を固定サイズのパッチに分割し、各パッチを1個のトークンとしてTransformerへ渡す画像認識モデルです。CNNの畳み込みを積み重ねる代わりに、パッチ列の各要素を同じ形式で処理します。

画像を一枚の絵として読むのではなく、同じ大きさのタイルを左上から順にカード化し、「このカードはどの位置にあったか」という札を添えて、文章のように読む構成です。タイル同士の関係はモデルが学習します。

なぜ必要か

自然言語では、単語をトークン列にしてTransformerへ入力します。ViTはこの考え方を画像へ移し、画像固有の畳み込み演算を前提にしない構成を作りました。入力画像を H×WH \times W、パッチの一辺を PP、チャンネル数を CC とすると、パッチ数は N=HW/P2N=HW/P^2 です。つまり、PP を小さくするほど細かい情報を残せますが、Transformerが扱う系列は長くなります。

原論文が比較している重要な点は、精度だけでなく帰納バイアスです。CNNは局所性、2次元の近傍構造、並進同変性を構造として持ちます。ViTはパッチを切る段階以外では画像固有の構造をほとんど仮定しないため、中規模データで強い正則化なしに学習するとCNN系に劣る結果になります。一方、原論文では14M〜300M画像規模のデータで事前学習すると、データからパターンを学ぶ利点が現れ、CNNを上回る方向へ変わると報告しています。これは「ViTは常にCNNより優れる」ではなく、データ量と事前学習を含めた比較です。

条件効きやすい性質
中規模データでの学習CNNの局所性・並進同変性
大規模データでの事前学習ViTが空間関係をデータから学ぶスケール

仕組み

画像を x∈RH×W×Cx \in \mathbb{R}^{H \times W \times C}、パッチ一辺を PP、埋め込み次元を DD とします。画像を N=HW/P2N=HW/P^2 個のパッチへ分け、各パッチを平坦化したベクトル xpi∈RP2Cx_p^i \in \mathbb{R}^{P^2C} として、学習可能な行列 E∈RP2C×DE \in \mathbb{R}^{P^2C \times D} で写像します。

zi=xpiE+ei,i=1,…,Nz_i = x_p^i E + e_i, \qquad i=1,\ldots,N

ziz_i は ii 番目パッチの埋め込み、eie_i はその位置の学習可能な位置埋め込みです。パッチを平坦化して線形写像する部分がパッチ埋め込みです。

さらに学習可能なクラストークン xclassx_{class} を先頭に追加します。入力列は

z0=[xclass;z1;…;zN]+Eposz_0 = [x_{class}; z_1;\ldots;z_N] + E_{pos}

です。EposE_{pos} は各位置に対応する位置埋め込み行列、z0z_0 はTransformerへ入る列です。エンコーダ後の先頭要素を分類ヘッドへ渡すため、パッチ全体を平均する設計とは切り分けます。位置埋め込みがないと、同じパッチ集合の順序を区別できません。

原論文では、初期のパッチ分割と位置埋め込み以外の空間関係は学習に委ねられると説明しています。解像度を変えるとパッチ数が変わるため、実装では事前学習済み位置埋め込みを2次元補間する扱いも登場します。

系列長は NN、クラストークン込みなら N+1N+1 です。自己注意のペア計算は系列長に対して二次に増えるため、同じ画像で PP を半分にすると NN は4倍になり、注意機構のペア数は概ね16倍になります。細かいパッチは表現上の利点と計算量を交換する選択です。

したがってパッチサイズは、単なる入力前処理の設定ではありません。大きくすれば一つのトークンが広い範囲を表し、系列を短くできます。小さくすれば系列長が増えて計算量も増えます。モデル名や設定を読むときは、画像解像度、パッチサイズ、クラストークン込みの系列長を順に確認します。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
画像をどうTransformerへ入力するか固定サイズのパッチを平坦化し、線形写像した列にする画素をそのまま1トークンにする/画像全体を1トークンにする
クラストークンの役割先頭の学習可能な要素を分類表現として使う各パッチのクラスを直接表す特殊トークンと考える
位置埋め込みの役割パッチの位置情報を列へ加えるパッチ内の局所性や並進同変性をCNNと同じ形で注入する
パッチサイズを小さくした影響パッチ数と系列長が 1/P21/P^2 に比例して増えるパッチ数も計算量も単純に2倍になる
CNNとのデータ依存の比較小規模では帰納バイアス、大規模事前学習ではViTのスケールが効くViTはデータ量に関係なくCNNを上回る

実装で確かめる

NumPyでパッチ数と、クラストークン込みの注意計算の大きさを確認します。ここでは正方形画像を仮定します。

import numpy as np

image_size = 224
for patch_size in (16, 32):
    n_patches = (image_size // patch_size) ** 2
    sequence_length = n_patches + 1  # class token
    pair_count = sequence_length ** 2
    print(patch_size, n_patches, sequence_length, pair_count)

実行結果は次のとおりです。

16 196 197 38809
32 49 50 2500

16×16パッチでは224×224画像が196個のパッチになり、系列長は197です。公式ドキュメントの patch16-224 は、パッチ解像度16、画像解像度224という読み方です。

実装で num_patches と系列長を混同しないでください。位置埋め込みや分類ヘッドが扱う列は、通常、パッチ数にクラストークン1個を足した長さです。画像サイズがパッチサイズで割り切れない場合の前処理も、モデル設定と合わせて確認します。

取り違えやすいもの

用語ViTとの切り分け
CNN畳み込みによる局所性・並進同変性を組み込む。ViTは画像をパッチ列にし、空間関係を主に学習する
パッチ埋め込み各パッチを平坦化して DD 次元へ写す入力部。Transformer本体や位置埋め込みそのものではない
クラストークン分類結果を集約するために先頭へ追加する学習可能なトークン。パッチの位置を表すものではない
位置埋め込み列の各要素へ位置情報を加えるもの。CNNの局所受容野や並進同変性を丸ごと再現するものではない
ハイブリッド構成CNNの特徴マップからパッチ列を作る構成。生画像パッチを直接入力する標準ViTとは入力部が異なる

想起チェック

ViTでは画像をどのような単位で系列化するか

固定サイズのパッチを切り出し、各パッチを平坦化して線形写像したベクトルを1トークンとして並べます。

クラストークンと位置埋め込みの役割を分けて説明すると

クラストークンは分類用の表現を集約する先頭要素です。位置埋め込みは各パッチが画像のどこにあったかを列へ伝えます。

パッチサイズを半分にすると系列長と注意計算はどうなるか

画像の縦横を固定すればパッチ数は4倍です。自己注意のペア計算は系列長の二次なので、クラストークンの影響を除けば概ね16倍になります。

ViTが小規模データでCNNに劣りうる理由は何か

CNNが持つ局所性や並進同変性をViTは強く仮定しないためです。原論文の比較では、十分な大規模事前学習によってこの差が逆転する結果も示されています。

出典