最近研究在自编码器,放一个复现的代码,移除了工程相关的代码,只保留了核心,有多卡accelerate就设置为True,没有就关了。

Decode 和 Encode 参考了stable diffusion的设计,Decode最后一层改成了方差和均值(也就是纯血VAE)特征图通过采样产生,再使用VQ量化特征图。图片最后还是有些胡,感觉是因为有些图像被压缩过,插值成256*256,或者jpeg格式的有损压缩导致了数据有噪声被学会了。

数据源:

Konachan动漫头像数据集_数据集-飞桨AI Studio星河社区

效果图

epoch 0 step 100

epoch 6 step 10000

epoch 50 step 85000epoch 100 176700

模型代码 

import math

import numpy as np
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F


class ConvBlock(nn.Module):
    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1, groups=1):
        super(ConvBlock, self).__init__()
        self.conv_block = nn.Sequential(
            nn.GroupNorm(groups, in_channels),
            nn.SiLU(),
            nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding),
        )

    def forward(self, x):
        return self.conv_block(x)


class ResnetBlock(nn.Module):
    def __init__(self, in_channels, out_channels, groups=32):
        super(ResnetBlock, self).__init__()
        self.conv_block = nn.Sequential(
            ConvBlock(in_channels, out_channels, groups=groups),
            ConvBlock(out_channels, out_channels, groups=groups),
        )
        if in_channels != out_channels:
            self.skip_conn = ConvBlock(in_channels, out_channels, kernel_size=1, padding=0, groups=groups)
        else:
            self.skip_conn = nn.Identity()

    def forward(self, x):
        return self.conv_block(x) + self.skip_conn(x)


class AttentionBlock(nn.Module):
    def __init__(self, in_channels, out_channels, groups=32):
        super(AttentionBlock, self).__init__()
        self.q_conv = ConvBlock(in_channels, out_channels, kernel_size=1, padding=0, groups=groups)
        self.k_conv = ConvBlock(in_channels, out_channels, kernel_size=1, padding=0, groups=groups)
        self.v_conv = ConvBlock(in_channels, out_channels, kernel_size=1, padding=0, groups=groups)
        self.out_conv = ConvBlock(out_channels, out_channels, kernel_size=1, padding=0, groups=groups)

        if in_channels != out_channels:
            self.skip_conn = ConvBlock(in_channels, out_channels, kernel_size=1, padding=0, groups=groups)
        else:
            self.skip_conn = nn.Identity()

    def forward(self, x):
        q = self.q_conv(x)
        k = self.k_conv(x)
        v = self.v_conv(x)

        attention = torch.einsum('bchw,bcHW->bhwHW', q, k)
        attention = attention / math.sqrt(q.shape[-1])
        attention = attention.softmax(dim=-1)

        out = torch.einsum('bhwHW,bcHW->bchw', attention, v)
        out = self.out_conv(out)

        return out + self.skip_conn(x)


class MiddleBlock(nn.Module):
    def __init__(self, in_channels, out_channels, groups=32):
        super(MiddleBlock, self).__init__()
        self.conv_block = nn.Sequential(
            ResnetBlock(in_channels, out_channels, groups=groups),
            AttentionBlock(out_channels, out_channels, groups=groups),
            ResnetBlock(out_channels, out_channels, groups=groups),
        )

    def forward(self, x):
        return self.conv_block(x)


