Commit bb5f15d1 authored by Myle Ott's avatar Myle Ott
Browse files

Iterate on need_attn and fix tests

parent 498a186d
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, need_attn=False):
    def forward(self, prev_output_tokens, encoder_out, incremental_state=None):
        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, need_attn=False):
    def forward(self, src_tokens, src_lengths, prev_output_tokens):
        encoder_out = self.encoder(src_tokens, src_lengths)
        decoder_out = self.decoder(prev_output_tokens, encoder_out, need_attn=need_attn)
        decoder_out = self.decoder(prev_output_tokens, encoder_out)
        return decoder_out

    def max_positions(self):
+6 −2
Changes for fairseq/models/fconv.py: 6 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -352,6 +352,7 @@ class FConvDecoder(FairseqIncrementalDecoder):
        self.dropout = dropout
        self.normalization_constant = normalization_constant
        self.left_pad = left_pad
        self.need_attn = True

        convolutions = extend_conv_spec(convolutions)
        in_channels = convolutions[0][0]
@@ -417,7 +418,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, need_attn=False):
    def forward(self, prev_output_tokens, encoder_out_dict=None, incremental_state=None):
        if encoder_out_dict is not None:
            encoder_out = encoder_out_dict['encoder_out']
            encoder_padding_mask = encoder_out_dict['encoder_padding_mask']
@@ -467,7 +468,7 @@ class FConvDecoder(FairseqIncrementalDecoder):

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

                if need_attn:
                if self.need_attn:
                    attn_scores = attn_scores / num_attn_layers
                    if avg_attn_scores is None:
                        avg_attn_scores = attn_scores
@@ -523,6 +524,9 @@ class FConvDecoder(FairseqIncrementalDecoder):
            state_dict['decoder.version'] = torch.Tensor([1])
        return state_dict

    def make_generation_fast_(self, need_attn=False, **kwargs):
        self.need_attn = need_attn

    def _embed_tokens(self, tokens, incremental_state):
        if incremental_state is not None:
            # keep only the last token for incremental forward pass
+6 −2
Changes for fairseq/models/fconv_self_att.py: 6 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -259,6 +259,7 @@ class FConvDecoder(FairseqDecoder):
        self.pretrained_decoder = trained_decoder
        self.dropout = dropout
        self.left_pad = left_pad
        self.need_attn = True
        in_channels = convolutions[0][0]

        def expand_bool_array(val):
@@ -352,7 +353,7 @@ class FConvDecoder(FairseqDecoder):

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

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

@@ -388,7 +389,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 self.need_attn:
                    if avg_attn_scores is None:
                        avg_attn_scores = attn_scores
                    else:
@@ -427,6 +428,9 @@ class FConvDecoder(FairseqDecoder):
        """Maximum output length supported by the decoder."""
        return self.embed_positions.max_positions()

    def make_generation_fast_(self, need_attn=False, **kwargs):
        self.need_attn = need_attn

    def _split_encoder_out(self, encoder_out):
        """Split and transpose encoder outputs."""
        # transpose only once to speed up attention layers
+6 −2
Changes for fairseq/models/lstm.py: 6 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -298,6 +298,7 @@ class LSTMDecoder(FairseqIncrementalDecoder):
        self.dropout_out = dropout_out
        self.hidden_size = hidden_size
        self.share_input_output_embed = share_input_output_embed
        self.need_attn = True

        num_embeddings = len(dictionary)
        padding_idx = dictionary.pad()
@@ -324,7 +325,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, need_attn=False):
    def forward(self, prev_output_tokens, encoder_out_dict, incremental_state=None):
        encoder_out = encoder_out_dict['encoder_out']
        encoder_padding_mask = encoder_out_dict['encoder_padding_mask']

@@ -395,7 +396,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) if need_attn else None
        attn_scores = attn_scores.transpose(0, 2) if self.need_attn else None

        # project back to size of vocabulary
        if hasattr(self, 'additional_fc'):
@@ -426,6 +427,9 @@ class LSTMDecoder(FairseqIncrementalDecoder):
        """Maximum output length supported by the decoder."""
        return int(1e5)  # an arbitrary large number

    def make_generation_fast_(self, need_attn=False, **kwargs):
        self.need_attn = need_attn


def Embedding(num_embeddings, embedding_dim, padding_idx):
    m = nn.Embedding(num_embeddings, embedding_dim, padding_idx=padding_idx)
Loading