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

torch.nn.Softmax等四类PyTorch softmax相关API的区别是什么?

PyTorch四个softmax相关接口区别及使用示例

核心差异总览

  • torch.nn.Softmax:是nn.Module的子类,属于可嵌套的网络层组件,需要先实例化再调用,适合放到nn.Sequential这类模块化搭建的网络结构中
  • torch.nn.functional.softmax:无状态的函数式接口,不需要实例化直接传参调用,适合前向计算逻辑灵活的场景
  • torch.softmax:PyTorch 1.2之后新增的张量方法/顶级函数,功能和torch.nn.functional.softmax完全一致,仅调用形式更灵活
  • torch.nn.functional.log_softmax:在softmax计算后自动做对数变换,等价于torch.log(torch.softmax(x))但做了数值稳定性优化,避免极小值取对数出现溢出问题,多用于多分类损失计算的前置步骤

所有softmax类接口都必须指定*dim参数*,明确沿哪个维度做概率归一化,否则会触发警告甚至得到错误结果。

逐接口使用示例

前置准备代码:

import torch
import torch.nn as nn
import torch.nn.functional as F

# 测试张量,shape(2,3)代表2个样本、3个类别
x = torch.randn(2, 3)

1. torch.nn.Softmax使用示例

# 先实例化层,指定沿类别维度(dim=1)归一化
softmax_layer = nn.Softmax(dim=1)
output = softmax_layer(x)

# 验证每个样本的归一化和为1
print(output.sum(dim=1)) # 输出接近[1., 1.]

典型使用场景:作为层组件加入网络序列:

class Classifier(nn.Module):
    def __init__(self):
        super().__init__()
        self.layers = nn.Sequential(
            nn.Linear(128, 3),
            nn.Softmax(dim=1)
        )
    def forward(self, x):
        return self.layers(x)

2. torch.nn.functional.softmax使用示例

# 直接调用函数,同步传入输入和dim参数
output = F.softmax(x, dim=1)
print(output.sum(dim=1)) # 输出接近[1., 1.]

典型使用场景:forward函数中需要灵活插入计算逻辑时使用,无需提前初始化,更轻便。

3. torch.softmax使用示例

# 调用方式1:作为torch顶级函数调用
output1 = torch.softmax(x, dim=1)
# 调用方式2:作为张量方法调用
output2 = x.softmax(dim=1)

# 验证和F.softmax结果完全一致
print(torch.allclose(output1, F.softmax(x, dim=1))) # 输出True
print(torch.allclose(output2, output1)) # 输出True

这个接口仅为调用便捷性设计,和F.softmax功能、性能没有区别。

4. torch.nn.functional.log_softmax使用示例

log_output = F.log_softmax(x, dim=1)
# 等价于以下写法,但内置实现用了log-sum-exp技巧,数值稳定性更高
equivalent_output = torch.log(torch.softmax(x, dim=1))
print(torch.allclose(log_output, equivalent_output, atol=1e-6)) # 绝大多数场景下为True

典型使用场景:配合nn.NLLLoss(负对数似然损失)实现多分类任务,注意nn.CrossEntropyLoss已经内置封装了log_softmax和NLLLoss,使用该损失时不需要手动调用log_softmax。

注意:dim参数要根据张量语义设置,比如时序任务中张量shape为(批次大小, 序列长度, 特征数),要对特征维度做归一化时dim要设为2,否则会得到错误的归一化结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 01:06:04