| 484 | self.generator.bias.data = weights.w['generator']['0.bias'] |
| 485 | |
| 486 | def forward(self, batch_size, beam_size, max_seq_len, memory, memory_seq_lens): |
| 487 | # nvtx.range_push("torch_decoding") |
| 488 | extended_memory = tile(memory, beam_size) |
| 489 | batchxbeam = extended_memory.size(0) |
| 490 | extended_memory = extended_memory.transpose(0, 1).contiguous() |
| 491 | |
| 492 | extended_memory_seq_lens = tile(memory_seq_lens, beam_size) |
| 493 | start_ids = extended_memory_seq_lens.new_full((batchxbeam,), self.start_id, dtype=torch.int64) |
| 494 | |
| 495 | initial_log_probs = extended_memory.new_full((beam_size,), -float("inf"), dtype=torch.float32) |
| 496 | initial_log_probs[0] = 0. |
| 497 | initial_log_probs = initial_log_probs.repeat(batch_size) |
| 498 | sequence_lengths = extended_memory_seq_lens.new_full((batchxbeam,), 0) |
| 499 | finished = extended_memory_seq_lens.new_full((batchxbeam,), 0, dtype=torch.bool) |
| 500 | |
| 501 | dtype_info = torch.finfo(extended_memory.dtype) |
| 502 | eos_max_prob = extended_memory.new_full((batchxbeam, self.vocab_size), dtype_info.min) |
| 503 | eos_max_prob[:, self.end_id] = dtype_info.max |
| 504 | |
| 505 | self.decoder.init_state(extended_memory, extended_memory, None) |
| 506 | word_ids = start_ids |
| 507 | cum_log_probs = initial_log_probs |
| 508 | |
| 509 | for step in range(max_seq_len): |
| 510 | if not torch.bitwise_not(finished).any(): |
| 511 | break |
| 512 | word_ids = word_ids.view(1, -1, 1) |
| 513 | dec_out, dec_attn = self.decoder(word_ids, extended_memory, memory_lengths=extended_memory_seq_lens, |
| 514 | step=step, decoding_max_seq_len=max_seq_len, sequence_lengths=sequence_lengths) |
| 515 | logits = self.generator(dec_out.squeeze(0)) |
| 516 | logits = torch.where(finished.view(-1, 1), eos_max_prob, logits).to(torch.float32) |
| 517 | log_probs = self.logsoftmax(logits.to(torch.float32)) |
| 518 | |
| 519 | total_probs = log_probs + torch.unsqueeze(cum_log_probs, 1) |
| 520 | total_probs = total_probs.view(-1, beam_size * self.vocab_size) |
| 521 | |
| 522 | # beamsearch |
| 523 | # _, sample_ids = torch.topk(total_probs, beam_size) |
| 524 | # sample_ids = sample_ids.view(-1) |
| 525 | |
| 526 | #diversesiblingsearch |
| 527 | sibling_score = torch.arange(1, beam_size+1).to(total_probs.dtype).to(extended_memory.device) * self.diversity_rate # [beam_size] |
| 528 | scores, ids = torch.topk(total_probs.view(-1, beam_size, self.vocab_size), beam_size) # [batch size, beam width, beam width] |
| 529 | scores = scores + sibling_score # [batch size, beam width, beam width] |
| 530 | scores = scores.view(-1, beam_size * beam_size) |
| 531 | ids = ids + torch.unsqueeze(torch.unsqueeze(torch.arange(0, beam_size).to(extended_memory.device) * self.vocab_size, 0), -1) |
| 532 | ids = ids.view(-1, beam_size * beam_size) |
| 533 | _, final_ids = torch.topk(scores, beam_size) # [batch size, beam size] |
| 534 | final_ids = final_ids.view(-1, 1) |
| 535 | batch_index = torch.arange(0, batch_size).to(extended_memory.device).view(-1, 1).repeat(1, beam_size).view(-1, 1) |
| 536 | index = torch.cat([batch_index, final_ids], 1) |
| 537 | sample_ids = gather_nd(ids, index) |
| 538 | |
| 539 | word_ids = sample_ids % self.vocab_size # [batch_size * beam_size] |
| 540 | beam_ids = sample_ids // self.vocab_size # [batch_size * beam_size] |
| 541 | beam_indices = (torch.arange(batchxbeam).to(extended_memory.device) // beam_size) * beam_size + beam_ids |
| 542 | |
| 543 | sequence_lengths = torch.where(finished, sequence_lengths, sequence_lengths + 1) |