0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

Deep Metric learningやってみる

0
Posted at

知らなかった、世の中にはDeep Metric leraningというものが存在するらしい

Deep Metric learning に対する Gemini回答

データ同士の「似ている・似ていない」関係が、そのまま空間上の「距離の近さ・遠さ」に反映されるような埋め込み空間(特徴空間)をディープラーニングで学習する手法

なのでその画像が何に分類されるかを直接的に導く画像分類(classification)とは異なるタスクぽい。

ユースケースとしては以下がありそう

  • 異常検知
  • 少量データでの分類

今回も親の顔よりもみたMNISTでやってみることにします。

進め方

  1. シンプルなCNNでベクトル出力するEmbedding Model構築
  2. 損失関数はArcFaceを採用(同じクラスはより近く、異なるクラスはより遠くなるよう学習しやすくなる関数)
  3. 1クラスに対して約100枚で学習、Modelパラメータ更新
  4. 学習を終えたModelで学習画像をベクトル変換→ベクトルDB構築
  5. 評価画像をModelでベクトル変換し、ベクトルDBとコサイン類似度で総当たりをし最も近しいクラス/画像を取得

結果

10問中、10問正解(まあMNISTだし、できなかったら泣く。まだやりようは沢山あるとはいえど。)

↓左が推論用画像で、右が登録画像
1.png

2.png

所感

  • 最新のEmbeddingModelもtimmから探して試したい
  • 少数画像でベクトルで比較するから、理論上より少数データで検証ができると思う。他のタスクで試したい。
  • やっぱりAI開発楽しい!

ソースコード

import os
from glob import glob

import numpy as np
import pandas as pd
from PIL import Image
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.metrics.pairwise import cosine_similarity

import torch
from torch import nn
from torch.nn import functional as F
from torchvision import transforms as T
from pytorch_metric_learning import losses
from torch.utils.data import Dataset, DataLoader


# ============================================================
# モデル定義。とにかくシンプルな構成にしてみた。
# 出力は「分類スコア」ではなく128次元のembeddingベクトル。
# forwardの最後でL2正規化(単位ベクトル化)しているので、
# 出力同士はコサイン類似度でそのまま距離比較できる。
# ============================================================
class EmbeddingModel(nn.Module):
    def __init__(self, embedding_dim: int = 128):
        super().__init__()
        self.backbone = nn.Sequential(
            nn.Conv2d(3, 16, kernel_size=3, stride=2, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1),
            nn.ReLU(inplace=True),
            nn.AdaptiveAvgPool2d(1)
        )
        self.fc = nn.Linear(64, embedding_dim)

    def forward(self, x):
        x = self.backbone(x)
        x = x.flatten(1)
        x = self.fc(x)
        return F.normalize(x, dim=1)

# ============================================================
# 前処理
# ============================================================
class Transform:
    def __init__(self):
        pass

    def get(self):
        return T.Compose([
            T.ToTensor(),
        ])

# ============================================================
# Dataset
# ファイル名(例: label3_xxx.jpg)からラベル番号を抽出
# ============================================================
class CustomDataset(Dataset):
    def __init__(self, device: str, img_paths: list):
        self.device = device
        self.img_paths = img_paths
        self.transform = Transform().get()

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

    def __getitem__(self, idx):
        img = self.img_paths[idx]
        label = os.path.basename(img)
        label = torch.tensor(
            int(
                label[:label.rfind('_')].replace('label', '')
            )
        ).to(self.device)
        img = self.transform(
            Image.open(img).convert('RGB')
        ).to(self.device)
        return img, label


class CustomDataLoader:
    def __init__(self, batch_size: int, device: str = 'mps'):
        self.batch_size = batch_size
        self.device = device

    def create(self, img_paths: list):
        dataset = CustomDataset(device=self.device, img_paths=img_paths)
        return DataLoader(dataset=dataset, batch_size=self.batch_size, shuffle=True, drop_last=True)


# ============================================================
# ハイパーパラメータ
# ============================================================
LR = 1e-3
EPOCH = 500
DEVICE = 'mps'
BATCH_SIZE = 32
NUM_CLASSES = 10
EMBEDDING_SIZE = 128
LABELS = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]

# モデルオブジェクト生成
model = EmbeddingModel().to(DEVICE)

