Spaces:
Running
Running
File size: 1,651 Bytes
17f1f54 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 | # coding: utf-8
import torch.nn as nn
from torch import Tensor
from helpers import freeze_params
from transformer_layers import TransformerEncoderLayer, PositionalEncoding
class Encoder(nn.Module):
def __init__(self,
hidden_size: int = 512,
ff_size: int = 2048,
num_layers: int = 2,
num_heads: int = 4,
dropout: float = 0.1,
emb_dropout: float = 0.1,
freeze: bool = False,
**kwargs):
super(Encoder, self).__init__()
self.layers = nn.ModuleList([
TransformerEncoderLayer(size=hidden_size, ff_size=ff_size,
num_heads=num_heads, dropout=dropout)
for _ in range(num_layers)])
self.layer_norm = nn.LayerNorm(hidden_size, eps=1e-6)
self.pe = PositionalEncoding(hidden_size)
self.emb_dropout = nn.Dropout(p=emb_dropout)
self._output_size = hidden_size
if freeze:
freeze_params(self)
def forward(self,
embed_src: Tensor,
src_length: Tensor,
mask: Tensor):
x = embed_src
# Add position encoding to word embeddings
x = self.pe(x)
# Add Dropout
x = self.emb_dropout(x)
# Apply each layer to the input
for layer in self.layers:
x = layer(x, mask)
return self.layer_norm(x)
def __repr__(self):
return "%s(num_layers=%r, num_heads=%r)" % (
self.__class__.__name__, len(self.layers),
self.layers[0].src_src_att.num_heads) |