class UpSample(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(UpSample, self).__init__()
        self.conv = nn.Conv2d(in_channels, out_channels, 3, 1, 1)

    def forward(self, x):
        x = nn.functional.interpolate(x, scale_factor=2)
        x = self.conv(x)
        return x


class DownSample(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(DownSample, self).__init__()
        self.conv = nn.Conv2d(in_channels, out_channels, 3, 2, 0)

    def forward(self, x):
        pad = (0, 1, 0, 1)
        x = F.pad(x, pad, mode='constant', value=0)
        x = self.conv(x)
        return x


class Encoder(nn.Module):
    def __init__(self, in_channels, out_channels, groups=32):
        super(Encoder, self).__init__()
        self.conv = nn.Conv2d(in_channels, 128, 3, 1, 1)
        self.down_block = self.create_down_block(128, 128, 1, groups=groups)
        self.down_block2 = self.create_down_block(128, 256, 2, groups=groups)
        self.down_block3 = self.create_down_block(256, 512, 2, groups=groups)
        self.down_block4 = self.create_down_block(512, 1024, 2, groups=groups)
        self.resnet_block = self.create_resnet_block(1024, 1024, 2, groups=groups)
        self.middle_block = MiddleBlock(1024, 1024, groups=groups)
        self.conv_block = ConvBlock(1024, out_channels, groups=groups)

    @staticmethod
    def create_down_block(in_channels, out_channels, num_blocks, groups=32):
        res_blocks = []
        for _ in range(num_blocks):
            res_blocks.append(ResnetBlock(in_channels, in_channels, groups=groups))
        res_blocks.append(DownSample(in_channels, out_channels))
        return nn.Sequential(*res_blocks)

    @staticmethod
    def create_resnet_block(in_channels, out_channels, num_blocks, groups=32):
        res_blocks = [ResnetBlock(in_channels, out_channels, groups=groups)]
        for _ in range(num_blocks - 1):
            res_blocks.append(ResnetBlock(out_channels, out_channels, groups=groups))
        return nn.Sequential(*res_blocks)

    def forward(self, x):
        x = self.conv(x)
        x = self.down_block(x)
        x = self.down_block2(x)
        x = self.down_block3(x)
        x = self.down_block4(x)
        x = self.resnet_block(x)
        x = self.middle_block(x)
        x = self.conv_block(x)
        return x


class Decoder(nn.Module):
    def __init__(self, in_channels, groups=32):
        super(Decoder, self).__init__()
        self.conv = nn.Conv2d(in_channels, 1024, 3, 1, 1)
        self.middle_block = MiddleBlock(1024, 1024, groups=groups)
        self.up_block = self.create_up_block(1024, 1024, 1, groups=groups)
        self.up_block2 = self.create_up_block(1024, 512, 2, groups=groups)
        self.up_block3 = self.create_up_block(512, 256, 3, groups=groups)
        self.up_block4 = self.create_up_block(256, 128, 3, groups=groups)
        self.resnet_block = self.create_resnet_block(128, 128, 2, groups=groups)
        self.conv_block = ConvBlock(128, 3, groups=groups)

    @staticmethod
    def create_up_block(in_channels, out_channels, num_blocks, groups=32):
        res_blocks = []
        for _ in range(num_blocks):
            res_blocks.append(ResnetBlock(in_channels, in_channels, groups=groups))
        res_blocks.append(UpSample(in_channels, out_channels))
        return nn.Sequential(*res_blocks)

    @staticmethod
    def create_resnet_block(in_channels, out_channels, num_blocks, groups=32):
        res_blocks = [ResnetBlock(in_channels, out_channels, groups=groups)]
        for _ in range(num_blocks - 1):
            res_blocks.append(ResnetBlock(out_channels, out_channels, groups=groups))
        return nn.Sequential(*res_blocks)

    def forward(self, x):
        x = self.conv(x)
        x = self.middle_block(x)
        x = self.up_block(x)
        x = self.up_block2(x)
        x = self.up_block3(x)
        x = self.up_block4(x)
        x = self.resnet_block(x)
        x = self.conv_block(x)
        return x

    def get_last_layer(self):
        return self.conv_block.conv_block[-1].weight


class DiagonalGaussianDistribution(object):
    def __init__(self, parameters, deterministic=False):
        self.parameters = parameters
        self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
        self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
        self.deterministic = deterministic
        self.std = torch.exp(0.5 * self.logvar)
        self.var = torch.exp(self.logvar)
        if self.deterministic:
            self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)

    def sample(self):
        x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
        return x

    def kl(self, other=None):
        if self.deterministic:
            return torch.Tensor([0.])
        else:
            if other is None:
                return 0.5 * torch.sum(torch.pow(self.mean, 2)
                                       + self.var - 1.0 - self.logvar,
                                       dim=[1, 2, 3])
            else:
                return 0.5 * torch.sum(
                    torch.pow(self.mean - other.mean, 2) / other.var
                    + self.var / other.var - 1.0 - self.logvar + other.logvar,
                    dim=[1, 2, 3])

    def nll(self, sample, dims=None):
        if dims is None:
            dims = [1, 2, 3]
        if self.deterministic:
            return torch.Tensor([0.])
        log_two_pi = np.log(2.0 * np.pi)
        return 0.5 * torch.sum(
            log_two_pi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
            dim=dims)

    def mode(self):
        return self.mean


class VectorQuantizer(nn.Module):
    """支持多卡EMA同步的向量量化层"""

    def __init__(self, num_embeddings, embedding_dim, beta=0.25, decay=0.99, epsilon=1e-5, ema=False):
        super().__init__()
        self.embedding_dim = embedding_dim
        self.num_embeddings = num_embeddings
        self.beta = beta
        self.decay = decay
        self.epsilon = epsilon
        self.ema = ema

        # 码本初始化
        self.embedding = nn.Embedding(self.num_embeddings, self.embedding_dim)
        self.embedding.weight.data.normal_()

        # EMA统计量
        self.register_buffer('_ema_cluster_size', torch.zeros(num_embeddings))
        self.register_buffer('_ema_w', self.embedding.weight.data.clone())

    @staticmethod
    def _all_reduce_tensor(tensor):
        """跨卡聚合张量"""
        tensor = tensor.clone().detach()
        dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
        return tensor

    @staticmethod
    def _broadcast_tensor(tensor):
        """广播主卡参数到所有卡"""
        dist.broadcast(tensor, src=0)
        return tensor

    def forward(self, z):
        # 形状变换
        z = z.permute(0, 2, 3, 1)  # [B, D, H, W] -> [B, H, W, D]
        z_flattened = z.reshape(-1, self.embedding_dim)

        # 计算码本距离
        distances = torch.cdist(z_flattened, self.embedding.weight, p=2.0) ** 2

        # 获取最近邻编码
        encoding_indices = torch.argmin(distances, dim=1)
        quantized = self.embedding(encoding_indices).view(z.shape)
        quantized = quantized.permute(0, 3, 1, 2)

        # 计算VQ损失
        vq_loss = self.beta * F.mse_loss(quantized.detach(), z.permute(0, 3, 1, 2))
        vq_loss = vq_loss + F.mse_loss(quantized, z.permute(0, 3, 1, 2).detach())

        # EMA 更新 (只在训练时执行)
        if self.training and self.ema:
            with torch.no_grad():
                # 生成one-hot编码 [N, num_embeddings]
                encodings = F.one_hot(encoding_indices, self.num_embeddings).float()

                # 跨卡聚合统计量
                cluster_size = encodings.sum(0)  # [num_embeddings]
                cluster_size = self._all_reduce_tensor(cluster_size)

                dw = torch.matmul(encodings.t(), z_flattened)  # [num_embeddings, dim]
                dw = self._all_reduce_tensor(dw)

                # 更新EMA统计量
                updated_ema_cluster_size = (self._ema_cluster_size * self.decay +
                                            (1 - self.decay) * cluster_size)

                # Laplace平滑
                n = torch.sum(updated_ema_cluster_size)
                updated_ema_cluster_size = (
                        (updated_ema_cluster_size + self.epsilon) /
                        (n + self.num_embeddings * self.epsilon) * n
                )

                updated_ema_w = (self._ema_w * self.decay +
                                 (1 - self.decay) * dw)

                # 主卡更新参数后广播到所有卡
                if dist.get_rank() == 0:
                    self._ema_cluster_size.copy_(updated_ema_cluster_size)
                    self._ema_w.copy_(updated_ema_w)

                # 确保所有卡使用相同的码本参数
                self._broadcast_tensor(self._ema_cluster_size)
                self._broadcast_tensor(self._ema_w)

                # 更新码本权重
                self.embedding.weight.data.copy_(
                    self._ema_w / (self._ema_cluster_size.unsqueeze(1) + 1e-6)
                )

        # 直通估计
        quantized = z.permute(0, 3, 1, 2) + (quantized - z.permute(0, 3, 1, 2)).detach()
        b, c, h, w = quantized.shape

        return quantized, vq_loss, encoding_indices.view(b, h, w)


class VAE(nn.Module):
    def __init__(self, in_channels, z_channels=4, embedding_dim=4, groups=32):
        super(VAE, self).__init__()
        self.encoder = Encoder(in_channels, z_channels * 2, groups=groups)
        self.decoder = Decoder(embedding_dim, groups=groups)
        self.quant_conv = nn.Conv2d(z_channels * 2, embedding_dim * 2, 1, 1, 0)
        self.post_quant_conv = nn.Conv2d(embedding_dim, z_channels, 1, 1, 0)

    def encode(self, x):
        h = self.encoder(x)
        moments = self.quant_conv(h)
        posterior = DiagonalGaussianDistribution(moments)
        out = posterior.sample()
        return out, posterior

    def decode(self, z):
        z = self.post_quant_conv(z)
        dec = self.decoder(z)
        return dec

    def forward(self, x):
        z, posterior = self.encode(x)
        dec = self.decode(z)
        return dec, posterior

    def generate(self, x):
        x = self.decoder(x)
        return x

    def get_last_layer(self):
        return self.decoder.get_last_layer()


class VQModel(nn.Module):
    def __init__(self, in_channels=3, groups=32, z_channels=4, embedding_dim=4, num_embeddings=8196, beta=0.25,
                 decay=0.99, epsilon=1e-5, use_ema=True):
        super(VQModel, self).__init__()
        self.encoder = Encoder(in_channels, z_channels, groups=groups)
        self.quant_conv = nn.Conv2d(z_channels, embedding_dim, 1, 1, 0)
        self.quantize = VectorQuantizer(num_embeddings,
                                        embedding_dim,
                                        ema=use_ema,
                                        beta=beta,
                                        decay=decay,
                                        epsilon=epsilon)
        self.decoder = Decoder(embedding_dim, groups=groups)
        self.post_quant_conv = nn.Conv2d(embedding_dim, z_channels, 1, 1, 0)

    def encode(self, x):
        h = self.encoder(x)
        h = self.quant_conv(h)
        quant, commit_loss, encoding_indices = self.quantize(h)
        return quant, commit_loss, encoding_indices

    def decode(self, quant, commit_loss=None):
        quant = self.post_quant_conv(quant)
        dec = self.decoder(quant)
        return dec, commit_loss

    def forward(self, x):
        quant, commit_loss, _ = self.encode(x)
        dec, commit_loss = self.decode(quant, commit_loss)
        return dec, commit_loss

    def generate(self, x):
        x = self.decoder(x)
        return x

    def get_last_layer(self):
        return self.decoder.get_last_layer()

    def calculate_balance_facter(self, perceptual_loss, gan_loss):
        last_layer = self.decoder.conv_block.conv_block[-1]
        last_layer_weight = last_layer.weight
        perceptual_loss_grads = torch.autograd.grad(perceptual_loss, last_layer_weight, retain_graph=True)[0]
        gan_loss_grads = torch.autograd.grad(gan_loss, last_layer_weight, retain_graph=True)[0]

        alpha = torch.norm(perceptual_loss_grads) / (torch.norm(gan_loss_grads) + 1e-4)
        alpha = torch.clamp(alpha, 0, 1e4).detach()
        return 0.8 * alpha

    def __getitem__(self, item):
        return self.quantize.embedding.weight[item]

    def tokenize(self, index):
        return self.quantize.embedding.weight[index]

损失函数
 

import torch
from taming.modules.losses.lpips import LPIPS
from taming.modules.losses.vqperceptual import adopt_weight, hinge_d_loss, vanilla_d_loss
from torch import nn


class LPIPSWithDiscriminator(nn.Module):
    def __init__(self, disc_start, discriminator, logvar_init=0.0, disc_factor=1.0, disc_weight=1.0,
                 perceptual_weight=1.0, disc_loss="hinge", codebook_weight=1.0):
        super().__init__()
        assert disc_loss in ["hinge", "vanilla"]
        self.perceptual_loss = LPIPS().eval()
        self.perceptual_weight = perceptual_weight
        self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
        # output log variance
        self.discriminator = discriminator
        self.discriminator_iter_start = disc_start
        self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
        self.disc_factor = disc_factor
        self.discriminator_weight = disc_weight
        self.codebook_weight = codebook_weight

    def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None):
        if last_layer is not None:
            nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
            g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
        else:
            nll_grads = torch.autograd.grad(nll_loss, self.last_layer[0], retain_graph=True)[0]
            g_grads = torch.autograd.grad(g_loss, self.last_layer[0], retain_graph=True)[0]

        d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
        d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
        d_weight = d_weight * self.discriminator_weight
        return d_weight

    def forward(self, generator_step, inputs, reconstructions, commit_loss, global_step, last_layer):
        rec_loss = torch.abs(inputs.contiguous() - reconstructions.contiguous())
        p_loss = self.perceptual_loss(inputs.contiguous(), reconstructions.contiguous())
        rec_loss = rec_loss + self.perceptual_weight * p_loss

        nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
        # nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
        nll_loss = nll_loss.mean()
        if generator_step:
            logits_fake = self.discriminator(reconstructions.contiguous())
            g_loss = -torch.mean(logits_fake)

            d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer)
            disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
            g_loss = d_weight * disc_factor * g_loss

            codebook_loss = commit_loss.mean() * self.codebook_weight
            loss = nll_loss + codebook_loss + g_loss
            return (loss,
                    d_weight.detach(),
                    self.logvar.detach(),
                    rec_loss.detach().mean(),
                    p_loss.detach().mean(),
                    g_loss.detach(),
                    codebook_loss.detach())

        else:
            logits_real = self.discriminator(inputs.contiguous().detach())
            logits_fake = self.discriminator(reconstructions.contiguous().detach())
            d_loss = self.disc_loss(logits_real, logits_fake)
            disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
            d_loss = disc_factor * d_loss
            return d_loss


