SW-MSAのAttention Maskについて
概要
Swin Transformer v1におけるSW-MSA(Shifted Window Multi-head Self-Attention)は、計算コストを抑えつつ長距離の依存関係を捉えるために導入されました。通常のW-MSA(Window-based Multi-head Self-Attention)は、重なりのないローカルなウィンドウ内のみでアテンションを計算するため、ウィンドウをまたぐパッチ間の関係性を計算できないという課題があります。これに対しSW-MSAは、ウィンドウの分割位置をずらす(シフトする)ことで、異なるウィンドウ間のパッチ同士を同一のウィンドウ内に配置し、アテンションを計算できるようにします。これにより、W-MSAの利点である計算の軽さを維持しながら、より大域的なパッチ間の関係性を効果的に計算することを目的としています。

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値)

これらの図を作成するために用いたコードを以下に記載します。
シフト前のパッチとシフト後のパッチの可視化コード
以下のコードでは、Swin Transformerのブロック構造を模したクラスを定義し、画像を入力した際のパッチの空間的配置がウィンドウのシフト処理によってどのように変化するかを可視化します。
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
import numpy as np
import seaborn as sns
import japanize_matplotlib
from matplotlib.colors import ListedColormap
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)
)
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 forward(self, x):
return x
if __name__ == "__main__":
B, C = 1, 4
H, W = 8, 8
window_size = 4
shift_size = 2
# データとラベルの準備
spatial_data = torch.arange(H * W, dtype=torch.float32).view(1, H, W, 1)
annot_orig = np.array([[f"V{r}H{c}" for c in range(W)] for r in range(H)])
# ウィンドウのIDグリッド (0:左上, 1:右上, 2:左下, 3:右下)
grid_ids = np.zeros((H, W))
grid_ids[0:window_size, 0:window_size] = 0
grid_ids[0:window_size, window_size:W] = 1
grid_ids[window_size:H, 0:window_size] = 2
grid_ids[window_size:H, window_size:W] = 3
# シフトされた後の状態を計算
shifted_grid_ids = np.roll(grid_ids, shift=(-shift_size, -shift_size), axis=(0, 1))
annot_shifted = np.roll(annot_orig, shift=(-shift_size, -shift_size), axis=(0, 1))
# カラーマップ設定
cmap_windows = ListedColormap(['#3498db', '#e74c3c', '#2ecc71', '#f1c40f'])
# --- 描画1: 画像パッチの空間的変化 ---
fig1, axes1 = plt.subplots(1, 2, figsize=(14, 6))
sns.heatmap(grid_ids, annot=annot_orig, fmt="", cmap=cmap_windows,
cbar=False, ax=axes1[0], annot_kws={"size": 9, "color": "white"}, alpha=0.8, linecolor='white', linewidths=0.5)
axes1[0].set_title("シフト前:ウィンドウ配置", fontsize=14)
axes1[0].axhline(4, color='black', lw=2)
axes1[0].axvline(4, color='black', lw=2)
sns.heatmap(shifted_grid_ids, annot=annot_shifted, fmt="", cmap=cmap_windows,
cbar=False, ax=axes1[1], annot_kws={"size": 9, "color": "white"}, alpha=0.8, linecolor='white', linewidths=0.5)
axes1[1].set_title(f"シフト後 (Shift={shift_size})", fontsize=14)
axes1[1].axhline(4, color='black', lw=2)
axes1[1].axvline(4, color='black', lw=2)
plt.suptitle("Swin Transformer: 窓の巡回シフトによる領域混在の可視化", fontsize=16)
plt.tight_layout()
plt.show()
このコードでは、まずSwinBlockクラス内でアテンションマスクの生成ロジックを実装しています。 クラスの初期化メソッドであるinitでは、shift_sizeが0より大きい場合(すなわちSW-MSAを適用する場合)にマスクテンソルを生成します。
マスクの生成手順は以下の通りです。
- img_maskというサイズが(1, H, W, 1)のテンソルを用意し、ウィンドウ境界とシフト領域に基づいてスライスを作成し、異なる領域ごとにユニークなインデックス(0から8)を割り振ります。
- window_partitionメソッドを用いて、画像をローカルウィンドウに分割します。この処理により、テンソルのサイズは(1, H, W, C)から(ウィンドウ数, window_size, window_size, C)へと変換されます。
- 分割されたウィンドウのマスクをフラット化し、mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)によって、ウィンドウ内のパッチ間でのインデックスの差分を算出します。
- この差分が0以外の箇所は「元々の画像において異なる領域(隣接していない領域)に属していたパッチ同士」であることを示すため、masked_fillを使用してその位置に負の大きな値である-100.0を代入します。差分が0の箇所には0.0を設定し、このテンソルをアテンションマスクであるattn_maskとして登録します。
mainブロックでは、8×8のグリッドに対して各ウィンドウ領域(0:左上, 1:右上, 2:左下, 3:右下)を表すgrid_idsを作成し、カスタムカラーマップであるcmap_windows(ListedColormap)で色分けしています。
そして、np.rollを用いてシフト量2で巡回シフトさせ、シフト前後で各ウィンドウの領域がどのように移動し、新たなウィンドウ内に混在するかをseaborn.heatmapを用いて可視化しています。これにより、シフト後の右下や右上などのウィンドウにおいて、異なる色の領域(本来隣接しない領域)が混在する様子が視覚的にわかりやすくなっています。

