【问题标题】:In pytorch, how to use the weight parameter in F.cross_entropy()?在pytorch中,如何使用F.cross_entropy()中的权重参数?
【发布时间】:2018-10-19 06:09:16
【问题描述】:

我正在尝试编写如下代码:

x = Variable(torch.Tensor([[1.0,2.0,3.0]]))
y = Variable(torch.LongTensor([1]))
w = torch.Tensor([1.0,1.0,1.0])
F.cross_entropy(x,y,w)
w = torch.Tensor([1.0,10.0,1.0])
F.cross_entropy(x,y,w)

然而,无论 w 是什么,交叉熵损失的输出总是 1.4076。 F.cross_entropy() 的权重参数背后是什么?如何正确使用?
我正在使用 pytorch 0.3

【问题讨论】:

    标签: deep-learning pytorch cross-entropy


    【解决方案1】:

    weight 参数用于根据目标类计算所有输入的加权结果。如果您只有一个输入或同一目标类的所有输入,weight 不会影响损失。

    查看不同目标类的 2 个输入的差异:

    import torch
    import torch.nn.functional as F
    from torch.autograd import Variable
    
    x = Variable(torch.Tensor([[1.0,2.0,3.0], [1.0,2.0,3.0]]))
    y = Variable(torch.LongTensor([1, 2]))
    w = torch.Tensor([1.0,1.0,1.0])
    res = F.cross_entropy(x,y,w)
    # 0.9076
    w = torch.Tensor([1.0,10.0,1.0])
    res = F.cross_entropy(x,y,w)
    # 1.3167
    

    【讨论】:

    • 我没有找到显示如何使用权重的实际表达式,而且 c++ 对我来说很难破译,有什么线索可以找到这些细节吗?
    • @pixelou:你可以在torch.nn.CrossEntropyLossdoc中找到weights的损失方程。否则你有 python 实现here
    • 我的错,我读这个问题太快了,虽然它是关于二进制 CE 的。不过,您仍然回答了我的问题;-):在您链接的 python 代码中,它显示后者在 pytorch 中使用逐点权重而不是类权重。
    • @pixelou:很高兴它仍然有帮助! :)
    猜你喜欢
    • 1970-01-01
    • 2019-05-01
    • 1970-01-01
    • 2020-11-09
    • 1970-01-01
    • 2018-09-01
    • 2021-08-16
    • 2018-03-14
    • 2019-06-24
    相关资源
    最近更新 更多