CAAE (Conditional Adversarial Autoencoder) による顔画像の年齢変換(PyTorch,Python 3.12 を使用)(Windows 上)

CAAE (Conditional Adversarial Autoencoder) は, Age Progression/Regression by Conditional Adversarial Autoencoder で提案されたモデルで, 1枚の顔画像から,さまざまな年齢の顔画像を合成する.

ここでは,PyTorch を用いて,CAAE のネットワーク構成(Encoder,Generator,画像に対する識別器,潜在変数に対する識別器)を実装し,UTKFace データセットを用いた学習とテストを行う.

参考文献: https://openaccess.thecvf.com/content_cvpr_2017/papers/Zhang_Age_ProgressionRegression_by_CVPR_2017_paper.pdf

参考にした TensorFlow 実装(ZZUTK/Face-Aging-CAAE)の Webページ: https://github.com/ZZUTK/Face-Aging-CAAE

前準備

Python 3.12 のインストール(Windows 上) [クリックして展開]

以下のいずれかの方法で Python 3.12 をインストールする。Python がインストール済みの場合、この手順は不要である。

方法1:winget によるインストール

管理者権限コマンドプロンプトで以下を実行する。管理者権限のコマンドプロンプトを起動するには、Windows キーまたはスタートメニューから「cmd」と入力し、表示された「コマンドプロンプト」を右クリックして「管理者として実行」を選択する。

winget install -e --id Python.Python.3.12 --scope machine --silent --accept-source-agreements --accept-package-agreements --override "/quiet InstallAllUsers=1 PrependPath=1 AssociateFiles=1 InstallLauncherAllUsers=1"

--scope machine を指定することで、システム全体(全ユーザー向け)にインストールされる。このオプションの実行には管理者権限が必要である。インストール完了後、コマンドプロンプトを再起動すると PATH が自動的に設定される。

方法2:インストーラーによるインストール

  1. Python 公式サイト(https://www.python.org/downloads/)にアクセスし、「Download Python 3.x.x」ボタンから Windows 用インストーラーをダウンロードする。
  2. ダウンロードしたインストーラーを実行する。
  3. 初期画面の下部に表示される「Add python.exe to PATH」に必ずチェックを入れてから「Customize installation」を選択する。このチェックを入れ忘れると、コマンドプロンプトから python コマンドを実行できない。
  4. 「Install Python 3.xx for all users」にチェックを入れ、「Install」をクリックする。

インストールの確認

コマンドプロンプトで以下を実行する。

python --version

バージョン番号(例:Python 3.12.x)が表示されればインストール成功である。「'python' は、内部コマンドまたは外部コマンドとして認識されていません。」と表示される場合は、インストールが正常に完了していない。

NVIDIA製GPUを使用する場合,NVIDIA ドライバのインストールが必要である.PyTorch の GPU 版パッケージには CUDA の実行時ライブラリが同梱されているため,NVIDIA CUDA ツールキットや cuDNN を別途システムにインストールする必要はない.

サイト内の関連ページ

PyTorch のインストール

コマンドプロンプト管理者として実行し,次のコマンドを実行する.

NVIDIA製GPU搭載のパソコンで,CUDA 12.6 を使用する場合.

python -m pip install -U torch torchvision --index-url https://download.pytorch.org/whl/cu126
python -m pip install -U numpy pillow

GPU を使用しない場合(CPU 版).

python -m pip install -U torch torchvision --index-url https://download.pytorch.org/whl/cpu
python -m pip install -U numpy pillow
使用しているGPUに対応するCUDAのバージョンは,PyTorch 公式サイトのインストール案内ページで確認できる.

UTKFace (Large Scale Face Dataset) のダウンロードと展開(解凍)

UTKFace (Large Scale Face Dataset) は,顔画像のデータセットである.
  1. Web ブラウザで次の URL を開く.

    https://susanqq.github.io/UTKFace/

  2. Aligned & Cropped Faces」データファイルを選ぶ.
  3. UTKFace.tar.gz」を選ぶ.

    別の方は使わない.

  4. ダウンロードするため,「ダウンロード (DOWNLOAD)」をクリックする.
  5. ダウンロードが始まるので確認する.
  6. ダウンロードしたファイルを展開(解凍)する.
    Windows での展開(解凍)に便利な 7-Zip: 別ページ »で説明

    tar.gz 形式ファイルを 7-Zip で展開(解凍)すると tar 形式ファイルができ, tar 形式ファイルを 7-Zip で展開(解凍)すると,画像ファイルの入ったディレクトリが得られる.

  7. 展開(解凍)してできたディレクトリ UTKFace を,%HOMEPATH%\caae\data\UTKFace となるように配置する.
    mkdir %HOMEPATH%\caae\data
    REM 展開してできた UTKFace ディレクトリを %HOMEPATH%\caae\data の下に移動する
    
  8. ディレクトリ UTKFace の下に多数の顔画像ファイルがあることを確認する.

CAAE のネットワーク構成

次のPython プログラムを,%HOMEPATH%\caae\models.py として保存する.

Encoder(顔画像を潜在ベクトルに変換),Generator(潜在ベクトルと年齢・性別のラベルから顔画像を生成),DiscriminatorZ(潜在ベクトルの分布を一様分布に近づける識別器),DiscriminatorImg(生成画像の写実性を判定する識別器)の4つのネットワークからなる.

import torch
import torch.nn as nn

N_AGE_GROUPS = 10   # 年齢を 10 の区分に分類
N_Z = 50            # 潜在ベクトルの次元数


class Encoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(3, 64, 5, stride=2, padding=2), nn.ReLU(inplace=True),
            nn.Conv2d(64, 128, 5, stride=2, padding=2), nn.ReLU(inplace=True),
            nn.Conv2d(128, 256, 5, stride=2, padding=2), nn.ReLU(inplace=True),
            nn.Conv2d(256, 512, 5, stride=2, padding=2), nn.ReLU(inplace=True),
            nn.Conv2d(512, 1024, 5, stride=2, padding=2), nn.ReLU(inplace=True),
        )
        self.fc = nn.Linear(1024 * 4 * 4, N_Z)

    def forward(self, x):
        x = self.conv(x)
        x = x.view(x.size(0), -1)
        return torch.tanh(self.fc(x))


