Commit 6ec5022e authored by Myle Ott's avatar Myle Ott
Browse files

Move reorder_encoder_out to FairseqEncoder and fix non-incremental decoding

parent e9967cd3
Loading
Loading
Loading
Loading
+6 −0
Changes for fairseq/models/composite_encoder.py: 6 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -26,6 +26,12 @@ class CompositeEncoder(FairseqEncoder):
            encoder_out[key] = self.encoders[key](src_tokens, src_lengths)
        return encoder_out

    def reorder_encoder_out(self, encoder_out, new_order):
        """Reorder encoder output according to new_order."""
        for key in self.encoders:
            encoder_out[key] = self.encoders[key].reorder_encoder_out(encoder_out[key], new_order)
        return encoder_out

    def max_positions(self):
        return min([self.encoders[key].max_positions() for key in self.encoders])

+4 −0
Changes for fairseq/models/fairseq_encoder.py: 4 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -18,6 +18,10 @@ class FairseqEncoder(nn.Module):
    def forward(self, src_tokens, src_lengths):
        raise NotImplementedError

    def reorder_encoder_out(self, encoder_out, new_order):
        """Reorder encoder output according to new_order."""
        raise NotImplementedError

    def max_positions(self):
        """Maximum input length supported by the encoder."""
        raise NotImplementedError
+0 −3
Changes for fairseq/models/fairseq_incremental_decoder.py: 0 added lines, 3 removed lines.
Original line number Diff line number Diff line
@@ -32,9 +32,6 @@ class FairseqIncrementalDecoder(FairseqDecoder):
                )
        self.apply(apply_reorder_incremental_state)

    def reorder_encoder_out(self, encoder_out, new_order):
        return encoder_out

    def set_beam_size(self, beam_size):
        """Sets the beam size in the decoder and all children."""
        if getattr(self, '_beam_size', -1) != beam_size:
+11 −6
Changes for fairseq/models/fconv.py: 11 added lines, 6 removed lines.
Original line number Diff line number Diff line
@@ -268,6 +268,17 @@ class FConvEncoder(FairseqEncoder):
            'encoder_padding_mask': encoder_padding_mask,  # B x T
        }

    def reorder_encoder_out(self, encoder_out_dict, new_order):
        if encoder_out_dict['encoder_out'] is not None:
            encoder_out_dict['encoder_out'] = (
                encoder_out_dict['encoder_out'][0].index_select(0, new_order),
                encoder_out_dict['encoder_out'][1].index_select(0, new_order),
            )
        if encoder_out_dict['encoder_padding_mask'] is not None:
            encoder_out_dict['encoder_padding_mask'] = \
                encoder_out_dict['encoder_padding_mask'].index_select(0, new_order)
        return encoder_out_dict

    def max_positions(self):
        """Maximum input length supported by the encoder."""
        return self.embed_positions.max_positions()
@@ -496,12 +507,6 @@ class FConvDecoder(FairseqIncrementalDecoder):
            encoder_out = tuple(eo.index_select(0, new_order) for eo in encoder_out)
            utils.set_incremental_state(self, incremental_state, 'encoder_out', encoder_out)

    def reorder_encoder_out(self, encoder_out_dict, new_order):
        if encoder_out_dict['encoder_padding_mask'] is not None:
            encoder_out_dict['encoder_padding_mask'] = \
                encoder_out_dict['encoder_padding_mask'].index_select(0, new_order)
        return encoder_out_dict

    def max_positions(self):
        """Maximum output length supported by the decoder."""
        return self.embed_positions.max_positions() if self.embed_positions is not None else float('inf')
+14 −19
Changes for fairseq/models/fconv_self_att.py: 14 added lines, 19 removed lines.
Original line number Diff line number Diff line
@@ -226,6 +226,19 @@ class FConvEncoder(FairseqEncoder):
            'encoder_out': (x, y),
        }

    def reorder_encoder_out(self, encoder_out_dict, new_order):
        encoder_out_dict['encoder_out'] = tuple(
            eo.index_select(0, new_order) for eo in encoder_out_dict['encoder_out']
        )

        if 'pretrained' in encoder_out_dict:
            encoder_out_dict['pretrained']['encoder_out'] = tuple(
                eo.index_select(0, new_order)
                for eo in encoder_out_dict['pretrained']['encoder_out']
            )

        return encoder_out_dict

    def max_positions(self):
        """Maximum input length supported by the encoder."""
        return self.embed_positions.max_positions()
@@ -409,30 +422,12 @@ class FConvDecoder(FairseqDecoder):
        else:
            return x, avg_attn_scores

    def reorder_incremental_state(self, incremental_state, new_order):
        """Reorder buffered internal state (for incremental generation)."""
        super().reorder_incremental_state(incremental_state, new_order)

    def reorder_encoder_out(self, encoder_out_dict, new_order):
        encoder_out_dict['encoder']['encoder_out'] = tuple(
            eo.index_select(0, new_order) for eo in encoder_out_dict['encoder']['encoder_out']
        )

        if 'pretrained' in encoder_out_dict:
            encoder_out_dict['pretrained']['encoder']['encoder_out'] = tuple(
                eo.index_select(0, new_order)
                for eo in encoder_out_dict['pretrained']['encoder']['encoder_out']
            )

        return encoder_out_dict

    def max_positions(self):
        """Maximum output length supported by the decoder."""
        return self.embed_positions.max_positions()

    def _split_encoder_out(self, encoder_out):
        """Split and transpose encoder outputs.
        """
        """Split and transpose encoder outputs."""
        # transpose only once to speed up attention layers
        encoder_a, encoder_b = encoder_out
        encoder_a = encoder_a.transpose(0, 1).contiguous()
Loading