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 = [], []