0
1

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

Python初心者の備忘録 #31 ~深層学習超入門編07~

0
Last updated at Posted at 2026-07-27

はじめに

今回私は最近はやりのchatGPTに興味を持ち、深層学習について学んでみたいと思い立ちました!
深層学習といえばPythonということなので、最終的にはPythonを使って深層学習ができるとこまでコツコツと学習していくことにしました。
ただ、勉強するだけではなく少しでもアウトプットをしようということで、備忘録として学習した内容をまとめていこうと思います。
この記事が少しでも誰かの糧になることを願っております!
※投稿主の環境はWindowsなのでMacの方は多少違う部分が出てくると思いますが、ご了承ください。
最初の記事:Python初心者の備忘録 #01
前の記事:Python初心者の備忘録 #30 ~深層学習超入門編06~
次の記事:まだ

今回はモデルの負荷Deta AugmentationTransfer LearningFine TuningAutoencoderTransposed ConvolutionSegmentetionについてまとめております。

■学習に使用している資料

Udemy:②米国AI開発者がやさしく教える深層学習超入門第二弾【Pythonで実践】

■FLOPsとメモリ使用量

  • NNの計算負荷とリソース消費量の指標
    • FLOPs:浮動小数点演算の回数を表す尺度で、予測時(順伝播)に必要な浮動小数点演算の回数を表す
    • メモリ使用量: モデルのパラメータや変数によって占有されるRAMの量
  • モバイルデバイスなどのリソースが限られた環境ではFLOPsやメモリ使用量が少ないモデルが望ましい
  • 論文でモデルの負荷を表す指標として、FLOPsやメモリ使用量がよく提示されている

FLOPs(Floarting Point OPerationS)

  • NNにおける予測時(順伝播)の浮動小数点演算の回数
  • 乗算のみカウントする場合MACs (Multiply-Accumulate Operation)という
    ※1MACs ≈ 2FLOPs

image.png

PythonでMACsを算出

  • 畳み込み層のMACsを計算する関数を作成する
    nn.Conv2dのインスタンスと入力shapeを引数にとり、.weight.shape()で重みのshapeを取得
  • 全結合層の MACs を計算する関数を作成する
    nn.Linearのインスタンスと入力shapeを引数にとる
  • 出力のshapeを計算する関数を作成する
    Layerインスタンスと入力tensorを作って順伝搬させることで計算する
import torch
from torch import nn

# 出力サイズを取得する関数
def calc_output_shape(layer, input_shape):
    input = torch.randn(input_shape)
    output = layer(input)
    return tuple(output.shape)

# 畳み込み層のMACs
def calc_macs_conv2d(layer, input_shape):
    b, in_ch, in_h, in_w= input_shape
    out_ch, _, f_h, f_w = layer.weight.shape

    _, _, out_h, out_w = calc_output_shape(layer, input_shape)

    macs = b * in_ch * out_ch * f_h * f_w * out_h * out_w
    return macs

# 全結合層のMACs
def calc_macs_linear(layer, input_shape):
    b, in_features= input_shape
    out_features, _ = layer.weight.shape

    macs = b * in_features * out_features
    return macs
    
input_shape = (1, 1, 128, 128)
X = torch.randn(input_shape)

# 畳み込み層
conv_layer = nn.Conv2d(1, 8, kernel_size=3)
calc_macs_conv2d(conv_layer, input_shape) # -> 1143072

# 全結合層
linear_layer = nn.Linear(64, 10)
input_shape = (1, 64)
calc_macs_linear(linear_layer, input_shape) # -> 640

ライブラリを使ってMACsを簡単に計算

  • thopライブラリを使用することで簡単にMACsを計算できる
  • thop.profile():
    • model引数にpytorchのモデルを渡す
    • inputs引数に入力tensorを(X, y)の形でtupleで渡す
      • 通常のmodelはXのみを入力と取るので、(X, )のようにして渡す
    • MACsおよびパラメータの総数を返す
import thop
# !pip install thop

# conv2d
macs, params = thop.profile(conv_layer, (X,))
macs # -> 1143072.0 スクラッチの関数で求めたMACsと等しい

# linear
X = torch.randn(input_shape)
macs, params = thop.profile(linear_layer, (X,)) 
macs # -> 640.0 スクラッチの関数で求めたMACsと等しい

▶メモリ使用量

  • 深層学習モデルの学習中に値を保持する必要がある(=メモリ使用)
  • メモリ使用量は主に以下の3つ

image.png

■Data Augmentation

  • 学習データのバリエーションを人工的に増やすことによってモデルの汎化性能を向上させる(過学習を抑制する)
  • 画像タスクで非常に一般的に使用されるテクニック

image.png