class Generator(nn.Module):
    def __init__(self):
        super().__init__()
        in_dim = N_Z + N_AGE_GROUPS + 1  # 潜在ベクトル + 年齢の one-hot + 性別
        self.fc = nn.Linear(in_dim, 1024 * 4 * 4)
        self.deconv = nn.Sequential(
            nn.ConvTranspose2d(1024, 512, 5, stride=2, padding=2, output_padding=1), nn.ReLU(inplace=True),
            nn.ConvTranspose2d(512, 256, 5, stride=2, padding=2, output_padding=1), nn.ReLU(inplace=True),
            nn.ConvTranspose2d(256, 128, 5, stride=2, padding=2, output_padding=1), nn.ReLU(inplace=True),
            nn.ConvTranspose2d(128, 64, 5, stride=2, padding=2, output_padding=1), nn.ReLU(inplace=True),
            nn.ConvTranspose2d(64, 3, 5, stride=2, padding=2, output_padding=1), nn.Tanh(),
        )

    def forward(self, z, age_onehot, gender):
        x = torch.cat([z, age_onehot, gender], dim=1)
        x = self.fc(x)
        x = x.view(x.size(0), 1024, 4, 4)
        return self.deconv(x)


class DiscriminatorZ(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(N_Z, 64), nn.ReLU(inplace=True),
            nn.Linear(64, 32), nn.ReLU(inplace=True),
            nn.Linear(32, 16), nn.ReLU(inplace=True),
            nn.Linear(16, 1), nn.Sigmoid(),
        )

    def forward(self, z):
        return self.net(z)


class DiscriminatorImg(nn.Module):
    def __init__(self):
        super().__init__()
        in_ch = 3 + N_AGE_GROUPS + 1
        self.conv = nn.Sequential(
            nn.Conv2d(in_ch, 16, 5, stride=2, padding=2), nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(16, 32, 5, stride=2, padding=2), nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(32, 64, 5, stride=2, padding=2), nn.LeakyReLU(0.2, inplace=True),
            nn.Conv2d(64, 128, 5, stride=2, padding=2), nn.LeakyReLU(0.2, inplace=True),
        )
        self.fc = nn.Sequential(
            nn.Linear(128 * 8 * 8, 1024), nn.LeakyReLU(0.2, inplace=True),
            nn.Linear(1024, 1), nn.Sigmoid(),
        )

    def forward(self, img, age_onehot, gender):
        n, _, h, w = img.size()
        age_map = age_onehot.view(n, N_AGE_GROUPS, 1, 1).expand(n, N_AGE_GROUPS, h, w)
        gender_map = gender.view(n, 1, 1, 1).expand(n, 1, h, w)
        x = torch.cat([img, age_map, gender_map], dim=1)
        x = self.conv(x)
        x = x.view(n, -1)
        return self.fc(x)

