INNER CODE UNIT · Python
mle_loss
yxuansu/SimCTG · dialogue_generation/simctgdialogue.py:61
mle_loss = train_fct(logits.view(-1, self.vocab_size), labels.view(-1))
norm_rep = last_hidden_states / last_hidden_states.norm(dim=2, keepdim=True)
cosine_scores = torch.matmul(norm_rep, norm_rep.transpose(1,2))
assert cosine_scores.size() == torch.Size([bsz, seqlen, seqlen])
cl_loss = contrastive_loss(margin, cosine_scores, input_ids, self.pad_token_id, prefix_len=0)
return mle_loss, cl_loss
def eval_loss(self, input_ids, labels):
bsz, seqlen = input_ids.size()
outputs = self.model(input_ids=input_ids, output_hidden_states=True)
logits = outputs.logits
assert logits.size() == torch.Size([bsz, seqlen, self.vocab_size])
last_hidden_states = outputs.hidden_states[-1]
assert last_hidden_states.size() == torch.Size([bsz, seqlen, self.embed_dim])
mle_loss = val_fct(logits.view(-1, self.vocab_size), labels.view(-1))
assert mle_loss.size() == torch.Size([bsz * seqlen])
mask_tmp = labels.masked_fill(~labels.eq(-100), 1.0)