▶よく使用されるData Augmentation

  • 画像の回転:画像をランダムな角度で回転させる
  • 画像のflip:水平(場合によっては垂直)に反転させる
  • cropping:画像からランダムに一部を切り取る
  • 色の調整:明るさ、コントラスト、彩度など色の特性を変更させる
  • ノイズの追加:画像にランダムなノイズを追加させる
  • ズーム:画像の一部を拡大/縮小させる
    ※複数のaugmentationをランダムに組み合わせるのが一般的で、学習時にのみ適用させ、検証/テストデータには使用しない

torchvision.transforms

  1. transforms.Composeクラスにさまざまなtransformのインスタンスのリストを渡す
    (Augmentationの種類を確認:https://pytorch.org/vision/0.15/transforms.html)
    • RandomHorizontalFlip:ランダムに画像を水平反転
    • RandomCrop:ランダムに画像を切り抜く
    • RandomRotation:ランダムに画像を回転
  2. Datasetクラスのtransform引数に、transforms.Composeのインスタンスを指定する
import
import numpy as np
from torchvision import transforms
from torch import optim
from torch.nn import functional as F
from torchvision.datasets import CIFAR10
from torch.utils.data import DataLoader
from torchvision.utils import make_grid
import matplotlib.pyplot as plt

%load_ext autoreload
%autoreload 2
import utils
# どのように加工するか指定
transform = transforms.Compose([
    transforms.RandomHorizontalFlip(),
    transforms.RandomCrop(32, padding=4),
    transforms.RandomRotation(30),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])

train_dataset = CIFAR10('./cifar10_data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=4)
Files already downloaded and verified
X, y = next(iter(train_loader))
X = make_grid(X).permute(1, 2, 0)
X = X / 2 + 0.5

# 描画
plt.imshow(X)

image.png

albumentations

  1. albumentarions.Composealbumentations.に続くクラスのインスタンスのリストを入れる
  2. tensorへの変換はalbumentations.pytorch.ToTensorV2を使用する
  3. wrapperクラスを作成
    transformを引数にとる
    __call__メソッドを作成し、transformを適用し画像とラベル(ラベルがない場合は画像のみ)を返す
    albumentations.Composeのインスタンスをcallする時にはnumpy arrayを引数とし、結果は[’image’]でアクセスする
  4. Datasetクラスのtransform引数にwrapperクラスのインスタンスを渡す
# 以下はOpenCVインストールに必要なコマンド
# !pip install albumentations
# !pip install opencv-python
#$sudo apt-get update
#$sudo apt-get install libgll-mesa-glx
import albumentations as A
from albumentations.pytorch import ToTensorV2
import cv2

# 条件を指定
transform = A.Compose([
    A.HorizontalFlip(),
    A.PadIfNeeded(min_height=40, min_width=40, border_mode=cv2.BORDER_CONSTANT),
    A.RandomCrop(height=32, width=32),
    A.Rotate(limit=30, border_mode=cv2.BORDER_CONSTANT),
    A.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]),
    ToTensorV2(),
])

# クラス作成
class AlbumentationsTransform:
    def __init__(self, transform):
        self.transform = transform

    def __call__(self, image, target=None):
        image = self.transform(image=np.array(image))['image']
        if target:
            return image, target
        else:
            return image

transform = AlbumentationsTransform(transform)
train_dataset = CIFAR10('./cifar10_data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=4)

X, y = next(iter(train_loader))
X = make_grid(X).permute(1, 2, 0)
X = X / 2 + 0.5

# 描画
plt.imshow(X)

image.png

imgaug

  1. imgaug.augmenters.Sequential()imgaug.augmenters.に続くaugmentation用のクラスのインスタンスのリストを渡しインスタンス作成
  2. wrapperクラスを作成
    • imgaug.augmenters.Sequential()のインスタンスを受け取り、インスタンス変数にする
    • __call__メソッドを作成、imgaug.augmenters.Sequentialのインスタンスに対して.augment_imageを実行し返す(引数にはnumpyを使用)
  3. torchvision.transforms.Composeの引数のリストにwrapperクラスのインスタンスを入れる
  4. Datasetのtransformに、torchvision.transforms.Composeのインスタンスを渡す
# !pip install imgaug
from imgaug import augmenters as iaa
import imgaug as ia

# 条件指定
transform_ia_seq = iaa.Sequential([
    iaa.Fliplr(0.5),
    iaa.Pad(px=4),
    iaa.CropToFixedSize(width=32, height=32),
    iaa.Affine(rotate=(-30, 30))
])

# クラス作成
class ImgAugTransform:
    def __init__(self, ia_seq):
        self.ia_seq = ia_seq

    def __call__(self, image, target=None):
        image = self.ia_seq.augment_image(np.array(image))
        if target:
            return image, target
        else:
            return image

transform_ia = ImgAugTransform(transform_ia_seq)
transform = transforms.Compose([
    transform_ia,
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])

train_dataset = CIFAR10('./cifar10_data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=4)
X, y = next(iter(train_loader))
X = make_grid(X).permute(1, 2, 0)
X = X / 2 + 0.5

# 描画
plt.imshow(X)

image.png

▶どのライブラリを使うべきか

  • すでに使っているものがあればそれを使う、そうでなければそれぞれの特徴を理解してやりたいことが実現できるライブラリを選択する

Pytorch:

  • ○ Pytorchとの統合が容易で簡単に始められる
  • × 機能が限られる

Albumentations:

  • ○ 高速で拡張機能が多様
  • ○ 様々なデータタイプに統一的なAPIを提供
  • × APIが少し複雑

imgaug:

  • ○ 豊富なAugmentation
  • × 場合によっては低速

▶Data Augmentationでモデル学習

# Augmentationの条件指定
# 学習用
transform = transforms.Compose([
    transforms.RandomHorizontalFlip(),
    transforms.RandomCrop(32, padding=4),
    transforms.RandomRotation(30),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])

# 検証用
transform_val = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),
])

