Swin Transformerとは?

(画像は、Google AI Studioの「Nano Banana 2」モデルを用いて作成されたものです)
Swin Transformerの概要
Swin Transformerは、Microsoft Research Asia(マイクロソフト・リサーチ・アジア:MSRA)の研究チームによって開発され、2021年に発表されました(論文はSwin Transformer: Hierarchical Vision Transformer using Shifted Windowsです)。 自然言語処理で大きな成果を上げたTransformerを画像認識に適用したVision Transformer(ViT)は、従来のCNNを凌ぐ精度を実現できるようになりました。しかし、通常のSelf-Attentionは画像内のすべてのパッチ(ピクセルの塊)同士の関連性をくまなく計算するため、画像のサイズ(画素数)が大きくなると、計算量とメモリ消費量が画素数の2乗に比例して爆発的に増えてしまうという課題がありました。
また、CNNは特徴マップを階層的にダウンサンプリングしていくことで、画像内に異なるサイズの物体が写っていても、浅い層での局所的な特徴から深い層での大域的な特徴までをマルチスケールで捉えることができます。これにより、物体検出や領域分割(セグメンテーション)を精度よく行うことが可能です。これに対して、従来のViTは特徴マップの解像度(サイズ)が全層を通じて単一(固定)であり、パッチサイズが固定されているため、様々なサイズの物体を認識することが難しく、物体検出や領域分割などの下流タスクへの適用に不向きであるというデメリットがありました。
Swin Transformerは、CNNのような「階層的な特徴マップの構築」と「アテンション領域の局所化(ウィンドウ制御)」を取り入れることで、計算コストを抑えつつ、Transformerにおいても高精度な物体検出や領域分割を実現できるようになりました。
Swin Transformerの処理概要

