INNER CODE UNIT · Python
base_mask
yxuansu/SimCTG · dialogue_generation/loss_func.py:28
base_mask = torch.ones(seqlen, seqlen) - torch.eye(seqlen, seqlen)
base_mask = base_mask.type(torch.FloatTensor)
bsz = len(valid_len_list)
for i in range(bsz):
one_base_mask = base_mask.clone()
one_valid_len = valid_len_list[i]
one_base_mask[:,one_valid_len:] = 0.
one_base_mask[one_valid_len:, :] = 0.
if prefix_len > 0:
one_base_mask[:prefix_len, :prefix_len] = 0.
res_list.append(one_base_mask)
res_mask = torch.stack(res_list, dim = 0)#torch.FloatTensor(res_list)
#print (res_mask)
assert res_mask.size() == torch.Size([bsz, seqlen, seqlen])
return res_mask
def contrastive_loss(margin, score_matrix, input_ids, pad_token_id, prefix_len=0):
'''