# データ準備(Aug有り)
train_dataset = CIFAR10('./cifar10_data', train=True, download=True, transform=transform)
val_dataset = CIFAR10('./cifar10_data', train=False, download=True, transform=transform_val)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=6)
val_loader = DataLoader(val_dataset, batch_size=128, num_workers=6)

# 学習(Aug有り)
conv_model = utils.get_conv_model(in_ch=3)
opt = optim.Adam(conv_model.parameters(), lr=0.03)
train_losses, val_losses, val_accuracies = utils.learn(conv_model, train_loader, val_loader, opt, F.cross_entropy, 5)
"""
epoch: 0: train error: 1.7928908734065492, validation error: 1.7756120389020895, validation accuracy: 0.3654074367088608
epoch: 1: train error: 1.6366446469445972, validation error: 1.499374205553079, validation accuracy: 0.45510284810126583
epoch: 2: train error: 1.5788945261474765, validation error: 1.4798819249189352, validation accuracy: 0.4609375
epoch: 3: train error: 1.5431847462568746, validation error: 1.466218373443507, validation accuracy: 0.46607990506329117
epoch: 4: train error: 1.5066445809800912, validation error: 1.3947291676002214, validation accuracy: 0.49139636075949367
"""

# データ準備(Aug無し)
train_dataset_no_aug = CIFAR10('./cifar10_data', train=True, download=True, transform=transform_val)
train_loader_no_aug = DataLoader(train_dataset_no_aug, batch_size=128, shuffle=True, num_workers=4)

# 学習(Aug無し)
conv_model_no_aug = utils.get_conv_model(in_ch=3)
opt_no_aug = optim.Adam(conv_model_no_aug.parameters(), lr=0.03)
train_losses, val_losses, val_accuracies = utils.learn(conv_model_no_aug, train_loader_no_aug, val_loader, opt_no_aug, F.cross_entropy, 5)
"""
epoch: 0: train error: 1.590190362137602, validation error: 1.4580148790456071, validation accuracy: 0.4622231012658228
epoch: 1: train error: 1.3713380400177158, validation error: 1.3607859158817726, validation accuracy: 0.5154272151898734
epoch: 2: train error: 1.2878778259772474, validation error: 1.2929550288598748, validation accuracy: 0.5323378164556962
epoch: 3: train error: 1.2372591385756002, validation error: 1.2776733697215212, validation accuracy: 0.5406447784810127
epoch: 4: train error: 1.1960674631016335, validation error: 1.2441376900371117, validation accuracy: 0.5599287974683544
"""

Data Augmentationを行わないもののほうがAccuracyが高く出ているが、Aug有りのほうはepoch毎に違う画像が対象となるので、学習の進みが遅くなりやすい
しかし、最終的なAccuracyはAug有りのほうが高く出やすい
※データのバリエーションが増えるので、汎化性能が上がり、過学習を抑えられる

■Transfer Learning(転移学習)とFine Tuning

  • 転移学習:学習済みのモデルを新しいタスクに転移(transfer)する手法
    例:ImageNetの画像分類用のモデルをMNISTの画像分類に転移学習する
    • 通常学習済みのモデルの重みを固定(学習させない)し、特徴量抽出器として使用する
    • 最終層を追加したり、学習済みモデルの最終層のみ重みを学習させる
  • Fine Tuning:転移学習の特定の手法で、学習済みのモデルの一部もしくは全ての重みを再学習する

image.png

なぜ全く関係のない画像で学習したモデルの重みをそのまま使えるのか?

  • 初期の層は一般的な視覚的特徴を抽出しているので、他の画像にも有効であることが多い
  • 人間が見たら全く違う画像でも、コンピュータが見ると視覚的構造が似ている場合がある

image.png
画像のLayer1やLayer2は、縦横斜め、曲線のエッジや色の差異というシンプルな特徴を抽出している
これぐらいの特徴(重み)はどの画像でも似たものになるので、転用できるという仕組み

