VQ-GAN复现
·
最近研究在自编码器,放一个复现的代码,移除了工程相关的代码,只保留了核心,有多卡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 85000
epoch 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)
更多推荐


所有评论(0)