【问题标题】:Pytorch TypeError - eq() received an invalid combination of argumentsPytorch TypeError - eq() 收到了无效的参数组合
【发布时间】:2020-08-27 05:19:34
【问题描述】:

我正在处理 BERT 的文本分类问题。在本地机器上训练时一切正常,但切换到服务器时,出现以下错误:

<ipython-input-28-508d35ac5f5f> in flat_accuracy(preds, labels)
      5     pred_flat = np.argmax(preds, axis=1).flatten()
      6     labels_flat = labels.flatten()
----> 7     return np.sum(pred_flat == labels_flat) / len(labels_flat)
      8 
      9 # Function to calculate the f1_score of our predictions vs labels

TypeError: eq() received an invalid combination of arguments - got (numpy.ndarray), but expected one of:
 * (Tensor other)
      didn't match because some of the arguments have invalid types: (numpy.ndarray)
 * (Number other)
      didn't match because some of the arguments have invalid types: (numpy.ndarray)

代码:

def flat_accuracy(preds, labels):
    pred_flat = np.argmax(preds, axis=1).flatten()
    labels_flat = labels.flatten()
    return np.sum(pred_flat == labels_flat) / len(labels_flat)

本地机器上的 Torch 版本:1.4.0

服务器上的 Torch 版本:1.3.1

任何帮助将不胜感激!

【问题讨论】:

    标签: numpy machine-learning pytorch tensor


    【解决方案1】:

    可能是您服务器上的 Torch 版本的 eq 实现不再允许您在 torch.Tensornp.ndarray 之间进行元素比较。您应该强制 pred_flat 成为 torch.Tensor,或强制 labels_flat 成为 numpy 数组。由于您在 return 语句中使用 np.sum 并且您只是返回一个标量值,所以我只是将所有内容移至 numpy,所以

    labels_flat = labels.numpy()
    

    但如果您使用 GPU,则可能需要调用 labels.cpu().numpy(),如果您要跟踪标签上的渐变,则可能需要 labels.detach().cpu().numpy()

    【讨论】:

      猜你喜欢
      • 2019-07-26
      • 2019-08-04
      • 1970-01-01
      • 2019-06-12
      • 2018-12-05
      • 1970-01-01
      • 1970-01-01
      • 2016-10-16
      • 2021-08-10
      相关资源
      最近更新 更多