AI・機械学習

Grad-CAM++をPyTorchで実装!画像分類AIの判定根拠を可視化する方法

AI・機械学習

前回の概要編では、Grad-CAM++がGrad-CAMと比べてどのような場面で強みを発揮するのかを、実例を交えて解説しました。今回はその判断根拠を、実際にPythonのPyTorchで動かして確認していきます。

以下に紹介するコードを実行すると、このようなgrad-cam++の出力画像を得ることが出来ます。

なお、前回のGrad-CAM実装編を読んでいただいた方には特に朗報ですが、コード上の変更点はインポートするクラス名1行だけです。それ以外の処理(モデルの読み込み、前処理、可視化)は一切変わりません。まずはこの手軽さを体感してみてください。

環境準備

この記事はGoogle Colaboratory(以下Colab)での実行を前提としています。Colab未経験の方は、先にこちらの記事を参考にColabの使い方を押さえておいてください。

まずはデバイスの確認とGrad-CAM++の計算に使うライブラリをインストールします。

import torch

# GPUが使用可能かどうかを確認します

# Colabで正しくGPU設定ができていれば、「使用するデバイス: cuda」と表示されます

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

print(f"使用するデバイス: {device}")

# ライブラリのインストールをします。

!pip install grad-cam

次に、ここで1つ注意点があります。

このライブラリは、**PyPI上のパッケージ名がgrad-camである一方、インポート時の名前はpytorch_grad_cam**です。そのため、以下のようにパッケージ名の方でインストールしようとするとエラーになります。

# NG: これだとインストールできません
!pip install pytorch_grad_cam

筆者自身、記事の検証中にこの点でModuleNotFoundErrorに遭遇しました。地味にはまりやすいポイントなので、先にここで済ませておきましょう。

モデルとデータの準備

概要編と同じ画像・モデルを使用します。ライセンスの詳細は概要編の該当セクションを参照してください(本記事では重複を避けるため割愛します)。

import urllib.request

# サンプル画像をGitHubリポジトリから取得します

image_url = "https://github.com/curiositycreates/curiosity-creates-codes/blob/main/AI/XAI/datasets/chest_xray_pneumonia/xray_fn.png?raw=true"

local_img_path = "sample_chest_xray.jpg"

urllib.request.urlretrieve(image_url, local_img_path)

print("サンプル画像のダウンロードが完了しました。")

!pip install transformers
from transformers import AutoImageProcessor, AutoModelForImageClassification

model_name = "Aunsiels/resnet-pneumonia-detection"

processor = AutoImageProcessor.from_pretrained(model_name)

model = AutoModelForImageClassification.from_pretrained(model_name)

Hugging Faceのモデルは、Grad-CAM系のライブラリがそのままでは扱えない形(辞書型)で出力を返します。そこで、必要な部分(logits)だけを取り出す薄いラッパークラスを用意します。

import torch.nn as nn

class HuggingFaceModelWrapper(nn.Module):

    """

    Hugging Faceモデルの出力(辞書型)から、

    Grad-CAMライブラリが必要とするlogits(テンソル)だけを

    取り出して返すためのラッパークラス

    """

    def __init__(self, model):

        super().__init__()

        self.model = model

    def forward(self, x):

        # Grad-CAM内部からの呼び出しに対して、

        # 辞書型ではなくテンソルを直接返すようにする

        outputs = self.model(pixel_values=x)

        return outputs.logits

# モデルをGPUに転送し、推論モードに設定

model = model.to(device)

model.eval()

# ラッパーでモデルを包む

wrapped_model = HuggingFaceModelWrapper(model)

最後に、Grad-CAM++の計算対象となる層(target_layers)を指定します。モデルの層構造は以下のコードで確認できます。

# モデル内部の層構造を確認します

# 出力の中から、最終畳み込み層に相当する部分を探します

for name, module in model.named_modules():

    print(name)

出力結果の中から、最終畳み込み層に相当する部分を確認し、以下のように指定します。

# 確認した結果に基づき、target_layerを指定します

target_layers = [model.resnet.encoder.stages[-1].layers[-1]]

この手順自体はGrad-CAMのときと変わりません。「どの層の勾配情報を使うか」という指定は、Grad-CAMとGrad-CAM++で共通です。

Grad-CAM++本体の実装

準備が整ったので、いよいよGrad-CAM++本体を実装します。Grad-CAMとの違いは、次のインポート文だけです。

import numpy as np

from PIL import Image

# ------------------------------------------

# Grad-CAMとの違いはここだけです。

# 「GradCAM」ではなく「GradCAMPlusPlus」をインポートします。

# ------------------------------------------

from pytorch_grad_cam import GradCAMPlusPlus

from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget

