INNER CODE UNIT · Python
contrastive_loss
yxuansu/SimCTG · dialogue_generation/loss_func.py:44
def contrastive_loss(margin, score_matrix, input_ids, pad_token_id, prefix_len=0):
'''
margin: predefined margin to push similarity score away
score_matrix: bsz x seqlen x seqlen
input_ids: bsz x seqlen
pad_token_id: indicating which tokens are padding token
'''
bsz, seqlen, _ = score_matrix.size()
gold_score = torch.diagonal(score_matrix, offset=0, dim1=1, dim2=2) # bsz x seqlen
gold_score = torch.unsqueeze(gold_score, -1)
assert gold_score.size() == torch.Size([bsz, seqlen, 1])
difference_matrix = gold_score - score_matrix
assert difference_matrix.size() == torch.Size([bsz, seqlen, seqlen])
loss_matrix = margin - difference_matrix # bsz x seqlen x seqlen
loss_matrix = torch.nn.functional.relu(loss_matrix)
### input mask
input_mask = torch.ones_like(input_ids).type(torch.FloatTensor)