データセットの読み込み

次のPython プログラムを,%HOMEPATH%\caae\dataset.py として保存する.

UTKFace のファイル名([年齢]_[性別]_[race]_[日時].jpg)から年齢と性別のラベルを取り出す.

import os
from PIL import Image
import torch
from torch.utils.data import Dataset
import torchvision.transforms as transforms

N_AGE_GROUPS = 10

# 年齢区分の境界(0-5, 6-10, 11-15, 16-20, 21-30, 31-40, 41-50, 51-60, 61-70, 71以上)
AGE_BOUNDARIES = [5, 10, 15, 20, 30, 40, 50, 60, 70]


def age_to_group(age):
    for i, boundary in enumerate(AGE_BOUNDARIES):
        if age <= boundary:
            return i
    return N_AGE_GROUPS - 1


class UTKFaceDataset(Dataset):
    def __init__(self, root_dir, image_size=128):
        self.root_dir = root_dir
        self.filenames = []
        for filename in os.listdir(root_dir):
            if not filename.lower().endswith(('.jpg', '.jpeg', '.png')):
                continue
            parts = filename.split('_')
            if len(parts) < 2:
                continue
            self.filenames.append(filename)
        self.transform = transforms.Compose([
            transforms.Resize((image_size, image_size)),
            transforms.ToTensor(),
            transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]),
        ])

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

    def __getitem__(self, idx):
        filename = self.filenames[idx]
        parts = filename.split('_')
        age = int(parts[0])
        gender = int(parts[1])  # 0: 男性, 1: 女性
        age_group = age_to_group(age)

        img = Image.open(os.path.join(self.root_dir, filename)).convert('RGB')
        img = self.transform(img)

        return img, age_group, gender

学習

次のPython プログラムを,%HOMEPATH%\caae\train.py として保存する.

再構成損失(L1),画像に対する GAN 損失,潜在ベクトルに対する GAN 損失,Total Variation 損失を組み合わせて学習する.

import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision.utils import save_image

from models import Encoder, Generator, DiscriminatorZ, DiscriminatorImg, N_AGE_GROUPS, N_Z
from dataset import UTKFaceDataset


def total_variation_loss(img):
    diff_h = img[:, :, 1:, :] - img[:, :, :-1, :]
    diff_w = img[:, :, :, 1:] - img[:, :, :, :-1]
    return diff_h.abs().mean() + diff_w.abs().mean()


def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

    data_dir = os.path.expanduser(r"~\caae\data\UTKFace")
    output_dir = os.path.expanduser(r"~\caae\output")
    os.makedirs(output_dir, exist_ok=True)

    dataset = UTKFaceDataset(data_dir, image_size=128)
    dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=2, drop_last=True)

    netE = Encoder().to(device)
    netG = Generator().to(device)
    netDz = DiscriminatorZ().to(device)
    netDimg = DiscriminatorImg().to(device)

    optE = torch.optim.Adam(netE.parameters(), lr=0.0002, betas=(0.5, 0.999))
    optG = torch.optim.Adam(netG.parameters(), lr=0.0002, betas=(0.5, 0.999))
    optDz = torch.optim.Adam(netDz.parameters(), lr=0.0002, betas=(0.5, 0.999))
    optDimg = torch.optim.Adam(netDimg.parameters(), lr=0.0002, betas=(0.5, 0.999))

    bce = nn.BCELoss()
    l1 = nn.L1Loss()

    n_epochs = 50

    for epoch in range(n_epochs):
        for i, (img, age_group, gender) in enumerate(dataloader):
            batch_size = img.size(0)
            img = img.to(device)
            age_onehot = F.one_hot(age_group, N_AGE_GROUPS).float().to(device)
            gender = gender.float().view(-1, 1).to(device)

            real_label = torch.ones(batch_size, 1, device=device)
            fake_label = torch.zeros(batch_size, 1, device=device)

            # ---- Encoder,Generator の学習(再構成損失 + GAN損失) ----
            optE.zero_grad()
            optG.zero_grad()

            z = netE(img)
            reconst = netG(z, age_onehot, gender)

            loss_l1 = l1(reconst, img)
            loss_g_img = bce(netDimg(reconst, age_onehot, gender), real_label)
            loss_g_z = bce(netDz(z), real_label)
            loss_tv = total_variation_loss(reconst)

            loss_eg = loss_l1 + 0.0001 * loss_g_img + 0.01 * loss_g_z + loss_tv
            loss_eg.backward()
            optE.step()
            optG.step()

            # ---- DiscriminatorZ の学習 ----
            optDz.zero_grad()
            z_prior = torch.empty(batch_size, N_Z, device=device).uniform_(-1, 1)
            loss_dz = bce(netDz(z_prior), real_label) + bce(netDz(z.detach()), fake_label)
            loss_dz.backward()
            optDz.step()

            # ---- DiscriminatorImg の学習 ----
            optDimg.zero_grad()
            loss_dimg = bce(netDimg(img, age_onehot, gender), real_label) \
                + bce(netDimg(reconst.detach(), age_onehot, gender), fake_label)
            loss_dimg.backward()
            optDimg.step()

            if i % 50 == 0:
                print(f"epoch {epoch+1}/{n_epochs}, step {i}, "
                      f"L1: {loss_l1.item():.4f}, EG: {loss_eg.item():.4f}, "
                      f"Dz: {loss_dz.item():.4f}, Dimg: {loss_dimg.item():.4f}")

        # 1エポックごとに再構成画像を保存
        save_image(reconst[:8], os.path.join(output_dir, f"reconst_epoch{epoch+1:03d}.png"), normalize=True)

        # 10エポックごとにモデルを保存
        if (epoch + 1) % 10 == 0:
            torch.save(netE.state_dict(), os.path.join(output_dir, f"netE_{epoch+1:03d}.pth"))
            torch.save(netG.state_dict(), os.path.join(output_dir, f"netG_{epoch+1:03d}.pth"))
            torch.save(netDz.state_dict(), os.path.join(output_dir, f"netDz_{epoch+1:03d}.pth"))
            torch.save(netDimg.state_dict(), os.path.join(output_dir, f"netDimg_{epoch+1:03d}.pth"))