from pytorch_grad_cam.utils.image import show_cam_on_image

import matplotlib.pyplot as plt

あとの処理は、Grad-CAMのときとまったく同じ書き方です。

# ------------------------------------------

# 1. 画像の読み込み

# ------------------------------------------

image = Image.open(local_img_path).convert("RGB")

# ------------------------------------------

# 2. モデルによる推論

# ------------------------------------------

inputs = processor(images=image, return_tensors="pt")

inputs = {k: v.to(device) for k, v in inputs.items()}

with torch.no_grad():

    outputs = model(**inputs)

    logits = outputs.logits

# ソフトマックスで確率に変換し、予測クラスを確認

probs = torch.nn.functional.softmax(logits, dim=-1)

predicted_label_idx = torch.argmax(probs, dim=-1).item()

labels = model.config.id2label

predicted_label_name = labels[predicted_label_idx]

confidence = probs[0][predicted_label_idx].item()

print(f"AIの予測クラス: {predicted_label_name}")

print(f"確信度: {confidence:.2%}")

# ------------------------------------------

# 3. Grad-CAM++の計算

# ------------------------------------------

# 先ほど確認したtarget_layerを対象にGrad-CAM++を構築

# クラス名以外、Grad-CAMのときと書き方は同じです

cam = GradCAMPlusPlus(model=wrapped_model, target_layers=target_layers)

# 「予測されたクラス」を対象に、そのクラスの判定根拠を可視化する

targets = [ClassifierOutputTarget(predicted_label_idx)]

# Grad-CAM++を実行し、ヒートマップ(0〜1の範囲の重要度マップ)を取得

cam_output = cam(input_tensor=inputs["pixel_values"], targets=targets)

grayscale_cam = cam_output[0, :]

見ての通り、GradCAM(…)がGradCAMPlusPlus(…)に変わった以外、コードの構造は完全に同じです。

出力の確認

最後に、元画像とヒートマップを重ね合わせて表示します。

# ------------------------------------------

# 4. 元画像とヒートマップの重ね合わせ

# ------------------------------------------

# 元画像をモデルの入力サイズにリサイズし、0〜1の範囲に正規化

# 224を決め打ちにせず、processorの設定から入力サイズを取得します

# (別のモデルに差し替えた場合でも入力サイズのズレが起きないようにするためです)

input_size = processor.size.get("shortest_edge", processor.size.get("height", 224))

input_image_np = np.array(image.resize((input_size, input_size))) / 255.0

# ヒートマップを元画像に重ね合わせる

visualization = show_cam_on_image(input_image_np, grayscale_cam, use_rgb=True)

# ------------------------------------------

# 5. 結果の表示

# ------------------------------------------

plt.figure(figsize=(12, 6))

plt.subplot(1, 2, 1)

plt.imshow(image)

plt.title("Original Image")

plt.axis("off")

plt.subplot(1, 2, 2)

plt.imshow(visualization)

plt.title(f"Grad-CAM++ (Target: {predicted_label_name})")

plt.axis("off")

plt.tight_layout()

plt.show()

image_urlを書き換えれば、皆さんの手元にある別の画像でも同じように試すことができます。ただし、モデルを別のものに差し替える場合は、target_layersの指定先(最終畳み込み層の名前)がモデルごとに異なるため、model.named_modules()で確認し直す必要がある点に注意してください。

概要編で紹介した2つの実例(PNEUMONIA判定での複数領域検出、NORMAL誤判定での反応の分散)は、このコードでそのまま再現できます。

まとめ

今回はGrad-CAM++をPyTorchで実装し、実際に動かして結果を確認しました。

コード上の変更点は「インポートするクラス名1行」だけであり、Grad-CAMを実装したことがあれば、ほとんど追加の学習コストなしにGrad-CAM++を試すことができます。

概要編で触れた通り、Grad-CAM++は複数領域の検出や小さな対象の検出に強い一方、モデルの判定自体が曖昧な場合はヒートマップの反応も分散しやすいという実感も、実際にコードを動かしてみることで確認できました。ヒートマップの見た目だけで判定の正しさを判断せず、あくまで「モデルの挙動を理解するための補助情報」として活用していただければと思います。

次回は、Grad-CAMやGrad-CAM++とはさらに異なるアプローチを取る「Layer-CAM」を紹介する予定です。

Grad-CAM++と関連がある記事

Grad-CAM++とは何か? Grad-CAMとの違いを実例で解説|胸部X線AIで見る判断根拠の可視化と限界 その2

Grad-CAMとは?仕組みをわかりやすく解説|胸部X線AIで見る判断根拠の可視化と限界

Grad-CAMをPyTorchで実装!画像分類AIの判定根拠を可視化する方法

コメント

タイトルとURLをコピーしました