(画像は、Geminiを用いて作成されたものです)
Swin Transformerは上図のように、入力画像をパッチに分割した後に4つのステージの処理を実施し、Layer NormalizationとAverage Pooling、そして最終段の線形処理(分類処理の場合は分類ヘッド)を行って結果を出力します。
BasicLayerと記述しているブロックは、Swin Transformerの中核となる技術であるW-MSA(Window-based Multi-Head Attention)とSW-MSA(Shifted Window-based Multi-Head Attention)、およびPatch Mergingの処理から構成されます。 W-MSAでは画像内のすべてのパッチではなく、重複のない近接するパッチの領域(ウィンドウ)内でのみAttention処理を行い、画像の局所的な特徴を捉えるために適用されます。 一方でSW-MSAは、W-MSAで分割されたウィンドウの境界をまたぐように新しいウィンドウを構成してAttention処理を行うことで、W-MSAよりも広い大域的な特徴(ウィンドウ間の情報伝達)を捉えるために適用されます。
各ステージの最後には、Patch Mergingによって特徴マップの縦と横のサイズを1/2に縮小し、チャネル数が2倍になるように集約します(ただし、最終ステージのみチャネル数は維持されます)。Patch Mergingを適用することで、より小さな特徴マップへと情報が集約されるため、後続のステージのW-MSAやSW-MSAは画像のさらに広い範囲の特徴を扱うことが可能になります。 また、W-MSAもSW-MSAも近接するウィンドウの範囲内にあるパッチ間でのみAttention処理を行うため、すべてのパッチ間でAttentionを行う通常のVision Transformerよりも計算量が少なく、より高解像度な画像を扱うことができます。
さらに詳しくSwin Transformerの処理を理解するために、コードサンプルを用いて確認してみましょう。
Swin Transformerの実装(概念的なシンプルな実装)
以下で解説するモデルは実際のSwin Transformerに対して、相対位置バイアスの有無やパッチサイズ、ウィンドウサイズの差異がありますが、 原理を理解するために重要な要素のみに限定したSwinToyClassifierを実装し、0から学習・評価してみましょう。
本記事の実装コードは、Microsoftの公式 Swin-Transformer リポジトリ(MIT License / Copyright (c) Microsoft Corporation)のコードを参考に、解説用に簡略化したものです。
SwinToyClassifierの実装
まずはSwin Transformerによる分類器の全体像を理解するために、パッチ埋め込みからパッチマージング、W-MSA/SW-MSA、Global Average Pooling、分類ヘッドなどの大まかな処理の流れを、SwinToyClassifierクラスを用いて説明します。 Swin Transformerにおいて重要な処理を行うステージクラス(上述のBasicLayerに相当するSwinStage)や、各ブロックの内部処理については後ほど順に解説します。
import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader, random_split
class SwinToyClassifier(nn.Module):
def __init__(self, img_size=224, patch_size=2, in_chans=3, embed_dim=32, num_classes=10):
super().__init__()
# パッチ埋め込み (224x224 -> 112x112)
self.patch_embed = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
# 各ステージの input_resolution を 224x224 入力に合わせてスケールアップ
# Stage 1: 入力 112x112, 出力 56x56 (Patch Merging後)
self.stage1 = SwinStage(dim=embed_dim, depth=2, num_heads=2, window_size=4, input_resolution=(112, 112), downsample=True)
# Stage 2: 入力 56x56, 出力 28x28
self.stage2 = SwinStage(dim=embed_dim*2, depth=2, num_heads=4, window_size=4, input_resolution=(56, 56), downsample=True)
# Stage 3: 入力 28x28, 出力 14x14
self.stage3 = SwinStage(dim=embed_dim*4, depth=2, num_heads=8, window_size=4, input_resolution=(28, 28), downsample=True)
# Stage 4: 入力 14x14, 出力 14x14 (これ以上小さくしないので downsample=False)
# ※ 14x14の解像度に対して window_size=2 は綺麗に割り切れるため、元コードのウィンドウサイズ2を維持しています。
self.stage4 = SwinStage(dim=embed_dim*8, depth=2, num_heads=16, window_size=2, input_resolution=(14, 14), downsample=False)
# 分類ヘッド
self.norm = nn.LayerNorm(embed_dim * 8)
self.head = nn.Linear(embed_dim * 8, num_classes)
def forward(self, x):
# 入力画像 x の形状: [Batch, 3, 224, 224]
# 0. パッチ埋め込み層
x = self.patch_embed(x).permute(0, 2, 3, 1)
# 形状変化: [B, 3, 224, 224] -> [B, 32, 112, 112] -> [B, 112, 112, 32]
# 1. Stage 1 (SwinBlock x 2 + PatchMerging)
x = self.stage1(x)
# ブロック処理中 : [B, 112, 112, 32] (解像度・チャネル維持)
# Merging処理後 : [B, 56, 56, 64] (解像度半分 H,W / 2、チャネル2倍 C * 2)
# 2. Stage 2 (SwinBlock x 2 + PatchMerging)
x = self.stage2(x)
# ブロック処理中 : [B, 56, 56, 64]
# Merging処理後 : [B, 28, 28, 128]
# 3. Stage 3 (SwinBlock x 2 + PatchMerging)
x = self.stage3(x)
# ブロック処理中 : [B, 28, 28, 128]
# Merging処理後 : [B, 14, 14, 256]
# 4. Stage 4 (SwinBlock x 2、PatchMergingなし)
x = self.stage4(x)
# ブロック処理後 : [B, 14, 14, 256] (downsample=False のため形状維持)
# 5. Global Average Pooling と 分類ヘッド
x = x.mean(dim=[1, 2]) # 空間方向(H, W)の平均を取る -> [B, 256]
x = self.norm(x)
return self.head(x) # 最終出力: [B, 10] (10クラス分類)
SwinToyClassifierは、Swin Transformerの構造を模した簡易的な画像分類器です。入力されたテンソルが各ステージを経てどのように形状(次元)を変化させ、最終的な予測値に変換されるかを表しています。
まず、self.patch_embedというnn.Conv2d層を用いて、入力画像(224x224)をパッチに分割しながら特徴空間へと埋め込みます。kernel_size=patch_sizeおよびstride=patch_size(ここでは2)とすることで、入力画像の解像度を112x112に縮小しつつチャネル数をembed_dim(32)に拡張します。その後、permute(0, 2, 3, 1)によってテンソルの軸を並び替え、チャネルラストの形状である[Batch, Height, Width, Channel]([B, 112, 112, 32])に変換します。これはSwin Transformerの各処理ブロックがこのフォーマットで特徴マップを扱うためです。
次に、モデルは4つのSwinStageを順番に通過します。各ステージの役割は以下の通りです。
- self.stage1: 112x112の解像度を入力とし、2層のブロック処理を行った後、ステージの終わりにPatchMerging(downsample=True)を適用して、形状を[B, 56, 56, 64]へと変換します(解像度は縦横1/2、チャネル数は2倍)。
- self.stage2: 56x56の解像度を入力として処理し、同様にPatchMergingによって形状を[B, 28, 28, 128]に変換します。
- self.stage3: 28x28の解像度を入力とし、PatchMergingを経て形状を[B, 14, 14, 256]に変換します。
- self.stage4: 14x14の解像度を入力とします。これ以上のダウンサンプリングは行わないためdownsample=Falseを設定し、最終ステージとしてのブロック処理を適用して形状は[B, 14, 14, 256]のまま維持します。
最後に、分類のための後処理を行います。x.mean(dim=[1, 2])によって空間方向(HおよびW)の平均を計算するGlobal Average Poolingを実行し、形状を[B, 256]のベクトルに圧縮します。これをself.norm(nn.LayerNorm)で正規化した後、全結合層であるself.head(nn.Linear)に入力して10クラス分類のための予測ロジット[B, 10]を出力します。
Swin Transformer Stageの実装
次にSwin Transformerの各Stageの処理を行うSwinStageを実装します。
SwinBlockとPatchMergingの具体的な処理は後ほど説明します。
class SwinStage(nn.Module):
def __init__(self, dim, depth, num_heads, window_size, input_resolution, downsample=True):
super().__init__()
self.blocks = nn.ModuleList([
SwinBlock(dim=dim, num_heads=num_heads, window_size=window_size,
input_resolution=input_resolution,
shift_size=0 if (i % 2 == 0) else window_size // 2)
for i in range(depth)
])
self.downsample = PatchMerging(dim) if downsample else None
def forward(self, x):
for blk in self.blocks:
x = blk(x)
if self.downsample is not None:
x = self.downsample(x)
return x
SwinStageは、Swin Transformerにおける一つの階層(ステージ)を構築するモジュールです。複数のSwinBlockと、オプションのPatchMerging層を統合する役割を担っています。
コンストラクタ(init)では、与えられたdepth(ブロックの層数)の数だけSwinBlockを作成し、nn.ModuleListに格納しています。Swin Transformerの核心的な特徴は、ウィンドウの分割境界における情報伝達を可能にするため、通常のウィンドウ分割(W-MSA)と、ウィンドウをずらして分割する(SW-MSA)を交互に適用することです。コード内では、ループのインデックスiが偶数のときはshift_size=0(通常のW-MSA)とし、奇数のときはshift_size = window_size // 2(SW-MSA)に設定することで、この交互の処理を自動的に定義しています。
さらに、downsample=Trueが指定されている場合は、ステージの出力において解像度を縮小しつつチャネル数を拡張するためのPatchMerging層をself.downsampleとして定義します。
フォワードパス(forward)では、まずself.blocksに格納された各SwinBlockにテンソルを順次通していきます。これにより、通常のアテンションとずらしたウィンドウでのアテンションが交互に計算され、局所的および大域的な特徴が抽出されます。すべてのブロックの処理が完了した後、self.downsampleが定義されていればPatchMergingを適用してテンソル形状を縮小させ、次のステージへと渡すための出力を生成します。
Swin Transformer Blockの実装
Swin Transformer Stageの中で出現したSwinBlockを実装します。
SwinBlockはW-MSAとSW-MSAの処理を行うブロックです。
class SwinBlock(nn.Module):
def __init__(self, dim, num_heads, input_resolution, window_size=4, shift_size=0):
super().__init__()
self.input_resolution = input_resolution
self.window_size = window_size
self.shift_size = shift_size
self.num_heads = num_heads
self.norm1 = nn.LayerNorm(dim)
self.attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
self.mlp = nn.Sequential(
nn.Linear(dim, dim * 4),
nn.GELU(),
nn.Linear(dim * 4, dim)
)
# Attention Mask の事前計算
if self.shift_size > 0:
H, W = self.input_resolution
img_mask = torch.zeros((1, H, W, 1))
h_slices = (slice(0, -self.window_size),
slice(-self.window_size, -self.shift_size),
slice(-self.shift_size, None))
w_slices = (slice(0, -self.window_size),
slice(-self.window_size, -self.shift_size),
slice(-self.shift_size, None))
cnt = 0
for h in h_slices:
for w in w_slices:
img_mask[:, h, w, :] = cnt
cnt += 1
mask_windows = self.window_partition(img_mask)
mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
self.register_buffer("attn_mask", attn_mask)
else:
self.attn_mask = None
def window_partition(self, x):
B, H, W, C = x.shape
x = x.view(B, H // self.window_size, self.window_size, W // self.window_size, self.window_size, C)
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, self.window_size, self.window_size, C)
return windows
def window_reverse(self, windows, H, W):
B = int(windows.shape[0] / (H * W / self.window_size / self.window_size))
x = windows.view(B, H // self.window_size, W // self.window_size, self.window_size, self.window_size, -1)
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
return x
def forward(self, x):
B, H, W, C = x.shape
shortcut = x
x = self.norm1(x)
if self.shift_size > 0:
shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
else:
shifted_x = x
x_windows = self.window_partition(shifted_x)
x_windows_flat = x_windows.view(-1, self.window_size * self.window_size, C)
if self.attn_mask is not None:
nW = x_windows.shape[0] // B
mask_for_mha = self.attn_mask.repeat(B, 1, 1)
mask_for_mha = mask_for_mha.repeat_interleave(self.num_heads, dim=0)
attn_windows, _ = self.attn(x_windows_flat, x_windows_flat, x_windows_flat, attn_mask=mask_for_mha)
else:
attn_windows, _ = self.attn(x_windows_flat, x_windows_flat, x_windows_flat)
attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
shifted_x = self.window_reverse(attn_windows, H, W)
if self.shift_size > 0:
x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
else:
x = shifted_x
x = shortcut + x
x = x + self.mlp(self.norm2(x))
return x
SwinBlockは、Swin Transformerで局所的な特徴を抽出する基本構成ブロックであり、アテンション窓をずらさないW-MSAと、窓をずらして適用するSW-MSAの両方に対応しています。
コンストラクタ(init)では、ウィンドウ分割のヘルパーメソッドであるwindow_partitionを定義し、テンソルを重ね合わせのないローカルウィンドウに分割します。また、分割されたウィンドウを元の画像レイアウトに復元する逆変換のためのwindow_reverseも定義しています。
SW-MSA(shift_size > 0)を処理する場合、ウィンドウを巡回シフトした際に生じる「本来画像内で隣接していないパッチ同士が同じウィンドウに配置される」という問題を解決する必要があります。このため、ここではアテンションマスク(attn_mask)の事前計算を行っています。マスク生成処理では、スライス(h_slices, w_slices)によってエリアごとに異なるIDを割り振ったマスク画像を作成し、window_partitionで切り出します。そして、異なるエリアID同士の関連度スコアを非常に大きな負の値(-100.0)でマスクし、同じエリア内のみでアテンションが機能するように設計し、register_bufferを用いてモデルのバッファに保存します。
フォワードパス(forward)では、入力テンソルに対して以下の手順で処理を実行します:
- 入力テンソルの残差接続(shortcut)を保持し、self.norm1(nn.LayerNorm)を適用します。
- shift_size > 0(SW-MSA)の場合、torch.rollを用いて特徴マップを左上に巡回シフトします。
- window_partitionによって特徴マップをアテンション窓単位に分割し、アテンション計算のために次元をフラット化します。
- attn_maskが存在する場合は、アテンションスコアが異なるエリア間で干渉しないようにマスクを適用し、self.attn(nn.MultiheadAttention)でSelf-Attentionを計算します。
- アテンションを適用後、window_reverseで元のテンソル空間解像度へと復元します。
- シフトされていた場合はtorch.rollで位置を元に戻します。
- 残差を足し合わせ、さらにself.norm2によるLayer Normalizationと、2層の全結合層からなるself.mlpを適用して残差を加算します。
(補足)SW-MSAのAttention Maskの計算
SW-MSAのAttention Maskの計算は複雑なため、直感的には理解しにくいかもしれません。ここでは例として、H=8, W=8, window_size=4, shift=2の場合を考えてみます。シフト前とシフト後のパッチの配置は下図のようになります。
SW-MSAでは、このシフトされたパッチ間でAttentionを計算することで、W-MSAのウィンドウ間の関係性から、より大域的なウィンドウ間の関係性を計算しますが、縦・横にそれぞれ2シフトすると、
シフト後のパッチ配置の右上のウィンドウのように、V2H7とV2H0、V3H7とV3H0という本来画像内で隣接しないパッチ同士が同じウィンドウに配置されてしまいます。
このため、これらのパッチに負の大きな値である-100のMaskを適用することで(ソフトマックス関数を通す際にexp(-100)というほぼ0の値になるため)、本来隣接しないパッチ同士でAttentionを計算されないようになります。
各ウィンドウごとのAttention Maskの結果を下図に示します。(赤が0、青が-100のMask値)
(上記の可視化コードはSW-MSAのAttention Maskについてを参照)
Patch Mergingの実装
Swin Transformer Stageの中で出現したPatchMergingを実装します。 PatchMergingは、各Stageから重要な情報を抽出し、より小さなサイズの特徴マップに集約します。
class PatchMerging(nn.Module):
def __init__(self, dim):
super().__init__()
self.norm = nn.LayerNorm(4 * dim)
self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
def forward(self, x):
# xの形状: [B, H, W, C]
x0 = x[:, 0::2, 0::2, :] # 左上
x1 = x[:, 1::2, 0::2, :] # 左下
x2 = x[:, 0::2, 1::2, :] # 右上
x3 = x[:, 1::2, 1::2, :] # 右下
x = torch.cat([x0, x1, x2, x3], dim=-1) # [B, H/2, W/2, 4*C]
x = self.norm(x)
x = self.reduction(x) # [B, H/2, W/2, 2*C]
return x
PatchMergingは、特徴マップの空間解像度(縦・横)を半分にし、チャネル数を2倍にするためのダウンサンプリング層です。CNNにおける Pooling(プーリング)層やストライド付き畳み込み層と同様の役割を担っています。
フォワードパス(forward)では、入力テンソル [Batch, Height, Width, Channel] に対し、インデックススライス 0::2 と 1::2 を用いて、2×2ピクセルのグリッド内における左上(x0)、左下(x1)、右上(x2)、右下(x3)のピクセルを間引いて取り出します。これにより、元の特徴マップの空間解像度が縦横それぞれ半分になった4つのテンソル(それぞれ [B, H/2, W/2, C])が作成されます。
取り出された4つのテンソルをチャネル次元(dim=-1)に沿って結合(torch.cat)することで、空間解像度が半分になりチャネル数が4倍になった [B, H/2, W/2, 4*C] のテンソルを得ます。その後、self.norm(nn.LayerNorm)で正規化を行い、バイアスを無効にした全結合層(self.reduction)によってチャネル数を 4 * dim から 2 * dim に射影し、次のステージへ引き渡すための出力([B, H/2, W/2, 2*C])を生成します。
以上でSwin Transformerの分類器モデルの実装は完了です。
このモデルを学習させる具体的なコードを以下で解説します。
モデルの学習と検証にはCIFAR-10データセットを使用します。
本記事の学習および検証で使用する「CIFAR-10」データセットは、研究目的、教育目的、および個人目的での利用が広く許可されています。詳細は記事の末尾に記載しています。
CIFAR-10データセットは、Alex Krizhevsky氏、Vinod Nair氏、Geoffrey Hinton氏によって作成・提供されています。
モデルの学習
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"使用デバイス: {device}")
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
])
full_train_dataset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
test_dataset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)
train_size = 45000
val_size = 5000
train_dataset, val_dataset = random_split(full_train_dataset, [train_size, val_size])
# 【重要】224x224は画像が非常に大きく、GPUメモリを大量に消費します。
train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2)
val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, num_workers=2)
test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False, num_workers=2)
model = SwinToyClassifier().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=0.001, weight_decay=0.01)
def run_epoch(model, loader, criterion, optimizer=None, is_train=True):
model.train() if is_train else model.eval()
running_loss, correct, total = 0.0, 0, 0
with torch.set_grad_enabled(is_train):
for images, labels in loader:
images, labels = images.to(device), labels.to(device)
if is_train: optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
if is_train:
loss.backward()
optimizer.step()
running_loss += loss.item() * images.size(0)
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
return running_loss / total, 100.0 * correct / total
print("\n--- 224x224版 フルステージ 学習・検証フェーズ開始 ---")
num_epochs = 20
for epoch in range(num_epochs):
train_loss, train_acc = run_epoch(model, train_loader, criterion, optimizer, is_train=True)
val_loss, val_acc = run_epoch(model, val_loader, criterion, is_train=False)
print(f"Epoch [{epoch+1}/{num_epochs}] Train Loss: {train_loss:.4f} Acc: {train_acc:.2f}% | Val Loss: {val_loss:.4f} Acc: {val_acc:.2f}%")
print("\n--- 最終テストフェーズ開始 ---")
test_loss, test_acc = run_epoch(model, test_loader, criterion, is_train=False)
print(f"最終テスト結果 -> Test Loss: {test_loss:.4f} | Test Accuracy: {test_acc:.2f}%")
ここでは、実装した SwinToyClassifier を用いてCIFAR-10データセットの画像分類タスクを学習・評価するためのパイプラインを実装しています。
前処理(transform)では、CIFAR-10に含まれる元々32x32ピクセルの画像を、定義したSwin Transformerモデルの入力解像度仕様に合わせて transforms.Resize((224, 224)) で拡大しています。その後、ピクセル値をテンソルに変換し、一般的なImageNet等と同じ標準的な平均値と標準偏差で正規化を行っています。
データローダ設定では、CIFAR-10の学習データセットを訓練用に45,000枚、検証用に5,000枚へと random_split で分割し、それぞれ DataLoader を通してバッチサイズ16で読み込みます。解像度を224x224に拡大したことでアテンション計算とメモリ消費が劇的に増加するため、GPUのメモリ不足(OOM)を防ぐ目的でバッチサイズを小さめに調整しています。
最適化器には重み減衰(L2正則化)を適用した optim.AdamW を採用し、損失関数には多クラス分類問題に適した nn.CrossEntropyLoss を使用しています。
エポックごとのループを制御する run_epoch 関数では、is_train フラグに従ってモデルを訓練モード(model.train())と評価モード(model.eval())に切り替えています。評価時は不要なメモリ消費を抑えるために torch.set_grad_enabled(is_train) で勾配計算の有無を制御します。訓練時には、各ステップで optimizer.zero_grad() を呼んで勾配を初期化した後に損失の逆伝播(loss.backward())を行い、最適化器(optimizer.step())でモデルのパラメータを更新します。
ここでは、CIFAR-10データセットを用いた学習および検証プログラムをGoogle Colaboratoryなどの小規模なリソース環境でテスト実行した際の動作検証ログ(出力結果)を示しています。 入力解像度が224x224と大きくアテンション計算量が多い点に加え、今回の概念モデルには実際のモデルで精度向上に寄与する相対位置バイアス(Relative Position Bias)などが組み込まれていないため、実用的な時間内での性能到達は困難です。そのため、本コードではパイプラインが正しく機能するかどうかの確認を主目的とし、3エポックのみの簡易的な学習を行っています。 出力は以下のようになります(Colab ProのA100での実行結果)。
Epoch [1/3] Train Loss: 1.9983 Acc: 24.22% | Val Loss: 1.8702 Acc: 30.26%
Epoch [2/3] Train Loss: 2.0248 Acc: 23.48% | Val Loss: 2.1222 Acc: 21.62%
Epoch [3/3] Train Loss: 1.9670 Acc: 26.18% | Val Loss: 1.9502 Acc: 26.88%
最終テスト結果 -> Test Loss: 1.9530 | Test Accuracy: 26.51%
出力結果から、エポックの進行に伴って損失(Loss)が計算され、最終テストにおける精度(Accuracy)が算出される一連の処理の流れが正常に実行できていることが確認できます。
実際のHugging Faceモデルを用いた推論テスト
ここまではSwin Transformerの仕組みを理解するためにシンプルな概念モデルをゼロから実装してきましたが、実際の開発ではHugging Faceのtransformersライブラリを利用することで、事前学習済みの強力なモデルを数行のコードで呼び出して利用することができます。
以下のコードでは、ImageNet-1Kデータセットで事前学習された本物のSwin Transformerモデルをロードし、テスト用に入力した任意の画像がどのカテゴリに分類されるかを推論してみます。
import torch
import requests
import matplotlib.pyplot as plt
from PIL import Image
from transformers import AutoImageProcessor, AutoModelForImageClassification
# 1. フリー素材サイトの画像URLを指定
# 例として、Pixabayのフリー素材(猫の画像)の直接リンクを使用しています。
# ここをご自身の好きな画像のURLに差し替えてください。
image_url = "https://cdn.pixabay.com/photo/2014/11/30/14/11/cat-551554_1280.jpg"
print("画像をダウンロードしています...")
# 2. サーバーにBotと判定されて弾かれないよう、User-Agentを偽装してリクエスト
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/114.0.0.0 Safari/537.36"
}
response = requests.get(image_url, headers=headers, stream=True)
response.raise_for_status() # URLが無効な場合はエラーを出す
image = Image.open(response.raw).convert("RGB")
# 3. Swin Transformer v1 (Tiny) モデルとプロセッサのロード
model_name = "microsoft/swin-tiny-patch4-window7-224"
print(f"\nモデル '{model_name}' をロードしています...")
processor = AutoImageProcessor.from_pretrained(model_name)
model = AutoModelForImageClassification.from_pretrained(model_name)
# 4. 前処理と推論
inputs = processor(images=image, return_tensors="pt")
with torch.no_grad():
outputs = model(**inputs)
# 5. 予測結果の取得
logits = outputs.logits
predicted_class_idx = logits.argmax(-1).item()
predicted_label = model.config.id2label[predicted_class_idx]
# 結果のコンソール出力
print("-" * 40)
print(f"Swinの予測結果 (ImageNet-1K): {predicted_label}")
print("-" * 40)
# 6. サンプル画像の可視化
plt.figure(figsize=(6, 6))
plt.imshow(image)
plt.title(f"Predicted: {predicted_label}", fontsize=14)
plt.axis("off")
plt.tight_layout()
plt.show()
ここでは、Hugging Faceの transformers ライブラリを利用して、事前学習済みの本物のSwin Transformerモデルをロードし、インターネットからダウンロードした未知の猫の画像を分類するプログラムを実装しています。
画像取得処理では、requests ライブラリを使用してフリー素材サイトから直接画像をダウンロードしています。一部のサーバーによるアクセス制限を回避するために User-Agent ヘッダーを設定してブラウザ経由のアクセスを模倣しています。
モデルのロードと前処理では、ImageNet-1Kデータセットで事前学習された microsoft/swin-tiny-patch4-window7-224 を指定しています。このモデルに対応する AutoImageProcessor は、入力された画像をモデルが必要とする224x224ピクセルへのリサイズや正規化を自動で適用し、PyTorchテンソル(pt)として inputs を構成します。
推論と予測結果の取得では、torch.no_grad() を使用して無駄な勾配計算を排除した上でモデルにデータを入力し、クラスごとの予測スコア(ロジット)を取得します。得られたロジットの最大値のインデックス(argmax(-1))をモデルの設定情報に含まれるラベル辞書(model.config.id2label)と照合することで、最終的な予測結果として「Egyptian cat(エジプト猫)」などのクラス名を抽出します。最後に、matplotlib を用いて画像と予測ラベルを可視化しています。
画像をダウンロードしています...
モデル 'microsoft/swin-tiny-patch4-window7-224' をロードしています...
Loading weights: 100%
221/221 [00:00<00:00, 3500.93it/s]
----------------------------------------
Swinの予測結果 (ImageNet-1K): Egyptian cat
----------------------------------------