PythonでTransfer Learning

  • torchvision.modelsに学習済みのモデルが多く提供されている
    (https://pytorch.org/vision/stable/models.html)
    1. モデルをロードし、最終層を新しいタスク向けに書き換える
      (クラス数の変更など)
    2. 最終層以外の重みを凍結する
      • parameterオブジェクトに対して.requires_grad = Falseを指定する
    3. 手元のデータで学習ループを回す

※GPUを使用したいので、Google Colabで試してください

import torch
from torch.nn import functional as F
import torchvision.models as models
from torchvision.models import ResNet18_Weights
import torchvision.transforms as transforms
from torchvision.datasets import CIFAR10
from torch.utils.data import DataLoader
from torch import optim, nn

from google.colab import drive
drive.mount('/content/drive')
import sys
sys.path.append('/content/drive/My Drive/Colab Notebooks')

%load_ext autoreload
%autoreload 2
import utils

# GPUの設定
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")

model = models.resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)

num_classes = 10
model.fc = nn.Linear(512, num_classes)

for name, param in model.named_parameters():
    if not name.startswith('fc'):
        param.requires_grad = False

# Augの条件設定
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# データ準備
train_dataset = CIFAR10(root='./cifar10_data', train=True, download=True, transform=transform)
val_dataset = CIFAR10(root='./cifar10_data', train=False, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=1024, shuffle=True, num_workers=4)
val_loader = DataLoader(val_dataset, batch_size=1024, shuffle=False, num_workers=4)

# モデルをGPUに移動
model = model.to(device)

# 学習
opt = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)

# resnetは大きいのでCUDAが必要
train_losses, val_losses, val_accuracies, losses_in_epoch = utils.learn(model, train_loader, val_loader, opt, F.cross_entropy, 3)

■AutoencoderとTransposed Convolution

▶Autoencoder

  • エンコーダとデコーダから構成されるモデル
    • エンコーダ(encoder):入力データを低次元の潜在空間(latent space)にマッピングする(圧縮するイメージ)
    • デコーダ(decoder):エンコーダで得た潜在空間を元の入力データと同じ空間に逆変換する(圧縮を元に戻すイメージ)
  • 次元削減やノイズ除去など多岐にわたるタスクに利用され、近年流行の生成AIにも使われる

image.png

PytorchでAutoencoder

  • 2層の畳み込みによるencoderと2層の転置畳み込みによるdecoder
    • encoder: conv -> relu -> pooling -> conv -> relu -> pooling
    • decoder: tconv -> relu -> tconv -> sigmoid
      • transposed convolutionにはnn.ConvTranspose2dを使用する
      • 最終層は0~1のpixel値にするため、sigmoid関数を使用する
    • 損失にはMSEを使用する
  • 実際にMNISTデータセットで学習し、再構築してみる
import
import numpy as np
import torch
from torchvision import transforms
from torch import optim, nn
from torch.nn import functional as F
from torchvision.datasets import CIFAR10, MNIST
from torch.utils.data import DataLoader
from torchvision.utils import make_grid
import matplotlib.pyplot as plt

%load_ext autoreload
%autoreload 2
import utils
class ConvAutoencoder(nn.Module):
    def __init__(self):
        super().__init__()
        # encoder 
        self.conv1 = nn.Conv2d(1, 16, 3, padding=1)
        self.conv2 = nn.Conv2d(16, 4, 3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)

        # decoder
        self.t_conv1 = nn.ConvTranspose2d(4, 16, 2, stride=2)
        self.t_conv2 = nn.ConvTranspose2d(16, 1, 2, stride=2)

    def forward(self, X):
        # encoder
        X = self.pool(F.relu(self.conv1(X)))
        X = self.pool(F.relu(self.conv2(X)))
        
        # decoder
        X = F.relu(self.t_conv1(X))
        X = F.sigmoid(self.t_conv2(X))
        return X
        
# データ準備
transform = transforms.ToTensor()
train_dataset = MNIST('./mnist_data', train=True, download=True, transform=transform)
val_dataset = MNIST('./mnist_data', train=False, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=128)

# 学習
model = ConvAutoencoder()
opt = optim.Adam(model.parameters(), lr=0.001)
num_epochs = 10
# 学習ループ (ここでは簡易的なループを使用する)
for epoch in range(1, num_epochs+1):
    train_loss = 0.0
    for X, _ in train_loader:
        opt.zero_grad()
        outputs = model(X)
        loss = F.mse_loss(outputs, X)
        loss.backward()
        opt.step()
        train_loss += loss.item()*X.shape[0]
    train_loss = train_loss / len(train_loader)
    print(f'Epoch: {epoch}: train loss: {train_loss}')

# 検証データで実際に予測してみる
val_loader = DataLoader(val_dataset, batch_size=8)
images, labels = next(iter(val_loader))
outputs = model(images)

# 予測結果可視化
images_grid = make_grid(images)
outputs_grid = make_grid(outputs)

# 元の画像
plt.imshow(torch.permute(images_grid, (1, 2, 0)))

image.png

# 出力結果
plt.imshow(torch.permute(outputs_grid, (1, 2, 0)))

image.png

