MCPcopy Create free account
hub / github.com/NVIDIA/FasterTransformer / forward

Method forward

examples/pytorch/decoding/utils/decoding.py:486–564  ·  view source on GitHub ↗
(self, batch_size, beam_size, max_seq_len, memory, memory_seq_lens)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 8

gather_ndFunction · 0.85
init_stateMethod · 0.80
anyMethod · 0.80
map_stateMethod · 0.80
finalizeFunction · 0.70
tileFunction · 0.50
sizeMethod · 0.45
toMethod · 0.45

Tested by

no test coverage detected