現(xiàn)變分自編碼器:從原理到實(shí)戰(zhàn),掌握生成模型核心)
1. 先搞清楚VAE到底解決了什么問(wèn)題以及它和普通自編碼器的核心區(qū)別如果你正在找PyTorch實(shí)現(xiàn)變分自編碼器的教程大概率是想用它來(lái)生成新數(shù)據(jù)比如生成新的人臉、手寫(xiě)數(shù)字或者某種風(fēng)格的圖片。但很多人一開(kāi)始會(huì)把它和普通的自編碼器搞混結(jié)果代碼跑通了生成的效果卻一塌糊涂或者根本沒(méi)法用。這里最關(guān)鍵的區(qū)別在于普通自編碼器AE學(xué)的是“壓縮和重建”而變分自編碼器VAE學(xué)的是“數(shù)據(jù)的概率分布”。普通自編碼器就像一個(gè)記憶力超強(qiáng)的學(xué)生你把一張貓的圖片輸入給它它壓縮成一個(gè)編碼潛在向量然后再盡力還原成原來(lái)的貓圖輸出。它還原得越好說(shuō)明壓縮編碼越有效。但問(wèn)題是這個(gè)學(xué)生只記住了你給它的那些具體圖片。你讓它“畫(huà)一張你沒(méi)見(jiàn)過(guò)的貓”它就懵了因?yàn)樗鼘W(xué)到的編碼空間可能是支離破碎、不連續(xù)的從一個(gè)編碼跳到另一個(gè)編碼生成的圖片可能毫無(wú)意義。VAE要解決的就是這個(gè)“生成新數(shù)據(jù)”的問(wèn)題。它不直接把輸入壓縮成一個(gè)固定的編碼點(diǎn)而是壓縮成一個(gè)概率分布——通常是一個(gè)高斯分布用均值mean和方差log_var來(lái)表示。然后從這個(gè)分布中采樣得到一個(gè)編碼點(diǎn)再用這個(gè)點(diǎn)去解碼生成圖片。這個(gè)“采樣”步驟是VAE的靈魂它強(qiáng)制模型學(xué)習(xí)一個(gè)連續(xù)、平滑的潛在空間。在這個(gè)空間里你稍微改變一下編碼值生成的圖片也會(huì)平滑地變化比如從微笑的臉變成嚴(yán)肅的臉。更重要的是你可以從這個(gè)分布的任何地方采樣理論上都能生成一個(gè)合理的新數(shù)據(jù)。所以如果你用PyTorch實(shí)現(xiàn)VAE目標(biāo)絕不是讓重建損失降到零那會(huì)過(guò)擬合而是要在重建精度和潛在空間的規(guī)整性KL散度之間找到一個(gè)平衡。這個(gè)平衡點(diǎn)才是VAE能穩(wěn)定生成新樣本的關(guān)鍵。2. 動(dòng)手前的環(huán)境準(zhǔn)備與核心概念拆解在開(kāi)始敲代碼之前有兩件事必須明確你的PyTorch環(huán)境和VAE模型里那幾個(gè)關(guān)鍵張量的維度。很多人在這一步?jīng)]理清后面維度對(duì)不上報(bào)錯(cuò)能找半天。2.1 PyTorch環(huán)境別在版本問(wèn)題上栽跟頭輸入材料里提到了大量關(guān)于PyTorch安裝、版本的問(wèn)題這確實(shí)是第一道坎。我的建議是優(yōu)先使用Conda管理環(huán)境。這能最大程度避免包沖突。創(chuàng)建一個(gè)專用于本實(shí)驗(yàn)的環(huán)境conda create -n pytorch-vae python3.9 conda activate pytorch-vae根據(jù)你的顯卡選擇安裝命令。去PyTorch官網(wǎng)pytorch.org用它的安裝命令生成器最穩(wěn)妥。比如對(duì)于CUDA 11.8的顯卡# 這是一個(gè)示例請(qǐng)以官網(wǎng)最新命令為準(zhǔn) pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118如果沒(méi)有GPU就用CPU版本。不要盲目追求最新版本特別是如果你的CUDA驅(qū)動(dòng)比較老。輸入材料里提到的“pytorch 2.5 is required but found”這種錯(cuò)誤就是版本不匹配的典型。驗(yàn)證安裝。跑一個(gè)簡(jiǎn)單的導(dǎo)入和CUDA檢查import torch print(torch.__version__) print(torch.cuda.is_available()) # 如果有GPU應(yīng)該返回True環(huán)境搞定后我們來(lái)看VAE模型里數(shù)據(jù)是怎么流動(dòng)的。2.2 理解數(shù)據(jù)流從圖片到分布再到新圖片假設(shè)我們處理的是28x28的灰度手寫(xiě)數(shù)字圖片MNIST數(shù)據(jù)集批次大小batch_size設(shè)為64。輸入x- 形狀為[64, 1, 28, 28]的張量。編碼器Encoder通過(guò)幾層卷積或全連接網(wǎng)絡(luò)把x映射到潛在空間。輸出不再是單個(gè)向量而是兩個(gè)向量mu(均值): 形狀[64, latent_dim]比如latent_dim20。log_var(對(duì)數(shù)方差): 形狀同樣是[64, 20]。用對(duì)數(shù)方差是為了訓(xùn)練穩(wěn)定性。重參數(shù)化技巧Reparameterization Trick這是VAE訓(xùn)練的核心。我們不能直接采樣因?yàn)椴蓸硬僮鞑豢蓪?dǎo)。所以用這個(gè)技巧std torch.exp(0.5 * log_var) # 計(jì)算標(biāo)準(zhǔn)差 eps torch.randn_like(std) # 從標(biāo)準(zhǔn)正態(tài)分布采樣噪聲 z mu eps * std # 得到最終的潛在編碼z這樣z的形狀也是[64, 20]并且梯度可以沿著mu和log_var回傳。解碼器Decoder將采樣得到的z([64, 20]) 通過(guò)反卷積或全連接網(wǎng)絡(luò)重建出圖片x_recon形狀恢復(fù)為[64, 1, 28, 28]。整個(gè)過(guò)程中維度必須嚴(yán)格對(duì)齊。編碼器最后的線性層輸出大小要等于latent_dim * 2因?yàn)橐瑫r(shí)輸出mu和log_var。解碼器的第一層線性輸入大小要等于latent_dim。3. 用PyTorch一步步搭建VAE模型理論清楚了我們開(kāi)始用PyTorch的nn.Module來(lái)搭建模型。我會(huì)把每個(gè)模塊拆開(kāi)講并解釋為什么這么設(shè)計(jì)。3.1 定義編碼器網(wǎng)絡(luò)編碼器的目的是把高維圖片壓縮成潛在分布的參數(shù)。對(duì)于MNIST這種小圖片用全連接網(wǎng)絡(luò)就夠簡(jiǎn)單直觀。import torch import torch.nn as nn import torch.nn.functional as F class Encoder(nn.Module): def __init__(self, input_dim784, hidden_dim400, latent_dim20): super(Encoder, self).__init__() # 將28x28784的圖片展平 self.fc1 nn.Linear(input_dim, hidden_dim) self.fc_mu nn.Linear(hidden_dim, latent_dim) # 輸出均值mu self.fc_logvar nn.Linear(hidden_dim, latent_dim) # 輸出對(duì)數(shù)方差log_var def forward(self, x): # x: [batch_size, 1, 28, 28] h x.view(x.size(0), -1) # 展平: [batch_size, 784] h F.relu(self.fc1(h)) # [batch_size, 400] mu self.fc_mu(h) # [batch_size, latent_dim] log_var self.fc_logvar(h) # [batch_size, latent_dim] return mu, log_var為什么用兩個(gè)獨(dú)立的線性層輸出mu和log_var因?yàn)榫岛头讲钍欠植嫉膬蓚€(gè)獨(dú)立參數(shù)讓網(wǎng)絡(luò)各自學(xué)習(xí)更靈活。隱藏層維度hidden_dim400是一個(gè)常用起點(diǎn)你可以根據(jù)任務(wù)調(diào)整。3.2 實(shí)現(xiàn)重參數(shù)化采樣層這個(gè)層本身沒(méi)有可學(xué)習(xí)參數(shù)它只是一個(gè)計(jì)算步驟但必須繼承nn.Module以便整合到模型里。class Reparameterization(nn.Module): def forward(self, mu, log_var): 根據(jù)均值mu和對(duì)數(shù)方差log_var采樣得到潛在編碼z。 使用重參數(shù)化技巧保證梯度可傳。 std torch.exp(0.5 * log_var) # 標(biāo)準(zhǔn)差 eps torch.randn_like(std) # 標(biāo)準(zhǔn)正態(tài)噪聲 z mu eps * std return z關(guān)鍵點(diǎn)torch.randn_like(std)確保噪聲eps和std在同一設(shè)備上CPU/GPU并且形狀一致。這是采樣多樣性的來(lái)源。3.3 定義解碼器網(wǎng)絡(luò)解碼器負(fù)責(zé)將采樣得到的低維編碼z還原成原始圖片。class Decoder(nn.Module): def __init__(self, latent_dim20, hidden_dim400, output_dim784): super(Decoder, self).__init__() self.fc1 nn.Linear(latent_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, output_dim) def forward(self, z): # z: [batch_size, latent_dim] h F.relu(self.fc1(z)) # [batch_size, 400] recon torch.sigmoid(self.fc2(h)) # [batch_size, 784] # 將輸出重塑為圖片形狀例如 [batch_size, 1, 28, 28] return recon.view(-1, 1, 28, 28)為什么最后用Sigmoid激活函數(shù)因?yàn)镸NIST圖片像素值被歸一化到[0,1]區(qū)間Sigmoid能將輸出約束在同一范圍方便用BCE損失計(jì)算重建誤差。3.4 組裝完整的VAE模型現(xiàn)在把編碼器、采樣層和解碼器串起來(lái)。class VAE(nn.Module): def __init__(self, input_dim784, hidden_dim400, latent_dim20): super(VAE, self).__init__() self.encoder Encoder(input_dim, hidden_dim, latent_dim) self.reparameterize Reparameterization() self.decoder Decoder(latent_dim, hidden_dim, input_dim) def forward(self, x): mu, log_var self.encoder(x) z self.reparameterize(mu, log_var) x_recon self.decoder(z) return x_recon, mu, log_var模型的前向傳播返回三個(gè)值重建圖片x_recon、均值mu和對(duì)數(shù)方差log_var。后兩者用于計(jì)算KL散度損失。4. 設(shè)計(jì)損失函數(shù)與訓(xùn)練循環(huán)VAE的損失函數(shù)是理解其工作的重中之重。它由兩部分組成分別對(duì)應(yīng)兩個(gè)目標(biāo)。4.1 分解損失函數(shù)重建損失 KL散度def loss_function(recon_x, x, mu, log_var): recon_x: 重建的圖片 x: 原始圖片 mu: 潛在空間均值 log_var: 潛在空間對(duì)數(shù)方差 # 1. 重建損失 (Reconstruction Loss) # 使用二元交叉熵因?yàn)橄袼刂翟?-1之間。也可以用MSE。 BCE F.binary_cross_entropy(recon_x.view(-1, 784), x.view(-1, 784), reductionsum) # 2. KL散度損失 (KL Divergence Loss) # KL散度衡量學(xué)到的分布與標(biāo)準(zhǔn)正態(tài)分布的差異。 # 公式: -0.5 * sum(1 log_var - mu^2 - exp(log_var)) KLD -0.5 * torch.sum(1 log_var - mu.pow(2) - log_var.exp()) # 總損失是兩者之和 total_loss BCE KLD return total_loss, BCE, KLD為什么損失要相加BCE迫使模型好好重建圖片KLD迫使?jié)撛诜植紂(z|x)接近標(biāo)準(zhǔn)正態(tài)分布p(z)。KLD的作用是“正則化”防止編碼器為了完美重建而把方差學(xué)成0那就退化成普通AE了。兩者之間的權(quán)重通常是1:1這也是最常用的VAE目標(biāo)函數(shù)。有些變體會(huì)調(diào)整這個(gè)權(quán)重如β-VAE。4.2 構(gòu)建完整的訓(xùn)練流程有了模型和損失就可以寫(xiě)訓(xùn)練循環(huán)了。這里我給出一個(gè)最小化的、但包含關(guān)鍵步驟的訓(xùn)練循環(huán)。import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader # 1. 數(shù)據(jù)準(zhǔn)備 transform transforms.Compose([transforms.ToTensor()]) train_dataset datasets.MNIST(./data, trainTrue, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) # 2. 初始化模型、優(yōu)化器 device torch.device(cuda if torch.cuda.is_available() else cpu) model VAE().to(device) optimizer optim.Adam(model.parameters(), lr1e-3) # 3. 訓(xùn)練循環(huán) num_epochs 20 model.train() for epoch in range(num_epochs): train_loss 0 train_bce 0 train_kld 0 for batch_idx, (data, _) in enumerate(train_loader): data data.to(device) optimizer.zero_grad() # 前向傳播 recon_batch, mu, log_var model(data) # 計(jì)算損失 loss, bce, kld loss_function(recon_batch, data, mu, log_var) # 反向傳播與優(yōu)化 loss.backward() optimizer.step() train_loss loss.item() train_bce bce.item() train_kld kld.item() # 打印每個(gè)epoch的平均損失 avg_loss train_loss / len(train_loader.dataset) avg_bce train_bce / len(train_loader.dataset) avg_kld train_kld / len(train_loader.dataset) print(fEpoch {epoch1:3d}, Total Loss: {avg_loss:.4f}, BCE: {avg_bce:.4f}, KLD: {avg_kld:.4f})訓(xùn)練時(shí)要注意觀察初期BCE會(huì)快速下降KLD會(huì)上升。隨著訓(xùn)練進(jìn)行兩者會(huì)達(dá)到一個(gè)動(dòng)態(tài)平衡。如果KLD一直非常小比如接近0說(shuō)明模型可能沒(méi)用到潛在空間的隨機(jī)性生成能力會(huì)弱。如果KLD太大重建圖片會(huì)非常模糊。5. 模型評(píng)估與生成新樣本訓(xùn)練完成后我們?cè)趺粗滥P秃貌缓霉饪磽p失下降不夠必須直觀地看生成效果。5.1 重建效果評(píng)估首先看看模型重建輸入圖片的能力。這能直接反映BCE損失是否有效。import matplotlib.pyplot as plt import numpy as np def visualize_reconstruction(model, data_loader, device, num_examples8): model.eval() with torch.no_grad(): data, _ next(iter(data_loader)) data data.to(device) recon, _, _ model(data) # 將張量轉(zhuǎn)回CPU和numpy用于繪圖 data data.cpu().numpy() recon recon.cpu().numpy() fig, axes plt.subplots(2, num_examples, figsize(num_examples*2, 4)) for i in range(num_examples): axes[0, i].imshow(data[i].squeeze(), cmapgray) axes[0, i].axis(off) axes[1, i].imshow(recon[i].squeeze(), cmapgray) axes[1, i].axis(off) axes[0, 0].set_ylabel(Original) axes[1, 0].set_ylabel(Reconstructed) plt.show() # 使用測(cè)試集 test_loader DataLoader(datasets.MNIST(./data, trainFalse, transformtransform), batch_size64) visualize_reconstruction(model, test_loader, device)理想情況下重建的圖片應(yīng)該和原圖非常接近但允許有輕微模糊這是VAE的特性。如果重建圖完全無(wú)法辨認(rèn)回去檢查模型結(jié)構(gòu)或損失計(jì)算。5.2 潛在空間插值與隨機(jī)生成這才是VAE的真正價(jià)值所在。我們可以探索其學(xué)習(xí)到的連續(xù)潛在空間。隨機(jī)生成直接從標(biāo)準(zhǔn)正態(tài)分布N(0, I)中采樣z丟給解碼器。def generate_random_samples(model, latent_dim, device, num_samples64): model.eval() with torch.no_grad(): # 從標(biāo)準(zhǔn)正態(tài)分布采樣 z torch.randn(num_samples, latent_dim).to(device) samples model.decoder(z) samples samples.cpu().numpy() # 繪制生成的圖片 fig, axes plt.subplots(8, 8, figsize(12, 12)) for i, ax in enumerate(axes.flat): ax.imshow(samples[i].squeeze(), cmapgray) ax.axis(off) plt.show() generate_random_samples(model, latent_dim20, devicedevice)如果生成的數(shù)字大部分清晰可辨且多樣性好0-9都有說(shuō)明模型學(xué)到的潛在空間質(zhì)量很高。潛在空間插值在兩個(gè)真實(shí)圖片對(duì)應(yīng)的潛在編碼之間進(jìn)行線性插值觀察生成圖片的平滑過(guò)渡。def interpolate(model, data_loader, device, index1, index2, steps10): model.eval() with torch.no_grad(): # 獲取兩幅真實(shí)圖片 data, _ next(iter(data_loader)) img1, img2 data[index1:index11], data[index2:index21] img1, img2 img1.to(device), img2.to(device) # 獲取它們的潛在編碼 mu mu1, _ model.encoder(img1) mu2, _ model.encoder(img2) # 線性插值 interpolations [] for alpha in np.linspace(0, 1, steps): z alpha * mu2 (1 - alpha) * mu1 recon model.decoder(z) interpolations.append(recon.cpu()) # 繪制插值序列 fig, axes plt.subplots(1, steps, figsize(steps*2, 2)) for i, ax in enumerate(axes): ax.imshow(interpolations[i].squeeze(), cmapgray) ax.axis(off) plt.show() # 例如在測(cè)試集中選兩個(gè)不同數(shù)字的圖片進(jìn)行插值 interpolate(model, test_loader, device, index10, index210)如果插值過(guò)程生成的圖片變化是連續(xù)的、有意義的比如從“0”平滑地變成“8”而不是出現(xiàn)亂碼或突變那就證明VAE成功學(xué)習(xí)到了一個(gè)連續(xù)且有語(yǔ)義的潛在空間。6. 從MNIST到更復(fù)雜任務(wù)卷積VAE與調(diào)參實(shí)戰(zhàn)上面的全連接VAE對(duì)于MNIST是夠用的但面對(duì)更復(fù)雜的圖片如CelebA人臉、CIFAR-10就需要更強(qiáng)大的編碼器-解碼器。卷積神經(jīng)網(wǎng)絡(luò)是自然的選擇。6.1 構(gòu)建卷積VAE (ConvVAE)用卷積層替換全連接層能更好地捕捉圖像的空間局部特征。class ConvVAE(nn.Module): def __init__(self, latent_dim128, img_channels1): super(ConvVAE, self).__init__() # 編碼器: 使用卷積層下采樣 self.encoder nn.Sequential( nn.Conv2d(img_channels, 32, kernel_size4, stride2, padding1), # [B, 32, 14, 14] nn.ReLU(), nn.Conv2d(32, 64, kernel_size4, stride2, padding1), # [B, 64, 7, 7] nn.ReLU(), nn.Conv2d(64, 128, kernel_size3, stride2, padding1), # [B, 128, 4, 4] nn.ReLU(), nn.Flatten(), # [B, 128*4*42048] ) self.fc_mu nn.Linear(2048, latent_dim) self.fc_logvar nn.Linear(2048, latent_dim) # 解碼器: 使用轉(zhuǎn)置卷積層上采樣 self.decoder_fc nn.Linear(latent_dim, 2048) self.decoder nn.Sequential( nn.Unflatten(1, (128, 4, 4)), # [B, 128, 4, 4] nn.ConvTranspose2d(128, 64, kernel_size3, stride2, padding1, output_padding1), # [B, 64, 7, 7] nn.ReLU(), nn.ConvTranspose2d(64, 32, kernel_size4, stride2, padding1, output_padding1), # [B, 32, 14, 14] nn.ReLU(), nn.ConvTranspose2d(32, img_channels, kernel_size4, stride2, padding1, output_padding1), # [B, 1, 28, 28] nn.Sigmoid() ) def encode(self, x): h self.encoder(x) mu self.fc_mu(h) log_var self.fc_logvar(h) return mu, log_var def decode(self, z): h self.decoder_fc(z) recon self.decoder(h) return recon def forward(self, x): mu, log_var self.encode(x) z self.reparameterize(mu, log_var) recon self.decode(z) return recon, mu, log_var def reparameterize(self, mu, log_var): # 同上 std torch.exp(0.5 * log_var) eps torch.randn_like(std) return mu eps * std注意維度計(jì)算構(gòu)建卷積VAE時(shí)最麻煩的是確保編碼器最后的Flatten維度和解碼器最初的Unflatten維度能對(duì)上。你需要根據(jù)輸入圖片尺寸、卷積核、步長(zhǎng)和填充仔細(xì)計(jì)算特征圖的大小。上面的例子是針對(duì)28x28輸入設(shè)計(jì)的。6.2 訓(xùn)練卷積VAE的關(guān)鍵調(diào)參點(diǎn)換用更復(fù)雜的模型和更大的數(shù)據(jù)集如CelebA時(shí)訓(xùn)練策略需要調(diào)整。學(xué)習(xí)率與優(yōu)化器Adam優(yōu)化器依然是不錯(cuò)的選擇但學(xué)習(xí)率可能需要調(diào)低例如3e-4或1e-4??梢耘浜蠈W(xué)習(xí)率調(diào)度器如ReduceLROnPlateau在損失平臺(tái)期時(shí)降低學(xué)習(xí)率。批次大小Batch Size在GPU顯存允許的情況下使用更大的批次大小如128, 256有助于穩(wěn)定訓(xùn)練尤其是對(duì)KLD項(xiàng)。潛在維度Latent Dimlatent_dim是一個(gè)關(guān)鍵超參數(shù)。對(duì)于MNIST20維可能就夠了。對(duì)于復(fù)雜的人臉數(shù)據(jù)集可能需要128、256甚至更高。維度太低模型表達(dá)能力不足圖片模糊維度太高訓(xùn)練困難且KLD可能難以優(yōu)化。損失權(quán)重β經(jīng)典的VAE使用BCE KLD。β-VAE通過(guò)引入權(quán)重β 1來(lái)增大KLD的權(quán)重即BCE β * KLD這能迫使模型學(xué)習(xí)到更解耦、更具解釋性的潛在因子例如一個(gè)維度控制笑容一個(gè)維度控制發(fā)型。但β太大會(huì)嚴(yán)重?fù)p害重建質(zhì)量。這是一個(gè)需要權(quán)衡的旋鈕。梯度裁剪對(duì)于深層卷積VAE梯度爆炸有時(shí)會(huì)發(fā)生。在optimizer.step()之前加入torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)可以穩(wěn)定訓(xùn)練。6.3 監(jiān)控訓(xùn)練與調(diào)試不要只盯著總損失。在TensorBoard或WB等工具中同時(shí)記錄BCE Loss和KLD Loss。如果BCE一直很高KLD很快降到0模型可能忽略了潛在變量退化為普通自編碼器。嘗試減小KLD項(xiàng)的權(quán)重β 1或者檢查重參數(shù)化采樣是否被正確應(yīng)用eps是否參與了計(jì)算圖。如果KLD一直很高BCE降不下來(lái)模型可能過(guò)于關(guān)注讓分布規(guī)整而忽略了重建。這會(huì)導(dǎo)致生成圖片非常模糊。嘗試增大KLD項(xiàng)的權(quán)重β 1或者增加latent_dim給模型更多表達(dá)空間。訓(xùn)練震蕩劇烈降低學(xué)習(xí)率或使用梯度裁剪。7. 避坑指南與進(jìn)階思考根據(jù)我自己的實(shí)測(cè)經(jīng)驗(yàn)新手在實(shí)現(xiàn)和訓(xùn)練VAE時(shí)最容易在以下幾個(gè)地方踩坑。7.1 輸入數(shù)據(jù)歸一化坑點(diǎn)輸入圖片像素值范圍是[0, 255]但損失函數(shù)如BCE默認(rèn)期望輸入在[0, 1]區(qū)間。直接輸入會(huì)導(dǎo)致數(shù)值不穩(wěn)定損失爆炸。解決務(wù)必使用transforms.ToTensor()它會(huì)將PIL圖像或NumPy數(shù)組轉(zhuǎn)換為[C, H, W]形狀的Torch張量并自動(dòng)縮放到[0.0, 1.0]。對(duì)于其他數(shù)據(jù)集確保進(jìn)行類似的歸一化。7.2 損失函數(shù)中的“reduction”參數(shù)坑點(diǎn)在F.binary_cross_entropy中reductionmean和reductionsum效果大不相同?!甿ean’是除以批次內(nèi)總像素?cái)?shù)‘sum’是直接求和。如果使用‘mean’那么BCE和KLD通常也是求和的量級(jí)可能不匹配需要手動(dòng)調(diào)整權(quán)重。建議按照原始VAE論文和大多數(shù)實(shí)現(xiàn)對(duì)BCE和KLD都使用reductionsum然后除以批次大小batch_size來(lái)求平均。這樣兩者是天然可加的。我上面給出的loss_function正是這么做的。7.3 潛在空間坍縮Posterior Collapse現(xiàn)象在訓(xùn)練某些更復(fù)雜的VAE變體如用于文本的VAE或當(dāng)解碼器過(guò)于強(qiáng)大時(shí)可能會(huì)發(fā)生KLD損失迅速變?yōu)?編碼器輸出的log_var變得非常負(fù)方差接近0。這意味著編碼器完全忽略了輸入潛在變量z沒(méi)有攜帶任何信息。緩解策略KL退火KL Annealing在訓(xùn)練初期將KLD項(xiàng)的權(quán)重從0線性增加到1給解碼器時(shí)間先學(xué)會(huì)重建。使用更弱的解碼器例如減少解碼器的層數(shù)或神經(jīng)元數(shù)量。調(diào)整模型架構(gòu)如使用殘差連接、更細(xì)致的歸一化層。7.4 VAE的局限性及與GAN的對(duì)比VAE生成圖片的清晰度通常不如GAN生成對(duì)抗網(wǎng)絡(luò)。這是因?yàn)閂AE的優(yōu)化目標(biāo)是最大化證據(jù)下界ELBO它傾向于生成“平均化”、“保守”的結(jié)果以避免在KLD懲罰下偏離先驗(yàn)分布太遠(yuǎn)。所以VAE生成的圖片往往偏模糊。何時(shí)選擇VAE需要學(xué)習(xí)一個(gè)結(jié)構(gòu)化的潛在空間并能夠進(jìn)行插值、屬性操作等。需要同時(shí)具備編碼推理和解碼生成能力。訓(xùn)練相對(duì)穩(wěn)定不像GAN那樣容易模式崩潰。對(duì)生成圖像的極致逼真度要求不是最高。何時(shí)選擇GAN首要目標(biāo)是生成盡可能逼真、清晰的圖像。不需要對(duì)潛在空間有精確的編碼/解碼能力。在實(shí)際項(xiàng)目中VAE和GAN也常被結(jié)合形成VAE-GAN這類混合模型以期兼得兩者之長(zhǎng)。7.5 在生產(chǎn)環(huán)境中的考量如果你打算將訓(xùn)練好的VAE模型部署用于推理例如一個(gè)在線圖像生成服務(wù)需要注意分離編碼和解碼訓(xùn)練時(shí)forward返回三者部署時(shí)你可能只需要編碼或解碼功能??梢詫ncode和decode方法單獨(dú)暴露出來(lái)。關(guān)閉梯度計(jì)算與啟用eval模式推理時(shí)使用with torch.no_grad():和model.eval()來(lái)節(jié)省內(nèi)存和計(jì)算資源并固定Dropout和BatchNorm層的行為。模型導(dǎo)出考慮使用torch.jit.script或torch.jit.trace將模型序列化為TorchScript或者使用ONNX格式以便在不依賴Python環(huán)境的其他平臺(tái)上部署。性能監(jiān)控監(jiān)控生成圖片的質(zhì)量分布、潛在空間采樣點(diǎn)的分布是否仍接近標(biāo)準(zhǔn)正態(tài)防止模型在線上服務(wù)一段時(shí)間后發(fā)生漂移。最后我建議你把VAE當(dāng)作理解生成模型的一個(gè)絕佳起點(diǎn)。它的數(shù)學(xué)思想優(yōu)美實(shí)現(xiàn)相對(duì)直接是通向更復(fù)雜生成模型如擴(kuò)散模型的重要基石。動(dòng)手實(shí)現(xiàn)一遍把損失曲線畫(huà)出來(lái)看看潛在空間插值的效果遠(yuǎn)比讀十篇理論文章收獲更大。先從MNIST上的全連接VAE跑通再挑戰(zhàn)卷積VAE和更復(fù)雜的數(shù)據(jù)集這個(gè)學(xué)習(xí)路徑會(huì)扎實(shí)很多。