▶Transposed Convolution(転置畳み込み層)

  • 画像のアップサンプリング(拡大)を行う処理で、通常の畳み込みがデータをダウンサンプリングするのに対して逆のことをしているイメージ
  • 画像生成や、セグメンテーションなどで頻繁に使用される

image.png

転置畳み込み(アップサンプリング)のイメージ
image.png

Transposed Convolutionにおけるstride

  • 指定したstride分、拡張する

image.png

Transposed Convolutionにおけるpadding

  • paddingは出力結果から間引くイメージで、通常の畳み込みとは逆の効果
  • Transposed Convolutionでは、出力結果が小さくなる -> 出力サイズを調整するために使用する

image.png

Transposed Convolutionをスクラッチ実装

  • 引数:入力tensor重みtensorstridepadding
    • 入力tensorは[batch_size, in_ch, h_in, w_in]
    • 重みのshapeは[in_ch, out_ch, f_h, f_w]
  • 順伝播のみで、逆伝播は考慮しない
  • biasは無視してよい
  • kernel size > strideを前提とする
def conv_transpose2d(input, weight, stride=1, padding=0):
    batch_size, ch_in, h_in, w_in = input.shape
    _, ch_out, f_h, f_w = weight.shape

    h_out = stride * (h_in - 1) + f_h - 2*padding
    w_out = stride * (w_in - 1) + f_w - 2*padding

    output = torch.zeros((batch_size, ch_out, h_out, w_out))

    for b in range(batch_size):
        for i in range(ch_in):
            for o in range(ch_out):
                for h in range(h_in):
                    for w in range(w_in):
                        h_start = h * stride - padding
                        w_start = w * stride - padding
                        for f_h_idx in range(f_h):
                            for f_w_idx in range(f_w):
                            
                                # 出力位置の計算
                                h_out_idx = h_start + f_h_idx
                                w_out_idx = w_start + f_w_idx
    
                                if 0 <= h_out_idx < h_out and 0 <= w_out_idx < w_out:
                                    output[b, o, h_out_idx, w_out_idx] += input[b, i, h, w] * weight[i, o, f_h_idx, f_w_idx]

    return output    

PytorchでTransposed Convolution

  • nn.ConvTranspose2d()で、転置畳み込み層を使用することができる
    • 基本的な引数はConv2dと同じ
    • 重みのshapeは[in_ch, out_ch, f_h, f_w]であることに注意
input = torch.randn(1, 3, 5, 5)
convt_layer = nn.ConvTranspose2d(3, 4, kernel_size=3, stride=2, padding=2, bias=False) # スクラッチではbias処理を無視しているので、こちらでもbias=Flase
weight = convt_layer.weight
output_scratch = conv_transpose2d(input, weight, stride=2, padding=2)
output = convt_layer(input)

# スクラッチ実装が正しかったかライブラリと比較
torch.allclose(output, output_scratch)  # -> True? Flase?

■セグメンテーション(Segmentetion)

  • 画像処理タスクの一つで、画像を複数のセグメント(部分)に分けるタスクで、各ピクセルに対してラベルをつける

image.png

▶Semantic Segmentation vs Instance Segmentation

  • Semantic Segmentation:同じクラスであれば同じラベルをマークする(インスタンス間の区別はしない)
  • Instance Segmentation:同じクラスでもインスタンスが別なら区別する

image.png

セグメンテーションの応用例

セグメンテーションは多くの分野で使われている

  • 医用画像解析:疾患箇所をセグメンテーションする
  • 自動運転:車両や歩行者、障害物などを識別する
  • ロボティクス:周囲の物体を識別する
  • 農業:作物の健康状態を監視する
    etc…

image.png

▶深層学習によるセグメンテーション

  • 深層学習の技術を用いてセグメンテーションタスクを行うことができる
  • アーキテクチャの基本はencoder-decoderの形
    encoderで画像を低次元の特徴空間に圧縮し、decoderでupsamplingし元の画像サイズに戻す
  • 従来の画像処理(閾値処理やエッジ検出)に比べ精度が高くロバストに動作するのが一般的
  • U-Net、Mask R-CNN、DeepLabなど様々なアーキテクチャが開発されている

▶U-Net

  • セグメンテーションを行う深層学習モデルとして非常に人気
    元は細胞の画像のセグメンテーションとして発表されたが、あらゆる画像で高い精度を記録している
  • encoderとdecoderの各層間にskip connectionを作り、decoderでの特徴マップの再構築を助ける
    • エッジや細部の情報をよく捉えることができる

image.png

PytorchでUNetを実装

  • 原論文の図を参考にPytorchでU-Netのクラスを作成する
  • convolutionのブロックを重ねる
    • ブロック:conv->relu->conv->relu
    • poolingによりdown sampling
    • transposed convolutionによりup sampling
  • encoderの各ブロックの層の出力を保持し、decoderの各ブロックの層につなげる(skip connection)
  • 最終層の活性化関数は損失関数側で処理するので、クラス数を出力サイズとする畳み込み層を使えばOK
    ※細かい実装は気にせず、U-Netの全体の流れを捉えられれば良い
