Commit 89e19d42 authored by Alexei Baevski's avatar Alexei Baevski Committed by Myle Ott
Browse files

disable printing alignment by default (for perf) and add a flag to enable it

parent f472d141
Loading
Loading
Loading
Loading
+1 −1
Changes for fairseq/models/fairseq_incremental_decoder.py: 1 added line, 1 removed line.
Original line number Diff line number Diff line
@@ -14,7 +14,7 @@ class FairseqIncrementalDecoder(FairseqDecoder):
    def __init__(self, dictionary):
        super().__init__(dictionary)

    def forward(self, prev_output_tokens, encoder_out, incremental_state=None):
    def forward(self, prev_output_tokens, encoder_out, incremental_state=None, need_attn=False):
        raise NotImplementedError

    def reorder_incremental_state(self, incremental_state, new_order):
+2 −2
Changes for fairseq/models/fairseq_model.py: 2 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -104,9 +104,9 @@ class FairseqModel(BaseFairseqModel):
        assert isinstance(self.encoder, FairseqEncoder)
        assert isinstance(self.decoder, FairseqDecoder)

    def forward(self, src_tokens, src_lengths, prev_output_tokens):
    def forward(self, src_tokens, src_lengths, prev_output_tokens, need_attn):
        encoder_out = self.encoder(src_tokens, src_lengths)
        decoder_out = self.decoder(prev_output_tokens, encoder_out)
        decoder_out = self.decoder(prev_output_tokens, encoder_out, need_attn=need_attn)
        return decoder_out

    def max_positions(self):
+3 −1
Changes for fairseq/models/fconv.py: 3 added lines, 1 removed line.
Original line number Diff line number Diff line
@@ -417,7 +417,7 @@ class FConvDecoder(FairseqIncrementalDecoder):
            else:
                self.fc3 = Linear(out_embed_dim, num_embeddings, dropout=dropout)

    def forward(self, prev_output_tokens, encoder_out_dict=None, incremental_state=None):
    def forward(self, prev_output_tokens, encoder_out_dict=None, incremental_state=None, need_attn=False):
        if encoder_out_dict is not None:
            encoder_out = encoder_out_dict['encoder_out']
            encoder_padding_mask = encoder_out_dict['encoder_padding_mask']
@@ -466,6 +466,8 @@ class FConvDecoder(FairseqIncrementalDecoder):
                x = self._transpose_if_training(x, incremental_state)

                x, attn_scores = attention(x, target_embedding, (encoder_a, encoder_b), encoder_padding_mask)

                if need_attn:
                    attn_scores = attn_scores / num_attn_layers
                    if avg_attn_scores is None:
                        avg_attn_scores = attn_scores
+2 −1
Changes for fairseq/models/fconv_self_att.py: 2 added lines, 1 removed line.
Original line number Diff line number Diff line
@@ -352,7 +352,7 @@ class FConvDecoder(FairseqDecoder):

            self.pretrained_decoder.fc2.register_forward_hook(save_output())

    def forward(self, prev_output_tokens, encoder_out_dict):
    def forward(self, prev_output_tokens, encoder_out_dict, need_attn=False):
        encoder_out = encoder_out_dict['encoder']['encoder_out']
        trained_encoder_out = encoder_out_dict['pretrained'] if self.pretrained else None

@@ -388,6 +388,7 @@ class FConvDecoder(FairseqDecoder):
                r = x
                x, attn_scores = attention(attproj(x) + target_embedding, encoder_a, encoder_b)
                x = x + r
                if need_attn:
                    if avg_attn_scores is None:
                        avg_attn_scores = attn_scores
                    else:
+2 −2
Changes for fairseq/models/lstm.py: 2 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -320,7 +320,7 @@ class LSTMDecoder(FairseqIncrementalDecoder):
        if not self.share_input_output_embed:
            self.fc_out = Linear(out_embed_dim, num_embeddings, dropout=dropout_out)

    def forward(self, prev_output_tokens, encoder_out_dict, incremental_state=None):
    def forward(self, prev_output_tokens, encoder_out_dict, incremental_state=None, need_attn=False):
        encoder_out = encoder_out_dict['encoder_out']
        encoder_padding_mask = encoder_out_dict['encoder_padding_mask']

@@ -391,7 +391,7 @@ class LSTMDecoder(FairseqIncrementalDecoder):
        x = x.transpose(1, 0)

        # srclen x tgtlen x bsz -> bsz x tgtlen x srclen
        attn_scores = attn_scores.transpose(0, 2)
        attn_scores = attn_scores.transpose(0, 2) if need_attn else None

        # project back to size of vocabulary
        if hasattr(self, 'additional_fc'):
Loading