INNER CODE UNIT · Python

accuracy

mlwithme/BertWithPretrained · Tasks/TaskForChineseNER.py:58

def accuracy(logits, y_true, ignore_idx=-100):
    """
    :param logits:  [src_len,batch_size,num_labels]
    :param y_true:  [src_len,batch_size]
    :param ignore_idx: 默认情况为-100
    :return:
    e.g.
    y_true = torch.tensor([[-100, 0, 0, 1, -100],
                       [-100, 2, 0, -100, -100]]).transpose(0, 1)
    logits = torch.tensor([[[0.5, 0.1, 0.2], [0.5, 0.4, 0.1], [0.7, 0.2, 0.3], [0.5, 0.7, 0.2], [0.1, 0.2, 0.5]],
                           [[0.3, 0.2, 0.5], [0.7, 0.2, 0.4], [0.8, 0.1, 0.3], [0.9, 0.2, 0.1], [0.1, 0.5, 0.2]]])
    logits = logits.transpose(0, 1)
    print(accuracy(logits, y_true, -100)) # (0.8, 4, 5)
    """
    y_pred = logits.transpose(0, 1).argmax(axis=2).reshape(-1).tolist()
    # 将 [src_len,batch_size,num_labels] 转成 [batch_size, src_len,num_labels]
    y_true = y_true.transpose(0, 1).reshape(-1).tolist()
    real_pred, real_true = [], []

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…