import
from tqdm import tqdm
import numpy as np
import matplotlib.pyplot as plt
from skimage import color
import torch
from torch import nn, optim
from torch.nn import functional as F
import torchvision
from torchvision.datasets import VOCSegmentation
from torchvision import transforms
from torch.utils.data import DataLoader
class UNet(nn.Module):

  def __init__(self, in_ch, num_classes, transposed=True):
    super().__init__()

    self.dconv_down1 = self._double_conv(in_ch, 64)
    self.dconv_down2 = self._double_conv(64, 128)
    self.dconv_down3 = self._double_conv(128, 256)
    self.dconv_down4 = self._double_conv(256, 512)
    self.dconv_down5 = self._double_conv(512, 1024)

    self.maxpool = nn.MaxPool2d(2)

    self.upconv4 = self._up_conv(1024, 512, transposed=transposed)
    self.upconv3 = self._up_conv(512, 256, transposed=transposed)
    self.upconv2 = self._up_conv(256, 128, transposed=transposed)
    self.upconv1 = self._up_conv(128, 64, transposed=transposed)

    self.dconv_up4 = self._double_conv(1024, 512)
    self.dconv_up3 = self._double_conv(512, 256)
    self.dconv_up2 = self._double_conv(256, 128)
    self.dconv_up1 = self._double_conv(128, 64)

    self.conv_last = nn.Conv2d(64, num_classes, 1)


  def _double_conv(self, in_ch, out_ch):
    return nn.Sequential(
        nn.Conv2d(in_ch, out_ch, 3, padding=1), # 現論文ではpadding=0だが、サイズが変わらないようにpadding=1に設定
        nn.ReLU(),
        nn.Conv2d(out_ch, out_ch, 3, padding=1),# 現論文ではpadding=0だが、サイズが変わらないようにpadding=1に設定
        nn.ReLU(),
    )

  def _up_conv(self, in_ch, out_ch, transposed=True):
    if transposed:
      #.ConvTranspose2d()ではチェッカーボードのようなアーティファクトが発生する
      return nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)
    else:
      # 最近では.Upsample()を使うことが増えている
      return nn.Sequential(
          nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
          nn.Conv2d(in_ch, out_ch, 1))

  def forward(self, X):
    conv1 = self.dconv_down1(X)
    X = self.maxpool(conv1)

    conv2 = self.dconv_down2(X)
    X = self.maxpool(conv2)

    conv3 = self.dconv_down3(X)
    X = self.maxpool(conv3)

    conv4 = self.dconv_down4(X)
    X = self.maxpool(conv4)

    X = self.dconv_down5(X)

    X = self.upconv4(X)
    X = self.dconv_up4(torch.cat([X, conv4], dim=1))
    X = self.upconv3(X)
    X = self.dconv_up3(torch.cat([X, conv3], dim=1))
    X = self.upconv2(X)
    X = self.dconv_up2(torch.cat([X, conv2], dim=1))
    X = self.upconv1(X)
    X = self.dconv_up1(torch.cat([X, conv1], dim=1))

    out = self.conv_last(X)

    return out

# 入力サイズは2^nを想定
X = torch.randn(1, 3, 256, 256)
model = UNet(3, 10, transposed=False)
output = model(X)
# output.shape

▶セグメンテーションの損失関数

  • 最終層から出力された特徴マップ(入力画像と同サイズ)の各ピクセルに対して交差エントロピーを計算する
    ※ダイス損失など、他の損失関数を用いることもある
  • クラスによる出現頻度に差が大きくある場合はクラス別に重み付けをすることが多い

image.png

▶IoU(Intersection Over Union)

  • 2つの領域の重なり具合を計算するために使用される
    -> 真の領域と予測した領域との重なり具合を測定することでセグメンテーションの結果を評価する

image.png

▶Pascal VOCデータセット

-セグメンテーションのデータセットとしてよく使われる

  • torchvision.datasets.VOCSegmentation
    【必要なtransform】
    • (256, 256)にresize
    • ToTensor
    • Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    • target(ラベル)には255を掛けてlong型にする
  • 可視化の際には正規化を戻し0~255にする
    • 正規化に使った平均、stdを使う
    • 0~1にクリップしたのちに0~255にrescaleする
  • ラベル毎に色を適用する
    • skimage.color.label2rgb()
def mask_to_tensor(mask):
  return (transforms.ToTensor()(mask) * 255).long()

# 画像用のtransform
transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# mask用のtransform
target_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    mask_to_tensor,
])

trainset = VOCSegmentation(root='./voc_data', year='2012', image_set='train', download=True, transform=transform, target_transform=target_transform)
valset = VOCSegmentation(root='./voc_data', year='2012', image_set='val', download=True, transform=transform, target_transform=target_transform)
trainloader = DataLoader(trainset, batch_size=32, shuffle=True)
valloader = DataLoader(valset, batch_size=32)

