You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.15 02:30:38