FastAI卷积VAE损失函数梯度无法计算,训练一步参数变NaN
问题描述
我将PyTorch官方VAE示例修改为卷积网络后,尝试在FastAI框架中实现,代码如下:
import torch import torch.nn as nn import torch.nn.functional as F from fastai.vision.all import * class convVAE(nn.Module): def __init__(self, dim_z=20): super(convVAE, self).__init__() self.cv1 = nn.Conv2d(1, 32, 3, stride=2) self.cv2 = nn.Conv2d(32, 64, 3, stride=2) self.fc31 = nn.Linear(2304, dim_z) self.fc32 = nn.Linear(2304, dim_z) self.fc4 = nn.Linear(dim_z, 2304) self.cv5 = nn.ConvTranspose2d(64, 32, 3, stride=2) self.cv6 = nn.ConvTranspose2d(32, 1, 3, stride=2, output_padding=1) def encode(self, x): h1 = F.leaky_relu(self.cv1(x)) h2 = F.leaky_relu(self.cv2(h1)).view(-1, 2304) return self.fc31(h2), self.fc32(h2) def reparameterize(self, mu, logvar): std = torch.exp(0.5*logvar) eps = torch.randn_like(std) return mu + eps*std def decode(self, z): h5 = F.leaky_relu(self.fc4(z)).view(-1, 64, 6, 6) h6 = F.leaky_relu(self.cv5(h5)) return torch.sigmoid(self.cv6(h6)) def forward(self, x): mu, logvar = self.encode(x) z = self.reparameterize(mu, logvar) return self.decode(z).view(-1, 784), mu, logvar def get_loss(res,y): y_hat, mu, logvar = res BCE = F.binary_cross_entropy( y.view(-1, 784), y_hat, reduction='sum') KLD = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) return BCE + KLD block = DataBlock( blocks=(ImageBlock(cls=PILImageBW),ImageBlock(cls=PILImageBW)), get_items=get_image_files, splitter=RandomSplitter(valid_pct=0.2, seed=42), get_y=(lambda x: x), batch_tfms=aug_transforms(mult=2., do_flip=False)) path = untar_data(URLs.MNIST) loaders = block.dataloaders(path/"training",num_workers=0,bs=32) loaders.train.show_batch(max_n=4, nrows=1) mdl = convVAE(5) learn = Learner(loaders, mdl, loss_func = convVAE.get_loss) learn.fit(1, cbs=ShortEpochCallback())
当前遇到的问题:损失函数可计算(数值量级约为1e6),但梯度无法正常传导,仅训练一步后所有参数就变为NaN。该模型与损失函数在原生PyTorch环境中可正常运行。
解决方案
1. 修正数据增强的像素范围
aug_transforms(mult=2.)会将图像像素值从[0,1]放大到[0,2],而模型输出经过sigmoid限制在[0,1],这会导致BCE损失计算时目标值超出输出范围,直接引发梯度爆炸。修改数据增强参数:
batch_tfms=aug_transforms(do_flip=False)
2. 调整损失函数为平均损失
原损失函数使用reduction='sum'会导致损失值过大(1e6量级),进而引发梯度爆炸。改为平均损失,同时适配FastAI的损失函数参数逻辑:
def vae_loss(preds, target): y_hat, mu, logvar = preds # 保持目标与输出形状一致,避免额外view操作 BCE = F.binary_cross_entropy(y_hat, target, reduction='mean') KLD = -0.5 * torch.mean(1 + logvar - mu.pow(2) - logvar.exp()) return BCE + KLD
同时修改模型forward方法,返回与输入形状一致的输出,无需展平:
def forward(self, x): mu, logvar = self.encode(x) z = self.reparameterize(mu, logvar) # 输出形状保持(batch_size,1,28,28),与目标匹配 return self.decode(z), mu, logvar
3. 初始化模型参数
FastAI默认的参数初始化可能不如原生PyTorch适配VAE,手动初始化卷积和线性层参数,避免梯度异常:
def init_weights(m): if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d, nn.Linear)): nn.init.xavier_normal_(m.weight) if m.bias is not None: nn.init.constant_(m.bias, 0) mdl = convVAE(5) mdl.apply(init_weights)
4. 指定合适的学习率
FastAI默认学习率可能过高,初始化Learner时显式设置合理的学习率:
learn = Learner(loaders, mdl, loss_func=vae_loss, lr=1e-3)
验证修改后的完整代码
import torch import torch.nn as nn import torch.nn.functional as F from fastai.vision.all import * class convVAE(nn.Module): def __init__(self, dim_z=20): super(convVAE, self).__init__() self.cv1 = nn.Conv2d(1, 32, 3, stride=2) self.cv2 = nn.Conv2d(32, 64, 3, stride=2) self.fc31 = nn.Linear(2304, dim_z) self.fc32 = nn.Linear(2304, dim_z) self.fc4 = nn.Linear(dim_z, 2304) self.cv5 = nn.ConvTranspose2d(64, 32, 3, stride=2) self.cv6 = nn.ConvTranspose2d(32, 1, 3, stride=2, output_padding=1) def encode(self, x): h1 = F.leaky_relu(self.cv1(x)) h2 = F.leaky_relu(self.cv2(h1)).view(-1, 2304) return self.fc31(h2), self.fc32(h2) def reparameterize(self, mu, logvar): std = torch.exp(0.5*logvar) eps = torch.randn_like(std) return mu + eps*std def decode(self, z): h5 = F.leaky_relu(self.fc4(z)).view(-1, 64, 6, 6) h6 = F.leaky_relu(self.cv5(h5)) return torch.sigmoid(self.cv6(h6)) def forward(self, x): mu, logvar = self.encode(x) z = self.reparameterize(mu, logvar) return self.decode(z), mu, logvar def vae_loss(preds, target): y_hat, mu, logvar = preds BCE = F.binary_cross_entropy(y_hat, target, reduction='mean') KLD = -0.5 * torch.mean(1 + logvar - mu.pow(2) - logvar.exp()) return BCE + KLD def init_weights(m): if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d, nn.Linear)): nn.init.xavier_normal_(m.weight) if m.bias is not None: nn.init.constant_(m.bias, 0) block = DataBlock( blocks=(ImageBlock(cls=PILImageBW),ImageBlock(cls=PILImageBW)), get_items=get_image_files, splitter=RandomSplitter(valid_pct=0.2, seed=42), get_y=(lambda x: x), batch_tfms=aug_transforms(do_flip=False)) path = untar_data(URLs.MNIST) loaders = block.dataloaders(path/"training",num_workers=0,bs=32) loaders.train.show_batch(max_n=4, nrows=1) mdl = convVAE(5) mdl.apply(init_weights) learn = Learner(loaders, mdl, loss_func=vae_loss, lr=1e-3) learn.fit(5, cbs=ShortEpochCallback())
内容的提问来源于stack exchange,提问作者Ben K
相关产品推荐
相关产品推荐

