zoukankan      html  css  js  c++  java
  • Pytorch 分类问题输出结果的数据整理方式:_, predicted = torch.max(outputs.data, 1)

    使用pytorch的小伙伴们,一定看过下面这段代码

    _, predicted = torch.max(outputs.data, 1)

    那么,这里的 下划线_ 表示什么意思?

    首先,torch.max()这个函数返回的是两个值,第一个值是具体的value(我们用下划线_表示),第二个值是value所在的index(也就是predicted)。

    那么,这个 下划线_ 表示的就是具体的value,也就是输出的最大值。那么为什么用 下划线_,可不可以用其他的变量名称来代替,比如x?答案自然是可以的。

    那么为什么这里选择用这么特殊的下划线?有没有特殊含义?这是因为我们不关心最大值是什么,而关心最大值对应的index是什么,所以选用下划线代表不需要用到的变量。比如在图像分类任务中,值所对应的index就对应着相应的类别class,当我们只关心网络预测的类别是什么,而不关心该类别的预测概率是多少时,就选择使用下划线_。

    其次,这里的数字1表示什么意思?

    数字1其实可以写为dim=1,这里简写为1,python也可以自动识别,dim=1表示输出所在行的最大值,若改写成dim=0则输出所在列的最大值。比如说测试集有10个数据,那么训练好的网络将会预测这10个数据,得到一个10×2的矩阵(假设是二分类问题),比如说预测结果是下面这个矩阵。

    在这里插入图片描述

    这里的数字就是,网络预测为对应类别的概率,而行代表样本、列代表类别,所以这里应该用dim=1,因为你需要输出的每个样本的预测类别。

    原文链接:https://blog.csdn.net/weixin_48249563/article/details/111387501




    如果这篇文章帮助到了你,你可以请作者喝一杯咖啡

  • 相关阅读:
    冒泡排序(可传函数参数)
    字符串转化成整数
    遍历文档内容,得到HTML层级结构
    跨域通信问题
    矩阵的快速幂
    CSS tricks
    牛客赛马网笔试攻略
    项目中遇到的问题
    Mysql远程链接访问权限设置
    恢复delete删除的数据
  • 原文地址:https://www.cnblogs.com/sddai/p/14447586.html
Copyright © 2011-2022 走看看