まとめ
本記事では、コンピュータビジョンにおける革新的なバックボーンであるSwin Transformerについて、その背景から実装までを詳しく解説しました。
記事を通じて、以下の内容を学習・実践しました。
- Swin Transformerのアーキテクチャの特徴: CNNのような階層的な特徴マップ構築と、Self-Attentionの計算量を削減するウィンドウ処理(W-MSA/SW-MSA)の重要性を理解しました。
- アテンション制限とシフトの仕組み: 重ね合わせのない局所アテンション(W-MSA)に加え、境界をまたいで情報を伝えるシフトウィンドウ(SW-MSA)と、そのアテンションマスク計算の意図を学びました。
- 概念モデルの実装: PyTorchを用いて、画像パッチ分割、W-MSA/SW-MSA、Patch Mergingなどの内部モジュールを再現し、CIFAR-10データセットを用いた画像分類パイプラインの動作を検証しました。
- Hugging Faceによる画像分類推論: 事前学習済みSwin-Tinyモデルを活用し、実際の画像データを素早く分類してクラス予測結果を可視化する実装を体験しました。
Swin Transformerは、従来のVision Transformerが抱えていた計算コストの課題をクリアし、分類タスクのみならず物体検出やセグメンテーションなど幅広いビジョンタスクで強力な性能を発揮します。本記事で学んだ基本構造を元に、ぜひ様々なビジョンアプリケーションの開発に応用してみてください。
出典・ライセンスについて
- CIFAR-10データセット: Alex Krizhevsky氏、Vinod Nair氏、Geoffrey Hinton氏によって作成・提供されている CIFAR-10 データセットを使用しています。
- Swin Transformerモデル: Microsoftが公開している Swin Transformer モデル(MITライセンス)を使用しています。
- 推論テスト用の画像: Pixabayが提供する商用利用可能なフリー素材を使用しています。
本記事の文章・構成の一部に生成AIを使用しています。