Initial commit
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Common ROSVOT layers and utilities."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Conformer layers for ROSVOT."""
|
||||
@@ -0,0 +1,96 @@
|
||||
from torch import nn
|
||||
from .espnet_positional_embedding import RelPositionalEncoding, ScaledPositionalEncoding, PositionalEncoding
|
||||
from .espnet_transformer_attn import RelPositionMultiHeadedAttention, MultiHeadedAttention
|
||||
from .layers import Swish, ConvolutionModule, EncoderLayer, MultiLayeredConv1d
|
||||
from ..layers import Embedding
|
||||
|
||||
|
||||
class ConformerLayers(nn.Module):
|
||||
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
|
||||
use_last_norm=True, save_hidden=False):
|
||||
super().__init__()
|
||||
self.use_last_norm = use_last_norm
|
||||
self.layers = nn.ModuleList()
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
|
||||
self.pos_embed = RelPositionalEncoding(hidden_size, dropout)
|
||||
self.encoder_layers = nn.ModuleList([EncoderLayer(
|
||||
hidden_size,
|
||||
RelPositionMultiHeadedAttention(num_heads, hidden_size, 0.0),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
ConvolutionModule(hidden_size, kernel_size, Swish()),
|
||||
dropout,
|
||||
) for _ in range(num_layers)])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(hidden_size)
|
||||
else:
|
||||
self.layer_norm = nn.Linear(hidden_size, hidden_size)
|
||||
self.save_hidden = save_hidden
|
||||
if save_hidden:
|
||||
self.hiddens = []
|
||||
|
||||
def forward(self, x, padding_mask=None):
|
||||
"""
|
||||
|
||||
:param x: [B, T, H]
|
||||
:param padding_mask: [B, T]
|
||||
:return: [B, T, H]
|
||||
"""
|
||||
self.hiddens = []
|
||||
nonpadding_mask = x.abs().sum(-1) > 0
|
||||
x = self.pos_embed(x)
|
||||
for l in self.encoder_layers:
|
||||
x, mask = l(x, nonpadding_mask[:, None, :])
|
||||
if self.save_hidden:
|
||||
self.hiddens.append(x[0])
|
||||
x = x[0]
|
||||
x = self.layer_norm(x) * nonpadding_mask.float()[:, :, None]
|
||||
return x
|
||||
|
||||
class FastConformerLayers(ConformerLayers):
|
||||
def __init__(self, hidden_size, num_layers, kernel_size=9, dropout=0.0, num_heads=4,
|
||||
use_last_norm=True, save_hidden=False):
|
||||
super(ConformerLayers, self).__init__()
|
||||
self.use_last_norm = use_last_norm
|
||||
self.layers = nn.ModuleList()
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (hidden_size, hidden_size * 4, 1, dropout)
|
||||
self.pos_embed = PositionalEncoding(hidden_size, dropout)
|
||||
self.encoder_layers = nn.ModuleList([EncoderLayer(
|
||||
hidden_size,
|
||||
MultiHeadedAttention(num_heads, hidden_size, 0.0, flash=True),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
ConvolutionModule(hidden_size, kernel_size, Swish()),
|
||||
dropout,
|
||||
) for _ in range(num_layers)])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(hidden_size)
|
||||
else:
|
||||
self.layer_norm = nn.Linear(hidden_size, hidden_size)
|
||||
self.save_hidden = save_hidden
|
||||
if save_hidden:
|
||||
self.hiddens = []
|
||||
|
||||
class ConformerEncoder(ConformerLayers):
|
||||
def __init__(self, hidden_size, dict_size, num_layers=None):
|
||||
conformer_enc_kernel_size = 9
|
||||
super().__init__(hidden_size, num_layers, conformer_enc_kernel_size)
|
||||
self.embed = Embedding(dict_size, hidden_size, padding_idx=0)
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
|
||||
:param src_tokens: [B, T]
|
||||
:return: [B x T x C]
|
||||
"""
|
||||
x = self.embed(x) # [B, T, H]
|
||||
x = super(ConformerEncoder, self).forward(x)
|
||||
return x
|
||||
|
||||
|
||||
class ConformerDecoder(ConformerLayers):
|
||||
def __init__(self, hidden_size, num_layers):
|
||||
conformer_dec_kernel_size = 9
|
||||
super().__init__(hidden_size, num_layers, conformer_dec_kernel_size)
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
class PositionalEncoding(torch.nn.Module):
|
||||
"""Positional encoding.
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
reverse (bool): Whether to reverse the input position.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000, reverse=False):
|
||||
"""Construct an PositionalEncoding object."""
|
||||
super(PositionalEncoding, self).__init__()
|
||||
self.d_model = d_model
|
||||
self.reverse = reverse
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.pe = None
|
||||
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
|
||||
|
||||
def extend_pe(self, x):
|
||||
"""Reset the positional encodings."""
|
||||
if self.pe is not None:
|
||||
if self.pe.size(1) >= x.size(1):
|
||||
if self.pe.dtype != x.dtype or self.pe.device != x.device:
|
||||
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
|
||||
return
|
||||
pe = torch.zeros(x.size(1), self.d_model)
|
||||
if self.reverse:
|
||||
position = torch.arange(
|
||||
x.size(1) - 1, -1, -1.0, dtype=torch.float32
|
||||
).unsqueeze(1)
|
||||
else:
|
||||
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.d_model, 2, dtype=torch.float32)
|
||||
* -(math.log(10000.0) / self.d_model)
|
||||
)
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
pe = pe.unsqueeze(0)
|
||||
self.pe = pe.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
"""Add positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class ScaledPositionalEncoding(PositionalEncoding):
|
||||
"""Scaled positional encoding module.
|
||||
See Sec. 3.2 https://arxiv.org/abs/1809.08895
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(d_model=d_model, dropout_rate=dropout_rate, max_len=max_len)
|
||||
self.alpha = torch.nn.Parameter(torch.tensor(1.0))
|
||||
|
||||
def reset_parameters(self):
|
||||
"""Reset parameters."""
|
||||
self.alpha.data = torch.tensor(1.0)
|
||||
|
||||
def forward(self, x):
|
||||
"""Add positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x + self.alpha * self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class RelPositionalEncoding(PositionalEncoding):
|
||||
"""Relative positional encoding module.
|
||||
See : Appendix B in https://arxiv.org/abs/1901.02860
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(d_model, dropout_rate, max_len, reverse=True)
|
||||
|
||||
def forward(self, x):
|
||||
"""Compute positional encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
torch.Tensor: Positional embedding tensor (1, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale
|
||||
pos_emb = self.pe[:, : x.size(1)]
|
||||
return self.dropout(x), self.dropout(pos_emb)
|
||||
+198
@@ -0,0 +1,198 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Multi-Head Attention layer definition."""
|
||||
|
||||
from packaging import version
|
||||
import math
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class MultiHeadedAttention(nn.Module):
|
||||
"""Multi-Head Attention layer.
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate, flash=False):
|
||||
"""Construct an MultiHeadedAttention object."""
|
||||
super(MultiHeadedAttention, self).__init__()
|
||||
assert n_feat % n_head == 0
|
||||
# We assume d_v always equals d_k
|
||||
self.d_k = n_feat // n_head
|
||||
self.h = n_head
|
||||
self.linear_q = nn.Linear(n_feat, n_feat)
|
||||
self.linear_k = nn.Linear(n_feat, n_feat)
|
||||
self.linear_v = nn.Linear(n_feat, n_feat)
|
||||
self.linear_out = nn.Linear(n_feat, n_feat)
|
||||
self.attn = None
|
||||
self.dropout = nn.Dropout(p=dropout_rate)
|
||||
self.dropout_rate = dropout_rate
|
||||
self.flash = flash
|
||||
|
||||
def forward_qkv(self, query, key, value):
|
||||
"""Transform query, key and value.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
Returns:
|
||||
torch.Tensor: Transformed query tensor (#batch, n_head, time1, d_k).
|
||||
torch.Tensor: Transformed key tensor (#batch, n_head, time2, d_k).
|
||||
torch.Tensor: Transformed value tensor (#batch, n_head, time2, d_k).
|
||||
"""
|
||||
n_batch = query.size(0)
|
||||
q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)
|
||||
k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)
|
||||
v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)
|
||||
q = q.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
k = k.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
v = v.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def forward_attention(self, value, scores, mask):
|
||||
"""Compute attention context vector.
|
||||
Args:
|
||||
value (torch.Tensor): Transformed value (#batch, n_head, time2, d_k).
|
||||
scores (torch.Tensor): Attention score (#batch, n_head, time1, time2).
|
||||
mask (torch.Tensor): Mask (#batch, 1, time2) or (#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Transformed value (#batch, time1, d_model)
|
||||
weighted by the attention score (#batch, time1, time2).
|
||||
"""
|
||||
n_batch = value.size(0)
|
||||
if mask is not None:
|
||||
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
|
||||
min_value = float(
|
||||
numpy.finfo(torch.tensor(0, dtype=scores.dtype).numpy().dtype).min
|
||||
)
|
||||
scores = scores.masked_fill(mask, min_value)
|
||||
self.attn = torch.softmax(scores, dim=-1).masked_fill(
|
||||
mask, 0.0
|
||||
) # (batch, head, time1, time2)
|
||||
else:
|
||||
self.attn = torch.softmax(scores, dim=-1) # (batch, head, time1, time2)
|
||||
|
||||
p_attn = self.dropout(self.attn)
|
||||
x = torch.matmul(p_attn, value) # (batch, head, time1, d_k)
|
||||
x = (
|
||||
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
|
||||
) # (batch, time1, d_model)
|
||||
|
||||
return self.linear_out(x) # (batch, time1, d_model)
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Compute scaled dot product attention.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
if version.parse(torch.__version__) >= version.parse("2.0") and self.flash:
|
||||
n_batch = value.size(0)
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=mask.unsqueeze(1) if mask is not None else None, dropout_p=self.dropout_rate)
|
||||
x = (
|
||||
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
|
||||
) # (batch, time1, d_model)
|
||||
return self.linear_out(x)
|
||||
else:
|
||||
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
|
||||
class RelPositionMultiHeadedAttention(MultiHeadedAttention):
|
||||
"""Multi-Head Attention layer with relative position encoding.
|
||||
Paper: https://arxiv.org/abs/1901.02860
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate):
|
||||
"""Construct an RelPositionMultiHeadedAttention object."""
|
||||
super().__init__(n_head, n_feat, dropout_rate)
|
||||
# linear transformation for positional ecoding
|
||||
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
|
||||
# these two learnable bias are used in matrix c and matrix d
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_u)
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_v)
|
||||
|
||||
def rel_shift(self, x, zero_triu=False):
|
||||
"""Compute relative positinal encoding.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, size).
|
||||
zero_triu (bool): If true, return the lower triangular part of the matrix.
|
||||
Returns:
|
||||
torch.Tensor: Output tensor.
|
||||
"""
|
||||
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
|
||||
x_padded = torch.cat([zero_pad, x], dim=-1)
|
||||
|
||||
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
|
||||
x = x_padded[:, :, 1:].view_as(x)
|
||||
|
||||
if zero_triu:
|
||||
ones = torch.ones((x.size(2), x.size(3)))
|
||||
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
|
||||
|
||||
return x
|
||||
|
||||
def forward(self, query, key, value, pos_emb, mask):
|
||||
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
pos_emb (torch.Tensor): Positional embedding tensor (#batch, time2, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
q = q.transpose(1, 2) # (batch, time1, head, d_k)
|
||||
|
||||
n_batch_pos = pos_emb.size(0)
|
||||
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
|
||||
p = p.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
# compute attention score
|
||||
# first compute matrix a and matrix c
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
# (batch, head, time1, time2)
|
||||
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
|
||||
|
||||
# compute matrix b and matrix d
|
||||
# (batch, head, time1, time2)
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
|
||||
matrix_bd = self.rel_shift(matrix_bd)
|
||||
|
||||
scores = (matrix_ac + matrix_bd) / math.sqrt(
|
||||
self.d_k
|
||||
) # (batch, head, time1, time2)
|
||||
|
||||
return self.forward_attention(v, scores, mask)
|
||||
@@ -0,0 +1,260 @@
|
||||
from torch import nn
|
||||
import torch
|
||||
|
||||
from ..layers import LayerNorm
|
||||
|
||||
|
||||
class ConvolutionModule(nn.Module):
|
||||
"""ConvolutionModule in Conformer model.
|
||||
Args:
|
||||
channels (int): The number of channels of conv layers.
|
||||
kernel_size (int): Kernerl size of conv layers.
|
||||
"""
|
||||
|
||||
def __init__(self, channels, kernel_size, activation=nn.ReLU(), bias=True):
|
||||
"""Construct an ConvolutionModule object."""
|
||||
super(ConvolutionModule, self).__init__()
|
||||
# kernerl_size should be a odd number for 'SAME' padding
|
||||
assert (kernel_size - 1) % 2 == 0
|
||||
|
||||
self.pointwise_conv1 = nn.Conv1d(
|
||||
channels,
|
||||
2 * channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=bias,
|
||||
)
|
||||
self.depthwise_conv = nn.Conv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
groups=channels,
|
||||
bias=bias,
|
||||
)
|
||||
self.norm = nn.BatchNorm1d(channels)
|
||||
self.pointwise_conv2 = nn.Conv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=bias,
|
||||
)
|
||||
self.activation = activation
|
||||
|
||||
def forward(self, x):
|
||||
"""Compute convolution module.
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, channels).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time, channels).
|
||||
"""
|
||||
# exchange the temporal dimension and the feature dimension
|
||||
x = x.transpose(1, 2)
|
||||
|
||||
# GLU mechanism
|
||||
x = self.pointwise_conv1(x) # (batch, 2*channel, dim)
|
||||
x = nn.functional.glu(x, dim=1) # (batch, channel, dim)
|
||||
|
||||
# 1D Depthwise Conv
|
||||
x = self.depthwise_conv(x)
|
||||
x = self.activation(self.norm(x))
|
||||
|
||||
x = self.pointwise_conv2(x)
|
||||
|
||||
return x.transpose(1, 2)
|
||||
|
||||
|
||||
class MultiLayeredConv1d(torch.nn.Module):
|
||||
"""Multi-layered conv1d for Transformer block.
|
||||
This is a module of multi-leyered conv1d designed
|
||||
to replace positionwise feed-forward network
|
||||
in Transforner block, which is introduced in
|
||||
`FastSpeech: Fast, Robust and Controllable Text to Speech`_.
|
||||
.. _`FastSpeech: Fast, Robust and Controllable Text to Speech`:
|
||||
https://arxiv.org/pdf/1905.09263.pdf
|
||||
"""
|
||||
|
||||
def __init__(self, in_chans, hidden_chans, kernel_size, dropout_rate):
|
||||
"""Initialize MultiLayeredConv1d module.
|
||||
Args:
|
||||
in_chans (int): Number of input channels.
|
||||
hidden_chans (int): Number of hidden channels.
|
||||
kernel_size (int): Kernel size of conv1d.
|
||||
dropout_rate (float): Dropout rate.
|
||||
"""
|
||||
super(MultiLayeredConv1d, self).__init__()
|
||||
self.w_1 = torch.nn.Conv1d(
|
||||
in_chans,
|
||||
hidden_chans,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
)
|
||||
self.w_2 = torch.nn.Conv1d(
|
||||
hidden_chans,
|
||||
in_chans,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
)
|
||||
self.dropout = torch.nn.Dropout(dropout_rate)
|
||||
|
||||
def forward(self, x):
|
||||
"""Calculate forward propagation.
|
||||
Args:
|
||||
x (torch.Tensor): Batch of input tensors (B, T, in_chans).
|
||||
Returns:
|
||||
torch.Tensor: Batch of output tensors (B, T, hidden_chans).
|
||||
"""
|
||||
x = torch.relu(self.w_1(x.transpose(-1, 1))).transpose(-1, 1)
|
||||
return self.w_2(self.dropout(x).transpose(-1, 1)).transpose(-1, 1)
|
||||
|
||||
|
||||
class Swish(torch.nn.Module):
|
||||
"""Construct an Swish object."""
|
||||
|
||||
def forward(self, x):
|
||||
"""Return Swich activation function."""
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
class EncoderLayer(nn.Module):
|
||||
"""Encoder layer module.
|
||||
Args:
|
||||
size (int): Input dimension.
|
||||
self_attn (torch.nn.Module): Self-attention module instance.
|
||||
`MultiHeadedAttention` or `RelPositionMultiHeadedAttention` instance
|
||||
can be used as the argument.
|
||||
feed_forward (torch.nn.Module): Feed-forward module instance.
|
||||
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
|
||||
can be used as the argument.
|
||||
feed_forward_macaron (torch.nn.Module): Additional feed-forward module instance.
|
||||
`PositionwiseFeedForward`, `MultiLayeredConv1d`, or `Conv1dLinear` instance
|
||||
can be used as the argument.
|
||||
conv_module (torch.nn.Module): Convolution module instance.
|
||||
`ConvlutionModule` instance can be used as the argument.
|
||||
dropout_rate (float): Dropout rate.
|
||||
normalize_before (bool): Whether to use layer_norm before the first block.
|
||||
concat_after (bool): Whether to concat attention layer's input and output.
|
||||
if True, additional linear will be applied.
|
||||
i.e. x -> x + linear(concat(x, att(x)))
|
||||
if False, no additional linear will be applied. i.e. x -> x + att(x)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
self_attn,
|
||||
feed_forward,
|
||||
feed_forward_macaron,
|
||||
conv_module,
|
||||
dropout_rate,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
):
|
||||
"""Construct an EncoderLayer object."""
|
||||
super(EncoderLayer, self).__init__()
|
||||
self.self_attn = self_attn
|
||||
self.feed_forward = feed_forward
|
||||
self.feed_forward_macaron = feed_forward_macaron
|
||||
self.conv_module = conv_module
|
||||
self.norm_ff = LayerNorm(size) # for the FNN module
|
||||
self.norm_mha = LayerNorm(size) # for the MHA module
|
||||
if feed_forward_macaron is not None:
|
||||
self.norm_ff_macaron = LayerNorm(size)
|
||||
self.ff_scale = 0.5
|
||||
else:
|
||||
self.ff_scale = 1.0
|
||||
if self.conv_module is not None:
|
||||
self.norm_conv = LayerNorm(size) # for the CNN module
|
||||
self.norm_final = LayerNorm(size) # for the final output of the block
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.size = size
|
||||
self.normalize_before = normalize_before
|
||||
self.concat_after = concat_after
|
||||
if self.concat_after:
|
||||
self.concat_linear = nn.Linear(size + size, size)
|
||||
|
||||
def forward(self, x_input, mask, cache=None):
|
||||
"""Compute encoded features.
|
||||
Args:
|
||||
x_input (Union[Tuple, torch.Tensor]): Input tensor w/ or w/o pos emb.
|
||||
- w/ pos emb: Tuple of tensors [(#batch, time, size), (1, time, size)].
|
||||
- w/o pos emb: Tensor (#batch, time, size).
|
||||
mask (torch.Tensor): Mask tensor for the input (#batch, time).
|
||||
cache (torch.Tensor): Cache tensor of the input (#batch, time - 1, size).
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time, size).
|
||||
torch.Tensor: Mask tensor (#batch, time).
|
||||
"""
|
||||
if isinstance(x_input, tuple):
|
||||
x, pos_emb = x_input[0], x_input[1]
|
||||
else:
|
||||
x, pos_emb = x_input, None
|
||||
|
||||
# whether to use macaron style
|
||||
if self.feed_forward_macaron is not None:
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_ff_macaron(x)
|
||||
x = residual + self.ff_scale * self.dropout(self.feed_forward_macaron(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_ff_macaron(x)
|
||||
|
||||
# multi-headed self-attention module
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_mha(x)
|
||||
|
||||
if cache is None:
|
||||
x_q = x
|
||||
else:
|
||||
assert cache.shape == (x.shape[0], x.shape[1] - 1, self.size)
|
||||
x_q = x[:, -1:, :]
|
||||
residual = residual[:, -1:, :]
|
||||
mask = None if mask is None else mask[:, -1:, :]
|
||||
|
||||
if pos_emb is not None:
|
||||
x_att = self.self_attn(x_q, x, x, pos_emb, mask)
|
||||
else:
|
||||
x_att = self.self_attn(x_q, x, x, mask)
|
||||
|
||||
if self.concat_after:
|
||||
x_concat = torch.cat((x, x_att), dim=-1)
|
||||
x = residual + self.concat_linear(x_concat)
|
||||
else:
|
||||
x = residual + self.dropout(x_att)
|
||||
if not self.normalize_before:
|
||||
x = self.norm_mha(x)
|
||||
|
||||
# convolution module
|
||||
if self.conv_module is not None:
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_conv(x)
|
||||
x = residual + self.dropout(self.conv_module(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_conv(x)
|
||||
|
||||
# feed forward module
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm_ff(x)
|
||||
x = residual + self.ff_scale * self.dropout(self.feed_forward(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm_ff(x)
|
||||
|
||||
if self.conv_module is not None:
|
||||
x = self.norm_final(x)
|
||||
|
||||
if cache is not None:
|
||||
x = torch.cat([cache, x], dim=1)
|
||||
|
||||
if pos_emb is not None:
|
||||
return (x, pos_emb), mask
|
||||
|
||||
return x, mask
|
||||
@@ -0,0 +1,175 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .layers import LayerNorm, Embedding
|
||||
|
||||
class LambdaLayer(nn.Module):
|
||||
def __init__(self, lambd):
|
||||
super(LambdaLayer, self).__init__()
|
||||
self.lambd = lambd
|
||||
|
||||
def forward(self, x):
|
||||
return self.lambd(x)
|
||||
|
||||
def init_weights_func(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv1d") != -1:
|
||||
torch.nn.init.xavier_uniform_(m.weight)
|
||||
|
||||
def get_norm_builder(norm_type, channels, ln_eps=1e-6):
|
||||
if norm_type == 'bn':
|
||||
norm_builder = lambda: nn.BatchNorm1d(channels)
|
||||
elif norm_type == 'in':
|
||||
norm_builder = lambda: nn.InstanceNorm1d(channels, affine=True)
|
||||
elif norm_type == 'gn':
|
||||
norm_builder = lambda: nn.GroupNorm(8, channels)
|
||||
elif norm_type == 'ln':
|
||||
norm_builder = lambda: LayerNorm(channels, dim=1, eps=ln_eps)
|
||||
else:
|
||||
norm_builder = lambda: nn.Identity()
|
||||
return norm_builder
|
||||
|
||||
def get_act_builder(act_type):
|
||||
if act_type == 'gelu':
|
||||
act_builder = lambda: nn.GELU()
|
||||
elif act_type == 'relu':
|
||||
act_builder = lambda: nn.ReLU(inplace=True)
|
||||
elif act_type == 'leakyrelu':
|
||||
act_builder = lambda: nn.LeakyReLU(negative_slope=0.01, inplace=True)
|
||||
elif act_type == 'swish':
|
||||
act_builder = lambda: nn.SiLU(inplace=True)
|
||||
else:
|
||||
act_builder = lambda: nn.Identity()
|
||||
return act_builder
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
"""Implements conv->PReLU->norm n-times"""
|
||||
|
||||
def __init__(self, channels, kernel_size, dilation, n=2, norm_type='bn', dropout=0.0,
|
||||
c_multiple=2, ln_eps=1e-12, act_type='gelu'):
|
||||
super(ResidualBlock, self).__init__()
|
||||
|
||||
norm_builder = get_norm_builder(norm_type, channels, ln_eps)
|
||||
act_builder = get_act_builder(act_type)
|
||||
|
||||
self.blocks = [
|
||||
nn.Sequential(
|
||||
norm_builder(),
|
||||
nn.Conv1d(channels, c_multiple * channels, kernel_size, dilation=dilation,
|
||||
padding=(dilation * (kernel_size - 1)) // 2),
|
||||
LambdaLayer(lambda x: x * kernel_size ** -0.5),
|
||||
act_builder(),
|
||||
nn.Conv1d(c_multiple * channels, channels, 1, dilation=dilation),
|
||||
)
|
||||
for i in range(n)
|
||||
]
|
||||
|
||||
self.blocks = nn.ModuleList(self.blocks)
|
||||
self.dropout = dropout
|
||||
|
||||
def forward(self, x):
|
||||
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
|
||||
for b in self.blocks:
|
||||
x_ = b(x)
|
||||
if self.dropout > 0 and self.training:
|
||||
x_ = F.dropout(x_, self.dropout, training=self.training)
|
||||
x = x + x_
|
||||
x = x * nonpadding
|
||||
return x
|
||||
|
||||
|
||||
class ConvBlocks(nn.Module):
|
||||
"""Decodes the expanded phoneme encoding into spectrograms"""
|
||||
|
||||
def __init__(self, hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5,
|
||||
init_weights=True, is_BTC=True, num_layers=None, post_net_kernel=3, act_type='gelu'):
|
||||
super(ConvBlocks, self).__init__()
|
||||
self.is_BTC = is_BTC
|
||||
if num_layers is not None:
|
||||
dilations = [1] * num_layers
|
||||
self.res_blocks = nn.Sequential(
|
||||
*[ResidualBlock(hidden_size, kernel_size, d,
|
||||
n=layers_in_block, norm_type=norm_type, c_multiple=c_multiple,
|
||||
dropout=dropout, ln_eps=ln_eps, act_type=act_type)
|
||||
for d in dilations],
|
||||
)
|
||||
norm = get_norm_builder(norm_type, hidden_size, ln_eps)()
|
||||
self.last_norm = norm
|
||||
self.post_net1 = nn.Conv1d(hidden_size, out_dims, kernel_size=post_net_kernel,
|
||||
padding=post_net_kernel // 2)
|
||||
if init_weights:
|
||||
self.apply(init_weights_func)
|
||||
|
||||
def forward(self, x, nonpadding=None):
|
||||
"""
|
||||
|
||||
:param x: [B, T, H]
|
||||
:return: [B, T, H]
|
||||
"""
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
if nonpadding is None:
|
||||
nonpadding = (x.abs().sum(1) > 0).float()[:, None, :]
|
||||
elif self.is_BTC:
|
||||
nonpadding = nonpadding.transpose(1, 2)
|
||||
x = self.res_blocks(x) * nonpadding
|
||||
x = self.last_norm(x) * nonpadding
|
||||
x = self.post_net1(x) * nonpadding
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
|
||||
|
||||
class TextConvEncoder(ConvBlocks):
|
||||
def __init__(self, dict_size, hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5, init_weights=True, num_layers=None, post_net_kernel=3):
|
||||
super().__init__(hidden_size, out_dims, dilations, kernel_size,
|
||||
norm_type, layers_in_block, c_multiple,
|
||||
dropout, ln_eps, init_weights, num_layers=num_layers,
|
||||
post_net_kernel=post_net_kernel)
|
||||
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
|
||||
self.embed_scale = math.sqrt(hidden_size)
|
||||
|
||||
def forward(self, txt_tokens):
|
||||
"""
|
||||
|
||||
:param txt_tokens: [B, T]
|
||||
:return: {
|
||||
'encoder_out': [B x T x C]
|
||||
}
|
||||
"""
|
||||
x = self.embed_scale * self.embed_tokens(txt_tokens)
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class ConditionalConvBlocks(ConvBlocks):
|
||||
def __init__(self, hidden_size, c_cond, c_out, dilations, kernel_size,
|
||||
norm_type='ln', layers_in_block=2, c_multiple=2,
|
||||
dropout=0.0, ln_eps=1e-5, init_weights=True, is_BTC=True, num_layers=None):
|
||||
super().__init__(hidden_size, c_out, dilations, kernel_size,
|
||||
norm_type, layers_in_block, c_multiple,
|
||||
dropout, ln_eps, init_weights, is_BTC=False, num_layers=num_layers)
|
||||
self.g_prenet = nn.Conv1d(c_cond, hidden_size, 3, padding=1)
|
||||
self.is_BTC_ = is_BTC
|
||||
if init_weights:
|
||||
self.g_prenet.apply(init_weights_func)
|
||||
|
||||
def forward(self, x, cond, nonpadding=None):
|
||||
if self.is_BTC_:
|
||||
x = x.transpose(1, 2)
|
||||
cond = cond.transpose(1, 2)
|
||||
if nonpadding is not None:
|
||||
nonpadding = nonpadding.transpose(1, 2)
|
||||
if nonpadding is None:
|
||||
nonpadding = x.abs().sum(1)[:, None]
|
||||
x = x + self.g_prenet(cond)
|
||||
x = x * nonpadding
|
||||
x = super(ConditionalConvBlocks, self).forward(x) # input needs to be BTC
|
||||
if self.is_BTC_:
|
||||
x = x.transpose(1, 2)
|
||||
return x
|
||||
@@ -0,0 +1,85 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
|
||||
class LayerNorm(torch.nn.LayerNorm):
|
||||
"""Layer normalization module.
|
||||
:param int nout: output dim size
|
||||
:param int dim: dimension to be normalized
|
||||
"""
|
||||
|
||||
def __init__(self, nout, dim=-1, eps=1e-5):
|
||||
"""Construct an LayerNorm object."""
|
||||
super(LayerNorm, self).__init__(nout, eps=eps)
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
"""Apply layer normalization.
|
||||
:param torch.Tensor x: input tensor
|
||||
:return: layer normalized tensor
|
||||
:rtype torch.Tensor
|
||||
"""
|
||||
if self.dim == -1:
|
||||
return super(LayerNorm, self).forward(x)
|
||||
return super(LayerNorm, self).forward(x.transpose(1, -1)).transpose(1, -1)
|
||||
|
||||
|
||||
class Reshape(nn.Module):
|
||||
def __init__(self, *args):
|
||||
super(Reshape, self).__init__()
|
||||
self.shape = args
|
||||
|
||||
def forward(self, x):
|
||||
return x.view(self.shape)
|
||||
|
||||
|
||||
class Permute(nn.Module):
|
||||
def __init__(self, *args):
|
||||
super(Permute, self).__init__()
|
||||
self.args = args
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(self.args)
|
||||
|
||||
|
||||
def Linear(in_features, out_features, bias=True, init_type='xavier'):
|
||||
m = nn.Linear(in_features, out_features, bias)
|
||||
if init_type == 'xavier':
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
elif init_type == 'kaiming':
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_in')
|
||||
if bias:
|
||||
nn.init.constant_(m.bias, 0.)
|
||||
return m
|
||||
|
||||
|
||||
def Embedding(num_embeddings, embedding_dim, padding_idx=None, init_type='normal'):
|
||||
m = nn.Embedding(num_embeddings, embedding_dim, padding_idx=padding_idx)
|
||||
if init_type == 'normal':
|
||||
nn.init.normal_(m.weight, mean=0, std=embedding_dim ** -0.5)
|
||||
elif init_type == 'kaiming':
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_in')
|
||||
if padding_idx is not None:
|
||||
nn.init.constant_(m.weight[padding_idx], 0)
|
||||
return m
|
||||
|
||||
|
||||
class GradientReverseFunction(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, input, coeff=1.):
|
||||
ctx.coeff = coeff
|
||||
output = input * 1.0
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output.neg() * ctx.coeff, None
|
||||
|
||||
|
||||
class GRL(nn.Module):
|
||||
def __init__(self):
|
||||
super(GRL, self).__init__()
|
||||
|
||||
def forward(self, *input):
|
||||
return GradientReverseFunction.apply(*input)
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .layers import Embedding
|
||||
|
||||
|
||||
def convert_pad_shape(pad_shape):
|
||||
l = pad_shape[::-1]
|
||||
pad_shape = [item for sublist in l for item in sublist]
|
||||
return pad_shape
|
||||
|
||||
|
||||
def shift_1d(x):
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
|
||||
return x
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, hidden_channels, filter_channels, n_heads, n_layers, kernel_size=1, p_dropout=0.,
|
||||
window_size=None, block_length=None, pre_ln=False, **kwargs):
|
||||
super().__init__()
|
||||
self.hidden_channels = hidden_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.n_heads = n_heads
|
||||
self.n_layers = n_layers
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.window_size = window_size
|
||||
self.block_length = block_length
|
||||
self.pre_ln = pre_ln
|
||||
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
self.attn_layers = nn.ModuleList()
|
||||
self.norm_layers_1 = nn.ModuleList()
|
||||
self.ffn_layers = nn.ModuleList()
|
||||
self.norm_layers_2 = nn.ModuleList()
|
||||
for i in range(self.n_layers):
|
||||
self.attn_layers.append(
|
||||
MultiHeadAttention(hidden_channels, hidden_channels, n_heads, window_size=window_size,
|
||||
p_dropout=p_dropout, block_length=block_length))
|
||||
self.norm_layers_1.append(LayerNorm(hidden_channels))
|
||||
self.ffn_layers.append(
|
||||
FFN(hidden_channels, hidden_channels, filter_channels, kernel_size, p_dropout=p_dropout))
|
||||
self.norm_layers_2.append(LayerNorm(hidden_channels))
|
||||
if pre_ln:
|
||||
self.last_ln = LayerNorm(hidden_channels)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
|
||||
for i in range(self.n_layers):
|
||||
x = x * x_mask
|
||||
x_ = x
|
||||
if self.pre_ln:
|
||||
x = self.norm_layers_1[i](x)
|
||||
y = self.attn_layers[i](x, x, attn_mask)
|
||||
y = self.drop(y)
|
||||
x = x_ + y
|
||||
if not self.pre_ln:
|
||||
x = self.norm_layers_1[i](x)
|
||||
|
||||
x_ = x
|
||||
if self.pre_ln:
|
||||
x = self.norm_layers_2[i](x)
|
||||
y = self.ffn_layers[i](x, x_mask)
|
||||
y = self.drop(y)
|
||||
x = x_ + y
|
||||
if not self.pre_ln:
|
||||
x = self.norm_layers_2[i](x)
|
||||
if self.pre_ln:
|
||||
x = self.last_ln(x)
|
||||
x = x * x_mask
|
||||
return x
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self, channels, out_channels, n_heads, window_size=None, heads_share=True, p_dropout=0.,
|
||||
block_length=None, proximal_bias=False, proximal_init=False):
|
||||
super().__init__()
|
||||
assert channels % n_heads == 0
|
||||
|
||||
self.channels = channels
|
||||
self.out_channels = out_channels
|
||||
self.n_heads = n_heads
|
||||
self.window_size = window_size
|
||||
self.heads_share = heads_share
|
||||
self.block_length = block_length
|
||||
self.proximal_bias = proximal_bias
|
||||
self.p_dropout = p_dropout
|
||||
self.attn = None
|
||||
|
||||
self.k_channels = channels // n_heads
|
||||
self.conv_q = nn.Conv1d(channels, channels, 1)
|
||||
self.conv_k = nn.Conv1d(channels, channels, 1)
|
||||
self.conv_v = nn.Conv1d(channels, channels, 1)
|
||||
if window_size is not None:
|
||||
n_heads_rel = 1 if heads_share else n_heads
|
||||
rel_stddev = self.k_channels ** -0.5
|
||||
self.emb_rel_k = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
|
||||
self.emb_rel_v = nn.Parameter(torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels) * rel_stddev)
|
||||
self.conv_o = nn.Conv1d(channels, out_channels, 1)
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
nn.init.xavier_uniform_(self.conv_q.weight)
|
||||
nn.init.xavier_uniform_(self.conv_k.weight)
|
||||
if proximal_init:
|
||||
self.conv_k.weight.data.copy_(self.conv_q.weight.data)
|
||||
self.conv_k.bias.data.copy_(self.conv_q.bias.data)
|
||||
nn.init.xavier_uniform_(self.conv_v.weight)
|
||||
|
||||
def forward(self, x, c, attn_mask=None):
|
||||
q = self.conv_q(x)
|
||||
k = self.conv_k(c)
|
||||
v = self.conv_v(c)
|
||||
|
||||
x, self.attn = self.attention(q, k, v, mask=attn_mask)
|
||||
|
||||
x = self.conv_o(x)
|
||||
return x
|
||||
|
||||
def attention(self, query, key, value, mask=None):
|
||||
# reshape [b, d, t] -> [b, n_h, t, d_k]
|
||||
b, d, t_s, t_t = (*key.size(), query.size(2))
|
||||
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
|
||||
key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
|
||||
value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
|
||||
|
||||
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.k_channels)
|
||||
if self.window_size is not None:
|
||||
assert t_s == t_t, "Relative attention is only available for self-attention."
|
||||
key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
|
||||
rel_logits = self._matmul_with_relative_keys(query, key_relative_embeddings)
|
||||
rel_logits = self._relative_position_to_absolute_position(rel_logits)
|
||||
scores_local = rel_logits / math.sqrt(self.k_channels)
|
||||
scores = scores + scores_local
|
||||
if self.proximal_bias:
|
||||
assert t_s == t_t, "Proximal bias is only available for self-attention."
|
||||
scores = scores + self._attention_bias_proximal(t_s).to(device=scores.device, dtype=scores.dtype)
|
||||
if mask is not None:
|
||||
scores = scores.masked_fill(mask == 0, -1e4)
|
||||
if self.block_length is not None:
|
||||
block_mask = torch.ones_like(scores).triu(-self.block_length).tril(self.block_length)
|
||||
scores = scores * block_mask + -1e4 * (1 - block_mask)
|
||||
p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
|
||||
p_attn = self.drop(p_attn)
|
||||
output = torch.matmul(p_attn, value)
|
||||
if self.window_size is not None:
|
||||
relative_weights = self._absolute_position_to_relative_position(p_attn)
|
||||
value_relative_embeddings = self._get_relative_embeddings(self.emb_rel_v, t_s)
|
||||
output = output + self._matmul_with_relative_values(relative_weights, value_relative_embeddings)
|
||||
output = output.transpose(2, 3).contiguous().view(b, d, t_t) # [b, n_h, t_t, d_k] -> [b, d, t_t]
|
||||
return output, p_attn
|
||||
|
||||
def _matmul_with_relative_values(self, x, y):
|
||||
"""
|
||||
x: [b, h, l, m]
|
||||
y: [h or 1, m, d]
|
||||
ret: [b, h, l, d]
|
||||
"""
|
||||
ret = torch.matmul(x, y.unsqueeze(0))
|
||||
return ret
|
||||
|
||||
def _matmul_with_relative_keys(self, x, y):
|
||||
"""
|
||||
x: [b, h, l, d]
|
||||
y: [h or 1, m, d]
|
||||
ret: [b, h, l, m]
|
||||
"""
|
||||
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
|
||||
return ret
|
||||
|
||||
def _get_relative_embeddings(self, relative_embeddings, length):
|
||||
max_relative_position = 2 * self.window_size + 1
|
||||
# Pad first before slice to avoid using cond ops.
|
||||
pad_length = max(length - (self.window_size + 1), 0)
|
||||
slice_start_position = max((self.window_size + 1) - length, 0)
|
||||
slice_end_position = slice_start_position + 2 * length - 1
|
||||
if pad_length > 0:
|
||||
padded_relative_embeddings = F.pad(
|
||||
relative_embeddings,
|
||||
convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]))
|
||||
else:
|
||||
padded_relative_embeddings = relative_embeddings
|
||||
used_relative_embeddings = padded_relative_embeddings[:, slice_start_position:slice_end_position]
|
||||
return used_relative_embeddings
|
||||
|
||||
def _relative_position_to_absolute_position(self, x):
|
||||
"""
|
||||
x: [b, h, l, 2*l-1]
|
||||
ret: [b, h, l, l]
|
||||
"""
|
||||
batch, heads, length, _ = x.size()
|
||||
# Concat columns of pad to shift from relative to absolute indexing.
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
|
||||
|
||||
# Concat extra elements so to add up to shape (len+1, 2*len-1).
|
||||
x_flat = x.view([batch, heads, length * 2 * length])
|
||||
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [0, length - 1]]))
|
||||
|
||||
# Reshape and slice out the padded elements.
|
||||
x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[:, :, :length, length - 1:]
|
||||
return x_final
|
||||
|
||||
def _absolute_position_to_relative_position(self, x):
|
||||
"""
|
||||
x: [b, h, l, l]
|
||||
ret: [b, h, l, 2*l-1]
|
||||
"""
|
||||
batch, heads, length, _ = x.size()
|
||||
# padd along column
|
||||
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]]))
|
||||
x_flat = x.view([batch, heads, length ** 2 + length * (length - 1)])
|
||||
# add 0's in the beginning that will skew the elements after reshape
|
||||
x_flat = F.pad(x_flat, convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
|
||||
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
|
||||
return x_final
|
||||
|
||||
def _attention_bias_proximal(self, length):
|
||||
"""Bias for self-attention to encourage attention to close positions.
|
||||
Args:
|
||||
length: an integer scalar.
|
||||
Returns:
|
||||
a Tensor with shape [1, 1, length, length]
|
||||
"""
|
||||
r = torch.arange(length, dtype=torch.float32)
|
||||
diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
|
||||
return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
|
||||
|
||||
|
||||
class FFN(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, filter_channels, kernel_size, p_dropout=0., activation=None):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.activation = activation
|
||||
|
||||
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size // 2)
|
||||
self.conv_2 = nn.Conv1d(filter_channels, out_channels, 1)
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
x = self.conv_1(x * x_mask)
|
||||
if self.activation == "gelu":
|
||||
x = x * torch.sigmoid(1.702 * x)
|
||||
else:
|
||||
x = torch.relu(x)
|
||||
x = self.drop(x)
|
||||
x = self.conv_2(x * x_mask)
|
||||
return x * x_mask
|
||||
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
def __init__(self, channels, eps=1e-4):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
self.eps = eps
|
||||
|
||||
self.gamma = nn.Parameter(torch.ones(channels))
|
||||
self.beta = nn.Parameter(torch.zeros(channels))
|
||||
|
||||
def forward(self, x):
|
||||
n_dims = len(x.shape)
|
||||
mean = torch.mean(x, 1, keepdim=True)
|
||||
variance = torch.mean((x - mean) ** 2, 1, keepdim=True)
|
||||
|
||||
x = (x - mean) * torch.rsqrt(variance + self.eps)
|
||||
|
||||
shape = [1, -1] + [1] * (n_dims - 2)
|
||||
x = x * self.gamma.view(*shape) + self.beta.view(*shape)
|
||||
return x
|
||||
|
||||
|
||||
class ConvReluNorm(nn.Module):
|
||||
def __init__(self, in_channels, hidden_channels, out_channels, kernel_size, n_layers, p_dropout):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.hidden_channels = hidden_channels
|
||||
self.out_channels = out_channels
|
||||
self.kernel_size = kernel_size
|
||||
self.n_layers = n_layers
|
||||
self.p_dropout = p_dropout
|
||||
assert n_layers > 1, "Number of layers should be larger than 0."
|
||||
|
||||
self.conv_layers = nn.ModuleList()
|
||||
self.norm_layers = nn.ModuleList()
|
||||
self.conv_layers.append(nn.Conv1d(in_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
|
||||
self.norm_layers.append(LayerNorm(hidden_channels))
|
||||
self.relu_drop = nn.Sequential(
|
||||
nn.ReLU(),
|
||||
nn.Dropout(p_dropout))
|
||||
for _ in range(n_layers - 1):
|
||||
self.conv_layers.append(nn.Conv1d(hidden_channels, hidden_channels, kernel_size, padding=kernel_size // 2))
|
||||
self.norm_layers.append(LayerNorm(hidden_channels))
|
||||
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
|
||||
self.proj.weight.data.zero_()
|
||||
self.proj.bias.data.zero_()
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
x_org = x
|
||||
for i in range(self.n_layers):
|
||||
x = self.conv_layers[i](x * x_mask)
|
||||
x = self.norm_layers[i](x)
|
||||
x = self.relu_drop(x)
|
||||
x = x_org + self.proj(x)
|
||||
return x * x_mask
|
||||
|
||||
|
||||
class RelTransformerEncoder(nn.Module):
|
||||
def __init__(self,
|
||||
n_vocab,
|
||||
out_channels,
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout=0.0,
|
||||
window_size=4,
|
||||
block_length=None,
|
||||
prenet=True,
|
||||
pre_ln=True,
|
||||
):
|
||||
|
||||
super().__init__()
|
||||
|
||||
self.n_vocab = n_vocab
|
||||
self.out_channels = out_channels
|
||||
self.hidden_channels = hidden_channels
|
||||
self.filter_channels = filter_channels
|
||||
self.n_heads = n_heads
|
||||
self.n_layers = n_layers
|
||||
self.kernel_size = kernel_size
|
||||
self.p_dropout = p_dropout
|
||||
self.window_size = window_size
|
||||
self.block_length = block_length
|
||||
self.prenet = prenet
|
||||
if n_vocab > 0:
|
||||
self.emb = Embedding(n_vocab, hidden_channels, padding_idx=0)
|
||||
|
||||
if prenet:
|
||||
self.pre = ConvReluNorm(hidden_channels, hidden_channels, hidden_channels,
|
||||
kernel_size=5, n_layers=3, p_dropout=0)
|
||||
self.encoder = Encoder(
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout,
|
||||
window_size=window_size,
|
||||
block_length=block_length,
|
||||
pre_ln=pre_ln,
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask=None):
|
||||
if self.n_vocab > 0:
|
||||
x_lengths = (x > 0).long().sum(-1)
|
||||
x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]
|
||||
else:
|
||||
x_lengths = (x.abs().sum(-1) > 0).long().sum(-1)
|
||||
x = torch.transpose(x, 1, -1) # [b, h, t]
|
||||
x_mask = torch.unsqueeze(sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)
|
||||
|
||||
if self.prenet:
|
||||
x = self.pre(x, x_mask)
|
||||
x = self.encoder(x, x_mask)
|
||||
return x.transpose(1, 2)
|
||||
@@ -0,0 +1,261 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class PreNet(nn.Module):
|
||||
def __init__(self, in_dims, fc1_dims=256, fc2_dims=128, dropout=0.5):
|
||||
super().__init__()
|
||||
self.fc1 = nn.Linear(in_dims, fc1_dims)
|
||||
self.fc2 = nn.Linear(fc1_dims, fc2_dims)
|
||||
self.p = dropout
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.p, training=self.training)
|
||||
x = self.fc2(x)
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.p, training=self.training)
|
||||
return x
|
||||
|
||||
|
||||
class HighwayNetwork(nn.Module):
|
||||
def __init__(self, size):
|
||||
super().__init__()
|
||||
self.W1 = nn.Linear(size, size)
|
||||
self.W2 = nn.Linear(size, size)
|
||||
self.W1.bias.data.fill_(0.)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.W1(x)
|
||||
x2 = self.W2(x)
|
||||
g = torch.sigmoid(x2)
|
||||
y = g * F.relu(x1) + (1. - g) * x
|
||||
return y
|
||||
|
||||
|
||||
class BatchNormConv(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel, relu=True):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv1d(in_channels, out_channels, kernel, stride=1, padding=kernel // 2, bias=False)
|
||||
self.bnorm = nn.BatchNorm1d(out_channels)
|
||||
self.relu = relu
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = F.relu(x) if self.relu is True else x
|
||||
return self.bnorm(x)
|
||||
|
||||
|
||||
class ConvNorm(torch.nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size=1, stride=1,
|
||||
padding=None, dilation=1, bias=True, w_init_gain='linear'):
|
||||
super(ConvNorm, self).__init__()
|
||||
if padding is None:
|
||||
assert (kernel_size % 2 == 1)
|
||||
padding = int(dilation * (kernel_size - 1) / 2)
|
||||
|
||||
self.conv = torch.nn.Conv1d(in_channels, out_channels,
|
||||
kernel_size=kernel_size, stride=stride,
|
||||
padding=padding, dilation=dilation,
|
||||
bias=bias)
|
||||
|
||||
torch.nn.init.xavier_uniform_(
|
||||
self.conv.weight, gain=torch.nn.init.calculate_gain(w_init_gain))
|
||||
|
||||
def forward(self, signal):
|
||||
conv_signal = self.conv(signal)
|
||||
return conv_signal
|
||||
|
||||
|
||||
class CBHG(nn.Module):
|
||||
def __init__(self, K, in_channels, channels, proj_channels, num_highways):
|
||||
super().__init__()
|
||||
|
||||
# List of all rnns to call `flatten_parameters()` on
|
||||
self._to_flatten = []
|
||||
|
||||
self.bank_kernels = [i for i in range(1, K + 1)]
|
||||
self.conv1d_bank = nn.ModuleList()
|
||||
for k in self.bank_kernels:
|
||||
conv = BatchNormConv(in_channels, channels, k)
|
||||
self.conv1d_bank.append(conv)
|
||||
|
||||
self.maxpool = nn.MaxPool1d(kernel_size=2, stride=1, padding=1)
|
||||
|
||||
self.conv_project1 = BatchNormConv(len(self.bank_kernels) * channels, proj_channels[0], 3)
|
||||
self.conv_project2 = BatchNormConv(proj_channels[0], proj_channels[1], 3, relu=False)
|
||||
|
||||
# Fix the highway input if necessary
|
||||
if proj_channels[-1] != channels:
|
||||
self.highway_mismatch = True
|
||||
self.pre_highway = nn.Linear(proj_channels[-1], channels, bias=False)
|
||||
else:
|
||||
self.highway_mismatch = False
|
||||
|
||||
self.highways = nn.ModuleList()
|
||||
for i in range(num_highways):
|
||||
hn = HighwayNetwork(channels)
|
||||
self.highways.append(hn)
|
||||
|
||||
self.rnn = nn.GRU(channels, channels, batch_first=True, bidirectional=True)
|
||||
self._to_flatten.append(self.rnn)
|
||||
|
||||
# Avoid fragmentation of RNN parameters and associated warning
|
||||
self._flatten_parameters()
|
||||
|
||||
def forward(self, x):
|
||||
# Although we `_flatten_parameters()` on init, when using DataParallel
|
||||
# the model gets replicated, making it no longer guaranteed that the
|
||||
# weights are contiguous in GPU memory. Hence, we must call it again
|
||||
self._flatten_parameters()
|
||||
|
||||
# Save these for later
|
||||
residual = x
|
||||
seq_len = x.size(-1)
|
||||
conv_bank = []
|
||||
|
||||
# Convolution Bank
|
||||
for conv in self.conv1d_bank:
|
||||
c = conv(x) # Convolution
|
||||
conv_bank.append(c[:, :, :seq_len])
|
||||
|
||||
# Stack along the channel axis
|
||||
conv_bank = torch.cat(conv_bank, dim=1)
|
||||
|
||||
# dump the last padding to fit residual
|
||||
x = self.maxpool(conv_bank)[:, :, :seq_len]
|
||||
|
||||
# Conv1d projections
|
||||
x = self.conv_project1(x)
|
||||
x = self.conv_project2(x)
|
||||
|
||||
# Residual Connect
|
||||
x = x + residual
|
||||
|
||||
# Through the highways
|
||||
x = x.transpose(1, 2)
|
||||
if self.highway_mismatch is True:
|
||||
x = self.pre_highway(x)
|
||||
for h in self.highways:
|
||||
x = h(x)
|
||||
|
||||
# And then the RNN
|
||||
x, _ = self.rnn(x)
|
||||
return x
|
||||
|
||||
def _flatten_parameters(self):
|
||||
"""Calls `flatten_parameters` on all the rnns used by the WaveRNN. Used
|
||||
to improve efficiency and avoid PyTorch yelling at us."""
|
||||
[m.flatten_parameters() for m in self._to_flatten]
|
||||
|
||||
|
||||
class TacotronEncoder(nn.Module):
|
||||
def __init__(self, embed_dims, num_chars, cbhg_channels, K, num_highways, dropout):
|
||||
super().__init__()
|
||||
self.embedding = nn.Embedding(num_chars, embed_dims)
|
||||
self.pre_net = PreNet(embed_dims, embed_dims, embed_dims, dropout=dropout)
|
||||
self.cbhg = CBHG(K=K, in_channels=cbhg_channels, channels=cbhg_channels,
|
||||
proj_channels=[cbhg_channels, cbhg_channels],
|
||||
num_highways=num_highways)
|
||||
self.proj_out = nn.Linear(cbhg_channels * 2, cbhg_channels)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.embedding(x)
|
||||
x = self.pre_net(x)
|
||||
x.transpose_(1, 2)
|
||||
x = self.cbhg(x)
|
||||
x = self.proj_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class RNNEncoder(nn.Module):
|
||||
def __init__(self, num_chars, embedding_dim, n_convolutions=3, kernel_size=5):
|
||||
super(RNNEncoder, self).__init__()
|
||||
self.embedding = nn.Embedding(num_chars, embedding_dim, padding_idx=0)
|
||||
convolutions = []
|
||||
for _ in range(n_convolutions):
|
||||
conv_layer = nn.Sequential(
|
||||
ConvNorm(embedding_dim,
|
||||
embedding_dim,
|
||||
kernel_size=kernel_size, stride=1,
|
||||
padding=int((kernel_size - 1) / 2),
|
||||
dilation=1, w_init_gain='relu'),
|
||||
nn.BatchNorm1d(embedding_dim))
|
||||
convolutions.append(conv_layer)
|
||||
self.convolutions = nn.ModuleList(convolutions)
|
||||
|
||||
self.lstm = nn.LSTM(embedding_dim, int(embedding_dim / 2), 1,
|
||||
batch_first=True, bidirectional=True)
|
||||
|
||||
def forward(self, x):
|
||||
input_lengths = (x > 0).sum(-1)
|
||||
input_lengths = input_lengths.cpu().numpy()
|
||||
|
||||
x = self.embedding(x)
|
||||
x = x.transpose(1, 2) # [B, H, T]
|
||||
for conv in self.convolutions:
|
||||
x = F.dropout(F.relu(conv(x)), 0.5, self.training) + x
|
||||
x = x.transpose(1, 2) # [B, T, H]
|
||||
|
||||
# pytorch tensor are not reversible, hence the conversion
|
||||
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
|
||||
|
||||
self.lstm.flatten_parameters()
|
||||
outputs, _ = self.lstm(x)
|
||||
outputs, _ = nn.utils.rnn.pad_packed_sequence(outputs, batch_first=True)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
class DecoderRNN(torch.nn.Module):
|
||||
def __init__(self, hidden_size, decoder_rnn_dim, dropout):
|
||||
super(DecoderRNN, self).__init__()
|
||||
self.in_conv1d = nn.Sequential(
|
||||
torch.nn.Conv1d(
|
||||
in_channels=hidden_size,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=9, padding=4,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv1d(
|
||||
in_channels=hidden_size,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=9, padding=4,
|
||||
),
|
||||
)
|
||||
self.ln = nn.LayerNorm(hidden_size)
|
||||
if decoder_rnn_dim == 0:
|
||||
decoder_rnn_dim = hidden_size * 2
|
||||
self.rnn = torch.nn.LSTM(
|
||||
input_size=hidden_size,
|
||||
hidden_size=decoder_rnn_dim,
|
||||
num_layers=1,
|
||||
batch_first=True,
|
||||
bidirectional=True,
|
||||
dropout=dropout
|
||||
)
|
||||
self.rnn.flatten_parameters()
|
||||
self.conv1d = torch.nn.Conv1d(
|
||||
in_channels=decoder_rnn_dim * 2,
|
||||
out_channels=hidden_size,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
input_masks = x.abs().sum(-1).ne(0).data[:, :, None]
|
||||
input_lengths = input_masks.sum([-1, -2])
|
||||
input_lengths = input_lengths.cpu().numpy()
|
||||
|
||||
x = self.in_conv1d(x.transpose(1, 2)).transpose(1, 2)
|
||||
x = self.ln(x)
|
||||
x = nn.utils.rnn.pack_padded_sequence(x, input_lengths, batch_first=True, enforce_sorted=False)
|
||||
self.rnn.flatten_parameters()
|
||||
x, _ = self.rnn(x) # [B, T, C]
|
||||
x, _ = nn.utils.rnn.pad_packed_sequence(x, batch_first=True)
|
||||
x = x * input_masks
|
||||
pre_mel = self.conv1d(x.transpose(1, 2)).transpose(1, 2) # [B, T, C]
|
||||
pre_mel = pre_mel * input_masks
|
||||
return pre_mel
|
||||
@@ -0,0 +1,751 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import Parameter, Linear
|
||||
from .layers import LayerNorm, Embedding
|
||||
from ...utils.nn.seq_utils import (
|
||||
get_incremental_state,
|
||||
set_incremental_state,
|
||||
softmax,
|
||||
make_positions,
|
||||
)
|
||||
import torch.nn.functional as F
|
||||
|
||||
DEFAULT_MAX_SOURCE_POSITIONS = 2000
|
||||
DEFAULT_MAX_TARGET_POSITIONS = 2000
|
||||
|
||||
|
||||
class SinusoidalPositionalEmbedding(nn.Module):
|
||||
"""This module produces sinusoidal positional embeddings of any length.
|
||||
|
||||
Padding symbols are ignored.
|
||||
"""
|
||||
|
||||
def __init__(self, embedding_dim, padding_idx, init_size=1024):
|
||||
super().__init__()
|
||||
self.embedding_dim = embedding_dim
|
||||
self.padding_idx = padding_idx
|
||||
self.weights = SinusoidalPositionalEmbedding.get_embedding(
|
||||
init_size,
|
||||
embedding_dim,
|
||||
padding_idx,
|
||||
)
|
||||
self.register_buffer('_float_tensor', torch.FloatTensor(1))
|
||||
|
||||
@staticmethod
|
||||
def get_embedding(num_embeddings, embedding_dim, padding_idx=None):
|
||||
"""Build sinusoidal embeddings.
|
||||
|
||||
This matches the implementation in tensor2tensor, but differs slightly
|
||||
from the description in Section 3.5 of "Attention Is All You Need".
|
||||
"""
|
||||
half_dim = embedding_dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = torch.exp(torch.arange(half_dim, dtype=torch.float) * -emb)
|
||||
emb = torch.arange(num_embeddings, dtype=torch.float).unsqueeze(1) * emb.unsqueeze(0)
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1).view(num_embeddings, -1)
|
||||
if embedding_dim % 2 == 1:
|
||||
# zero pad
|
||||
emb = torch.cat([emb, torch.zeros(num_embeddings, 1)], dim=1)
|
||||
if padding_idx is not None:
|
||||
emb[padding_idx, :] = 0
|
||||
return emb
|
||||
|
||||
def forward(self, input, incremental_state=None, timestep=None, positions=None, **kwargs):
|
||||
"""Input is expected to be of size [bsz x seqlen]."""
|
||||
bsz, seq_len = input.shape[:2]
|
||||
max_pos = self.padding_idx + 1 + seq_len
|
||||
if self.weights is None or max_pos > self.weights.size(0):
|
||||
# recompute/expand embeddings if needed
|
||||
self.weights = SinusoidalPositionalEmbedding.get_embedding(
|
||||
max_pos,
|
||||
self.embedding_dim,
|
||||
self.padding_idx,
|
||||
)
|
||||
self.weights = self.weights.to(self._float_tensor)
|
||||
|
||||
if incremental_state is not None:
|
||||
# positions is the same for every token when decoding a single step
|
||||
pos = timestep.view(-1)[0] + 1 if timestep is not None else seq_len
|
||||
return self.weights[self.padding_idx + pos, :].expand(bsz, 1, -1)
|
||||
|
||||
positions = make_positions(input, self.padding_idx) if positions is None else positions
|
||||
return self.weights.index_select(0, positions.view(-1)).view(bsz, seq_len, -1).detach()
|
||||
|
||||
def max_positions(self):
|
||||
"""Maximum number of supported positions."""
|
||||
return int(1e5) # an arbitrary large number
|
||||
|
||||
|
||||
class TransformerFFNLayer(nn.Module):
|
||||
def __init__(self, hidden_size, filter_size, padding="SAME", kernel_size=1, dropout=0., act='gelu'):
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.dropout = dropout
|
||||
self.act = act
|
||||
if padding == 'SAME':
|
||||
self.ffn_1 = nn.Conv1d(hidden_size, filter_size, kernel_size, padding=kernel_size // 2)
|
||||
elif padding == 'LEFT':
|
||||
self.ffn_1 = nn.Sequential(
|
||||
nn.ConstantPad1d((kernel_size - 1, 0), 0.0),
|
||||
nn.Conv1d(hidden_size, filter_size, kernel_size)
|
||||
)
|
||||
self.ffn_2 = Linear(filter_size, hidden_size)
|
||||
|
||||
def forward(self, x, incremental_state=None):
|
||||
# x: T x B x C
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_input' in saved_state:
|
||||
prev_input = saved_state['prev_input']
|
||||
x = torch.cat((prev_input, x), dim=0)
|
||||
x = x[-self.kernel_size:]
|
||||
saved_state['prev_input'] = x
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
x = self.ffn_1(x.permute(1, 2, 0)).permute(2, 0, 1)
|
||||
x = x * self.kernel_size ** -0.5
|
||||
|
||||
if incremental_state is not None:
|
||||
x = x[-1:]
|
||||
if self.act == 'gelu':
|
||||
x = F.gelu(x)
|
||||
if self.act == 'relu':
|
||||
x = F.relu(x)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = self.ffn_2(x)
|
||||
return x
|
||||
|
||||
def _get_input_buffer(self, incremental_state):
|
||||
return get_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'f',
|
||||
) or {}
|
||||
|
||||
def _set_input_buffer(self, incremental_state, buffer):
|
||||
set_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'f',
|
||||
buffer,
|
||||
)
|
||||
|
||||
def clear_buffer(self, incremental_state):
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_input' in saved_state:
|
||||
del saved_state['prev_input']
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
|
||||
class MultiheadAttention(nn.Module):
|
||||
def __init__(self, embed_dim, num_heads, kdim=None, vdim=None, dropout=0., bias=True,
|
||||
add_bias_kv=False, add_zero_attn=False, self_attention=False,
|
||||
encoder_decoder_attention=False):
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.kdim = kdim if kdim is not None else embed_dim
|
||||
self.vdim = vdim if vdim is not None else embed_dim
|
||||
self.qkv_same_dim = self.kdim == embed_dim and self.vdim == embed_dim
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.dropout = dropout
|
||||
self.head_dim = embed_dim // num_heads
|
||||
assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
|
||||
self.scaling = self.head_dim ** -0.5
|
||||
|
||||
self.self_attention = self_attention
|
||||
self.encoder_decoder_attention = encoder_decoder_attention
|
||||
|
||||
assert not self.self_attention or self.qkv_same_dim, 'Self-attention requires query, key and ' \
|
||||
'value to be of the same size'
|
||||
|
||||
if self.qkv_same_dim:
|
||||
self.in_proj_weight = Parameter(torch.Tensor(3 * embed_dim, embed_dim))
|
||||
else:
|
||||
self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim))
|
||||
self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim))
|
||||
self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim))
|
||||
|
||||
if bias:
|
||||
self.in_proj_bias = Parameter(torch.Tensor(3 * embed_dim))
|
||||
else:
|
||||
self.register_parameter('in_proj_bias', None)
|
||||
|
||||
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
|
||||
|
||||
if add_bias_kv:
|
||||
self.bias_k = Parameter(torch.Tensor(1, 1, embed_dim))
|
||||
self.bias_v = Parameter(torch.Tensor(1, 1, embed_dim))
|
||||
else:
|
||||
self.bias_k = self.bias_v = None
|
||||
|
||||
self.add_zero_attn = add_zero_attn
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
self.enable_torch_version = False
|
||||
if hasattr(F, "multi_head_attention_forward"):
|
||||
self.enable_torch_version = True
|
||||
else:
|
||||
self.enable_torch_version = False
|
||||
self.last_attn_probs = None
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.qkv_same_dim:
|
||||
nn.init.xavier_uniform_(self.in_proj_weight)
|
||||
else:
|
||||
nn.init.xavier_uniform_(self.k_proj_weight)
|
||||
nn.init.xavier_uniform_(self.v_proj_weight)
|
||||
nn.init.xavier_uniform_(self.q_proj_weight)
|
||||
|
||||
nn.init.xavier_uniform_(self.out_proj.weight)
|
||||
if self.in_proj_bias is not None:
|
||||
nn.init.constant_(self.in_proj_bias, 0.)
|
||||
nn.init.constant_(self.out_proj.bias, 0.)
|
||||
if self.bias_k is not None:
|
||||
nn.init.xavier_normal_(self.bias_k)
|
||||
if self.bias_v is not None:
|
||||
nn.init.xavier_normal_(self.bias_v)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query, key, value,
|
||||
key_padding_mask=None,
|
||||
incremental_state=None,
|
||||
need_weights=True,
|
||||
static_kv=False,
|
||||
attn_mask=None,
|
||||
before_softmax=False,
|
||||
need_head_weights=False,
|
||||
enc_dec_attn_constraint_mask=None,
|
||||
reset_attn_weight=None
|
||||
):
|
||||
"""Input shape: Time x Batch x Channel
|
||||
|
||||
Args:
|
||||
key_padding_mask (ByteTensor, optional): mask to exclude
|
||||
keys that are pads, of shape `(batch, src_len)`, where
|
||||
padding elements are indicated by 1s.
|
||||
need_weights (bool, optional): return the attention weights,
|
||||
averaged over heads (default: False).
|
||||
attn_mask (ByteTensor, optional): typically used to
|
||||
implement causal attention, where the mask prevents the
|
||||
attention from looking forward in time (default: None).
|
||||
before_softmax (bool, optional): return the raw attention
|
||||
weights and values before the attention softmax.
|
||||
need_head_weights (bool, optional): return the attention
|
||||
weights for each head. Implies *need_weights*. Default:
|
||||
return the average attention weights over all heads.
|
||||
"""
|
||||
if need_head_weights:
|
||||
need_weights = True
|
||||
|
||||
tgt_len, bsz, embed_dim = query.size()
|
||||
assert embed_dim == self.embed_dim
|
||||
assert list(query.size()) == [tgt_len, bsz, embed_dim]
|
||||
if self.enable_torch_version and incremental_state is None and not static_kv and reset_attn_weight is None:
|
||||
if self.qkv_same_dim:
|
||||
return F.multi_head_attention_forward(query, key, value,
|
||||
self.embed_dim, self.num_heads,
|
||||
self.in_proj_weight,
|
||||
self.in_proj_bias, self.bias_k, self.bias_v,
|
||||
self.add_zero_attn, self.dropout,
|
||||
self.out_proj.weight, self.out_proj.bias,
|
||||
self.training, key_padding_mask, need_weights,
|
||||
attn_mask)
|
||||
else:
|
||||
return F.multi_head_attention_forward(query, key, value,
|
||||
self.embed_dim, self.num_heads,
|
||||
torch.empty([0]),
|
||||
self.in_proj_bias, self.bias_k, self.bias_v,
|
||||
self.add_zero_attn, self.dropout,
|
||||
self.out_proj.weight, self.out_proj.bias,
|
||||
self.training, key_padding_mask, need_weights,
|
||||
attn_mask, use_separate_proj_weight=True,
|
||||
q_proj_weight=self.q_proj_weight,
|
||||
k_proj_weight=self.k_proj_weight,
|
||||
v_proj_weight=self.v_proj_weight)
|
||||
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_key' in saved_state:
|
||||
# previous time steps are cached - no need to recompute
|
||||
# key and value if they are static
|
||||
if static_kv:
|
||||
assert self.encoder_decoder_attention and not self.self_attention
|
||||
key = value = None
|
||||
else:
|
||||
saved_state = None
|
||||
|
||||
if self.self_attention:
|
||||
# self-attention
|
||||
q, k, v = self.in_proj_qkv(query)
|
||||
elif self.encoder_decoder_attention:
|
||||
# encoder-decoder attention
|
||||
q = self.in_proj_q(query)
|
||||
if key is None:
|
||||
assert value is None
|
||||
k = v = None
|
||||
else:
|
||||
k = self.in_proj_k(key)
|
||||
v = self.in_proj_v(key)
|
||||
|
||||
else:
|
||||
q = self.in_proj_q(query)
|
||||
k = self.in_proj_k(key)
|
||||
v = self.in_proj_v(value)
|
||||
q *= self.scaling
|
||||
|
||||
if self.bias_k is not None:
|
||||
assert self.bias_v is not None
|
||||
k = torch.cat([k, self.bias_k.repeat(1, bsz, 1)])
|
||||
v = torch.cat([v, self.bias_v.repeat(1, bsz, 1)])
|
||||
if attn_mask is not None:
|
||||
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
|
||||
if key_padding_mask is not None:
|
||||
key_padding_mask = torch.cat(
|
||||
[key_padding_mask, key_padding_mask.new_zeros(key_padding_mask.size(0), 1)], dim=1)
|
||||
|
||||
q = q.contiguous().view(tgt_len, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
if k is not None:
|
||||
k = k.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
if v is not None:
|
||||
v = v.contiguous().view(-1, bsz * self.num_heads, self.head_dim).transpose(0, 1)
|
||||
|
||||
if saved_state is not None:
|
||||
# saved states are stored with shape (bsz, num_heads, seq_len, head_dim)
|
||||
if 'prev_key' in saved_state:
|
||||
prev_key = saved_state['prev_key'].view(bsz * self.num_heads, -1, self.head_dim)
|
||||
if static_kv:
|
||||
k = prev_key
|
||||
else:
|
||||
k = torch.cat((prev_key, k), dim=1)
|
||||
if 'prev_value' in saved_state:
|
||||
prev_value = saved_state['prev_value'].view(bsz * self.num_heads, -1, self.head_dim)
|
||||
if static_kv:
|
||||
v = prev_value
|
||||
else:
|
||||
v = torch.cat((prev_value, v), dim=1)
|
||||
if 'prev_key_padding_mask' in saved_state and saved_state['prev_key_padding_mask'] is not None:
|
||||
prev_key_padding_mask = saved_state['prev_key_padding_mask']
|
||||
if static_kv:
|
||||
key_padding_mask = prev_key_padding_mask
|
||||
else:
|
||||
key_padding_mask = torch.cat((prev_key_padding_mask, key_padding_mask), dim=1)
|
||||
|
||||
saved_state['prev_key'] = k.view(bsz, self.num_heads, -1, self.head_dim)
|
||||
saved_state['prev_value'] = v.view(bsz, self.num_heads, -1, self.head_dim)
|
||||
saved_state['prev_key_padding_mask'] = key_padding_mask
|
||||
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
src_len = k.size(1)
|
||||
|
||||
# This is part of a workaround to get around fork/join parallelism
|
||||
# not supporting Optional types.
|
||||
if key_padding_mask is not None and key_padding_mask.shape == torch.Size([]):
|
||||
key_padding_mask = None
|
||||
|
||||
if key_padding_mask is not None:
|
||||
assert key_padding_mask.size(0) == bsz
|
||||
assert key_padding_mask.size(1) == src_len
|
||||
|
||||
if self.add_zero_attn:
|
||||
src_len += 1
|
||||
k = torch.cat([k, k.new_zeros((k.size(0), 1) + k.size()[2:])], dim=1)
|
||||
v = torch.cat([v, v.new_zeros((v.size(0), 1) + v.size()[2:])], dim=1)
|
||||
if attn_mask is not None:
|
||||
attn_mask = torch.cat([attn_mask, attn_mask.new_zeros(attn_mask.size(0), 1)], dim=1)
|
||||
if key_padding_mask is not None:
|
||||
key_padding_mask = torch.cat(
|
||||
[key_padding_mask, torch.zeros(key_padding_mask.size(0), 1).type_as(key_padding_mask)], dim=1)
|
||||
|
||||
attn_weights = torch.bmm(q, k.transpose(1, 2))
|
||||
attn_weights = self.apply_sparse_mask(attn_weights, tgt_len, src_len, bsz)
|
||||
|
||||
assert list(attn_weights.size()) == [bsz * self.num_heads, tgt_len, src_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
if len(attn_mask.shape) == 2:
|
||||
attn_mask = attn_mask.unsqueeze(0)
|
||||
elif len(attn_mask.shape) == 3:
|
||||
attn_mask = attn_mask[:, None].repeat([1, self.num_heads, 1, 1]).reshape(
|
||||
bsz * self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights + attn_mask
|
||||
|
||||
if enc_dec_attn_constraint_mask is not None: # bs x head x L_kv
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights.masked_fill(
|
||||
enc_dec_attn_constraint_mask.unsqueeze(2).bool(),
|
||||
-1e8,
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
if key_padding_mask is not None:
|
||||
# don't attend to padding symbols
|
||||
attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights.masked_fill(
|
||||
key_padding_mask.unsqueeze(1).unsqueeze(2),
|
||||
-1e8,
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
|
||||
|
||||
attn_logits = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
|
||||
|
||||
if before_softmax:
|
||||
return attn_weights, v
|
||||
|
||||
attn_weights_float = softmax(attn_weights, dim=-1)
|
||||
attn_weights = attn_weights_float.type_as(attn_weights)
|
||||
attn_probs = F.dropout(attn_weights_float.type_as(attn_weights), p=self.dropout, training=self.training)
|
||||
|
||||
if reset_attn_weight is not None:
|
||||
if reset_attn_weight:
|
||||
self.last_attn_probs = attn_probs.detach()
|
||||
else:
|
||||
assert self.last_attn_probs is not None
|
||||
attn_probs = self.last_attn_probs
|
||||
attn = torch.bmm(attn_probs, v)
|
||||
assert list(attn.size()) == [bsz * self.num_heads, tgt_len, self.head_dim]
|
||||
attn = attn.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)
|
||||
attn = self.out_proj(attn)
|
||||
|
||||
if need_weights:
|
||||
attn_weights = attn_weights_float.view(bsz, self.num_heads, tgt_len, src_len).transpose(1, 0)
|
||||
if not need_head_weights:
|
||||
# average attention weights over heads
|
||||
attn_weights = attn_weights.mean(dim=0)
|
||||
else:
|
||||
attn_weights = None
|
||||
|
||||
return attn, (attn_weights, attn_logits)
|
||||
|
||||
def in_proj_qkv(self, query):
|
||||
return self._in_proj(query).chunk(3, dim=-1)
|
||||
|
||||
def in_proj_q(self, query):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(query, end=self.embed_dim)
|
||||
else:
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[:self.embed_dim]
|
||||
return F.linear(query, self.q_proj_weight, bias)
|
||||
|
||||
def in_proj_k(self, key):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(key, start=self.embed_dim, end=2 * self.embed_dim)
|
||||
else:
|
||||
weight = self.k_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[self.embed_dim:2 * self.embed_dim]
|
||||
return F.linear(key, weight, bias)
|
||||
|
||||
def in_proj_v(self, value):
|
||||
if self.qkv_same_dim:
|
||||
return self._in_proj(value, start=2 * self.embed_dim)
|
||||
else:
|
||||
weight = self.v_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
if bias is not None:
|
||||
bias = bias[2 * self.embed_dim:]
|
||||
return F.linear(value, weight, bias)
|
||||
|
||||
def _in_proj(self, input, start=0, end=None):
|
||||
weight = self.in_proj_weight
|
||||
bias = self.in_proj_bias
|
||||
weight = weight[start:end, :]
|
||||
if bias is not None:
|
||||
bias = bias[start:end]
|
||||
return F.linear(input, weight, bias)
|
||||
|
||||
def _get_input_buffer(self, incremental_state):
|
||||
return get_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'attn_state',
|
||||
) or {}
|
||||
|
||||
def _set_input_buffer(self, incremental_state, buffer):
|
||||
set_incremental_state(
|
||||
self,
|
||||
incremental_state,
|
||||
'attn_state',
|
||||
buffer,
|
||||
)
|
||||
|
||||
def apply_sparse_mask(self, attn_weights, tgt_len, src_len, bsz):
|
||||
return attn_weights
|
||||
|
||||
def clear_buffer(self, incremental_state=None):
|
||||
if incremental_state is not None:
|
||||
saved_state = self._get_input_buffer(incremental_state)
|
||||
if 'prev_key' in saved_state:
|
||||
del saved_state['prev_key']
|
||||
if 'prev_value' in saved_state:
|
||||
del saved_state['prev_value']
|
||||
self._set_input_buffer(incremental_state, saved_state)
|
||||
|
||||
|
||||
class EncSALayer(nn.Module):
|
||||
def __init__(self, c, num_heads, dropout, attention_dropout=0.1,
|
||||
relu_dropout=0.1, kernel_size=9, padding='SAME', act='gelu'):
|
||||
super().__init__()
|
||||
self.c = c
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
if num_heads > 0:
|
||||
self.layer_norm1 = LayerNorm(c)
|
||||
self.self_attn = MultiheadAttention(
|
||||
self.c, num_heads, self_attention=True, dropout=attention_dropout, bias=False)
|
||||
self.layer_norm2 = LayerNorm(c)
|
||||
self.ffn = TransformerFFNLayer(
|
||||
c, 4 * c, kernel_size=kernel_size, dropout=relu_dropout, padding=padding, act=act)
|
||||
|
||||
def forward(self, x, encoder_padding_mask=None, **kwargs):
|
||||
layer_norm_training = kwargs.get('layer_norm_training', None)
|
||||
if layer_norm_training is not None:
|
||||
self.layer_norm1.training = layer_norm_training
|
||||
self.layer_norm2.training = layer_norm_training
|
||||
if self.num_heads > 0:
|
||||
residual = x
|
||||
x = self.layer_norm1(x)
|
||||
x, _, = self.self_attn(
|
||||
query=x,
|
||||
key=x,
|
||||
value=x,
|
||||
key_padding_mask=encoder_padding_mask
|
||||
)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
|
||||
|
||||
residual = x
|
||||
x = self.layer_norm2(x)
|
||||
x = self.ffn(x)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
x = x * (1 - encoder_padding_mask.float()).transpose(0, 1)[..., None]
|
||||
return x
|
||||
|
||||
|
||||
class DecSALayer(nn.Module):
|
||||
def __init__(self, c, num_heads, dropout, attention_dropout=0.1, relu_dropout=0.1,
|
||||
kernel_size=9, act='gelu'):
|
||||
super().__init__()
|
||||
self.c = c
|
||||
self.dropout = dropout
|
||||
self.layer_norm1 = LayerNorm(c)
|
||||
self.self_attn = MultiheadAttention(
|
||||
c, num_heads, self_attention=True, dropout=attention_dropout, bias=False
|
||||
)
|
||||
self.layer_norm2 = LayerNorm(c)
|
||||
self.encoder_attn = MultiheadAttention(
|
||||
c, num_heads, encoder_decoder_attention=True, dropout=attention_dropout, bias=False,
|
||||
)
|
||||
self.layer_norm3 = LayerNorm(c)
|
||||
self.ffn = TransformerFFNLayer(
|
||||
c, 4 * c, padding='LEFT', kernel_size=kernel_size, dropout=relu_dropout, act=act)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
encoder_out=None,
|
||||
encoder_padding_mask=None,
|
||||
incremental_state=None,
|
||||
self_attn_mask=None,
|
||||
self_attn_padding_mask=None,
|
||||
attn_out=None,
|
||||
reset_attn_weight=None,
|
||||
**kwargs,
|
||||
):
|
||||
layer_norm_training = kwargs.get('layer_norm_training', None)
|
||||
if layer_norm_training is not None:
|
||||
self.layer_norm1.training = layer_norm_training
|
||||
self.layer_norm2.training = layer_norm_training
|
||||
self.layer_norm3.training = layer_norm_training
|
||||
residual = x
|
||||
x = self.layer_norm1(x)
|
||||
x, _ = self.self_attn(
|
||||
query=x,
|
||||
key=x,
|
||||
value=x,
|
||||
key_padding_mask=self_attn_padding_mask,
|
||||
incremental_state=incremental_state,
|
||||
attn_mask=self_attn_mask
|
||||
)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
|
||||
attn_logits = None
|
||||
if encoder_out is not None or attn_out is not None:
|
||||
residual = x
|
||||
x = self.layer_norm2(x)
|
||||
if encoder_out is not None:
|
||||
x, attn = self.encoder_attn(
|
||||
query=x,
|
||||
key=encoder_out,
|
||||
value=encoder_out,
|
||||
key_padding_mask=encoder_padding_mask,
|
||||
incremental_state=incremental_state,
|
||||
static_kv=True,
|
||||
enc_dec_attn_constraint_mask=get_incremental_state(self, incremental_state,
|
||||
'enc_dec_attn_constraint_mask'),
|
||||
reset_attn_weight=reset_attn_weight
|
||||
)
|
||||
attn_logits = attn[1]
|
||||
elif attn_out is not None:
|
||||
x = self.encoder_attn.in_proj_v(attn_out)
|
||||
if encoder_out is not None or attn_out is not None:
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
|
||||
residual = x
|
||||
x = self.layer_norm3(x)
|
||||
x = self.ffn(x, incremental_state=incremental_state)
|
||||
x = F.dropout(x, self.dropout, training=self.training)
|
||||
x = residual + x
|
||||
return x, attn_logits
|
||||
|
||||
def clear_buffer(self, input, encoder_out=None, encoder_padding_mask=None, incremental_state=None):
|
||||
self.encoder_attn.clear_buffer(incremental_state)
|
||||
self.ffn.clear_buffer(incremental_state)
|
||||
|
||||
def set_buffer(self, name, tensor, incremental_state):
|
||||
return set_incremental_state(self, incremental_state, name, tensor)
|
||||
|
||||
|
||||
class TransformerEncoderLayer(nn.Module):
|
||||
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
self.op = EncSALayer(
|
||||
hidden_size, num_heads, dropout=dropout,
|
||||
attention_dropout=0.0, relu_dropout=dropout,
|
||||
kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
return self.op(x, **kwargs)
|
||||
|
||||
|
||||
class TransformerDecoderLayer(nn.Module):
|
||||
def __init__(self, hidden_size, dropout, kernel_size=9, num_heads=2):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.num_heads = num_heads
|
||||
self.op = DecSALayer(
|
||||
hidden_size, num_heads, dropout=dropout,
|
||||
attention_dropout=0.0, relu_dropout=dropout,
|
||||
kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
return self.op(x, **kwargs)
|
||||
|
||||
def clear_buffer(self, *args):
|
||||
return self.op.clear_buffer(*args)
|
||||
|
||||
def set_buffer(self, *args):
|
||||
return self.op.set_buffer(*args)
|
||||
|
||||
|
||||
class FFTBlocks(nn.Module):
|
||||
def __init__(self, hidden_size, num_layers, ffn_kernel_size=9, dropout=0.0,
|
||||
num_heads=2, use_pos_embed=True, use_last_norm=True,
|
||||
use_pos_embed_alpha=True):
|
||||
super().__init__()
|
||||
self.num_layers = num_layers
|
||||
embed_dim = self.hidden_size = hidden_size
|
||||
self.dropout = dropout
|
||||
self.use_pos_embed = use_pos_embed
|
||||
self.use_last_norm = use_last_norm
|
||||
if use_pos_embed:
|
||||
self.max_source_positions = DEFAULT_MAX_TARGET_POSITIONS
|
||||
self.padding_idx = 0
|
||||
self.pos_embed_alpha = nn.Parameter(torch.Tensor([1])) if use_pos_embed_alpha else 1
|
||||
self.embed_positions = SinusoidalPositionalEmbedding(
|
||||
embed_dim, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
self.layers.extend([
|
||||
TransformerEncoderLayer(self.hidden_size, self.dropout,
|
||||
kernel_size=ffn_kernel_size, num_heads=num_heads)
|
||||
for _ in range(self.num_layers)
|
||||
])
|
||||
if self.use_last_norm:
|
||||
self.layer_norm = nn.LayerNorm(embed_dim)
|
||||
else:
|
||||
self.layer_norm = None
|
||||
|
||||
def forward(self, x, padding_mask=None, attn_mask=None, return_hiddens=False):
|
||||
"""
|
||||
:param x: [B, T, C]
|
||||
:param padding_mask: [B, T]
|
||||
:return: [B, T, C] or [L, B, T, C]
|
||||
"""
|
||||
padding_mask = x.abs().sum(-1).eq(0).data if padding_mask is None else padding_mask
|
||||
nonpadding_mask_TB = 1 - padding_mask.transpose(0, 1).float()[:, :, None] # [T, B, 1]
|
||||
if self.use_pos_embed:
|
||||
positions = self.pos_embed_alpha * self.embed_positions(x[..., 0])
|
||||
x = x + positions
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
# B x T x C -> T x B x C
|
||||
x = x.transpose(0, 1) * nonpadding_mask_TB
|
||||
hiddens = []
|
||||
for layer in self.layers:
|
||||
x = layer(x, encoder_padding_mask=padding_mask, attn_mask=attn_mask) * nonpadding_mask_TB
|
||||
hiddens.append(x)
|
||||
if self.use_last_norm:
|
||||
x = self.layer_norm(x) * nonpadding_mask_TB
|
||||
if return_hiddens:
|
||||
x = torch.stack(hiddens, 0) # [L, T, B, C]
|
||||
x = x.transpose(1, 2) # [L, B, T, C]
|
||||
else:
|
||||
x = x.transpose(0, 1) # [B, T, C]
|
||||
return x
|
||||
|
||||
|
||||
class FastSpeechEncoder(FFTBlocks):
|
||||
def __init__(self, dict_size, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2,
|
||||
dropout=0.0):
|
||||
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads,
|
||||
use_pos_embed=False, dropout=dropout) # use_pos_embed_alpha for compatibility
|
||||
self.embed_tokens = Embedding(dict_size, hidden_size, 0)
|
||||
self.embed_scale = math.sqrt(hidden_size)
|
||||
self.padding_idx = 0
|
||||
self.embed_positions = SinusoidalPositionalEmbedding(
|
||||
hidden_size, self.padding_idx, init_size=DEFAULT_MAX_TARGET_POSITIONS,
|
||||
)
|
||||
|
||||
def forward(self, txt_tokens, attn_mask=None):
|
||||
"""
|
||||
|
||||
:param txt_tokens: [B, T]
|
||||
:return: {
|
||||
'encoder_out': [B x T x C]
|
||||
}
|
||||
"""
|
||||
encoder_padding_mask = txt_tokens.eq(self.padding_idx).data
|
||||
x = self.forward_embedding(txt_tokens) # [B, T, H]
|
||||
if self.num_layers > 0:
|
||||
x = super(FastSpeechEncoder, self).forward(x, encoder_padding_mask, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
def forward_embedding(self, txt_tokens):
|
||||
# embed tokens and positions
|
||||
x = self.embed_scale * self.embed_tokens(txt_tokens)
|
||||
positions = self.embed_positions(txt_tokens)
|
||||
x = x + positions
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
return x
|
||||
|
||||
|
||||
class FastSpeechDecoder(FFTBlocks):
|
||||
def __init__(self, hidden_size=256, num_layers=4, kernel_size=9, num_heads=2):
|
||||
super().__init__(hidden_size, num_layers, kernel_size, num_heads=num_heads)
|
||||
@@ -0,0 +1,109 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from packaging import version
|
||||
|
||||
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
|
||||
n_channels_int = n_channels[0]
|
||||
in_act = input_a + input_b
|
||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||
acts = t_act * s_act
|
||||
return acts
|
||||
|
||||
jit_fused_add_tanh_sigmoid_multiply = fused_add_tanh_sigmoid_multiply
|
||||
|
||||
def script_function():
|
||||
if version.parse(torch.__version__) >= version.parse('2.0'):
|
||||
global jit_fused_add_tanh_sigmoid_multiply
|
||||
jit_fused_add_tanh_sigmoid_multiply = torch.jit.script(fused_add_tanh_sigmoid_multiply)
|
||||
|
||||
|
||||
class WN(torch.nn.Module):
|
||||
def __init__(self, hidden_size, kernel_size, dilation_rate, n_layers, c_cond=0,
|
||||
p_dropout=0, share_cond_layers=False, is_BTC=False):
|
||||
super(WN, self).__init__()
|
||||
assert (kernel_size % 2 == 1)
|
||||
assert (hidden_size % 2 == 0)
|
||||
self.is_BTC = is_BTC
|
||||
self.hidden_size = hidden_size
|
||||
self.kernel_size = kernel_size
|
||||
self.dilation_rate = dilation_rate
|
||||
self.n_layers = n_layers
|
||||
self.gin_channels = c_cond
|
||||
self.p_dropout = p_dropout
|
||||
self.share_cond_layers = share_cond_layers
|
||||
|
||||
self.in_layers = torch.nn.ModuleList()
|
||||
self.res_skip_layers = torch.nn.ModuleList()
|
||||
self.drop = nn.Dropout(p_dropout)
|
||||
|
||||
if c_cond != 0 and not share_cond_layers:
|
||||
cond_layer = torch.nn.Conv1d(c_cond, 2 * hidden_size * n_layers, 1)
|
||||
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name='weight')
|
||||
|
||||
for i in range(n_layers):
|
||||
dilation = dilation_rate ** i
|
||||
padding = int((kernel_size * dilation - dilation) / 2)
|
||||
in_layer = torch.nn.Conv1d(hidden_size, 2 * hidden_size, kernel_size,
|
||||
dilation=dilation, padding=padding)
|
||||
in_layer = torch.nn.utils.weight_norm(in_layer, name='weight')
|
||||
self.in_layers.append(in_layer)
|
||||
|
||||
# last one is not necessary
|
||||
if i < n_layers - 1:
|
||||
res_skip_channels = 2 * hidden_size
|
||||
else:
|
||||
res_skip_channels = hidden_size
|
||||
|
||||
res_skip_layer = torch.nn.Conv1d(hidden_size, res_skip_channels, 1)
|
||||
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name='weight')
|
||||
self.res_skip_layers.append(res_skip_layer)
|
||||
|
||||
script_function()
|
||||
|
||||
def forward(self, x, nonpadding=None, cond=None):
|
||||
if self.is_BTC:
|
||||
x = x.transpose(1, 2)
|
||||
cond = cond.transpose(1, 2) if cond is not None else None
|
||||
nonpadding = nonpadding.transpose(1, 2) if nonpadding is not None else None
|
||||
if nonpadding is None:
|
||||
nonpadding = 1
|
||||
output = torch.zeros_like(x)
|
||||
n_channels_tensor = torch.IntTensor([self.hidden_size])
|
||||
|
||||
if cond is not None and not self.share_cond_layers:
|
||||
cond = self.cond_layer(cond)
|
||||
|
||||
for i in range(self.n_layers):
|
||||
x_in = self.in_layers[i](x)
|
||||
x_in = self.drop(x_in)
|
||||
if cond is not None:
|
||||
cond_offset = i * 2 * self.hidden_size
|
||||
cond_l = cond[:, cond_offset:cond_offset + 2 * self.hidden_size, :]
|
||||
else:
|
||||
cond_l = torch.zeros_like(x_in)
|
||||
|
||||
if version.parse(torch.__version__) >= version.parse('2.0'):
|
||||
acts = jit_fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
|
||||
else:
|
||||
acts = fused_add_tanh_sigmoid_multiply(x_in, cond_l, n_channels_tensor)
|
||||
|
||||
res_skip_acts = self.res_skip_layers[i](acts)
|
||||
if i < self.n_layers - 1:
|
||||
x = (x + res_skip_acts[:, :self.hidden_size, :]) * nonpadding
|
||||
output = output + res_skip_acts[:, self.hidden_size:, :]
|
||||
else:
|
||||
output = output + res_skip_acts
|
||||
output = output * nonpadding
|
||||
if self.is_BTC:
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
def remove_weight_norm(self):
|
||||
def remove_weight_norm(m):
|
||||
try:
|
||||
nn.utils.remove_weight_norm(m)
|
||||
except ValueError: # this module didn't have weight norm
|
||||
return
|
||||
|
||||
self.apply(remove_weight_norm)
|
||||
Reference in New Issue
Block a user