Commit 81e99d8d authored by Myle Ott's avatar Myle Ott
Browse files

Flake8

parent 6f96ad78
Loading
Loading
Loading
Loading
+2 −5
Changes for fairseq/models/lstm.py: 2 added lines, 5 removed lines.
Original line number Diff line number Diff line
@@ -117,8 +117,7 @@ class LSTMEncoder(FairseqEncoder):
        self.padding_idx = dictionary.pad()
        self.embed_tokens = Embedding(num_embeddings, embed_dim, self.padding_idx)
        if embed_dict:
            self.embed_tokens = utils.load_embedding(
                embed_dict, self.dictionary, self.embed_tokens)
            self.embed_tokens = utils.load_embedding(embed_dict, self.dictionary, self.embed_tokens)

        self.lstm = LSTM(
            input_size=embed_dim,
@@ -246,9 +245,7 @@ class LSTMDecoder(FairseqIncrementalDecoder):
        padding_idx = dictionary.pad()
        self.embed_tokens = Embedding(num_embeddings, embed_dim, padding_idx)
        if embed_dict:
            self.embed_tokens = utils.load_embedding(
                embed_dict, self.dictionary, self.embed_tokens)

            self.embed_tokens = utils.load_embedding(embed_dict, self.dictionary, self.embed_tokens)

        self.layers = nn.ModuleList([
            LSTMCell(
+3 −0
Changes for fairseq/utils.py: 3 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -263,6 +263,7 @@ def print_embed_overlap(embed_dict, vocab_dict):
    overlap = len(embed_keys & vocab_keys)
    print("| Found {}/{} types in embedding file.".format(overlap, len(vocab_dict)))


def parse_embedding(embed_path):
    """Parse embedding text file into a dictionary of word and embedding tensors.

@@ -282,6 +283,7 @@ def parse_embedding(embed_path):
            embed_dict[pieces[0]] = torch.Tensor([float(weight) for weight in pieces[1:]])
    return embed_dict


def load_embedding(embed_dict, vocab, embedding):
    for idx in range(len(vocab)):
        token = vocab[idx]
@@ -289,6 +291,7 @@ def load_embedding(embed_dict, vocab, embedding):
            embedding.weight.data[idx] = embed_dict[token]
    return embedding


def replace_unk(hypo_str, src_str, alignment, align_dict, unk):
    from fairseq import tokenizer
    # Tokens are strings here
+2 −2

File changed.

Contains only whitespace changes.