# 正規化を元に戻す
image, mask = trainset[0]
mean = np.array([0.485, 0.456, 0.406])
std = np.array([0.229, 0.224, 0.225])
image_denormalized = image.numpy() * std[:, None, None] + mean[:, None, None]
image_clipped = np.clip(image_denormalized, 0, 1) # 0~1にclipping
image_rescaled = (image_clipped * 255).astype(np.uint8)

# 描画
# plt.imshow(image_rescaled.transpose(1, 2, 0))
# plt.imshow(mask.permute(1, 2, 0)) # 0~255の256段階なので,普通に画像を表示するだけだとあまりマスクの色に違いがみれない
colored_mask = color.label2rgb(mask[0].numpy(), image_rescaled.transpose(1, 2, 0))
plt.imshow(colored_mask)

image.png

U-Netを実際に学習

  • PASCALデータセットから2クラス問題を作成する
    例: person(ID=15) vs 背景(ID=0)
    • wrapperクラスを作成し、personが映っていない画像は排除し、他のマスクの値は0にする
  • Data Augmentationを使用する
  • クラス別の重みを作成する
    • 学習データにおける各クラスの出現頻度をカウントし,ベクトル化する
    • nn.CrossEntropyLoss()のweight引数に渡す
  • セグメンテーションの結果を描画
2クラス分類のデータセットを作る
# 21クラス分類だと学習に時間がかかるため、ここではシンプルな2クラス分類(person vs 背景)にする
class CustomVOCSegmentation(VOCSegmentation):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.person_class_id = 15
        self.data, self.targets = self.filter_dataset()

    def __len__(self):
        return len(self.data)

    def __getitem__(self, index):
        img = self.data[index]
        mask = self.targets[index]
        return img, mask

    def filter_dataset(self):
        new_data = []
        new_targets = []
        for i in range(super().__len__()):
            img, mask = super().__getitem__(i)
            mask = (mask == self.person_class_id).long()
            if torch.sum(mask) > 0:
                new_data.append(img)
                new_targets.append(mask)
        return new_data, new_targets

# Data Augmetnationを含む前処理
train_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.RandomCrop(256, padding=4),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

train_target_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.RandomCrop(256, padding=4),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10, fill=(0,)), # 余白は背景とする (255をセットし,損失関数でignore_index=255にするのでもOK)
    mask_to_tensor,
    transforms.Lambda(lambda x: x.squeeze(0)) # DataLoaderから取得する時点で[b, 1, h, w]ではなく[b, h, w]にする (train loopで処理してもよい)
])

# validationではdata augmentaitonをしないので,別途用意
val_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

val_target_transform = transforms.Compose([
    transforms.Resize((256, 256)),
    mask_to_tensor,
    transforms.Lambda(lambda x: x.squeeze(0))
])
trainset_person = CustomVOCSegmentation(root='./voc_data', year='2012', image_set='train', download=True, transform=train_transform, target_transform=train_target_transform)
valset_person = CustomVOCSegmentation(root='./voc_data', year='2012', image_set='val', download=True, transform=val_transform, target_transform=val_target_transform)
trainloader = DataLoader(trainset_person, batch_size=4, shuffle=True, num_workers=4) # 有料版では複数のスレッドを使用可能
valloader = DataLoader(valset_person, batch_size=4, num_workers=4)
クラスの重み計算
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# クラスの重みつけ計算
num_classes = 2
class_sample_counts = torch.zeros(num_classes, dtype=torch.int64)
for _, masks in trainloader:
    for mask in masks:
        mask = mask[mask != 255] # すでにDatasetのWrapperクラスで対処ずみだが,21クラス分類にする場合は必要
        class_sample_counts += torch.bincount(mask.flatten(), minlength=num_classes)
class_sample_counts = class_sample_counts.float() + 1e-5
# クラスの出現頻度の逆数を重みにする
class_weights = 1. / class_sample_counts
class_weights = class_weights / class_weights.sum()
class_weights = class_weights.to(device)
モデル、Optimizer、損失関数定義
model = UNet(in_ch=3, num_classes=num_classes)
model = model.to(device)
opt = optim.Adam(model.parameters(), lr=0.0001)
loss_func = nn.CrossEntropyLoss(ignore_index=255, weight=class_weights) # すでにDatasetのWrapperクラスで対処ずみだが,21クラス分類にする場合はignoreする必要がある
学習ループ
num_epochs = 30
save_interval = 10  # 10エポック毎にモデルを評価/保存

