讲师中心 微信公众号
AI工具推荐 视频效率加速

如何在Python中自定义PyTorch的Loss损失函数与反向传播?

大芳小哥_6982

大芳小哥_6982

发布时间:2026-09-03 09:22:24

|

592人浏览过

|

来源于php中文网

原创

PyTorch自定义Loss必须继承nn.Module或torch.autograd.Function才能参与反向传播;普通函数因不构建计算图而无法求导;继承nn.Module需在forward中实现可导运算并调用super().__init__();需避免.item()、.numpy()等中断梯度操作;复杂场景(如不可导操作、梯度定制)应使用autograd.Function,显式定义forward与backward;务必验证loss.grad_fn非None且参数梯度不为None。

如何在python中自定义pytorch的loss损失函数与反向传播?

PyTorch自定义Loss必须继承torch.nn.Module或重写forward

直接写一个普通函数(比如def my_loss(y_pred, y_true): return torch.mean((y_pred - y_true) ** 2))虽然能算出数值,但无法参与反向传播——因为没构建计算图。PyTorch的自动求导只追踪nn.Module子类或torch.autograd.Function定义的前向/反向逻辑。

最稳妥的做法是继承nn.Module,在forward里实现计算,并确保所有张量操作都是可导的:

class MyMSELoss(torch.nn.Module):
    def __init__(self, reduction='mean'):
        super().__init__()
        self.reduction = reduction
<pre class="brush:php;toolbar:false;">def forward(self, input, target):
    loss = (input - target) ** 2
    if self.reduction == 'mean':
        return loss.mean()
    elif self.reduction == 'sum':
        return loss.sum()
    else:
        return loss

  • 别漏掉super().__init__(),否则self.reduction不会被注册为模块属性
  • reduction参数要自己实现,PyTorch不会自动帮你处理
  • 避免在forward里用.item()、.numpy()或print()等中断梯度的操作

需要控制反向传播细节时,用torch.autograd.Function

当Loss里包含不可导操作(如argmax)、想手动覆写梯度、或需节省显存(比如梯度裁剪嵌入Loss),就得用autograd.Function。它要求显式定义forward和backward,且backward的输入输出必须与forward的输出/输入一一对应。

立即学习“Python免费学习笔记(深入)”;

Python Use Agent
Python Use Agent

智能执行Python任务,自动生成、执行代码并反馈结果,无需额外配置,兼容旧命令。

下载

例如实现一个带梯度缩放的L1 Loss:

class ScaledL1Loss(torch.autograd.Function):
    @staticmethod
    def forward(ctx, input, target, scale=1.0):
        ctx.save_for_backward(input, target)
        ctx.scale = scale
        return torch.mean(torch.abs(input - target))
<pre class="brush:php;toolbar:false;">@staticmethod
def backward(ctx, grad_output):
    input, target = ctx.saved_tensors
    grad_input = grad_output * ctx.scale * torch.sign(input - target) / input.numel()
    return grad_input, None, None  # input有梯度,target和scale没有

  • ctx.save_for_backward()保存中间变量供backward使用;不能直接存Python标量(如scale),要存张量或用ctx.xxx存
  • backward返回的梯度顺序必须和forward参数顺序一致;不需要梯度的参数返回None
  • 调用时要用ScaledL1Loss.apply(input, target, 0.1),不是ScaledL1Loss()(…)

验证自定义Loss是否真正参与反向传播

光看loss值下降不够,得确认梯度确实回传到了模型参数。常见错误是Loss返回了标量但没连上模型输出,或用了.detach()、with torch.no_grad():包裹。

  • 检查Loss输出是否有grad_fn:loss.grad_fn不为None才说明在计算图中
  • 运行后立刻打印某层权重的梯度:model.layer.weight.grad应为非None的Tensor
  • 如果loss.backward()后所有grad都是None,大概率是Loss里用了.item()、numpy(),或输入input本身requires_grad=False
  • 注意input和target的device和dtype要一致,否则可能静默失败(如float32 vs float64)

多输出Loss或带辅助项时,别忘了梯度归一化

比如在目标Loss外加一个正则项:total_loss = task_loss + 0.01 * reg_loss。如果reg_loss量级远大于task_loss(如L2 norm比交叉熵大几个数量级),梯度会严重偏向正则项,训练失衡。

  • 建议对每项Loss单独.mean()或.sum()后再加权,而不是直接加未归一化的张量
  • 调试时可分别打印task_loss.item()和reg_loss.item(),观察量级是否合理
  • 若用autograd.Function实现复合Loss,backward里返回的梯度也要按相同系数缩放,否则梯度更新步长错乱

最难察觉的是Loss里隐式改变了tensor的requires_grad状态,比如通过torch.where选值后忘记保留梯度路径,或者用torch.max取index再索引——这些操作容易切断计算图。动手前先用torch.autograd.gradcheck做数值梯度校验。

热门AI工具

更多
Lovart
Lovart Hot

一款面向视觉设计创作的AI设计平台,可通过智能体和画布工作流辅助制作海报、Logo、网页、PPT及其他视觉内容。

Seko
Seko Hot

一款AI视频创作工具,主要用于商汤科技推出的创编一体的AI短视频创作Agent,适合需要提升相关任务效率的用户。

DeepSeek

DeepSeek是一款面向对话、写作、编程和推理场景的AI大模型工具。

UP简历
UP简历 Hot

一款AI办公效率工具,主要用于基于AI技术的免费在线简历制作工具,适合需要提升相关任务效率的用户。

音述AI
音述AI Hot

一款AI音频处理工具,主要用于音述AI是一个以“用声音述说故事”为核心的 AI 音乐创作与声音分享社区,适合需要提升相关任务效率的用户。

SkildArt
SkildArt Hot

SkildArt是一款AI文本写作工具,一站式 AI 视觉创作平台。

WorkBuddy

一款AI办公效率工具,主要用于腾讯云推出的AI原生桌面智能体工作台,适合需要提升相关任务效率的用户。

豆包大模型

豆包大模型是一款由字节跳动推出的企业级大语言模型服务平台。

咔片AIPPT

一款在线AI演示文稿制作工具,可根据主题和内容需求辅助生成PPT结构与页面,提高演示材料制作效率。

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

1611

2023.07.20

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

3904

2023.07.25

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

1609

2023.07.31

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

22517

2023.08.03

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2767

2023.08.04

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2807

2023.08.04

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

1123

2023.08.11

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

596

2023.08.10

PDF转图片方法
PDF转图片方法

需要把 PDF 页面用于上传、预览、分享或图片归档时,PDF 转图片方法专题整理 JPG/PNG 格式选择、逐页导出、清晰度设置、批量下载和结果检查等流程,帮助用户稳定完成 PDF 图片化处理。

0

2026.09.30

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程
关于我们 免责申明 举报中心 意见反馈 讲师合作 广告合作 最新更新
php中文网:公益在线php培训,帮助PHP学习者快速成长!
关注服务号
PHP中文网订阅号
每天精选资源文章推送

Copyright 2014-2026 https://www.php.cn/ All Rights Reserved | php.cn