Commit 9998bbfa authored by Myle Ott's avatar Myle Ott Committed by Facebook Github Bot
Browse files

Merge internal changes

Summary: Pull Request resolved: https://github.com/pytorch/fairseq/pull/505

Differential Revision: D14110201

Pulled By: myleott

fbshipit-source-id: 099ce61fa386c016f3a1d1815c6fe1a9a6c9005d
parent 184629a7
Loading
Loading
Loading
Loading
+0 −4
Original line number Diff line number Diff line
@@ -153,10 +153,6 @@ class BacktranslationDataset(FairseqDataset):
        """Just use the tgt dataset ordered_indices"""
        return self.tgt_dataset.ordered_indices()

    def valid_size(self, index, max_positions):
        """Just use the tgt dataset size"""
        return self.tgt_dataset.valid_size(index, max_positions)

    def size(self, index):
        """Return an example's size as a float or tuple. This value is used
        when filtering a dataset with ``--max-positions``.
+3 −1
Original line number Diff line number Diff line
@@ -30,9 +30,10 @@ class CountingIterator(object):
        self.iterable = iterable
        self.count = 0
        self.itr = iter(self)
        self.len = len(iterable)

    def __len__(self):
        return len(self.iterable)
        return self.len

    def __iter__(self):
        for x in self.iterable:
@@ -49,6 +50,7 @@ class CountingIterator(object):
    def skip(self, num_to_skip):
        """Fast-forward the iterator by skipping *num_to_skip* elements."""
        next(itertools.islice(self.itr, num_to_skip, num_to_skip), None)
        self.len -= num_to_skip
        return self


+0 −7
Original line number Diff line number Diff line
@@ -104,13 +104,6 @@ class RoundRobinZipDatasets(FairseqDataset):
        """Ordered indices for batching."""
        return np.arange(len(self))

    def valid_size(self, index, max_positions):
        """Check if an example's size is valid according to max_positions."""
        return all(
            dataset.valid_size(self._map_index(key, index), max_positions[key])
            for key, dataset in self.datasets.items()
        )

    @property
    def supports_prefetch(self):
        return all(
+6 −2
Original line number Diff line number Diff line
@@ -37,8 +37,12 @@ class FairseqDecoder(nn.Module):
        """Get normalized probabilities (or log probs) from a net's output."""

        if hasattr(self, 'adaptive_softmax') and self.adaptive_softmax is not None:
            assert sample is not None and 'target' in sample
            out = self.adaptive_softmax.get_log_prob(net_output[0], sample['target'])
            if sample is not None:
                assert 'target' in sample
                target = sample['target']
            else:
                target = None
            out = self.adaptive_softmax.get_log_prob(net_output[0], target=target)
            return out.exp_() if not log_probs else out

        logits = net_output[0].float()
+12 −3
Original line number Diff line number Diff line
@@ -67,8 +67,8 @@ class MultilingualTransformerModel(FairseqMultiModel):
        if not hasattr(args, 'max_target_positions'):
            args.max_target_positions = 1024

        src_langs = [lang_pair.split('-')[0] for lang_pair in args.lang_pairs]
        tgt_langs = [lang_pair.split('-')[1] for lang_pair in args.lang_pairs]
        src_langs = [lang_pair.split('-')[0] for lang_pair in task.lang_pairs]
        tgt_langs = [lang_pair.split('-')[1] for lang_pair in task.lang_pairs]

        if args.share_encoders:
            args.share_encoder_embeddings = True
@@ -158,12 +158,21 @@ class MultilingualTransformerModel(FairseqMultiModel):
            shared_decoder = get_decoder(tgt_langs[0])

        encoders, decoders = OrderedDict(), OrderedDict()
        for lang_pair, src, tgt in zip(args.lang_pairs, src_langs, tgt_langs):
        for lang_pair, src, tgt in zip(task.lang_pairs, src_langs, tgt_langs):
            encoders[lang_pair] = shared_encoder if shared_encoder is not None else get_encoder(src)
            decoders[lang_pair] = shared_decoder if shared_decoder is not None else get_decoder(tgt)

        return MultilingualTransformerModel(encoders, decoders)

    def load_state_dict(self, state_dict, strict=True):
        state_dict_subset = state_dict.copy()
        for k, v in state_dict.items():
            assert k.startswith('models.')
            lang_pair = k.split('.')[1]
            if lang_pair not in self.models:
                del state_dict_subset[k]
        super().load_state_dict(state_dict_subset, strict=strict)


@register_model_architecture('multilingual_transformer', 'multilingual_transformer')
def base_multilingual_architecture(args):
Loading