class LPIPSWithDiscriminatorKL(nn.Module):
    def __init__(self, disc_start, discriminator, logvar_init=0.0, kl_weight=1.0, disc_factor=1.0, disc_weight=1.0,
                 perceptual_weight=1.0, disc_loss="hinge"):
        super().__init__()
        assert disc_loss in ["hinge", "vanilla"]
        self.kl_weight = kl_weight
        self.perceptual_loss = LPIPS().eval()
        self.perceptual_weight = perceptual_weight
        # output log variance
        self.logvar = nn.Parameter(torch.ones(size=()) * logvar_init)
        self.discriminator = discriminator
        self.discriminator_iter_start = disc_start
        self.disc_loss = hinge_d_loss if disc_loss == "hinge" else vanilla_d_loss
        self.disc_factor = disc_factor
        self.discriminator_weight = disc_weight

    def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None):
        if last_layer is not None:
            nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0]
            g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
        else:
            nll_grads = torch.autograd.grad(nll_loss, self.last_layer[0], retain_graph=True)[0]
            g_grads = torch.autograd.grad(g_loss, self.last_layer[0], retain_graph=True)[0]

        d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
        d_weight = torch.clamp(d_weight, 0.0, 1e4).detach()
        d_weight = d_weight * self.discriminator_weight
        return d_weight

    def forward(self, inputs, reconstructions, posteriors, optimizer_idx, global_step, last_layer, weight=None):
        # 公共部分:重建损失和感知损失计算
        rec_loss = torch.abs(inputs.contiguous() - reconstructions.contiguous())
        p_loss = self.perceptual_loss(inputs.contiguous(), reconstructions.contiguous())
        rec_loss = rec_loss + self.perceptual_weight * p_loss

        # 公共部分:NLL损失计算
        nll_loss = rec_loss / torch.exp(self.logvar) + self.logvar
        weighted_nll_loss = nll_loss
        if weight is not None:
            weighted_nll_loss = weight * nll_loss
        weighted_nll_loss = torch.sum(weighted_nll_loss) / weighted_nll_loss.shape[0]
        nll_loss = torch.sum(nll_loss) / nll_loss.shape[0]
        kl_loss = posteriors.kl()
        kl_loss = torch.sum(kl_loss) / kl_loss.shape[0]

        if optimizer_idx == 0:
            # 生成器训练分支
            logits_fake = self.discriminator(reconstructions.contiguous())
            g_loss = -torch.mean(logits_fake)

            # 自适应权重计算
            d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer)
            disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)

            # 总损失组成
            total_loss = (
                    weighted_nll_loss +
                    self.kl_weight * kl_loss +
                    d_weight * disc_factor * g_loss
            )

            return (
                total_loss,
                self.logvar.detach(),
                kl_loss.detach(),
                nll_loss.detach(),
                rec_loss.detach().mean(),
                d_weight.detach(),
                p_loss.detach(),
                g_loss.detach()
            )

        elif optimizer_idx == 1:
            # 判别器训练分支
            # 停止生成器部分的梯度传播
            logits_real = self.discriminator(inputs.contiguous().detach())
            logits_fake = self.discriminator(reconstructions.contiguous().detach())

            # 判别器损失计算
            disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start)
            d_loss = disc_factor * self.disc_loss(logits_real, logits_fake)

            return d_loss

        else:
            raise ValueError(f"Invalid optimizer_idx: {optimizer_idx}")

 训练脚本

