Commit 8441cbf3 authored by Peng-Jen Chen's avatar Peng-Jen Chen Committed by Facebook Github Bot
Browse files

Manually port pull request 385

Summary:
Manually port fairinternal fairseq-py pull request #385 [1] to fbcode.

Resolve the merge conflict of removing fp16_trainer per offline discussion with Myle. Also updated codes to make generate.py works.

[1] https://github.com/fairinternal/fairseq-py/pull/385/commits/18fa6e154781cf0c4b1596429dba7e753a545069

Reviewed By: liezl200

Differential Revision: D10052908

fbshipit-source-id: c3c378d78dc1e9ac087c815f359e78c0048ff2f5
parent 0a628401
Loading
Loading
Loading
Loading
+2 −0
Original line number Diff line number Diff line
@@ -12,6 +12,7 @@ from .indexed_dataset import IndexedDataset, IndexedCachedDataset, IndexedInMemo
from .append_eos_dataset import AppendEosDataset
from .language_pair_dataset import LanguagePairDataset
from .monolingual_dataset import MonolingualDataset
from .round_robin_zip_datasets import RoundRobinZipDatasets
from .token_block_dataset import TokenBlockDataset

from .iterators import (
@@ -35,6 +36,7 @@ __all__ = [
    'IndexedRawTextDataset',
    'LanguagePairDataset',
    'MonolingualDataset',
    'RoundRobinZipDatasets',
    'ShardedIterator',
    'TokenBlockDataset',
]
+7 −0
Original line number Diff line number Diff line
@@ -87,6 +87,13 @@ def filter_by_size(indices, size_fn, max_positions, raise_exception=False):
    def check_size(idx):
        if isinstance(max_positions, float) or isinstance(max_positions, int):
            return size_fn(idx) <= max_positions
        elif isinstance(max_positions, dict):
            idx_size = size_fn(idx)
            assert isinstance(idx_size, dict)
            intersect_keys = set(max_positions.keys()) & set(idx_size.keys())
            return all(
                idx_size[key] <= max_positions[key] for key in intersect_keys
            )
        else:
            return all(a is None or b is None or a <= b
                       for a, b in zip(size_fn(idx), max_positions))
+107 −0
Original line number Diff line number Diff line
# Copyright (c) 2017-present, Facebook, Inc.
# All rights reserved.
#
# This source code is licensed under the license found in the LICENSE file in
# the root directory of this source tree. An additional grant of patent rights
# can be found in the PATENTS file in the same directory.

from collections import OrderedDict

import numpy as np

from . import FairseqDataset


class RoundRobinZipDatasets(FairseqDataset):
    """Zip multiple FairseqDatasets together, repeating shorter datasets in a
    round-robin fashion to match the length of the longest one.

    Args:
        datasets: a dictionary of FairseqDatasets
        eval_key: an optional key used at evaluation time that causes this
            instance to pass-through batches from `datasets[eval_key]`.
    """

    def __init__(self, datasets, eval_key=None):
        super().__init__()
        assert isinstance(datasets, OrderedDict)
        self.datasets = datasets
        self.eval_key = eval_key

        self.longest_dataset = None
        self.longest_dataset_key = None
        for key, dataset in datasets.items():
            assert isinstance(dataset, FairseqDataset)
            if self.longest_dataset is None or len(dataset) > len(self.longest_dataset):
                self.longest_dataset = dataset
                self.longest_dataset_key = key

        self._ordered_indices = OrderedDict([
            (key, dataset.ordered_indices())
            for key, dataset in datasets.items()
        ])

    def _map_index(self, key, index):
        return self._ordered_indices[key][index % len(self.datasets[key])]

    def __getitem__(self, index):
        if self.eval_key is None:
            return OrderedDict([
                (key, dataset[self._map_index(key, index)])
                for key, dataset in self.datasets.items()
            ])
        else:
            # at evaluation time it's useful to pass-through batches from a single key
            return self.datasets[self.eval_key][self._map_index(self.eval_key, index)]

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

    def collater(self, samples):
        """Merge a list of samples to form a mini-batch."""
        if self.eval_key is None:
            return OrderedDict([
                (key, dataset.collater([sample[key] for sample in samples]))
                for key, dataset in self.datasets.items()
            ])
        else:
            # at evaluation time it's useful to pass-through batches from a single key
            return self.datasets[self.eval_key].collater(samples)

    def get_dummy_batch(self, max_tokens, max_positions):
        if self.eval_key is None:
            # TODO should max_tokens be used independently for each batch like this?
            return OrderedDict([
                (key, dataset.get_dummy_batch(max_tokens, max_positions[key]))
                for key, dataset in self.datasets.items()
            ])
        else:
            # at evaluation time it's useful to return a single batch directly
            return self.datasets[self.eval_key].get_dummy_batch(max_tokens, max_positions[self.eval_key])

    def num_tokens(self, index):
        """Return an example's length (number of tokens), used for batching."""
        # TODO make it configurable whether to use max() or sum() here
        return max(
            dataset.num_tokens(self._map_index(key, index))
            for key, dataset in self.datasets.items()
        )

    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``."""
        return {
            key: dataset.size(self._map_index(key, index))
            for key, dataset in self.datasets.items()
        }

    def ordered_indices(self):
        """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()
        )