if __name__ == '__main__':
    main()

コマンドプロンプト管理者として実行し,次のコマンドを実行して学習を行う.

終了まで時間がかかるので待つ.

cd /d c:%HOMEPATH%\caae
python train.py

テスト(年齢変換画像の生成)

次のPython プログラムを,%HOMEPATH%\caae\test.py として保存する.

1枚の入力画像から,学習済みの Encoder と Generator を用いて,10段階の年齢区分すべての顔画像を生成する.

import os
import torch
import torch.nn.functional as F
from PIL import Image
import torchvision.transforms as transforms
from torchvision.utils import save_image

from models import Encoder, Generator, N_AGE_GROUPS


def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

    output_dir = os.path.expanduser(r"~\caae\output")
    test_image_path = os.path.expanduser(r"~\caae\data\test_image.jpg")
    gender = 0  # 0: 男性として生成,1: 女性として生成

    netE = Encoder().to(device).eval()
    netG = Generator().to(device).eval()
    netE.load_state_dict(torch.load(os.path.join(output_dir, "netE_050.pth"), map_location=device))
    netG.load_state_dict(torch.load(os.path.join(output_dir, "netG_050.pth"), map_location=device))

    transform = transforms.Compose([
        transforms.Resize((128, 128)),
        transforms.ToTensor(),
        transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]),
    ])

    img = Image.open(test_image_path).convert('RGB')
    img_tensor = transform(img).unsqueeze(0).to(device)

    with torch.no_grad():
        z = netE(img_tensor)
        z = z.repeat(N_AGE_GROUPS, 1)
        age_onehot = F.one_hot(torch.arange(N_AGE_GROUPS), N_AGE_GROUPS).float().to(device)
        gender_tensor = torch.full((N_AGE_GROUPS, 1), float(gender), device=device)
        generated = netG(z, age_onehot, gender_tensor)

    suffix = "male" if gender == 0 else "female"
    save_image(generated, os.path.join(output_dir, f"test_as_{suffix}.png"), normalize=True, nrow=N_AGE_GROUPS)
    print(f"Done! Results are saved as {output_dir}\\test_as_{suffix}.png")


if __name__ == '__main__':
    main()

判定したい顔画像を %HOMEPATH%\caae\data\test_image.jpg という名前で配置してから,次のコマンドを実行する.

cd /d c:%HOMEPATH%\caae
python test.py

実行結果として,10段階の年齢区分(0~5歳,6~10歳,11~15歳,16~20歳,21~30歳,31~40歳,41~50歳,51~60歳,61~70歳,71歳以上)の生成画像が1枚の画像として保存される.