Commit b59815bc authored by Angela Fan's avatar Angela Fan Committed by Myle Ott
Browse files

added multiscale gated self attention layer with multiple heads, and pretrained fusion models

parent 50931d69
Loading
Loading
Loading
Loading
+1 −0
Changes for fairseq/models/__init__.py: 1 added line, 0 removed lines.
Original line number Diff line number Diff line
@@ -12,6 +12,7 @@ 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 .composite_encoder import CompositeEncoder # noqa: F401

MODEL_REGISTRY = {}
ARCH_MODEL_REGISTRY = {}
+35 −0
Changes for fairseq/models/composite_encoder.py: 35 added lines, 0 removed lines.
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 . import FairseqEncoder


class CompositeEncoder(FairseqEncoder):
    """
    Encoder class that forwards on multiple encoders, for example for a fusion model or question-answering
    Accepts a dictionary of encoder, the first encoder's dictionary is used for initialization
    """

    def __init__(self, encoders):
        super().__init__(next(iter(encoders.values())).dictionary)
        self.encoders = encoders
        for key in self.encoders:
            self.add_module(key, self.encoders[key])

    def forward(self, src_tokens, src_lengths):
        encoder_out = {}
        for key in self.encoders:
            encoder_out[key] = self.encoders[key](src_tokens, src_lengths)
        return encoder_out

    def max_positions(self):
        return min([self.encoders[key].max_positions() for key in self.encoders])

    def upgrade_state_dict(self, state_dict):
        for key in self.encoders:
            self.encoders[key].upgrade_state_dict(state_dict)
        return state_dict
+502 −0

File added.

Preview size limit exceeded, changes collapsed.

+4 −0
Changes for fairseq/modules/__init__.py: 4 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -8,19 +8,23 @@
from .adaptive_softmax import AdaptiveSoftmax
from .beamable_mm import BeamableMM
from .conv_tbc import ConvTBC
from .downsampled_multihead_attention import DownsampledMultiHeadAttention
from .grad_multiply import GradMultiply
from .learned_positional_embedding import LearnedPositionalEmbedding
from .linearized_convolution import LinearizedConvolution
from .multihead_attention import MultiheadAttention
from .scalar_bias import ScalarBias
from .sinusoidal_positional_embedding import SinusoidalPositionalEmbedding

__all__ = [
    'AdaptiveSoftmax',
    'BeamableMM',
    'ConvTBC',
    'DownsampledMultiHeadAttention',
    'GradMultiply',
    'LearnedPositionalEmbedding',
    'LinearizedConvolution',
    'MultiheadAttention',
    'ScalarBias',
    'SinusoidalPositionalEmbedding',
]
+272 −0

File added.

Preview size limit exceeded, changes collapsed.

Loading