windowごとのAttention Maskの可視化コード
以下のコードでは、前述のクラスで生成された各ウィンドウ(Window 0からWindow 3)のアテンションマスクの値をヒートマップとして可視化します。これにより、どの位置のパッチ間でアテンションの計算がブロックされるかを視覚的に確認できます。
fig2, axes2 = plt.subplots(2, 2, figsize=(15, 13))
axes2 = axes2.flatten()
window_names = ["Window 0 (Top-Left)", "Window 1 (Top-Right)", "Window 2 (Bottom-Left)", "Window 3 (Bottom-Right)"]
for i in range(4):
mask = block_sw_msa.attn_mask[i].numpy()
indices = patch_indices_per_window[i]
# 1次元インデックス(0-63)から "V{r}H{c}" のラベルに変換
tick_labels = [f"V{idx // 8}H{idx % 8}" for idx in indices]
# マスクの描画
sns.heatmap(mask, cmap="coolwarm", vmin=-100, vmax=0, cbar=True, ax=axes2[i])
axes2[i].set_title(window_names[i], fontsize=14)
# 軸に座標ラベルをセット
axes2[i].set_xticks(np.arange(16) + 0.5)
axes2[i].set_xticklabels(tick_labels, rotation=90, fontsize=10)
axes2[i].set_yticks(np.arange(16) + 0.5)
axes2[i].set_yticklabels(tick_labels, rotation=0, fontsize=10)
axes2[i].set_xlabel("Key Patch (Attend To)", fontsize=12)
axes2[i].set_ylabel("Query Patch (Attend From)", fontsize=12)
fig2.suptitle("Attention Masks (0 = Compute Attention, -100 = Masked)", fontsize=16)
plt.tight_layout()
plt.show()
この可視化コードでは、分割された4つの各ウィンドウについて、生成されたアテンションマスクをヒートマップとしてプロットしています。
- 各ウィンドウに含まれる元のパッチインデックスをpatch_indices_per_windowから取得し、これを「V(縦の座標)H(横の座標)」の形式(例: V0H0など)の文字列ラベルに変換して軸のラベル(tick_labels)に設定しています。
- アテンションマスクの値(0または-100)をsns.heatmapによって色分けして描画します。カラーマップにはcoolwarmを用いており、マスクがかかっておらず通常通りアテンションが計算される箇所は赤色(0.0)、本来隣接しないパッチ同士のためアテンションがブロックされる箇所は青色(-100.0)で表示されます。これにより、シフトされた領域同士でのみアテンションの干渉を防ぐマスクが正しく構成されていることが視覚的に理解できます。

まとめ
本記事では、Swin Transformer v1の主要なコンポーネントであるSW-MSA(Shifted Window Multi-head Self-Attention)と、その中で不可欠なAttention Mask(アテンションマスク)の役割について解説しました。
W-MSAによる局所的な処理の限界を補うため、ウィンドウ位置をずらしてアテンションを計算するSW-MSAですが、そのシフト処理の過程で本来隣接しないパッチ同士が同じウィンドウに含まれてしまいます。この問題に対してアテンションマスク(値:-100.0)を適用し、ソフトマックス計算時に影響を実質ゼロにすることで、不自然なアテンションの発生を防いでいます。
Swin Transformerの持つ高い計算効率と、大域的なコンテキストのモデリング性能は、この精巧なマスク処理によって支えられています。可視化コードを通じて、テンソルの流れやマスクの構造を直感的に理解する一助となれば幸いです。
本記事の文章・構成の一部に生成AIを使用しています。