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

使用SpeechBrain预训练模型构造自定义损失时梯度不回传问题求助

问题描述

配置神经网络计算由SpeechBrain预训练模型定义的梯度时出现梯度回传失败问题,训练轮次间损失无变化,权重未得到更新。需求为仅将SpeechBrain预训练模型作为损失函数的组成部分,不对该预训练模型执行反向传播,仅需要获取梯度df(x)/dxi。
问题复现代码如下:

import torch.nn as nn
from pytorch_lightning import LightningModule
from speechbrain.pretrained import EncoderClassifier
import torch.nn as nnloss

from torch import optim

from torch.nn import functional as F
from torch.utils.data import DataLoader, random_split
from torchvision.datasets import MNIST
import os
from torchvision import datasets, transforms
import pytorch_lightning as pl


class LightningMNISTClassifier(pl.LightningModule):
  def __init__(self):
    super().__init__()
    # mnist images are (1, 28, 28) (channels, width, height)
    self.layer_1 = nn.Linear(28 * 28, 128)
    self.layer_2 = nn.Linear(128, 256)
    self.layer_3 = nn.Linear(256, 10)

  def forward(self, x):
      batch_size, channels, width, height = x.size()

      # (b, 1, 28, 28) -> (b, 1*28*28)
      x = x.view(batch_size, -1)

      # layer 1 (b, 1*28*28) -> (b, 128)
      x = self.layer_1(x)


      # layer 2 (b, 128) -> (b, 256)
      x = self.layer_2(x)


      # layer 3 (b, 256) -> (b, 10)
      x = self.layer_3(x)



      return x

  def cross_entropy_loss(self, logits, labels):
    return F.nll_loss(logits, labels)

  def customLoss(self, logits, target):
      transformed_logits = w2v.encode_batch(logits.flatten()).to('cuda').requires_grad_(True)
      transformed_target = w2v.encode_batch(target.flatten()).to('cuda')
      loss = mse(transformed_logits, transformed_target)
      return loss

  def training_step(self, train_batch, batch_idx):
      x, y = train_batch
      logits = self.forward(x)
      loss = self.customLoss(logits, x)
      self.log('train_loss', loss)
      return loss

  def validation_step(self, val_batch, batch_idx):
      x, y = val_batch
      logits = self.forward(x)
      loss = self.cross_entropy_loss(logits, y)
      self.log('val_loss', loss)

  def configure_optimizers(self):
      optimizer = optim.Adam(self.parameters(), lr=1e-3)
      return optimizer




if __name__=='__main__':
    # data
    # transforms for images
    transform = transforms.Compose([transforms.ToTensor(),
                                    transforms.Normalize((0.1307,), (0.3081,))])

    # prepare transforms standard to MNIST
    mnist_train = MNIST(os.getcwd(), train=True, download=True, transform=transform)
    mnist_test = MNIST(os.getcwd(), train=False, download=True, transform=transform)

    train_dataloader = DataLoader(mnist_train, batch_size=64)
    val_loader = DataLoader(mnist_test, batch_size=64)

    # functions for custom loss functions
    w2v = EncoderClassifier.from_hparams(source='speechbrain/spkrec-ecapa-voxceleb')  # this is by default on the cpu
    mse = nnloss.MSELoss()

    # train
    model = LightningMNISTClassifier()
    trainer = pl.Trainer()

    trainer.fit(model, train_dataloader, val_loader)

经排查,自定义损失函数中的transformed_logits变量没有grad_fn属性,计算图断裂导致梯度无法回传至上游主网络。

问题原因
  • SpeechBrain封装的encode_batch接口默认会在torch.no_grad()上下文下执行推理,直接截断计算图,输出的张量不会保留梯度关联
  • 代码中对transformed_logits手动调用requires_grad_(True)仅能给当前张量设置梯度属性,无法修复已经断裂的上游计算图,梯度依然无法回传给主网络层
  • 预训练模型默认加载在CPU上,调用.to('cuda')转移设备的操作会进一步断裂计算图
修复方案
  1. 加载预训练模型时先完成设备迁移、参数冻结,开启eval模式的同时保留计算图跟踪:
w2v = EncoderClassifier.from_hparams(source='speechbrain/spkrec-ecapa-voxceleb')
# 冻结预训练模型所有参数,不会参与梯度更新
for param in w2v.parameters():
    param.requires_grad = False
# 提前迁移到cuda,避免后续转设备断裂计算图
w2v = w2v.cuda()
# 开启eval模式,关闭dropout、BN等训练态逻辑,同时不影响计算图跟踪
w2v.eval()
  1. 修改自定义损失函数,弃用手动设置梯度、手动转设备的操作,如果encode_batch依然截断梯度,直接调用模型内部的前向逻辑:
def customLoss(self, logits, target):
    # 直接调用模型内部前向链路,不走封装的encode_batch避免梯度截断
    logits_input = logits.flatten().unsqueeze(1)
    target_input = target.flatten().unsqueeze(1)
    transformed_logits = w2v.mods.encoder(w2v.mods.mean_var_norm(logits_input))
    transformed_target = w2v.mods.encoder(w2v.mods.mean_var_norm(target_input))
    loss = mse(transformed_logits, transformed_target)
    return loss

修复后transformed_logits会正常携带grad_fn属性,梯度可以正常回传至主网络的全连接层,权重会正常更新。

内容的提问来源于stack exchange,提问作者mathemagier

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 10:54:05