INNER CODE UNIT · Python
compute_logits_and_hidden_states
yxuansu/SimCTG · dialogue_generation/simctgdialogue.py:46
def compute_logits_and_hidden_states(self, input_ids):
# used for advanced decoding
# input_ids: 1 x seqlen
outputs = self.model(input_ids=input_ids, output_hidden_states=True)
last_hidden_states = outputs.hidden_states[-1]
logits = outputs.logits
return last_hidden_states, logits
def forward(self, input_ids, labels, margin):
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 = train_fct(logits.view(-1, self.vocab_size), labels.view(-1))
norm_rep = last_hidden_states / last_hidden_states.norm(dim=2, keepdim=True)