+6 −1
Original line number Diff line number Diff line
@@ -12,7 +12,12 @@ import os
from .fairseq_decoder import FairseqDecoder  # noqa: F401
from .fairseq_encoder import FairseqEncoder  # noqa: F401
from .fairseq_incremental_decoder import FairseqIncrementalDecoder  # noqa: F401
from .fairseq_model import BaseFairseqModel, FairseqModel, FairseqLanguageModel  # noqa: F401
from .fairseq_model import (
    BaseFairseqModel,
    FairseqModel,  # noqa: F401
    FairseqMultiModel,  # noqa: F401
    FairseqLanguageModel,  # noqa: F401
)

from .composite_encoder import CompositeEncoder  # noqa: F401
from .distributed_fairseq_model import DistributedFairseqModel  # noqa: F401
+42 −0
Original line number Diff line number Diff line
@@ -168,6 +168,48 @@ class FairseqModel(BaseFairseqModel):
        return (self.encoder.max_positions(), self.decoder.max_positions())


class FairseqMultiModel(BaseFairseqModel):
    """Base class for combining multiple encoder-decoder models."""
    def __init__(self, encoders, decoders):
        super().__init__()
        assert encoders.keys() == decoders.keys()
        self.keys = list(encoders.keys())
        for key in self.keys:
            assert isinstance(encoders[key], FairseqEncoder)
            assert isinstance(decoders[key], FairseqDecoder)

        self.models = nn.ModuleDict({
            key: FairseqModel(encoders[key], decoders[key])
            for key in self.keys
        })

    def forward(self, src_tokens, src_lengths, prev_output_tokens):
        decoder_outs = {}
        for key in self.keys:
            encoder_out = self.models[key].encoder(src_tokens, src_lengths)
            decoder_outs[key] = self.models[key].decoder(prev_output_tokens, encoder_out)
        return decoder_outs

    def max_positions(self):
        """Maximum length supported by the model."""
        return {
            key: (self.models[key].encoder.max_positions(), self.models[key].decoder.max_positions())
            for key in self.keys
        }

    def max_decoder_positions(self):
        """Maximum length supported by the decoder."""
        return min(model.decoder.max_positions() for model in self.models.values())

    @property
    def encoder(self):
        return self.models[self.keys[0]].encoder

    @property
    def decoder(self):
        return self.models[self.keys[0]].decoder


class FairseqLanguageModel(BaseFairseqModel):
    """Base class for decoder-only models.

Loading