(
self,
tokens,
encoder_outs: List[Dict[str, List[Tensor]]],
incremental_states: List[Dict[str, Dict[str, Optional[Tensor]]]],
temperature: float = 1.0,
)
| 761 | |
| 762 | @torch.jit.export |
| 763 | def forward_decoder( |
| 764 | self, |
| 765 | tokens, |
| 766 | encoder_outs: List[Dict[str, List[Tensor]]], |
| 767 | incremental_states: List[Dict[str, Dict[str, Optional[Tensor]]]], |
| 768 | temperature: float = 1.0, |
| 769 | ): |
| 770 | log_probs = [] |
| 771 | avg_attn: Optional[Tensor] = None |
| 772 | encoder_out: Optional[Dict[str, List[Tensor]]] = None |
| 773 | for i, model in enumerate(self.models): |
| 774 | if self.has_encoder(): |
| 775 | encoder_out = encoder_outs[i] |
| 776 | # decode each model |
| 777 | if self.has_incremental_states(): |
| 778 | decoder_out = model.decoder.forward( |
| 779 | tokens, |
| 780 | encoder_out=encoder_out, |
| 781 | incremental_state=incremental_states[i], |
| 782 | ) |
| 783 | else: |
| 784 | if hasattr(model, "decoder"): |
| 785 | decoder_out = model.decoder.forward(tokens, encoder_out=encoder_out) |
| 786 | else: |
| 787 | decoder_out = model.forward(tokens) |
| 788 | |
| 789 | attn: Optional[Tensor] = None |
| 790 | decoder_len = len(decoder_out) |
| 791 | if decoder_len > 1 and decoder_out[1] is not None: |
| 792 | if isinstance(decoder_out[1], Tensor): |
| 793 | attn = decoder_out[1] |
| 794 | else: |
| 795 | attn_holder = decoder_out[1]["attn"] |
| 796 | if isinstance(attn_holder, Tensor): |
| 797 | attn = attn_holder |
| 798 | elif attn_holder is not None: |
| 799 | attn = attn_holder[0] |
| 800 | if attn is not None: |
| 801 | attn = attn[:, -1, :] |
| 802 | |
| 803 | decoder_out_tuple = ( |
| 804 | decoder_out[0][:, -1:, :].div_(temperature), |
| 805 | None if decoder_len <= 1 else decoder_out[1], |
| 806 | ) |
| 807 | probs = model.get_normalized_probs( |
| 808 | decoder_out_tuple, log_probs=True, sample=None |
| 809 | ) |
| 810 | probs = probs[:, -1, :] |
| 811 | if self.models_size == 1: |
| 812 | return probs, attn |
| 813 | |
| 814 | log_probs.append(probs) |
| 815 | if attn is not None: |
| 816 | if avg_attn is None: |
| 817 | avg_attn = attn |
| 818 | else: |
| 819 | avg_attn.add_(attn) |
| 820 |
no test coverage detected