【问题标题】:How do I extract indices of non-equivalent entries between two tensors?如何提取两个张量之间的非等价条目的索引?
【发布时间】:2020-03-03 02:14:22
【问题描述】:

我有一个具有 N 个对象类别的 N 个预测的张量,我还有另一个具有真实 N 个目标对象类别的张量。我想提取我的分类器预测错误的张量索引。

考虑以下两个张量定义为:

import torch
predictions = torch.tensor([ [0], [1], [1], [0], [0], [1] ])
target      = torch.tensor([ [0], [0], [1], [1], [0], [1] ])

我想找到一些可以传递这两个向量并返回类似index_diff = [1, 3] 的列表的函数。这个功能存在吗?我目前的想法是将这两个向量都转换为 numpy 数组,然后循环 N 次并比较每个索引处的每个条目,但这对我来说似乎有点迂回。有其他选择吗?

【问题讨论】:

    标签: python pytorch tensor


    【解决方案1】:

    类似

    index_diff = (predictions.flatten() != target.flatten()).nonzero().flatten()
    

    应该可以。

    【讨论】:

    • 成功了,谢谢!当张量的格式略有不同时,它甚至可以工作,例如:predictions = torch.tensor([ [0], [0] ])target = torch.tensor([0, 0])
    • 是的,.flatten() 将张量重塑为一维张量。您可以尝试删除它们以查看其行为方式。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-02-19
    • 2019-12-28
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多