import os

import numpy as np
import torch
import torchvision
from accelerate import DistributedDataParallelKwargs, Accelerator
from taming.modules.discriminator.model import NLayerDiscriminator, weights_init
from torch.optim import AdamW
from torch.utils.data import DataLoader
from torchvision.datasets import ImageFolder
from torchvision.transforms import transforms
from tqdm import tqdm

from modules.autoencoders.vae import VQModel
from modules.losses import LPIPSWithDiscriminator


def get_imagenet_dataloader(batch_size=32, data_path="datasets/faces", means=None, stds=None):
    if means is None:
        means = [0.485, 0.456, 0.406]
    if stds is None:
        stds = [0.229, 0.224, 0.225]
    # 数据加载
    transform = transforms.Compose([
        transforms.Resize(256),
        transforms.RandomCrop(256),
        transforms.RandomHorizontalFlip(),
        transforms.RandomVerticalFlip(),
        transforms.ToTensor(),
        transforms.Normalize(mean=means, std=stds)
    ])
    train_dataset = ImageFolder(data_path, transform=transform)

    return DataLoader(train_dataset, batch_size=batch_size, shuffle=True, pin_memory=True, num_workers=4)


def train(in_channels=3, disc_layers=3, use_actnorm=False, disc_start=25001, checkpoints_path='checkpoints/geo_50_9',
          data_path='datasets/geo_50_9',
          lr=4.5e-6, epochs=300, means=None, stds=None, batch_size=8,
          mixed_precision=False, monitor_path='monitors/geo_50_9', accelerate=True, disc_weight=0.75,
          accumulate_grad_batches=2, disc_ndf=32, num_embeddings=8196):
    means = [0.476, 0.476, 0.476] if means is None else means
    stds = [0.224, 0.224, 0.224] if stds is None else stds
    os.makedirs(checkpoints_path, exist_ok=True)
    os.makedirs(monitor_path, exist_ok=True)
    model = VQModel(num_embeddings=num_embeddings)
    dataloader = get_imagenet_dataloader(batch_size=batch_size,
                                         data_path=data_path,
                                         means=means,
                                         stds=stds)

    discriminator = NLayerDiscriminator(
        input_nc=in_channels,
        n_layers=disc_layers,
        ndf=disc_ndf,
        use_actnorm=use_actnorm).apply(weights_init)
    criterion = LPIPSWithDiscriminator(
        disc_start=disc_start,
        discriminator=discriminator,
        disc_factor=1.0,
        disc_weight=disc_weight,
        perceptual_weight=1.0,
        disc_loss="hinge")

    opt_ae = AdamW(model.parameters(), lr=lr, betas=(0.5, 0.9))
    opt_disc = AdamW(discriminator.parameters(), lr=lr, betas=(0.5, 0.9))

    step = 0
    start_epoch = 0
    all_steps = epochs * len(dataloader) // accumulate_grad_batches
    scheduler_ae = torch.optim.lr_scheduler.CosineAnnealingLR(opt_ae, all_steps, eta_min=1e-7)
    scheduler_disc = torch.optim.lr_scheduler.CosineAnnealingLR(opt_disc, all_steps, eta_min=1e-7)
    latest_checkpoint = os.path.join(checkpoints_path, "latest.pth")

    if os.path.exists(latest_checkpoint):
        state_dict = torch.load(latest_checkpoint)
        step = state_dict.get("step", 0)
        start_epoch = state_dict.get("epoch", 0)
        model.load_state_dict(state_dict.get("model", {}))
        discriminator.load_state_dict(state_dict.get("discriminator", {}))
        opt_ae.load_state_dict(state_dict.get("opt_ae", {}))
        opt_disc.load_state_dict(state_dict.get("opt_disc", {}))
        scheduler_ae.load_state_dict(state_dict.get("scheduler_ae", {}))
        scheduler_disc.load_state_dict(state_dict.get("scheduler_disc", {}))
    ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
    if accelerate:
        accelerator = Accelerator(mixed_precision='fp16' if mixed_precision else 'no',
                                  kwargs_handlers=[ddp_kwargs], gradient_accumulation_steps=accumulate_grad_batches)
        # 加速器
        model, criterion, opt_ae, opt_disc, dataloader, scheduler_ae, scheduler_disc = accelerator.prepare(
            model, criterion, opt_ae, opt_disc, dataloader, scheduler_ae, scheduler_disc)
        device = accelerator.device
    else:
        accelerator = None
        device = "cuda" if torch.cuda.is_available() else "cpu"
        model = model.to(device)
        criterion = criterion.to(device)

    means_tensor = torch.tensor(means).view(1, 3, 1, 1).to(device)
    stds_tensor = torch.tensor(stds).view(1, 3, 1, 1).to(device)
    disc_loss = torch.tensor(0.0).to(device)
    for epoch in range(start_epoch, epochs):
        # 仅主进程初始化进度条
        pbar = None
        if (accelerator and accelerator.is_main_process) or not accelerator:
            pbar = tqdm(total=len(dataloader), desc=f"Epoch {epoch}")
        for batch in dataloader:
            generator_step = ((step // accumulate_grad_batches) % 2) == 0
            x, _ = batch
            x = x.to(device)
            if generator_step:
                opt_ae.zero_grad(set_to_none=True)
            else:
                opt_disc.zero_grad(set_to_none=True)
            if generator_step:
                if accelerator is not None:
                    with accelerator.accumulate(model):
                        with accelerator.autocast():
                            dec, codebook_loss = model(x)
                            total_loss, d_weight, logvar, rec_loss, p_loss, g_loss, codebook_loss = criterion(
                                generator_step,
                                x,
                                dec,
                                codebook_loss,
                                step,
                                model.module.get_last_layer())
                            accelerator.backward(total_loss)
                            opt_ae.step()
                else:
                    with torch.amp.autocast(device_type=device.split(':')[0], enabled=mixed_precision):
                        dec, codebook_loss = model(x)
                        total_loss, d_weight, logvar, rec_loss, p_loss, g_loss, codebook_loss = criterion(
                            generator_step,
                            x,
                            dec,
                            codebook_loss,
                            step,
                            model.get_last_layer())
                        total_loss.backward()
                        if step % accumulate_grad_batches == 0:
                            opt_ae.step()
            else:
                if accelerator is not None:
                    with accelerator.accumulate(discriminator):
                        with accelerator.autocast():
                            dec, codebook_loss = model(x)
                            disc_loss = criterion(
                                generator_step,
                                x,
                                dec,
                                codebook_loss,
                                step,
                                model.module.get_last_layer())
                            accelerator.backward(disc_loss)
                            opt_disc.step()
                else:
                    with torch.amp.autocast(device_type=device.split(':')[0], enabled=mixed_precision):
                        dec, codebook_loss = model(x)
                        disc_loss = criterion(
                            generator_step,
                            x,
                            dec,
                            codebook_loss,
                            step,
                            model.get_last_layer())
                        disc_loss.backward()
                        if step % accumulate_grad_batches == 0:
                            opt_disc.step()
            scheduler_ae.step()
            scheduler_disc.step()

            if pbar is not None:
                pbar.update(1)
                pbar.set_postfix(
                    TotalLoss=np.round(total_loss.cpu().detach().numpy().item(), 5),
                    DiscLoss=np.round(disc_loss.cpu().detach().numpy().item(), 3),
                    PerceptualLoss=np.round(p_loss.cpu().numpy().item(), 5),
                    RecLoss=np.round(rec_loss.cpu().numpy().item(), 5),
                    GenLoss=np.round(g_loss.cpu().numpy().item(), 5),
                    CodebookLoss=np.round(codebook_loss.cpu().detach().numpy().item(), 5),
                    LearningRate=opt_ae.param_groups[0]['lr']
                )
                pbar.update(0)

            # 日志记录
            if step % 100 == 0:
                with torch.no_grad():
                    # 反标准化处理
                    fake_image = dec[:4].mul(stds_tensor).add(means_tensor).clamp(0, 1)
                    real_image = x[:4].mul(stds_tensor).add(means_tensor).clamp(0, 1)
                    real_fake_images = torch.cat([real_image, fake_image])
                    save_path = os.path.join(monitor_path, f"{epoch}_{step}.jpg")
                    if accelerator:
                        if accelerator.is_main_process:
                            torchvision.utils.save_image(real_fake_images, save_path, nrow=4)
                    else:
                        torchvision.utils.save_image(real_fake_images, save_path, nrow=4)

            step += 1
        # 保存模型
        if (accelerator and accelerator.is_main_process) or not accelerate:
            state_dict = {
                "model": accelerator.unwrap_model(model).state_dict() if accelerator else model.state_dict(),
                "discriminator": accelerator.unwrap_model(
                    discriminator).state_dict() if accelerator else discriminator.state_dict(),
                "opt_ae": opt_ae.state_dict(),
                "opt_disc": opt_disc.state_dict(),
                "scheduler_ae": scheduler_ae.state_dict(),
                "scheduler_disc": scheduler_disc.state_dict(),
                "step": step,
                "epoch": epoch + 1
            }
            epoch_checkpoint = os.path.join(checkpoints_path, f"epoch_{epoch}.pth")
            torch.save(state_dict, latest_checkpoint)
            torch.save(state_dict, epoch_checkpoint)


if __name__ == '__main__':
    means = [0.99601443, 0.84884399, 0.51539658]
    stds = [0.00807457, 0.1418739, 0.23776156]
    train(checkpoints_path='checkpoints/chemical_vqgan', data_path='datasets/filtered_folder', means=means, stds=stds,
          mixed_precision=True,
          batch_size=6,
          monitor_path='monitors/chemical_vqgan',
          accumulate_grad_batches=4, epochs=30, num_embeddings=4096, disc_start=25001)

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