# ============================================================
# データ読み込み・分割
# dfのカラムはpath(画像までのパス), label(ラベル。0~9の値が格納)とする
# test_df はラベルごとに1枚だけ残す
# ============================================================
df = pd.read_csv('./datasets/dataset.csv')
train_df, test_df = train_test_split(df, test_size=0.1, shuffle=True, stratify=df['label'])
train_dataloader = CustomDataLoader(batch_size=BATCH_SIZE, device=DEVICE).create(img_paths=train_df['path'].values)
# テストデータセットは各ラベル(0~9)の1枚ずつのデータセットとする
test_df = test_df.drop_duplicates(subset=['label']).sort_values(by='label')


# ============================================================
# 損失関数・最適化関数の定義
# ArcFaceLoss は内部に学習可能なクラス代表ベクトル(num_classes x embedding_size)
# を持つため、criterion.parameters() も optimizer に含める必要がある。
# ============================================================
criterion = losses.ArcFaceLoss(
    num_classes=NUM_CLASSES,
    embedding_size=EMBEDDING_SIZE
).to(DEVICE)

optimizer = torch.optim.AdamW(
    params=list(model.parameters()) + list(criterion.parameters()),
    lr=LR,
)


# ============================================================
# 学習ループ
# シンプルな学習ループ。検証データの活用, Early Stopping, 学習率スケジューラー等は今回無し。
# ============================================================
losses_history = []

model.train()
for epoch in range(1, EPOCH + 1):
    epoch_losses = []

    # ミニバッチ学習
    for imgs, labels in train_dataloader:
        optimizer.zero_grad()
        pred = model(imgs)
        loss = criterion(pred, labels)
        loss.backward()
        optimizer.step()
        epoch_losses.append(loss.item())

    epoch_loss = np.mean(epoch_losses)
    losses_history.append(epoch_loss)
    print(epoch, epoch_loss)


# ============================================================
# 推論・評価
# 学習済みモデルで train_df の全画像をベクトル化し、
# 「登録済みベクトルDB(vector_df)」として保持する。
# ============================================================
transform = Transform().get()

def load_img(path: str, device: str):
    return (
        transform(
            Image.open(path)
            .convert('RGB')
        )
        .unsqueeze(dim=0)
        .to(device)
    )

# ベクトルDB作成
vector_df = pd.DataFrame()
with torch.inference_mode():
    for label in LABELS:
        preds, paths = [], []
        for img in train_df[train_df['label'] == label]['path']:
            paths.append(img)
            img = load_img(path=img, device=DEVICE)
            preds.append(model(img))
        label_list = [label] * len(preds)

        a_vector_df = pd.DataFrame(
            data={
                'label': label_list,
                'vector': preds,
                'path': paths,
            }
        )
        if len(vector_df) == 0:
            vector_df = a_vector_df
        else:
            vector_df = pd.concat([vector_df, a_vector_df], axis=0)


# 可視化用コード
def plot(reg_path, reg_label, true_path, true_label):
    fig, ax = plt.subplots(dpi=120, ncols=2, figsize=(12, 5))

    ax[0].imshow(Image.open(true_path))
    ax[1].imshow(Image.open(reg_path))
    ax[0].axis('off')
    ax[1].axis('off')
    ax[0].set_title(f'True ({true_label})')
    ax[1].set_title(f'Registered ({reg_label})')

    plt.show()


# ============================================================
# 評価
# test_df(未知画像)を1枚ずつベクトル化し、vector_df内の全ベクトルと
# コサイン類似度を比較、最も近いものを予測ラベルとする。
# ============================================================
with torch.inference_mode():
    for _, test_ser in test_df.iterrows():
        img = load_img(path=test_ser['path'], device=DEVICE)
        pred = model(img)

        reg_label_box = []
        similarity_box = []
        for _, train_ser in vector_df.iterrows():
            reg_label_box.append(train_ser['label'])
            similarity_box.append(cosine_similarity(pred.to('cpu'), train_ser['vector'].to('cpu')))

        reg_label = reg_label_box[np.argmax(similarity_box)]

        print(f'True label is {test_ser["label"]} vs Prediction label is {reg_label}')
        reg_path = vector_df.iloc[np.argmax(similarity_box)]['path']
        plot(reg_path, reg_label, test_ser['path'], test_ser['label'])


以上、備忘録でした。読んでくださりありがとうございました。

0
0
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
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?