for epoch in range(num_epochs):
    model.train()
    running_loss = 0.0
    for images, masks in tqdm(trainloader, total=len(trainloader), desc="Training", leave=False):
        images = images.to(device)
        masks = masks.to(device)
        # lossが(b, h, w)しか受け付けない -> transformでtransforms.Lambda(lambda x: x.squeeze(0))を実施していれば不要
        # masks = masks.squeeze(1)
        opt.zero_grad()

        outputs = model(images)
        loss = loss_func(outputs, masks)
        loss.backward()
        opt.step()

        running_loss += loss.item()

    print(f'Epoch [{epoch + 1}/{num_epochs}], Training Loss: {running_loss / len(trainloader):.4f}')

    # 各10エポック毎にモデルの評価
    if (epoch + 1) % save_interval == 0:
        model.eval()
        val_loss = 0.0

        with torch.no_grad():
            for images, masks in valloader:
                images = images.to(device)
                masks = masks.to(device)
                # masks = masks.squeeze(1)

                outputs = model(images)
                loss = loss_func(outputs, masks)

                val_loss += loss.item()

        print(f'Epoch [{epoch + 1}/{num_epochs}], Validation Loss: {val_loss / len(valloader):.4f}')

        # モデルの保存
        torch.save(model.state_dict(), f'unet_epoch_{epoch+1}.pth')
モデルの予測結果可視化
# モデルの予測結果の描画
model.eval()
with torch.no_grad():
    images, masks = next(iter(valloader))
    images = images.to(device)
    masks = masks.to(device)
    outputs = model(images)
    # ひとまず最も値が大きいクラスを出力とするが,実際にはアプリケーションに応じて閾値を決める
    _, predicted_masks = torch.max(outputs, 1)

    # TensorをGPUからCPUに移動する
    images = images.cpu()
    predicted_masks = predicted_masks.cpu()
    masks = masks.cpu()

    index = 2

    image = images[index].permute(1, 2, 0)
    predicted_mask = predicted_masks[index]
    mask = masks[index]

    # Tensor -> Numpy Array
    image = image.numpy()
    predicted_mask = predicted_mask.numpy()
    mask = mask.numpy()

    fig, ax = plt.subplots(1, 3)

    # Plot image
    ax[0].imshow(image)
    ax[0].title.set_text('Original Image')

     # Plot mask
    ax[1].imshow(mask, cmap='gray')
    ax[1].title.set_text('Ground Truth Mask')

    # Plot prediction
    ax[2].imshow(predicted_mask, cmap='gray')
    ax[2].title.set_text('Predicted Mask')

    plt.show()

image.png

モデルの生の出力を可視化
pred_map = outputs[index, 1, :, :]
plt.imshow(pred_map.cpu())

黄色は人と認識しており値が高く、暗い部分は値が低く、人以外と認識している
image.png

VOCSegmentationクラスでは、imageとmaskのtransformが別々に実行されるため、それぞれのtransformが同期する様にtransformを処理するカスタムクラスが必要となる。
(たとえば、Albumentationを使用するなどでこの問題を回避することもできる)

実装例
import random
class CustomTransform:
    def __init__(self, image_transform, mask_transform=None):
        self.image_transform = image_transform
        self.mask_transform = mask_transform if mask_transform else image_transform
 
    def __call__(self, image, mask):
        seed = np.random.randint(2147483647)  # Random seedを生成
        random_state = np.random.RandomState(seed)  # Random stateを生成
 
        # PIL ImageをTensorに変換する前のTransformを作成
        if self.image_transform is not None:
            torch.manual_seed(seed) # 画像TransformのためのRandom seedをセット
            image = self.image_transform(image)
 
        # マスクに対するTransform
        if self.mask_transform is not None:
            torch.manual_seed(seed) # マスクTransformのためのRandom seedをセット
            mask = self.mask_transform(mask)
 
        return image, mask
 
class CustomVOCSegmentation(VOCSegmentation):
    def __init__(self, *args, transform, target_transform=None, **kwargs):
        super().__init__(*args, transform=None, target_transform=None, **kwargs) # 親クラスでのtransformを無効化
        self.custom_transform = CustomTransform(transform, target_transform) # imageとtargetのtransformを同期するためのカスタムのtransformを使用
        self.person_class_id = 15
        self.data, self.targets = self.filter_dataset()
 
    def __len__(self):
        return len(self.data)
 
    def __getitem__(self, index):
        img = self.data[index]
        mask = self.targets[index]
        # 毎epochで異なるtransformをするため,毎回変換をかける。
        img, mask = self.custom_transform(img, mask)
        # 人以外は0にする。
        mask = (mask == self.person_class_id).long()
        return img, mask
 
    def filter_dataset(self):
        new_data = []
        new_targets = []
        for i in range(super().__len__()):
            original_img, original_mask = super().__getitem__(i) # ここではVOCSemgmentationクラスによるtransformはされない
            img, mask = self.custom_transform(original_img, original_mask) # オリジナルの画像は保持しておく(__getitem__時に毎回custom_tansformを呼び出せるようにしたいので)
            mask = (mask == self.person_class_id).long()
            if torch.sum(mask) > 0:
                new_data.append(original_img)
                new_targets.append(original_mask)
        return new_data, new_targets

次の記事

まだ

0
1
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
1

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?