自编码器#
在训练CNN时,一个问题是我们需要大量的标注数据。以图像分类为例,我们需要将图像分成不同的类别,这通常需要手动完成。
然而,我们可能希望使用原始(未标注)数据来训练CNN特征提取器,这种方法称为自监督学习。在这种情况下,我们将使用训练图像作为网络的输入和输出。自动编码器的核心思想是,我们会有一个编码器网络,将输入图像转换为某种潜在空间(通常是一个较小尺寸的向量),然后通过解码器网络,其目标是重建原始图像。
由于我们训练自动编码器的目的是尽可能捕捉原始图像中的信息以实现准确的重建,网络会尝试找到输入图像的最佳嵌入来捕捉其含义。

图片来源:Keras博客
让我们为MNIST创建最简单的自动编码器吧!
import torch
import torchvision
import matplotlib.pyplot as plt
from torchvision import transforms
from torch import nn
from torch import optim
from tqdm import tqdm
import numpy as np
import torch.nn.functional as F
torch.manual_seed(42)
np.random.seed(42)定义训练参数并检查GPU是否可用:
device = 'cuda:0' if torch.cuda.is_available() else 'cpu'
train_size = 0.9
lr = 1e-3
eps = 1e-8
batch_size = 256
epochs = 30以下函数将加载MNIST数据集并应用指定的转换。它还会将其拆分为训练/测试数据集。
def mnist(train_part, transform=None):
dataset = torchvision.datasets.MNIST('.', download=True, transform=transform)
train_part = int(train_part * len(dataset))
train_dataset, test_dataset = torch.utils.data.random_split(dataset, [train_part, len(dataset) - train_part])
return train_dataset, test_dataset现在让我们加载数据集并为训练和测试定义数据加载器:
transform = transforms.Compose([transforms.ToTensor()])
train_dataset, test_dataset = mnist(train_size, transform)
train_dataloader = torch.utils.data.DataLoader(train_dataset, drop_last=True, batch_size=batch_size, shuffle=True)
test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=False)
dataloaders = (train_dataloader, test_dataloader)def plotn(n, data, noisy=False, super_res=None):
fig, ax = plt.subplots(1, n)
for i, z in enumerate(data):
if i == n:
break
preprocess = z[0].reshape(1, 28, 28) if z[0].shape[1] == 28 else z[0].reshape(1, 14, 14) if z[0].shape[1] == 14 else z[0]
if super_res is not None:
_transform = transforms.Resize((int(preprocess.shape[1] / super_res), int(preprocess.shape[2] / super_res)))
preprocess = _transform(preprocess)
if noisy:
shapes = list(preprocess.shape)
preprocess += noisify(shapes)
ax[i].imshow(preprocess[0])
plt.show()def noisify(shapes):
return np.random.normal(loc=0.5, scale=0.3, size=shapes)plotn(5, train_dataset)class Encoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=(3, 3), padding='same')
self.maxpool1 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv2 = nn.Conv2d(16, 8, kernel_size=(3, 3), padding='same')
self.maxpool2 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv3 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.maxpool3 = nn.MaxPool2d(kernel_size=(2, 2), padding=(1, 1))
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.maxpool1(self.relu(self.conv1(input)))
hidden2 = self.maxpool2(self.relu(self.conv2(hidden1)))
encoded = self.maxpool3(self.relu(self.conv3(hidden2)))
return encodedclass Decoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.upsample1 = nn.Upsample(scale_factor=(2, 2))
self.conv2 = nn.Conv2d(8, 8, kernel_size=(3, 3), padding='same')
self.upsample2 = nn.Upsample(scale_factor=(2, 2))
self.conv3 = nn.Conv2d(8, 16, kernel_size=(3, 3))
self.upsample3 = nn.Upsample(scale_factor=(2, 2))
self.conv4 = nn.Conv2d(16, 1, kernel_size=(3, 3), padding='same')
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.upsample1(self.relu(self.conv1(input)))
hidden2 = self.upsample2(self.relu(self.conv2(hidden1)))
hidden3 = self.upsample3(self.relu(self.conv3(hidden2)))
decoded = self.sigmoid(self.conv4(hidden3))
return decodedclass AutoEncoder(nn.Module):
def __init__(self, super_resolution=False):
super().__init__()
if not super_resolution:
self.encoder = Encoder()
else:
self.encoder = SuperResolutionEncoder()
self.decoder = Decoder()
def forward(self, input):
encoded = self.encoder(input)
decoded = self.decoder(encoded)
return decodedmodel = AutoEncoder().to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()def train(dataloaders, model, loss_fn, optimizer, epochs, device, noisy=None, super_res=None):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
for epoch in tqdm_iter:
model.train()
train_loss = 0.0
test_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
shapes = list(imgs.shape)
if super_res is not None:
shapes[2], shapes[3] = int(shapes[2] / super_res), int(shapes[3] / super_res)
_transform = transforms.Resize((shapes[2], shapes[3]))
imgs_transformed = _transform(imgs)
imgs_transformed = imgs_transformed.to(device)
imgs = imgs.to(device)
labels = labels.to(device)
if noisy is not None:
noisy_tensor = noisy[0]
else:
noisy_tensor = torch.zeros(tuple(shapes)).to(device)
if super_res is None:
imgs_noisy = imgs + noisy_tensor
else:
imgs_noisy = imgs_transformed + noisy_tensor
imgs_noisy = torch.clamp(imgs_noisy, 0., 1.)
preds = model(imgs_noisy)
loss = loss_fn(preds, imgs)
optimizer.zero_grad()
loss.backward()
optimizer.step()
train_loss += loss.item()
model.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
shapes = list(imgs.shape)
if super_res is not None:
shapes[2], shapes[3] = int(shapes[2] / super_res), int(shapes[3] / super_res)
_transform = transforms.Resize((shapes[2], shapes[3]))
imgs_transformed = _transform(imgs)
imgs_transformed = imgs_transformed.to(device)
imgs = imgs.to(device)
labels = labels.to(device)
if noisy is not None:
test_noisy_tensor = noisy[1]
else:
test_noisy_tensor = torch.zeros(tuple(shapes)).to(device)
if super_res is None:
imgs_noisy = imgs + test_noisy_tensor
else:
imgs_noisy = imgs_transformed + test_noisy_tensor
imgs_noisy = torch.clamp(imgs_noisy, 0., 1.)
preds = model(imgs_noisy)
loss = loss_fn(preds, imgs)
test_loss += loss.item()
train_loss /= len(train_dataloader)
test_loss /= len(test_dataloader)
tqdm_dct = {'train loss:': train_loss, 'test loss:': test_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()train(dataloaders, model, loss_fn, optimizer, epochs, device)100%|██████████| 30/30 [06:49<00:00, 13.65s/it, train loss:=0.104, test loss:=0.104]
model.eval()
predictions = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
predictions.append(model(data[0].to(device).unsqueeze(0)).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)任务 1:尝试使用非常小的潜在向量大小(例如 2)来训练自动编码器,并绘制与不同数字对应的点。提示:在卷积部分之后使用全连接的密集层,将向量大小减少到所需值。
任务 2:从不同的数字开始,获取它们的潜在空间表示,观察在潜在空间中添加一些噪声对生成数字的影响。
去噪#
自编码器可以被有效地用来从图像中去除噪声。为了训练一个去噪器,我们将从无噪声的图像开始,并人为地向它们添加噪声。然后,我们将带有噪声的图像作为输入,无噪声的图像作为输出,输入到自编码器中。
让我们看看这在 MNIST 数据集上的效果如何:
plotn(5, train_dataset, noisy=True)model = AutoEncoder().to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()noisy_tensor = torch.FloatTensor(noisify([256, 1, 28, 28])).to(device)
test_noisy_tensor = torch.FloatTensor(noisify([1, 1, 28, 28])).to(device)
noisy_tensors = (noisy_tensor, test_noisy_tensor)train(dataloaders, model, loss_fn, optimizer, 100, device, noisy=noisy_tensors)100%|██████████| 100/100 [22:29<00:00, 13.49s/it, train loss:=0.134, test loss:=0.133]
model.eval()
predictions = []
noise = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
shapes = data[0].shape
noisy_data = data[0] + test_noisy_tensor[0].detach().cpu()
noise.append(noisy_data)
predictions.append(model(noisy_data.to(device).unsqueeze(0)).detach().cpu())
plotn(plots, noise)
plotn(plots, predictions)练习: 观察在MNIST数字上训练的去噪器如何处理不同的图像。作为一个例子,你可以使用Fashion MNIST数据集,它具有相同的图像尺寸。请注意,去噪器仅在与其训练时相同类型的图像上效果良好(即输入数据的概率分布相同)。
超分辨率#
与去噪器类似,我们可以训练自动编码器来提高图像的分辨率。为了训练超分辨率网络,我们将从高分辨率图像开始,并自动将其缩小以生成网络输入。然后,我们将小尺寸图像作为输入,高分辨率图像作为输出,输入到自动编码器中。
为此,让我们在训练时将图像缩小到14x14。
super_res_koeff = 2.0
plotn(5, train_dataset, super_res=super_res_koeff)class SuperResolutionEncoder(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=(3, 3), padding='same')
self.maxpool1 = nn.MaxPool2d(kernel_size=(2, 2))
self.conv2 = nn.Conv2d(16, 8, kernel_size=(3, 3), padding='same')
self.maxpool2 = nn.MaxPool2d(kernel_size=(2, 2), padding=(1, 1))
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.maxpool1(self.relu(self.conv1(input)))
encoded = self.maxpool2(self.relu(self.conv2(hidden1)))
return encodedmodel = AutoEncoder(super_resolution=True).to(device)
optimizer = optim.Adam(model.parameters(), lr=lr, eps=eps)
loss_fn = nn.BCELoss()train(dataloaders, model, loss_fn, optimizer, epochs, device, super_res=2.0)100%|██████████| 30/30 [06:43<00:00, 13.47s/it, train loss:=0.102, test loss:=0.103]
model.eval()
predictions = []
plots = 5
shapes = test_dataset[0][0].shape
for i, data in enumerate(test_dataset):
if i == plots:
break
_transform = transforms.Resize((int(shapes[1] / super_res_koeff), int(shapes[2] / super_res_koeff)))
predictions.append(model(_transform(data[0]).to(device).unsqueeze(0)).detach().cpu())
plotn(plots, test_dataset, super_res=super_res_koeff)
plotn(plots, predictions)练习: 尝试在 CIFAR-10 上训练超分辨率网络,实现2倍和4倍的上采样。使用噪声作为4倍上采样模型的输入,并观察结果。
变分自编码器 (VAE)#
传统的自编码器通过某种方式降低输入数据的维度,从而提取输入图像的重要特征。然而,潜在向量往往缺乏明确的意义。换句话说,以 MNIST 数据集为例,想要弄清楚哪些数字对应于不同的潜在向量并不是一件容易的事,因为相近的潜在向量不一定对应于相同的数字。
另一方面,为了训练生成模型,理解潜在空间是更有帮助的。这一想法引出了变分自编码器(VAE)。
VAE 是一种能够学习预测潜在参数的统计分布(即所谓的潜在分布)的自编码器。例如,我们可以假设潜在向量服从分布 $N(\mathrm{z_mean},e^{\mathrm{z_log}})$,其中 $\mathrm{z_mean}, \mathrm{z_log} \in\mathbb{R}^d$。VAE 的编码器学习预测这些参数,然后解码器从该分布中随机采样一个向量来重建对象。
总结如下:
- 从输入向量中,我们预测
z_mean和z_log(我们预测的是标准差的对数,而不是标准差本身) - 我们从分布 $N(\mathrm{z_mean},e^{\mathrm{z_log_sigma}})$ 中采样一个向量
sample(z_val in code) - 解码器尝试使用
sample作为输入向量来解码原始图像
图片来源于 Isaak Dykeman 的这篇博客文章
class VAEEncoder(nn.Module):
def __init__(self, device):
super().__init__()
self.intermediate_dim = 512
self.latent_dim = 2
self.linear = nn.Linear(784, self.intermediate_dim)
self.z_mean = nn.Linear(self.intermediate_dim, self.latent_dim)
self.z_log = nn.Linear(self.intermediate_dim, self.latent_dim)
self.relu = nn.ReLU()
self.device = device
def forward(self, input):
bs = input.shape[0]
hidden = self.relu(self.linear(input))
z_mean = self.z_mean(hidden)
z_log = self.z_log(hidden)
eps = torch.FloatTensor(np.random.normal(size=(bs, self.latent_dim))).to(device)
z_val = z_mean + torch.exp(z_log) * eps
return z_mean, z_log, z_valclass VAEDecoder(nn.Module):
def __init__(self):
super().__init__()
self.intermediate_dim = 512
self.latent_dim = 2
self.linear = nn.Linear(self.latent_dim, self.intermediate_dim)
self.output = nn.Linear(self.intermediate_dim, 784)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden = self.relu(self.linear(input))
decoded = self.sigmoid(self.output(hidden))
return decodedclass VAEAutoEncoder(nn.Module):
def __init__(self, device):
super().__init__()
self.encoder = VAEEncoder(device)
self.decoder = VAEDecoder()
self.z_vals = None
def forward(self, input):
bs, c, h, w = input.shape[0], input.shape[1], input.shape[2], input.shape[3]
input = input.view(bs, -1)
encoded = self.encoder(input)
self.z_vals = encoded
decoded = self.decoder(encoded[2])
return decoded
def get_zvals(self):
return self.z_vals变分自编码器使用由两部分组成的复杂损失函数:
- 重建损失 是一种损失函数,用于衡量重建图像与目标图像的接近程度(可以是均方误差 MSE)。它与普通自编码器中的损失函数相同。
- KL 损失,确保潜在变量分布接近正态分布。它基于 Kullback-Leibler 散度 的概念——一种用于估计两个统计分布相似程度的度量方法。
def vae_loss(preds, targets, z_vals):
mse = nn.MSELoss()
reconstruction_loss = mse(preds, targets.view(targets.shape[0], -1)) * 784.0
temp = 1.0 + z_vals[1] - torch.square(z_vals[0]) - torch.exp(z_vals[1])
kl_loss = -0.5 * torch.sum(temp, axis=-1)
return torch.mean(reconstruction_loss + kl_loss)model = VAEAutoEncoder(device).to(device)
optimizer = optim.RMSprop(model.parameters(), lr=lr, eps=eps)def train_vae(dataloaders, model, optimizer, epochs, device):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
for epoch in tqdm_iter:
model.train()
train_loss = 0.0
test_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
imgs = imgs.to(device)
labels = labels.to(device)
preds = model(imgs)
z_vals = model.get_zvals()
loss = vae_loss(preds, imgs, z_vals)
optimizer.zero_grad()
loss.backward()
optimizer.step()
train_loss += loss.item()
model.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
imgs = imgs.to(device)
labels = labels.to(device)
preds = model(imgs)
z_vals = model.get_zvals()
loss = vae_loss(preds, imgs, z_vals)
test_loss += loss.item()
train_loss /= len(train_dataloader)
test_loss /= len(test_dataloader)
tqdm_dct = {'train loss:': train_loss, 'test loss:': test_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()train_vae(dataloaders, model, optimizer, epochs, device)100%|██████████| 30/30 [04:54<00:00, 9.83s/it, train loss:=35.1, test loss:=35.6]
model.eval()
predictions = []
plots = 5
for i, data in enumerate(test_dataset):
if i == plots:
break
predictions.append(model(data[0].to(device).unsqueeze(0)).view(1, 28, 28).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)任务: 在我们的样本中,我们已经训练了全连接VAE。现在从上面的传统自动编码器中提取CNN,并创建基于CNN的VAE。
对抗自编码器 (AAE)#
对抗自编码器(Adversarial Auto-Encoders)是生成对抗网络(Generative Adversarial Networks)和变分自编码器(Variational Auto-Encoders)的结合。
编码器将作为生成器,判别器将学习区分编码器输出的真实图像和生成的图像。编码器的输出是一个分布,解码器将尝试从这个输出中解码图像。
在这种方法中,我们有三个损失函数:来自GAN的生成器损失、判别器损失,以及来自VAE的重构损失。
图片来源于这篇博客文章,作者为Felipe Ducau
class AAEEncoder(nn.Module):
def __init__(self, input_dim, inter_dim, latent_dim):
super().__init__()
self.linear1 = nn.Linear(input_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, latent_dim)
self.relu = nn.ReLU()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
encoded = self.linear4(hidden3)
return encodedclass AAEDecoder(nn.Module):
def __init__(self, latent_dim, inter_dim, output_dim):
super().__init__()
self.linear1 = nn.Linear(latent_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, output_dim)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
decoded = self.sigmoid(self.linear4(hidden3))
return decodedclass AAEDiscriminator(nn.Module):
def __init__(self, latent_dim, inter_dim):
super().__init__()
self.latent_dim = latent_dim
self.inter_dim = inter_dim
self.linear1 = nn.Linear(latent_dim, inter_dim)
self.linear2 = nn.Linear(inter_dim, inter_dim)
self.linear3 = nn.Linear(inter_dim, inter_dim)
self.linear4 = nn.Linear(inter_dim, inter_dim)
self.linear5 = nn.Linear(inter_dim, 1)
self.relu = nn.ReLU()
self.sigmoid = nn.Sigmoid()
def forward(self, input):
hidden1 = self.relu(self.linear1(input))
hidden2 = self.relu(self.linear2(hidden1))
hidden3 = self.relu(self.linear3(hidden2))
hidden4 = self.relu(self.linear4(hidden3))
decoded = self.sigmoid(self.linear4(hidden4))
return decoded
def get_dims(self):
return self.latent_dim, self.inter_dim
input_dims = 784
inter_dims = 1000
latent_dims = 150aae_encoder = AAEEncoder(input_dims, inter_dims, latent_dims).to(device)
aae_decoder = AAEDecoder(latent_dims, inter_dims, input_dims).to(device)
aae_discriminator = AAEDiscriminator(latent_dims, int(inter_dims / 2)).to(device)lr = 1e-4
regularization_lr = 5e-5optim_encoder = optim.Adam(aae_encoder.parameters(), lr=lr)
optim_encoder_regularization = optim.Adam(aae_encoder.parameters(), lr=regularization_lr)
optim_decoder = optim.Adam(aae_decoder.parameters(), lr=lr)
optim_discriminator = optim.Adam(aae_discriminator.parameters(), lr=regularization_lr)def train_aae(dataloaders, models, optimizers, epochs, device):
tqdm_iter = tqdm(range(epochs))
train_dataloader, test_dataloader = dataloaders[0], dataloaders[1]
enc, dec, disc = models[0], models[1], models[2]
optim_enc, optim_enc_reg, optim_dec, optim_disc = optimizers[0], optimizers[1], optimizers[2], optimizers[3]
eps = 1e-9
for epoch in tqdm_iter:
enc.train()
dec.train()
disc.train()
train_reconst_loss = 0.0
train_disc_loss = 0.0
train_enc_loss = 0.0
test_reconst_loss = 0.0
test_disc_loss = 0.0
test_enc_loss = 0.0
for batch in train_dataloader:
imgs, labels = batch
imgs = imgs.view(imgs.shape[0], -1).to(device)
labels = labels.to(device)
enc.zero_grad()
dec.zero_grad()
disc.zero_grad()
encoded = enc(imgs)
decoded = dec(encoded)
reconstruction_loss = F.binary_cross_entropy(decoded, imgs)
reconstruction_loss.backward()
optim_enc.step()
optim_dec.step()
enc.eval()
latent_dim, disc_inter_dim = disc.get_dims()
real = torch.randn(imgs.shape[0], latent_dim).to(device)
disc_real = disc(real)
disc_fake = disc(enc(imgs))
disc_loss = -torch.mean(torch.log(disc_real + eps) + torch.log(1.0 - disc_fake + eps))
disc_loss.backward()
optim_dec.step()
enc.train()
disc_fake = disc(enc(imgs))
enc_loss = -torch.mean(torch.log(disc_fake + eps))
enc_loss.backward()
optim_enc_reg.step()
train_reconst_loss += reconstruction_loss.item()
train_disc_loss += disc_loss.item()
train_enc_loss += enc_loss.item()
enc.eval()
dec.eval()
disc.eval()
with torch.no_grad():
for batch in test_dataloader:
imgs, labels = batch
imgs = imgs.view(imgs.shape[0], -1).to(device)
labels = labels.to(device)
encoded = enc(imgs)
decoded = dec(encoded)
reconstruction_loss = F.binary_cross_entropy(decoded, imgs)
latent_dim, disc_inter_dim = disc.get_dims()
real = torch.randn(imgs.shape[0], latent_dim).to(device)
disc_real = disc(real)
disc_fake = disc(enc(imgs))
disc_loss = -torch.mean(torch.log(disc_real + eps) + torch.log(1.0 - disc_fake + eps))
disc_fake = disc(enc(imgs))
enc_loss = -torch.mean(torch.log(disc_fake + eps))
test_reconst_loss += reconstruction_loss.item()
test_disc_loss += disc_loss.item()
test_enc_loss += enc_loss.item()
train_reconst_loss /= len(train_dataloader)
train_disc_loss /= len(train_dataloader)
train_enc_loss /= len(train_dataloader)
test_reconst_loss /= len(test_dataloader)
test_disc_loss /= len(test_dataloader)
test_enc_loss /= len(test_dataloader)
tqdm_dct = {'train reconst loss:': train_reconst_loss, 'train disc loss:': train_disc_loss, 'train enc loss': train_enc_loss, \
'test reconst loss:': test_reconst_loss, 'test disc loss:': test_disc_loss, 'test enc loss': test_enc_loss}
tqdm_iter.set_postfix(tqdm_dct, refresh=True)
tqdm_iter.refresh()models = (aae_encoder, aae_decoder, aae_discriminator)
optimizers = (optim_encoder, optim_encoder_regularization, optim_decoder, optim_discriminator)train_aae(dataloaders, models, optimizers, epochs, device)100%|██████████| 30/30 [09:22<00:00, 18.75s/it, train reconst loss:=0.0919, train disc loss:=1.39, train enc loss=0.692, test reconst loss:=0.0945, test disc loss:=1.39, test enc loss=0.692]
aae_encoder.eval()
aae_decoder.eval()
predictions = []
plots = 10
for i, data in enumerate(test_dataset):
if i == plots:
break
pred = aae_decoder(aae_encoder(data[0].to(device).unsqueeze(0).view(1, 784)))
predictions.append(pred.view(1, 28, 28).detach().cpu())
plotn(plots, test_dataset)
plotn(plots, predictions)额外资料#
免责声明:
本文档使用AI翻译服务Co-op Translator进行翻译。尽管我们努力确保翻译的准确性,但请注意,自动翻译可能包含错误或不准确之处。原始语言的文档应被视为权威来源。对于关键信息,建议使用专业人工翻译。我们不对因使用此翻译而产生的任何误解或误读承担责任。