【问题标题】:How to understand "-input[range(target.shape[0]), target]"?如何理解“-input[range(target.shape[0]), target]”?
【发布时间】:2020-11-21 03:30:40
【问题描述】:

在the official example of PyTorch中,给出如下损失函数。

def nll(input, target):
    return -input[range(target.shape[0]), target].mean()

loss_func = nll

如何理解上述函数中“input[range(target.shape[0]), target]”的语法? “输入”有一个 torch.Size([64, 10]),“目标”有一个 torch.Size([64])。为什么在这里使用“range”函数?

【问题讨论】:

    标签: python pytorch


    【解决方案1】:

    范围函数用作创建从 0 到 64 的向量/列表/生成器的快捷方式。因此它本质上是 [0,1,2,...64] 的简写

    要明确这一点,您可以执行以下操作:

    def nll(input, target):
        minputlist = list(range(target.shape[0]))
        print(minputlist )
        return -input[minputlist, target].mean()
    

    【讨论】:

    • 为什么直接运行a = np.random.randn(3,2) b = np.random.randn(3) a[range(b.shape[0]),b] 会报错?
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2012-06-13
    • 2014-10-21
    • 2020-09-07
    • 1970-01-01
    • 2013-12-16
    • 2013-04-27
    • 2018-05-01
    相关资源
    最近更新 更多