Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled
Update API Documentation / build-api-docs (push) Has been cancelled
Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Unility functions for Transformer."""
|
||||
|
||||
import torch
|
||||
from funasr.models.transformer.utils.nets_utils import pad_list
|
||||
|
||||
|
||||
def add_sos_eos(ys_pad, sos, eos, ignore_id):
|
||||
"""Add <sos> and <eos> labels.
|
||||
|
||||
:param torch.Tensor ys_pad: batch of padded target sequences (B, Lmax)
|
||||
:param int sos: index of <sos>
|
||||
:param int eos: index of <eos>
|
||||
:param int ignore_id: index of padding
|
||||
:return: padded tensor (B, Lmax)
|
||||
:rtype: torch.Tensor
|
||||
:return: padded tensor (B, Lmax)
|
||||
:rtype: torch.Tensor
|
||||
"""
|
||||
|
||||
_sos = ys_pad.new([sos])
|
||||
_eos = ys_pad.new([eos])
|
||||
ys = [y[y != ignore_id] for y in ys_pad] # parse padded ys
|
||||
ys_in = [torch.cat([_sos, y], dim=0) for y in ys]
|
||||
ys_out = [torch.cat([y, _eos], dim=0) for y in ys]
|
||||
return pad_list(ys_in, eos), pad_list(ys_out, ignore_id)
|
||||
|
||||
def add_sos_and_eos(ys_pad, sos, eos, ignore_id):
|
||||
"""Add <sos> at the beginning and <eos> at the end (length + 2).
|
||||
|
||||
Unlike add_sos_eos which returns (ys_in, ys_out) separately,
|
||||
this returns a single sequence with both sos and eos added.
|
||||
|
||||
:param torch.Tensor ys_pad: batch of padded target sequences (B, Lmax)
|
||||
:param int sos: index of <sos>
|
||||
:param int eos: index of <eos>
|
||||
:param int ignore_id: index of padding
|
||||
:return: ys_in with sos prepended (B, Lmax+1)
|
||||
:return: ys with both sos and eos (B, Lmax+2)
|
||||
"""
|
||||
_sos = ys_pad.new([sos])
|
||||
_eos = ys_pad.new([eos])
|
||||
ys = [y[y != ignore_id] for y in ys_pad]
|
||||
ys_in = [torch.cat([_sos, y], dim=0) for y in ys]
|
||||
ys_both = [torch.cat([_sos, y, _eos], dim=0) for y in ys]
|
||||
return pad_list(ys_in, eos), pad_list(ys_both, ignore_id)
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Dynamic Convolution module."""
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
MIN_VALUE = float(numpy.finfo(numpy.float32).min)
|
||||
|
||||
|
||||
class DynamicConvolution(nn.Module):
|
||||
"""Dynamic Convolution layer.
|
||||
|
||||
This implementation is based on
|
||||
https://github.com/pytorch/fairseq/tree/master/fairseq
|
||||
|
||||
Args:
|
||||
wshare (int): the number of kernel of convolution
|
||||
n_feat (int): the number of features
|
||||
dropout_rate (float): dropout_rate
|
||||
kernel_size (int): kernel size (length)
|
||||
use_kernel_mask (bool): Use causal mask or not for convolution kernel
|
||||
use_bias (bool): Use bias term or not.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
wshare,
|
||||
n_feat,
|
||||
dropout_rate,
|
||||
kernel_size,
|
||||
use_kernel_mask=False,
|
||||
use_bias=False,
|
||||
):
|
||||
"""Construct Dynamic Convolution layer."""
|
||||
super(DynamicConvolution, self).__init__()
|
||||
|
||||
assert n_feat % wshare == 0
|
||||
self.wshare = wshare
|
||||
self.use_kernel_mask = use_kernel_mask
|
||||
self.dropout_rate = dropout_rate
|
||||
self.kernel_size = kernel_size
|
||||
self.attn = None
|
||||
|
||||
# linear -> GLU -- -> lightconv -> linear
|
||||
# \ /
|
||||
# Linear
|
||||
self.linear1 = nn.Linear(n_feat, n_feat * 2)
|
||||
self.linear2 = nn.Linear(n_feat, n_feat)
|
||||
self.linear_weight = nn.Linear(n_feat, self.wshare * 1 * kernel_size)
|
||||
nn.init.xavier_uniform(self.linear_weight.weight)
|
||||
self.act = nn.GLU()
|
||||
|
||||
# dynamic conv related
|
||||
self.use_bias = use_bias
|
||||
if self.use_bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(n_feat))
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Forward of 'Dynamic Convolution'.
|
||||
|
||||
This function takes query, key and value but uses only quert.
|
||||
This is just for compatibility with self-attention layer (attention.py)
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): (batch, time1, d_model) input tensor
|
||||
key (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
value (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
mask (torch.Tensor): (batch, time1, time2) mask
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): (batch, time1, d_model) output
|
||||
|
||||
"""
|
||||
# linear -> GLU -- -> lightconv -> linear
|
||||
# \ /
|
||||
# Linear
|
||||
x = query
|
||||
B, T, C = x.size()
|
||||
H = self.wshare
|
||||
k = self.kernel_size
|
||||
|
||||
# first liner layer
|
||||
x = self.linear1(x)
|
||||
|
||||
# GLU activation
|
||||
x = self.act(x)
|
||||
|
||||
# get kernel of convolution
|
||||
weight = self.linear_weight(x) # B x T x kH
|
||||
weight = F.dropout(weight, self.dropout_rate, training=self.training)
|
||||
weight = weight.view(B, T, H, k).transpose(1, 2).contiguous() # B x H x T x k
|
||||
weight_new = torch.zeros(B * H * T * (T + k - 1), dtype=weight.dtype)
|
||||
weight_new = weight_new.view(B, H, T, T + k - 1).fill_(float("-inf"))
|
||||
weight_new = weight_new.to(x.device) # B x H x T x T+k-1
|
||||
weight_new.as_strided((B, H, T, k), ((T + k - 1) * T * H, (T + k - 1) * T, T + k, 1)).copy_(
|
||||
weight
|
||||
)
|
||||
weight_new = weight_new.narrow(-1, int((k - 1) / 2), T) # B x H x T x T(k)
|
||||
if self.use_kernel_mask:
|
||||
kernel_mask = torch.tril(torch.ones(T, T, device=x.device)).unsqueeze(0)
|
||||
weight_new = weight_new.masked_fill(kernel_mask == 0.0, float("-inf"))
|
||||
weight_new = F.softmax(weight_new, dim=-1)
|
||||
self.attn = weight_new
|
||||
weight_new = weight_new.view(B * H, T, T)
|
||||
|
||||
# convolution
|
||||
x = x.transpose(1, 2).contiguous() # B x C x T
|
||||
x = x.view(B * H, int(C / H), T).transpose(1, 2)
|
||||
x = torch.bmm(weight_new, x) # BH x T x C/H
|
||||
x = x.transpose(1, 2).contiguous().view(B, C, T)
|
||||
|
||||
if self.use_bias:
|
||||
x = x + self.bias.view(1, -1, 1)
|
||||
x = x.transpose(1, 2) # B x T x C
|
||||
|
||||
if mask is not None and not self.use_kernel_mask:
|
||||
mask = mask.transpose(-1, -2)
|
||||
x = x.masked_fill(mask == 0, 0.0)
|
||||
|
||||
# second linear layer
|
||||
x = self.linear2(x)
|
||||
return x
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Dynamic 2-Dimensional Convolution module."""
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
MIN_VALUE = float(numpy.finfo(numpy.float32).min)
|
||||
|
||||
|
||||
class DynamicConvolution2D(nn.Module):
|
||||
"""Dynamic 2-Dimensional Convolution layer.
|
||||
|
||||
This implementation is based on
|
||||
https://github.com/pytorch/fairseq/tree/master/fairseq
|
||||
|
||||
Args:
|
||||
wshare (int): the number of kernel of convolution
|
||||
n_feat (int): the number of features
|
||||
dropout_rate (float): dropout_rate
|
||||
kernel_size (int): kernel size (length)
|
||||
use_kernel_mask (bool): Use causal mask or not for convolution kernel
|
||||
use_bias (bool): Use bias term or not.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
wshare,
|
||||
n_feat,
|
||||
dropout_rate,
|
||||
kernel_size,
|
||||
use_kernel_mask=False,
|
||||
use_bias=False,
|
||||
):
|
||||
"""Construct Dynamic 2-Dimensional Convolution layer."""
|
||||
super(DynamicConvolution2D, self).__init__()
|
||||
|
||||
assert n_feat % wshare == 0
|
||||
self.wshare = wshare
|
||||
self.use_kernel_mask = use_kernel_mask
|
||||
self.dropout_rate = dropout_rate
|
||||
self.kernel_size = kernel_size
|
||||
self.padding_size = int(kernel_size / 2)
|
||||
self.attn_t = None
|
||||
self.attn_f = None
|
||||
|
||||
# linear -> GLU -- -> lightconv -> linear
|
||||
# \ /
|
||||
# Linear
|
||||
self.linear1 = nn.Linear(n_feat, n_feat * 2)
|
||||
self.linear2 = nn.Linear(n_feat * 2, n_feat)
|
||||
self.linear_weight = nn.Linear(n_feat, self.wshare * 1 * kernel_size)
|
||||
nn.init.xavier_uniform(self.linear_weight.weight)
|
||||
self.linear_weight_f = nn.Linear(n_feat, kernel_size)
|
||||
nn.init.xavier_uniform(self.linear_weight_f.weight)
|
||||
self.act = nn.GLU()
|
||||
|
||||
# dynamic conv related
|
||||
self.use_bias = use_bias
|
||||
if self.use_bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(n_feat))
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Forward of 'Dynamic 2-Dimensional Convolution'.
|
||||
|
||||
This function takes query, key and value but uses only query.
|
||||
This is just for compatibility with self-attention layer (attention.py)
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): (batch, time1, d_model) input tensor
|
||||
key (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
value (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
mask (torch.Tensor): (batch, time1, time2) mask
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): (batch, time1, d_model) output
|
||||
|
||||
"""
|
||||
# linear -> GLU -- -> lightconv -> linear
|
||||
# \ /
|
||||
# Linear
|
||||
x = query
|
||||
B, T, C = x.size()
|
||||
H = self.wshare
|
||||
k = self.kernel_size
|
||||
|
||||
# first liner layer
|
||||
x = self.linear1(x)
|
||||
|
||||
# GLU activation
|
||||
x = self.act(x)
|
||||
|
||||
# convolution of frequency axis
|
||||
weight_f = self.linear_weight_f(x).view(B * T, 1, k) # B x T x k
|
||||
self.attn_f = weight_f.view(B, T, k).unsqueeze(1)
|
||||
xf = F.conv1d(x.view(1, B * T, C), weight_f, padding=self.padding_size, groups=B * T)
|
||||
xf = xf.view(B, T, C)
|
||||
|
||||
# get kernel of convolution
|
||||
weight = self.linear_weight(x) # B x T x kH
|
||||
weight = F.dropout(weight, self.dropout_rate, training=self.training)
|
||||
weight = weight.view(B, T, H, k).transpose(1, 2).contiguous() # B x H x T x k
|
||||
weight_new = torch.zeros(B * H * T * (T + k - 1), dtype=weight.dtype)
|
||||
weight_new = weight_new.view(B, H, T, T + k - 1).fill_(float("-inf"))
|
||||
weight_new = weight_new.to(x.device) # B x H x T x T+k-1
|
||||
weight_new.as_strided((B, H, T, k), ((T + k - 1) * T * H, (T + k - 1) * T, T + k, 1)).copy_(
|
||||
weight
|
||||
)
|
||||
weight_new = weight_new.narrow(-1, int((k - 1) / 2), T) # B x H x T x T(k)
|
||||
if self.use_kernel_mask:
|
||||
kernel_mask = torch.tril(torch.ones(T, T, device=x.device)).unsqueeze(0)
|
||||
weight_new = weight_new.masked_fill(kernel_mask == 0.0, float("-inf"))
|
||||
weight_new = F.softmax(weight_new, dim=-1)
|
||||
self.attn_t = weight_new
|
||||
weight_new = weight_new.view(B * H, T, T)
|
||||
|
||||
# convolution
|
||||
x = x.transpose(1, 2).contiguous() # B x C x T
|
||||
x = x.view(B * H, int(C / H), T).transpose(1, 2)
|
||||
x = torch.bmm(weight_new, x)
|
||||
x = x.transpose(1, 2).contiguous().view(B, C, T)
|
||||
|
||||
if self.use_bias:
|
||||
x = x + self.bias.view(1, -1, 1)
|
||||
x = x.transpose(1, 2) # B x T x C
|
||||
x = torch.cat((x, xf), -1) # B x T x Cx2
|
||||
|
||||
if mask is not None and not self.use_kernel_mask:
|
||||
mask = mask.transpose(-1, -2)
|
||||
x = x.masked_fill(mask == 0, 0.0)
|
||||
|
||||
# second linear layer
|
||||
x = self.linear2(x)
|
||||
return x
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Lightweight Convolution Module."""
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
MIN_VALUE = float(numpy.finfo(numpy.float32).min)
|
||||
|
||||
|
||||
class LightweightConvolution(nn.Module):
|
||||
"""Lightweight Convolution layer.
|
||||
|
||||
This implementation is based on
|
||||
https://github.com/pytorch/fairseq/tree/master/fairseq
|
||||
|
||||
Args:
|
||||
wshare (int): the number of kernel of convolution
|
||||
n_feat (int): the number of features
|
||||
dropout_rate (float): dropout_rate
|
||||
kernel_size (int): kernel size (length)
|
||||
use_kernel_mask (bool): Use causal mask or not for convolution kernel
|
||||
use_bias (bool): Use bias term or not.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
wshare,
|
||||
n_feat,
|
||||
dropout_rate,
|
||||
kernel_size,
|
||||
use_kernel_mask=False,
|
||||
use_bias=False,
|
||||
):
|
||||
"""Construct Lightweight Convolution layer."""
|
||||
super(LightweightConvolution, self).__init__()
|
||||
|
||||
assert n_feat % wshare == 0
|
||||
self.wshare = wshare
|
||||
self.use_kernel_mask = use_kernel_mask
|
||||
self.dropout_rate = dropout_rate
|
||||
self.kernel_size = kernel_size
|
||||
self.padding_size = int(kernel_size / 2)
|
||||
|
||||
# linear -> GLU -> lightconv -> linear
|
||||
self.linear1 = nn.Linear(n_feat, n_feat * 2)
|
||||
self.linear2 = nn.Linear(n_feat, n_feat)
|
||||
self.act = nn.GLU()
|
||||
|
||||
# lightconv related
|
||||
self.weight = nn.Parameter(torch.Tensor(self.wshare, 1, kernel_size).uniform_(0, 1))
|
||||
self.use_bias = use_bias
|
||||
if self.use_bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(n_feat))
|
||||
|
||||
# mask of kernel
|
||||
kernel_mask0 = torch.zeros(self.wshare, int(kernel_size / 2))
|
||||
kernel_mask1 = torch.ones(self.wshare, int(kernel_size / 2 + 1))
|
||||
self.kernel_mask = torch.cat((kernel_mask1, kernel_mask0), dim=-1).unsqueeze(1)
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Forward of 'Lightweight Convolution'.
|
||||
|
||||
This function takes query, key and value but uses only query.
|
||||
This is just for compatibility with self-attention layer (attention.py)
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): (batch, time1, d_model) input tensor
|
||||
key (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
value (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
mask (torch.Tensor): (batch, time1, time2) mask
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): (batch, time1, d_model) output
|
||||
|
||||
"""
|
||||
# linear -> GLU -> lightconv -> linear
|
||||
x = query
|
||||
B, T, C = x.size()
|
||||
H = self.wshare
|
||||
|
||||
# first liner layer
|
||||
x = self.linear1(x)
|
||||
|
||||
# GLU activation
|
||||
x = self.act(x)
|
||||
|
||||
# lightconv
|
||||
x = x.transpose(1, 2).contiguous().view(-1, H, T) # B x C x T
|
||||
weight = F.dropout(self.weight, self.dropout_rate, training=self.training)
|
||||
if self.use_kernel_mask:
|
||||
self.kernel_mask = self.kernel_mask.to(x.device)
|
||||
weight = weight.masked_fill(self.kernel_mask == 0.0, float("-inf"))
|
||||
weight = F.softmax(weight, dim=-1)
|
||||
x = F.conv1d(x, weight, padding=self.padding_size, groups=self.wshare).view(B, C, T)
|
||||
if self.use_bias:
|
||||
x = x + self.bias.view(1, -1, 1)
|
||||
x = x.transpose(1, 2) # B x T x C
|
||||
|
||||
if mask is not None and not self.use_kernel_mask:
|
||||
mask = mask.transpose(-1, -2)
|
||||
x = x.masked_fill(mask == 0, 0.0)
|
||||
|
||||
# second linear layer
|
||||
x = self.linear2(x)
|
||||
return x
|
||||
@@ -0,0 +1,120 @@
|
||||
"""Lightweight 2-Dimensional Convolution module."""
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
MIN_VALUE = float(numpy.finfo(numpy.float32).min)
|
||||
|
||||
|
||||
class LightweightConvolution2D(nn.Module):
|
||||
"""Lightweight 2-Dimensional Convolution layer.
|
||||
|
||||
This implementation is based on
|
||||
https://github.com/pytorch/fairseq/tree/master/fairseq
|
||||
|
||||
Args:
|
||||
wshare (int): the number of kernel of convolution
|
||||
n_feat (int): the number of features
|
||||
dropout_rate (float): dropout_rate
|
||||
kernel_size (int): kernel size (length)
|
||||
use_kernel_mask (bool): Use causal mask or not for convolution kernel
|
||||
use_bias (bool): Use bias term or not.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
wshare,
|
||||
n_feat,
|
||||
dropout_rate,
|
||||
kernel_size,
|
||||
use_kernel_mask=False,
|
||||
use_bias=False,
|
||||
):
|
||||
"""Construct Lightweight 2-Dimensional Convolution layer."""
|
||||
super(LightweightConvolution2D, self).__init__()
|
||||
|
||||
assert n_feat % wshare == 0
|
||||
self.wshare = wshare
|
||||
self.use_kernel_mask = use_kernel_mask
|
||||
self.dropout_rate = dropout_rate
|
||||
self.kernel_size = kernel_size
|
||||
self.padding_size = int(kernel_size / 2)
|
||||
|
||||
# linear -> GLU -> lightconv -> linear
|
||||
self.linear1 = nn.Linear(n_feat, n_feat * 2)
|
||||
self.linear2 = nn.Linear(n_feat * 2, n_feat)
|
||||
self.act = nn.GLU()
|
||||
|
||||
# lightconv related
|
||||
self.weight = nn.Parameter(torch.Tensor(self.wshare, 1, kernel_size).uniform_(0, 1))
|
||||
self.weight_f = nn.Parameter(torch.Tensor(1, 1, kernel_size).uniform_(0, 1))
|
||||
self.use_bias = use_bias
|
||||
if self.use_bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(n_feat))
|
||||
|
||||
# mask of kernel
|
||||
kernel_mask0 = torch.zeros(self.wshare, int(kernel_size / 2))
|
||||
kernel_mask1 = torch.ones(self.wshare, int(kernel_size / 2 + 1))
|
||||
self.kernel_mask = torch.cat((kernel_mask1, kernel_mask0), dim=-1).unsqueeze(1)
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Forward of 'Lightweight 2-Dimensional Convolution'.
|
||||
|
||||
This function takes query, key and value but uses only query.
|
||||
This is just for compatibility with self-attention layer (attention.py)
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): (batch, time1, d_model) input tensor
|
||||
key (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
value (torch.Tensor): (batch, time2, d_model) NOT USED
|
||||
mask (torch.Tensor): (batch, time1, time2) mask
|
||||
|
||||
Return:
|
||||
x (torch.Tensor): (batch, time1, d_model) output
|
||||
|
||||
"""
|
||||
# linear -> GLU -> lightconv -> linear
|
||||
x = query
|
||||
B, T, C = x.size()
|
||||
H = self.wshare
|
||||
|
||||
# first liner layer
|
||||
x = self.linear1(x)
|
||||
|
||||
# GLU activation
|
||||
x = self.act(x)
|
||||
|
||||
# convolution along frequency axis
|
||||
weight_f = F.softmax(self.weight_f, dim=-1)
|
||||
weight_f = F.dropout(weight_f, self.dropout_rate, training=self.training)
|
||||
weight_new = torch.zeros(B * T, 1, self.kernel_size, device=x.device, dtype=x.dtype).copy_(
|
||||
weight_f
|
||||
)
|
||||
xf = F.conv1d(
|
||||
x.view(1, B * T, C), weight_new, padding=self.padding_size, groups=B * T
|
||||
).view(B, T, C)
|
||||
|
||||
# lightconv
|
||||
x = x.transpose(1, 2).contiguous().view(-1, H, T) # B x C x T
|
||||
weight = F.dropout(self.weight, self.dropout_rate, training=self.training)
|
||||
if self.use_kernel_mask:
|
||||
self.kernel_mask = self.kernel_mask.to(x.device)
|
||||
weight = weight.masked_fill(self.kernel_mask == 0.0, float("-inf"))
|
||||
weight = F.softmax(weight, dim=-1)
|
||||
x = F.conv1d(x, weight, padding=self.padding_size, groups=self.wshare).view(B, C, T)
|
||||
if self.use_bias:
|
||||
x = x + self.bias.view(1, -1, 1)
|
||||
x = x.transpose(1, 2) # B x T x C
|
||||
x = torch.cat((x, xf), -1) # B x T x Cx2
|
||||
|
||||
if mask is not None and not self.use_kernel_mask:
|
||||
mask = mask.transpose(-1, -2)
|
||||
x = x.masked_fill(mask == 0, 0.0)
|
||||
|
||||
# second linear layer
|
||||
x = self.linear2(x)
|
||||
return x
|
||||
@@ -0,0 +1,52 @@
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Mask module."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def subsequent_mask(size, device="cpu", dtype=torch.bool):
|
||||
"""Create mask for subsequent steps (size, size).
|
||||
|
||||
:param int size: size of mask
|
||||
:param str device: "cpu" or "cuda" or torch.Tensor.device
|
||||
:param torch.dtype dtype: result dtype
|
||||
:rtype: torch.Tensor
|
||||
>>> subsequent_mask(3)
|
||||
[[1, 0, 0],
|
||||
[1, 1, 0],
|
||||
[1, 1, 1]]
|
||||
"""
|
||||
ret = torch.ones(size, size, device=device, dtype=dtype)
|
||||
return torch.tril(ret, out=ret)
|
||||
|
||||
|
||||
def target_mask(ys_in_pad, ignore_id):
|
||||
"""Create mask for decoder self-attention.
|
||||
|
||||
:param torch.Tensor ys_pad: batch of padded target sequences (B, Lmax)
|
||||
:param int ignore_id: index of padding
|
||||
:param torch.dtype dtype: result dtype
|
||||
:rtype: torch.Tensor (B, Lmax, Lmax)
|
||||
"""
|
||||
ys_mask = ys_in_pad != ignore_id
|
||||
m = subsequent_mask(ys_mask.size(-1), device=ys_mask.device).unsqueeze(0)
|
||||
return ys_mask.unsqueeze(-2) & m
|
||||
|
||||
|
||||
def vad_mask(size, vad_pos, device="cpu", dtype=torch.bool):
|
||||
"""Create mask for decoder self-attention.
|
||||
|
||||
:param int size: size of mask
|
||||
:param int vad_pos: index of vad index
|
||||
:param str device: "cpu" or "cuda" or torch.Tensor.device
|
||||
:param torch.dtype dtype: result dtype
|
||||
:rtype: torch.Tensor (B, Lmax, Lmax)
|
||||
"""
|
||||
ret = torch.ones(size, size, device=device, dtype=dtype)
|
||||
if vad_pos <= 0 or vad_pos >= size:
|
||||
return ret
|
||||
sub_corner = torch.zeros(vad_pos - 1, size - vad_pos, device=device, dtype=dtype)
|
||||
ret[0 : vad_pos - 1, vad_pos:] = sub_corner
|
||||
return ret
|
||||
@@ -0,0 +1,157 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Tomoki Hayashi
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Layer modules for FFT block in FastSpeech (Feed-forward Transformer)."""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
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 FsmnFeedForward(torch.nn.Module):
|
||||
"""Position-wise feed forward for FSMN blocks.
|
||||
|
||||
This is a module of multi-leyered conv1d designed
|
||||
to replace position-wise feed-forward network
|
||||
in FSMN block.
|
||||
"""
|
||||
|
||||
def __init__(self, in_chans, hidden_chans, out_chans, kernel_size, dropout_rate):
|
||||
"""Initialize FsmnFeedForward module.
|
||||
|
||||
Args:
|
||||
in_chans (int): Number of input channels.
|
||||
hidden_chans (int): Number of hidden channels.
|
||||
out_chans (int): Number of output channels.
|
||||
kernel_size (int): Kernel size of conv1d.
|
||||
dropout_rate (float): Dropout rate.
|
||||
|
||||
"""
|
||||
super(FsmnFeedForward, 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,
|
||||
out_chans,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
bias=False,
|
||||
)
|
||||
self.norm = torch.nn.LayerNorm(hidden_chans)
|
||||
self.dropout = torch.nn.Dropout(dropout_rate)
|
||||
|
||||
def forward(self, x, ilens=None):
|
||||
"""Calculate forward propagation.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Batch of input tensors (B, T, in_chans).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Batch of output tensors (B, T, out_chans).
|
||||
|
||||
"""
|
||||
x = torch.relu(self.w_1(x.transpose(-1, 1))).transpose(-1, 1)
|
||||
return self.w_2(self.norm(self.dropout(x)).transpose(-1, 1)).transpose(-1, 1), ilens
|
||||
|
||||
|
||||
class Conv1dLinear(torch.nn.Module):
|
||||
"""Conv1D + Linear for Transformer block.
|
||||
|
||||
A variant of MultiLayeredConv1d, which replaces second conv-layer to linear.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, in_chans, hidden_chans, kernel_size, dropout_rate):
|
||||
"""Initialize Conv1dLinear 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(Conv1dLinear, 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.Linear(hidden_chans, in_chans)
|
||||
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))
|
||||
@@ -0,0 +1,740 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""Network related utility tools."""
|
||||
|
||||
import logging
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def to_device(m, x):
|
||||
"""Send tensor into the device of the module.
|
||||
|
||||
Args:
|
||||
m (torch.nn.Module): Torch module.
|
||||
x (Tensor): Torch tensor.
|
||||
|
||||
Returns:
|
||||
Tensor: Torch tensor located in the same place as torch module.
|
||||
|
||||
"""
|
||||
if isinstance(m, torch.nn.Module):
|
||||
device = next(m.parameters()).device
|
||||
elif isinstance(m, torch.Tensor):
|
||||
device = m.device
|
||||
else:
|
||||
raise TypeError("Expected torch.nn.Module or torch.tensor, " f"bot got: {type(m)}")
|
||||
return x.to(device)
|
||||
|
||||
|
||||
def pad_list(xs, pad_value):
|
||||
"""Perform padding for the list of tensors.
|
||||
|
||||
Args:
|
||||
xs (List): List of Tensors [(T_1, `*`), (T_2, `*`), ..., (T_B, `*`)].
|
||||
pad_value (float): Value for padding.
|
||||
|
||||
Returns:
|
||||
Tensor: Padded tensor (B, Tmax, `*`).
|
||||
|
||||
Examples:
|
||||
>>> x = [torch.ones(4), torch.ones(2), torch.ones(1)]
|
||||
>>> x
|
||||
[tensor([1., 1., 1., 1.]), tensor([1., 1.]), tensor([1.])]
|
||||
>>> pad_list(x, 0)
|
||||
tensor([[1., 1., 1., 1.],
|
||||
[1., 1., 0., 0.],
|
||||
[1., 0., 0., 0.]])
|
||||
|
||||
"""
|
||||
n_batch = len(xs)
|
||||
max_len = max(x.size(0) for x in xs)
|
||||
pad = xs[0].new(n_batch, max_len, *xs[0].size()[1:]).fill_(pad_value)
|
||||
|
||||
for i in range(n_batch):
|
||||
pad[i, : xs[i].size(0)] = xs[i]
|
||||
|
||||
return pad
|
||||
|
||||
|
||||
def pad_list_all_dim(xs, pad_value):
|
||||
"""Perform padding for the list of tensors.
|
||||
|
||||
Args:
|
||||
xs (List): List of Tensors [(T_1, `*`), (T_2, `*`), ..., (T_B, `*`)].
|
||||
pad_value (float): Value for padding.
|
||||
|
||||
Returns:
|
||||
Tensor: Padded tensor (B, Tmax, `*`).
|
||||
|
||||
Examples:
|
||||
>>> x = [torch.ones(4), torch.ones(2), torch.ones(1)]
|
||||
>>> x
|
||||
[tensor([1., 1., 1., 1.]), tensor([1., 1.]), tensor([1.])]
|
||||
>>> pad_list(x, 0)
|
||||
tensor([[1., 1., 1., 1.],
|
||||
[1., 1., 0., 0.],
|
||||
[1., 0., 0., 0.]])
|
||||
|
||||
"""
|
||||
n_batch = len(xs)
|
||||
num_dim = len(xs[0].shape)
|
||||
max_len_all_dim = []
|
||||
for i in range(num_dim):
|
||||
max_len_all_dim.append(max(x.size(i) for x in xs))
|
||||
pad = xs[0].new(n_batch, *max_len_all_dim).fill_(pad_value)
|
||||
|
||||
for i in range(n_batch):
|
||||
if num_dim == 1:
|
||||
pad[i, : xs[i].size(0)] = xs[i]
|
||||
elif num_dim == 2:
|
||||
pad[i, : xs[i].size(0), : xs[i].size(1)] = xs[i]
|
||||
elif num_dim == 3:
|
||||
pad[i, : xs[i].size(0), : xs[i].size(1), : xs[i].size(2)] = xs[i]
|
||||
else:
|
||||
raise ValueError(
|
||||
"pad_list_all_dim only support 1-D, 2-D and 3-D tensors, not {}-D.".format(num_dim)
|
||||
)
|
||||
|
||||
return pad
|
||||
|
||||
|
||||
def make_pad_mask(lengths, xs=None, length_dim=-1, maxlen=None):
|
||||
"""Make mask tensor containing indices of padded part.
|
||||
|
||||
Args:
|
||||
lengths (LongTensor or List): Batch of lengths (B,).
|
||||
xs (Tensor, optional): The reference tensor.
|
||||
If set, masks will be the same shape as this tensor.
|
||||
length_dim (int, optional): Dimension indicator of the above tensor.
|
||||
See the example.
|
||||
|
||||
Returns:
|
||||
Tensor: Mask tensor containing indices of padded part.
|
||||
dtype=torch.uint8 in PyTorch 1.2-
|
||||
dtype=torch.bool in PyTorch 1.2+ (including 1.2)
|
||||
|
||||
Examples:
|
||||
With only lengths.
|
||||
|
||||
>>> lengths = [5, 3, 2]
|
||||
>>> make_pad_mask(lengths)
|
||||
masks = [[0, 0, 0, 0 ,0],
|
||||
[0, 0, 0, 1, 1],
|
||||
[0, 0, 1, 1, 1]]
|
||||
|
||||
With the reference tensor.
|
||||
|
||||
>>> xs = torch.zeros((3, 2, 4))
|
||||
>>> make_pad_mask(lengths, xs)
|
||||
tensor([[[0, 0, 0, 0],
|
||||
[0, 0, 0, 0]],
|
||||
[[0, 0, 0, 1],
|
||||
[0, 0, 0, 1]],
|
||||
[[0, 0, 1, 1],
|
||||
[0, 0, 1, 1]]], dtype=torch.uint8)
|
||||
>>> xs = torch.zeros((3, 2, 6))
|
||||
>>> make_pad_mask(lengths, xs)
|
||||
tensor([[[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1]],
|
||||
[[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1]],
|
||||
[[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1]]], dtype=torch.uint8)
|
||||
|
||||
With the reference tensor and dimension indicator.
|
||||
|
||||
>>> xs = torch.zeros((3, 6, 6))
|
||||
>>> make_pad_mask(lengths, xs, 1)
|
||||
tensor([[[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[1, 1, 1, 1, 1, 1]],
|
||||
[[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1]],
|
||||
[[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1]]], dtype=torch.uint8)
|
||||
>>> make_pad_mask(lengths, xs, 2)
|
||||
tensor([[[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1],
|
||||
[0, 0, 0, 0, 0, 1]],
|
||||
[[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1],
|
||||
[0, 0, 0, 1, 1, 1]],
|
||||
[[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1],
|
||||
[0, 0, 1, 1, 1, 1]]], dtype=torch.uint8)
|
||||
|
||||
"""
|
||||
if length_dim == 0:
|
||||
raise ValueError("length_dim cannot be 0: {}".format(length_dim))
|
||||
|
||||
if not isinstance(lengths, list):
|
||||
lengths = lengths.tolist()
|
||||
bs = int(len(lengths))
|
||||
if maxlen is None:
|
||||
if xs is None:
|
||||
maxlen = int(max(lengths))
|
||||
else:
|
||||
maxlen = xs.size(length_dim)
|
||||
else:
|
||||
assert xs is None
|
||||
assert maxlen >= int(max(lengths))
|
||||
|
||||
seq_range = torch.arange(0, maxlen, dtype=torch.int64)
|
||||
seq_range_expand = seq_range.unsqueeze(0).expand(bs, maxlen)
|
||||
seq_length_expand = seq_range_expand.new(lengths).unsqueeze(-1)
|
||||
mask = seq_range_expand >= seq_length_expand
|
||||
|
||||
if xs is not None:
|
||||
assert xs.size(0) == bs, (xs.size(0), bs)
|
||||
|
||||
if length_dim < 0:
|
||||
length_dim = xs.dim() + length_dim
|
||||
# ind = (:, None, ..., None, :, , None, ..., None)
|
||||
ind = tuple(slice(None) if i in (0, length_dim) else None for i in range(xs.dim()))
|
||||
mask = mask[ind].expand_as(xs).to(xs.device)
|
||||
return mask
|
||||
|
||||
|
||||
def make_non_pad_mask(lengths, xs=None, length_dim=-1):
|
||||
"""Make mask tensor containing indices of non-padded part.
|
||||
|
||||
Args:
|
||||
lengths (LongTensor or List): Batch of lengths (B,).
|
||||
xs (Tensor, optional): The reference tensor.
|
||||
If set, masks will be the same shape as this tensor.
|
||||
length_dim (int, optional): Dimension indicator of the above tensor.
|
||||
See the example.
|
||||
|
||||
Returns:
|
||||
ByteTensor: mask tensor containing indices of padded part.
|
||||
dtype=torch.uint8 in PyTorch 1.2-
|
||||
dtype=torch.bool in PyTorch 1.2+ (including 1.2)
|
||||
|
||||
Examples:
|
||||
With only lengths.
|
||||
|
||||
>>> lengths = [5, 3, 2]
|
||||
>>> make_non_pad_mask(lengths)
|
||||
masks = [[1, 1, 1, 1 ,1],
|
||||
[1, 1, 1, 0, 0],
|
||||
[1, 1, 0, 0, 0]]
|
||||
|
||||
With the reference tensor.
|
||||
|
||||
>>> xs = torch.zeros((3, 2, 4))
|
||||
>>> make_non_pad_mask(lengths, xs)
|
||||
tensor([[[1, 1, 1, 1],
|
||||
[1, 1, 1, 1]],
|
||||
[[1, 1, 1, 0],
|
||||
[1, 1, 1, 0]],
|
||||
[[1, 1, 0, 0],
|
||||
[1, 1, 0, 0]]], dtype=torch.uint8)
|
||||
>>> xs = torch.zeros((3, 2, 6))
|
||||
>>> make_non_pad_mask(lengths, xs)
|
||||
tensor([[[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0]],
|
||||
[[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0]],
|
||||
[[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0]]], dtype=torch.uint8)
|
||||
|
||||
With the reference tensor and dimension indicator.
|
||||
|
||||
>>> xs = torch.zeros((3, 6, 6))
|
||||
>>> make_non_pad_mask(lengths, xs, 1)
|
||||
tensor([[[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[0, 0, 0, 0, 0, 0]],
|
||||
[[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0]],
|
||||
[[1, 1, 1, 1, 1, 1],
|
||||
[1, 1, 1, 1, 1, 1],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0],
|
||||
[0, 0, 0, 0, 0, 0]]], dtype=torch.uint8)
|
||||
>>> make_non_pad_mask(lengths, xs, 2)
|
||||
tensor([[[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0],
|
||||
[1, 1, 1, 1, 1, 0]],
|
||||
[[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0],
|
||||
[1, 1, 1, 0, 0, 0]],
|
||||
[[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0],
|
||||
[1, 1, 0, 0, 0, 0]]], dtype=torch.uint8)
|
||||
|
||||
"""
|
||||
return ~make_pad_mask(lengths, xs, length_dim)
|
||||
|
||||
|
||||
def mask_by_length(xs, lengths, fill=0):
|
||||
"""Mask tensor according to length.
|
||||
|
||||
Args:
|
||||
xs (Tensor): Batch of input tensor (B, `*`).
|
||||
lengths (LongTensor or List): Batch of lengths (B,).
|
||||
fill (int or float): Value to fill masked part.
|
||||
|
||||
Returns:
|
||||
Tensor: Batch of masked input tensor (B, `*`).
|
||||
|
||||
Examples:
|
||||
>>> x = torch.arange(5).repeat(3, 1) + 1
|
||||
>>> x
|
||||
tensor([[1, 2, 3, 4, 5],
|
||||
[1, 2, 3, 4, 5],
|
||||
[1, 2, 3, 4, 5]])
|
||||
>>> lengths = [5, 3, 2]
|
||||
>>> mask_by_length(x, lengths)
|
||||
tensor([[1, 2, 3, 4, 5],
|
||||
[1, 2, 3, 0, 0],
|
||||
[1, 2, 0, 0, 0]])
|
||||
|
||||
"""
|
||||
assert xs.size(0) == len(lengths)
|
||||
ret = xs.data.new(*xs.size()).fill_(fill)
|
||||
for i, l in enumerate(lengths):
|
||||
ret[i, :l] = xs[i, :l]
|
||||
return ret
|
||||
|
||||
|
||||
def to_torch_tensor(x):
|
||||
"""Change to torch.Tensor or ComplexTensor from numpy.ndarray.
|
||||
|
||||
Args:
|
||||
x: Inputs. It should be one of numpy.ndarray, Tensor, ComplexTensor, and dict.
|
||||
|
||||
Returns:
|
||||
Tensor or ComplexTensor: Type converted inputs.
|
||||
|
||||
Examples:
|
||||
>>> xs = np.ones(3, dtype=np.float32)
|
||||
>>> xs = to_torch_tensor(xs)
|
||||
tensor([1., 1., 1.])
|
||||
>>> xs = torch.ones(3, 4, 5)
|
||||
>>> assert to_torch_tensor(xs) is xs
|
||||
>>> xs = {'real': xs, 'imag': xs}
|
||||
>>> to_torch_tensor(xs)
|
||||
ComplexTensor(
|
||||
Real:
|
||||
tensor([1., 1., 1.])
|
||||
Imag;
|
||||
tensor([1., 1., 1.])
|
||||
)
|
||||
|
||||
"""
|
||||
# If numpy, change to torch tensor
|
||||
if isinstance(x, np.ndarray):
|
||||
if x.dtype.kind == "c":
|
||||
# Dynamically importing because torch_complex requires python3
|
||||
from torch_complex.tensor import ComplexTensor
|
||||
|
||||
return ComplexTensor(x)
|
||||
else:
|
||||
return torch.from_numpy(x)
|
||||
|
||||
# If {'real': ..., 'imag': ...}, convert to ComplexTensor
|
||||
elif isinstance(x, dict):
|
||||
# Dynamically importing because torch_complex requires python3
|
||||
from torch_complex.tensor import ComplexTensor
|
||||
|
||||
if "real" not in x or "imag" not in x:
|
||||
raise ValueError("has 'real' and 'imag' keys: {}".format(list(x)))
|
||||
# Relative importing because of using python3 syntax
|
||||
return ComplexTensor(x["real"], x["imag"])
|
||||
|
||||
# If torch.Tensor, as it is
|
||||
elif isinstance(x, torch.Tensor):
|
||||
return x
|
||||
|
||||
else:
|
||||
error = (
|
||||
"x must be numpy.ndarray, torch.Tensor or a dict like "
|
||||
"{{'real': torch.Tensor, 'imag': torch.Tensor}}, "
|
||||
"but got {}".format(type(x))
|
||||
)
|
||||
try:
|
||||
from torch_complex.tensor import ComplexTensor
|
||||
except Exception:
|
||||
# If PY2
|
||||
raise ValueError(error)
|
||||
else:
|
||||
# If PY3
|
||||
if isinstance(x, ComplexTensor):
|
||||
return x
|
||||
else:
|
||||
raise ValueError(error)
|
||||
|
||||
|
||||
def get_subsample(train_args, mode, arch):
|
||||
"""Parse the subsampling factors from the args for the specified `mode` and `arch`.
|
||||
|
||||
Args:
|
||||
train_args: argument Namespace containing options.
|
||||
mode: one of ('asr', 'mt', 'st')
|
||||
arch: one of ('rnn', 'rnn-t', 'rnn_mix', 'rnn_mulenc', 'transformer')
|
||||
|
||||
Returns:
|
||||
np.ndarray / List[np.ndarray]: subsampling factors.
|
||||
"""
|
||||
if arch == "transformer":
|
||||
return np.array([1])
|
||||
|
||||
elif mode == "mt" and arch == "rnn":
|
||||
# +1 means input (+1) and layers outputs (train_args.elayer)
|
||||
subsample = np.ones(train_args.elayers + 1, dtype=np.int32)
|
||||
logging.warning("Subsampling is not performed for machine translation.")
|
||||
logging.info("subsample: " + " ".join([str(x) for x in subsample]))
|
||||
return subsample
|
||||
|
||||
elif (
|
||||
(mode == "asr" and arch in ("rnn", "rnn-t"))
|
||||
or (mode == "mt" and arch == "rnn")
|
||||
or (mode == "st" and arch == "rnn")
|
||||
):
|
||||
subsample = np.ones(train_args.elayers + 1, dtype=np.int32)
|
||||
if train_args.etype.endswith("p") and not train_args.etype.startswith("vgg"):
|
||||
ss = train_args.subsample.split("_")
|
||||
for j in range(min(train_args.elayers + 1, len(ss))):
|
||||
subsample[j] = int(ss[j])
|
||||
else:
|
||||
logging.warning(
|
||||
"Subsampling is not performed for vgg*. "
|
||||
"It is performed in max pooling layers at CNN."
|
||||
)
|
||||
logging.info("subsample: " + " ".join([str(x) for x in subsample]))
|
||||
return subsample
|
||||
|
||||
elif mode == "asr" and arch == "rnn_mix":
|
||||
subsample = np.ones(train_args.elayers_sd + train_args.elayers + 1, dtype=np.int32)
|
||||
if train_args.etype.endswith("p") and not train_args.etype.startswith("vgg"):
|
||||
ss = train_args.subsample.split("_")
|
||||
for j in range(min(train_args.elayers_sd + train_args.elayers + 1, len(ss))):
|
||||
subsample[j] = int(ss[j])
|
||||
else:
|
||||
logging.warning(
|
||||
"Subsampling is not performed for vgg*. "
|
||||
"It is performed in max pooling layers at CNN."
|
||||
)
|
||||
logging.info("subsample: " + " ".join([str(x) for x in subsample]))
|
||||
return subsample
|
||||
|
||||
elif mode == "asr" and arch == "rnn_mulenc":
|
||||
subsample_list = []
|
||||
for idx in range(train_args.num_encs):
|
||||
subsample = np.ones(train_args.elayers[idx] + 1, dtype=np.int32)
|
||||
if train_args.etype[idx].endswith("p") and not train_args.etype[idx].startswith("vgg"):
|
||||
ss = train_args.subsample[idx].split("_")
|
||||
for j in range(min(train_args.elayers[idx] + 1, len(ss))):
|
||||
subsample[j] = int(ss[j])
|
||||
else:
|
||||
logging.warning(
|
||||
"Encoder %d: Subsampling is not performed for vgg*. "
|
||||
"It is performed in max pooling layers at CNN.",
|
||||
idx + 1,
|
||||
)
|
||||
logging.info("subsample: " + " ".join([str(x) for x in subsample]))
|
||||
subsample_list.append(subsample)
|
||||
return subsample_list
|
||||
|
||||
else:
|
||||
raise ValueError("Invalid options: mode={}, arch={}".format(mode, arch))
|
||||
|
||||
|
||||
def rename_state_dict(old_prefix: str, new_prefix: str, state_dict: Dict[str, torch.Tensor]):
|
||||
"""Replace keys of old prefix with new prefix in state dict."""
|
||||
# need this list not to break the dict iterator
|
||||
old_keys = [k for k in state_dict if k.startswith(old_prefix)]
|
||||
if len(old_keys) > 0:
|
||||
logging.warning(f"Rename: {old_prefix} -> {new_prefix}")
|
||||
for k in old_keys:
|
||||
v = state_dict.pop(k)
|
||||
new_k = k.replace(old_prefix, new_prefix)
|
||||
state_dict[new_k] = v
|
||||
|
||||
|
||||
class Swish(torch.nn.Module):
|
||||
"""Swish activation definition.
|
||||
|
||||
Swish(x) = (beta * x) * sigmoid(x)
|
||||
where beta = 1 defines standard Swish activation.
|
||||
|
||||
References:
|
||||
https://arxiv.org/abs/2108.12943 / https://arxiv.org/abs/1710.05941v1.
|
||||
E-swish variant: https://arxiv.org/abs/1801.07145.
|
||||
|
||||
Args:
|
||||
beta: Beta parameter for E-Swish.
|
||||
(beta >= 1. If beta < 1, use standard Swish).
|
||||
use_builtin: Whether to use PyTorch function if available.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, beta: float = 1.0, use_builtin: bool = False) -> None:
|
||||
"""Initialize Swish.
|
||||
|
||||
Args:
|
||||
beta: TODO.
|
||||
use_builtin: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.beta = beta
|
||||
|
||||
if beta > 1:
|
||||
self.swish = lambda x: (self.beta * x) * torch.sigmoid(x)
|
||||
else:
|
||||
if use_builtin:
|
||||
self.swish = torch.nn.SiLU()
|
||||
else:
|
||||
self.swish = lambda x: x * torch.sigmoid(x)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Forward computation."""
|
||||
return self.swish(x)
|
||||
|
||||
|
||||
def get_activation(act):
|
||||
"""Return activation function."""
|
||||
|
||||
activation_funcs = {
|
||||
"hardtanh": torch.nn.Hardtanh,
|
||||
"tanh": torch.nn.Tanh,
|
||||
"relu": torch.nn.ReLU,
|
||||
"selu": torch.nn.SELU,
|
||||
"swish": Swish,
|
||||
}
|
||||
|
||||
return activation_funcs[act]()
|
||||
|
||||
|
||||
class TooShortUttError(Exception):
|
||||
"""Raised when the utt is too short for subsampling.
|
||||
|
||||
Args:
|
||||
message: Error message to display.
|
||||
actual_size: The size that cannot pass the subsampling.
|
||||
limit: The size limit for subsampling.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, actual_size: int, limit: int) -> None:
|
||||
"""Construct a TooShortUttError module."""
|
||||
super().__init__(message)
|
||||
|
||||
self.actual_size = actual_size
|
||||
self.limit = limit
|
||||
|
||||
|
||||
def check_short_utt(sub_factor: int, size: int) -> Tuple[bool, int]:
|
||||
"""Check if the input is too short for subsampling.
|
||||
|
||||
Args:
|
||||
sub_factor: Subsampling factor for Conv2DSubsampling.
|
||||
size: Input size.
|
||||
|
||||
Returns:
|
||||
: Whether an error should be sent.
|
||||
: Size limit for specified subsampling factor.
|
||||
|
||||
"""
|
||||
if sub_factor == 2 and size < 3:
|
||||
return True, 7
|
||||
elif sub_factor == 4 and size < 7:
|
||||
return True, 7
|
||||
elif sub_factor == 6 and size < 11:
|
||||
return True, 11
|
||||
|
||||
return False, -1
|
||||
|
||||
|
||||
def sub_factor_to_params(sub_factor: int, input_size: int) -> Tuple[int, int, int]:
|
||||
"""Get conv2D second layer parameters for given subsampling factor.
|
||||
|
||||
Args:
|
||||
sub_factor: Subsampling factor (1/X).
|
||||
input_size: Input size.
|
||||
|
||||
Returns:
|
||||
: Kernel size for second convolution.
|
||||
: Stride for second convolution.
|
||||
: Conv2DSubsampling output size.
|
||||
|
||||
"""
|
||||
if sub_factor == 2:
|
||||
return 3, 1, (((input_size - 1) // 2 - 2))
|
||||
elif sub_factor == 4:
|
||||
return 3, 2, (((input_size - 1) // 2 - 1) // 2)
|
||||
elif sub_factor == 6:
|
||||
return 5, 3, (((input_size - 1) // 2 - 2) // 3)
|
||||
else:
|
||||
raise ValueError("subsampling_factor parameter should be set to either 2, 4 or 6.")
|
||||
|
||||
|
||||
def make_chunk_mask(
|
||||
size: int,
|
||||
chunk_size: int,
|
||||
left_chunk_size: int = 0,
|
||||
device: torch.device = None,
|
||||
) -> torch.Tensor:
|
||||
"""Create chunk mask for the subsequent steps (size, size).
|
||||
|
||||
Reference: https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py
|
||||
|
||||
Args:
|
||||
size: Size of the source mask.
|
||||
chunk_size: Number of frames in chunk.
|
||||
left_chunk_size: Size of the left context in chunks (0 means full context).
|
||||
device: Device for the mask tensor.
|
||||
|
||||
Returns:
|
||||
mask: Chunk mask. (size, size)
|
||||
|
||||
"""
|
||||
mask = torch.zeros(size, size, device=device, dtype=torch.bool)
|
||||
|
||||
for i in range(size):
|
||||
if left_chunk_size < 0:
|
||||
start = 0
|
||||
else:
|
||||
start = max((i // chunk_size - left_chunk_size) * chunk_size, 0)
|
||||
|
||||
end = min((i // chunk_size + 1) * chunk_size, size)
|
||||
mask[i, start:end] = True
|
||||
|
||||
return ~mask
|
||||
|
||||
|
||||
def make_source_mask(lengths: torch.Tensor) -> torch.Tensor:
|
||||
"""Create source mask for given lengths.
|
||||
|
||||
Reference: https://github.com/k2-fsa/icefall/blob/master/icefall/utils.py
|
||||
|
||||
Args:
|
||||
lengths: Sequence lengths. (B,)
|
||||
|
||||
Returns:
|
||||
: Mask for the sequence lengths. (B, max_len)
|
||||
|
||||
"""
|
||||
max_len = lengths.max()
|
||||
batch_size = lengths.size(0)
|
||||
|
||||
expanded_lengths = torch.arange(max_len).expand(batch_size, max_len).to(lengths)
|
||||
|
||||
return expanded_lengths >= lengths.unsqueeze(1)
|
||||
|
||||
|
||||
def get_transducer_task_io(
|
||||
labels: torch.Tensor,
|
||||
encoder_out_lens: torch.Tensor,
|
||||
ignore_id: int = -1,
|
||||
blank_id: int = 0,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Get Transducer loss I/O.
|
||||
|
||||
Args:
|
||||
labels: Label ID sequences. (B, L)
|
||||
encoder_out_lens: Encoder output lengths. (B,)
|
||||
ignore_id: Padding symbol ID.
|
||||
blank_id: Blank symbol ID.
|
||||
|
||||
Returns:
|
||||
decoder_in: Decoder inputs. (B, U)
|
||||
target: Target label ID sequences. (B, U)
|
||||
t_len: Time lengths. (B,)
|
||||
u_len: Label lengths. (B,)
|
||||
|
||||
"""
|
||||
|
||||
def pad_list(labels: List[torch.Tensor], padding_value: int = 0):
|
||||
"""Create padded batch of labels from a list of labels sequences.
|
||||
|
||||
Args:
|
||||
labels: Labels sequences. [B x (?)]
|
||||
padding_value: Padding value.
|
||||
|
||||
Returns:
|
||||
labels: Batch of padded labels sequences. (B,)
|
||||
|
||||
"""
|
||||
batch_size = len(labels)
|
||||
|
||||
padded = (
|
||||
labels[0]
|
||||
.new(batch_size, max(x.size(0) for x in labels), *labels[0].size()[1:])
|
||||
.fill_(padding_value)
|
||||
)
|
||||
|
||||
for i in range(batch_size):
|
||||
padded[i, : labels[i].size(0)] = labels[i]
|
||||
|
||||
return padded
|
||||
|
||||
device = labels.device
|
||||
|
||||
labels_unpad = [y[y != ignore_id] for y in labels]
|
||||
blank = labels[0].new([blank_id])
|
||||
|
||||
decoder_in = pad_list(
|
||||
[torch.cat([blank, label], dim=0) for label in labels_unpad], blank_id
|
||||
).to(device)
|
||||
|
||||
target = pad_list(labels_unpad, blank_id).type(torch.int32).to(device)
|
||||
|
||||
encoder_out_lens = list(map(int, encoder_out_lens))
|
||||
t_len = torch.IntTensor(encoder_out_lens).to(device)
|
||||
|
||||
u_len = torch.IntTensor([y.size(0) for y in labels_unpad]).to(device)
|
||||
|
||||
return decoder_in, target, t_len, u_len
|
||||
|
||||
|
||||
def pad_to_len(t: torch.Tensor, pad_len: int, dim: int):
|
||||
"""Pad the tensor `t` at `dim` to the length `pad_len` with right padding zeros."""
|
||||
if t.size(dim) == pad_len:
|
||||
return t
|
||||
else:
|
||||
pad_size = list(t.shape)
|
||||
pad_size[dim] = pad_len - t.size(dim)
|
||||
return torch.cat([t, torch.zeros(*pad_size, dtype=t.dtype, device=t.device)], dim=dim)
|
||||
@@ -0,0 +1,137 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Repeat the same layer definition."""
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
import torch
|
||||
|
||||
|
||||
class MultiSequential(torch.nn.Sequential):
|
||||
"""Multi-input multi-output torch.nn.Sequential."""
|
||||
|
||||
def __init__(self, *args, layer_drop_rate=0.0):
|
||||
"""Initialize MultiSequential with layer_drop.
|
||||
|
||||
Args:
|
||||
layer_drop_rate (float): Probability of dropping out each fn (layer).
|
||||
|
||||
"""
|
||||
super(MultiSequential, self).__init__(*args)
|
||||
self.layer_drop_rate = layer_drop_rate
|
||||
|
||||
def forward(self, *args):
|
||||
"""Repeat."""
|
||||
_probs = torch.empty(len(self)).uniform_()
|
||||
for idx, m in enumerate(self):
|
||||
if not self.training or (_probs[idx] >= self.layer_drop_rate):
|
||||
args = m(*args)
|
||||
return args
|
||||
|
||||
|
||||
def repeat(N, fn, layer_drop_rate=0.0):
|
||||
"""Repeat module N times.
|
||||
|
||||
Args:
|
||||
N (int): Number of repeat time.
|
||||
fn (Callable): Function to generate module.
|
||||
layer_drop_rate (float): Probability of dropping out each fn (layer).
|
||||
|
||||
Returns:
|
||||
MultiSequential: Repeated model instance.
|
||||
|
||||
"""
|
||||
return MultiSequential(*[fn(n) for n in range(N)], layer_drop_rate=layer_drop_rate)
|
||||
|
||||
|
||||
class MultiBlocks(torch.nn.Module):
|
||||
"""MultiBlocks definition.
|
||||
Args:
|
||||
block_list: Individual blocks of the encoder architecture.
|
||||
output_size: Architecture output size.
|
||||
norm_class: Normalization module class.
|
||||
norm_args: Normalization module arguments.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
block_list: List[torch.nn.Module],
|
||||
output_size: int,
|
||||
norm_class: torch.nn.Module = LayerNorm,
|
||||
) -> None:
|
||||
"""Construct a MultiBlocks object."""
|
||||
super().__init__()
|
||||
|
||||
self.blocks = torch.nn.ModuleList(block_list)
|
||||
self.norm_blocks = norm_class(output_size)
|
||||
|
||||
self.num_blocks = len(block_list)
|
||||
|
||||
def reset_streaming_cache(self, left_context: int, device: torch.device) -> None:
|
||||
"""Initialize/Reset encoder streaming cache.
|
||||
Args:
|
||||
left_context: Number of left frames during chunk-by-chunk inference.
|
||||
device: Device to use for cache tensor.
|
||||
"""
|
||||
for idx in range(self.num_blocks):
|
||||
self.blocks[idx].reset_streaming_cache(left_context, device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
pos_enc: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
chunk_mask: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Forward each block of the encoder architecture.
|
||||
Args:
|
||||
x: MultiBlocks input sequences. (B, T, D_block_1)
|
||||
pos_enc: Positional embedding sequences.
|
||||
mask: Source mask. (B, T)
|
||||
chunk_mask: Chunk mask. (T_2, T_2)
|
||||
Returns:
|
||||
x: Output sequences. (B, T, D_block_N)
|
||||
"""
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
x, mask, pos_enc = block(x, pos_enc, mask, chunk_mask=chunk_mask)
|
||||
|
||||
x = self.norm_blocks(x)
|
||||
|
||||
return x
|
||||
|
||||
def chunk_forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
pos_enc: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
chunk_size: int = 0,
|
||||
left_context: int = 0,
|
||||
right_context: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Forward each block of the encoder architecture.
|
||||
Args:
|
||||
x: MultiBlocks input sequences. (B, T, D_block_1)
|
||||
pos_enc: Positional embedding sequences. (B, 2 * (T - 1), D_att)
|
||||
mask: Source mask. (B, T_2)
|
||||
left_context: Number of frames in left context.
|
||||
right_context: Number of frames in right context.
|
||||
Returns:
|
||||
x: MultiBlocks output sequences. (B, T, D_block_N)
|
||||
"""
|
||||
for block_idx, block in enumerate(self.blocks):
|
||||
x, pos_enc = block.chunk_forward(
|
||||
x,
|
||||
pos_enc,
|
||||
mask,
|
||||
chunk_size=chunk_size,
|
||||
left_context=left_context,
|
||||
right_context=right_context,
|
||||
)
|
||||
|
||||
x = self.norm_blocks(x)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,641 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Subsampling layer definition."""
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from funasr.models.transformer.embedding import PositionalEncoding
|
||||
import logging
|
||||
from funasr.models.scama.utils import sequence_mask
|
||||
from funasr.models.transformer.utils.nets_utils import sub_factor_to_params, pad_to_len
|
||||
from typing import Optional, Tuple, Union
|
||||
import math
|
||||
|
||||
|
||||
class TooShortUttError(Exception):
|
||||
"""Raised when the utt is too short for subsampling.
|
||||
|
||||
Args:
|
||||
message (str): Message for error catch
|
||||
actual_size (int): the short size that cannot pass the subsampling
|
||||
limit (int): the limit size for subsampling
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, message, actual_size, limit):
|
||||
"""Construct a TooShortUttError for error handler."""
|
||||
super().__init__(message)
|
||||
self.actual_size = actual_size
|
||||
self.limit = limit
|
||||
|
||||
|
||||
def check_short_utt(ins, size):
|
||||
"""Check if the utterance is too short for subsampling."""
|
||||
if isinstance(ins, Conv2dSubsampling2) and size < 3:
|
||||
return True, 3
|
||||
if isinstance(ins, Conv2dSubsampling) and size < 7:
|
||||
return True, 7
|
||||
if isinstance(ins, Conv2dSubsampling6) and size < 11:
|
||||
return True, 11
|
||||
if isinstance(ins, Conv2dSubsampling8) and size < 15:
|
||||
return True, 15
|
||||
return False, -1
|
||||
|
||||
|
||||
class Conv2dSubsampling(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling (to 1/4 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, pos_enc=None):
|
||||
"""Construct an Conv2dSubsampling object."""
|
||||
super(Conv2dSubsampling, self).__init__()
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
self.out = torch.nn.Sequential(
|
||||
torch.nn.Linear(odim * (((idim - 1) // 2 - 1) // 2), odim),
|
||||
pos_enc if pos_enc is not None else PositionalEncoding(odim, dropout_rate),
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 4.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 4.
|
||||
|
||||
"""
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
return x, x_mask[:, :, :-2:2][:, :, :-2:2]
|
||||
|
||||
def __getitem__(self, key):
|
||||
"""Get item.
|
||||
|
||||
When reset_parameters() is called, if use_scaled_pos_enc is used,
|
||||
return the positioning encoding.
|
||||
|
||||
"""
|
||||
if key != -1:
|
||||
raise NotImplementedError("Support only `-1` (for `reset_parameters`).")
|
||||
return self.out[key]
|
||||
|
||||
|
||||
class Conv2dSubsamplingPad(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling (to 1/4 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, pos_enc=None):
|
||||
"""Construct an Conv2dSubsampling object."""
|
||||
super(Conv2dSubsamplingPad, self).__init__()
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, odim, 3, 2, padding=(0, 0)),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 3, 2, padding=(0, 0)),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
self.out = torch.nn.Sequential(
|
||||
torch.nn.Linear(odim * (((idim - 1) // 2 - 1) // 2), odim),
|
||||
pos_enc if pos_enc is not None else PositionalEncoding(odim, dropout_rate),
|
||||
)
|
||||
self.pad_fn = torch.nn.ConstantPad1d((0, 4), 0.0)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 4.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 4.
|
||||
|
||||
"""
|
||||
x = x.transpose(1, 2)
|
||||
x = self.pad_fn(x)
|
||||
x = x.transpose(1, 2)
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
x_len = torch.sum(x_mask[:, 0, :], dim=-1)
|
||||
x_len = (x_len - 1) // 2 + 1
|
||||
x_len = (x_len - 1) // 2 + 1
|
||||
mask = sequence_mask(x_len, None, x_len.dtype, x[0].device)
|
||||
return x, mask[:, None, :]
|
||||
|
||||
def __getitem__(self, key):
|
||||
"""Get item.
|
||||
|
||||
When reset_parameters() is called, if use_scaled_pos_enc is used,
|
||||
return the positioning encoding.
|
||||
|
||||
"""
|
||||
if key != -1:
|
||||
raise NotImplementedError("Support only `-1` (for `reset_parameters`).")
|
||||
return self.out[key]
|
||||
|
||||
|
||||
class Conv2dSubsampling2(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling (to 1/2 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, pos_enc=None):
|
||||
"""Construct an Conv2dSubsampling2 object."""
|
||||
super(Conv2dSubsampling2, self).__init__()
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 3, 1),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
self.out = torch.nn.Sequential(
|
||||
torch.nn.Linear(odim * (((idim - 1) // 2 - 2)), odim),
|
||||
pos_enc if pos_enc is not None else PositionalEncoding(odim, dropout_rate),
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 2.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 2.
|
||||
|
||||
"""
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
return x, x_mask[:, :, :-2:2][:, :, :-2:1]
|
||||
|
||||
def __getitem__(self, key):
|
||||
"""Get item.
|
||||
|
||||
When reset_parameters() is called, if use_scaled_pos_enc is used,
|
||||
return the positioning encoding.
|
||||
|
||||
"""
|
||||
if key != -1:
|
||||
raise NotImplementedError("Support only `-1` (for `reset_parameters`).")
|
||||
return self.out[key]
|
||||
|
||||
|
||||
class Conv2dSubsampling6(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling (to 1/6 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, pos_enc=None):
|
||||
"""Construct an Conv2dSubsampling6 object."""
|
||||
super(Conv2dSubsampling6, self).__init__()
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 5, 3),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
self.out = torch.nn.Sequential(
|
||||
torch.nn.Linear(odim * (((idim - 1) // 2 - 2) // 3), odim),
|
||||
pos_enc if pos_enc is not None else PositionalEncoding(odim, dropout_rate),
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 6.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 6.
|
||||
|
||||
"""
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
return x, x_mask[:, :, :-2:2][:, :, :-4:3]
|
||||
|
||||
|
||||
class Conv2dSubsampling8(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling (to 1/8 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, pos_enc=None):
|
||||
"""Construct an Conv2dSubsampling8 object."""
|
||||
super(Conv2dSubsampling8, self).__init__()
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(odim, odim, 3, 2),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
self.out = torch.nn.Sequential(
|
||||
torch.nn.Linear(odim * ((((idim - 1) // 2 - 1) // 2 - 1) // 2), odim),
|
||||
pos_enc if pos_enc is not None else PositionalEncoding(odim, dropout_rate),
|
||||
)
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 8.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 8.
|
||||
|
||||
"""
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
return x, x_mask[:, :, :-2:2][:, :, :-2:2][:, :, :-2:2]
|
||||
|
||||
|
||||
class Conv1dSubsampling(torch.nn.Module):
|
||||
"""Convolutional 1D subsampling (to 1/2 length).
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
pos_enc (torch.nn.Module): Custom position encoding layer.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
idim,
|
||||
odim,
|
||||
kernel_size,
|
||||
stride,
|
||||
pad,
|
||||
tf2torch_tensor_name_prefix_torch: str = "stride_conv",
|
||||
tf2torch_tensor_name_prefix_tf: str = "seq2seq/proj_encoder/downsampling",
|
||||
):
|
||||
"""Initialize Conv1dSubsampling.
|
||||
|
||||
Args:
|
||||
idim: TODO.
|
||||
odim: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
stride: TODO.
|
||||
pad: TODO.
|
||||
tf2torch_tensor_name_prefix_torch: TODO.
|
||||
tf2torch_tensor_name_prefix_tf: TODO.
|
||||
"""
|
||||
super(Conv1dSubsampling, self).__init__()
|
||||
self.conv = torch.nn.Conv1d(idim, odim, kernel_size, stride)
|
||||
self.pad_fn = torch.nn.ConstantPad1d(pad, 0.0)
|
||||
self.stride = stride
|
||||
self.odim = odim
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self.odim
|
||||
|
||||
def forward(self, x, x_len):
|
||||
"""Subsample x."""
|
||||
x = x.transpose(1, 2) # (b, d ,t)
|
||||
x = self.pad_fn(x)
|
||||
# x = F.relu(self.conv(x))
|
||||
x = F.leaky_relu(self.conv(x), negative_slope=0.0)
|
||||
x = x.transpose(1, 2) # (b, t ,d)
|
||||
|
||||
if x_len is None:
|
||||
|
||||
return x, None
|
||||
x_len = (x_len - 1) // self.stride + 1
|
||||
return x, x_len
|
||||
|
||||
|
||||
class StreamingConvInput(torch.nn.Module):
|
||||
"""Streaming ConvInput module definition.
|
||||
Args:
|
||||
input_size: Input size.
|
||||
conv_size: Convolution size.
|
||||
subsampling_factor: Subsampling factor.
|
||||
vgg_like: Whether to use a VGG-like network.
|
||||
output_size: Block output dimension.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
conv_size: Union[int, Tuple],
|
||||
subsampling_factor: int = 4,
|
||||
vgg_like: bool = True,
|
||||
conv_kernel_size: int = 3,
|
||||
output_size: Optional[int] = None,
|
||||
) -> None:
|
||||
"""Construct a ConvInput object."""
|
||||
super().__init__()
|
||||
if vgg_like:
|
||||
if subsampling_factor == 1:
|
||||
conv_size1, conv_size2 = conv_size
|
||||
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(
|
||||
1,
|
||||
conv_size1,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(
|
||||
conv_size1,
|
||||
conv_size1,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((1, 2)),
|
||||
torch.nn.Conv2d(
|
||||
conv_size1,
|
||||
conv_size2,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(
|
||||
conv_size2,
|
||||
conv_size2,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((1, 2)),
|
||||
)
|
||||
|
||||
output_proj = conv_size2 * ((input_size // 2) // 2)
|
||||
|
||||
self.subsampling_factor = 1
|
||||
|
||||
self.stride_1 = 1
|
||||
|
||||
self.create_new_mask = self.create_new_vgg_mask
|
||||
|
||||
else:
|
||||
conv_size1, conv_size2 = conv_size
|
||||
|
||||
kernel_1 = int(subsampling_factor / 2)
|
||||
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(
|
||||
1,
|
||||
conv_size1,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(
|
||||
conv_size1,
|
||||
conv_size1,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((kernel_1, 2)),
|
||||
torch.nn.Conv2d(
|
||||
conv_size1,
|
||||
conv_size2,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(
|
||||
conv_size2,
|
||||
conv_size2,
|
||||
conv_kernel_size,
|
||||
stride=1,
|
||||
padding=(conv_kernel_size - 1) // 2,
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((2, 2)),
|
||||
)
|
||||
|
||||
output_proj = conv_size2 * ((input_size // 2) // 2)
|
||||
|
||||
self.subsampling_factor = subsampling_factor
|
||||
|
||||
self.create_new_mask = self.create_new_vgg_mask
|
||||
|
||||
self.stride_1 = kernel_1
|
||||
|
||||
else:
|
||||
if subsampling_factor == 1:
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, conv_size, 3, [1, 2], [1, 0]),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(conv_size, conv_size, conv_kernel_size, [1, 2], [1, 0]),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
|
||||
output_proj = conv_size * (((input_size - 1) // 2 - 1) // 2)
|
||||
|
||||
self.subsampling_factor = subsampling_factor
|
||||
self.kernel_2 = conv_kernel_size
|
||||
self.stride_2 = 1
|
||||
|
||||
self.create_new_mask = self.create_new_conv2d_mask
|
||||
|
||||
else:
|
||||
kernel_2, stride_2, conv_2_output_size = sub_factor_to_params(
|
||||
subsampling_factor,
|
||||
input_size,
|
||||
)
|
||||
|
||||
self.conv = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, conv_size, 3, 2, [1, 0]),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(
|
||||
conv_size, conv_size, kernel_2, stride_2, [(kernel_2 - 1) // 2, 0]
|
||||
),
|
||||
torch.nn.ReLU(),
|
||||
)
|
||||
|
||||
output_proj = conv_size * conv_2_output_size
|
||||
|
||||
self.subsampling_factor = subsampling_factor
|
||||
self.kernel_2 = kernel_2
|
||||
self.stride_2 = stride_2
|
||||
|
||||
self.create_new_mask = self.create_new_conv2d_mask
|
||||
|
||||
self.vgg_like = vgg_like
|
||||
self.min_frame_length = 7
|
||||
|
||||
if output_size is not None:
|
||||
self.output = torch.nn.Linear(output_proj, output_size)
|
||||
self.output_size = output_size
|
||||
else:
|
||||
self.output = None
|
||||
self.output_size = output_proj
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, mask: Optional[torch.Tensor], chunk_size: Optional[torch.Tensor]
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Encode input sequences.
|
||||
Args:
|
||||
x: ConvInput input sequences. (B, T, D_feats)
|
||||
mask: Mask of input sequences. (B, 1, T)
|
||||
Returns:
|
||||
x: ConvInput output sequences. (B, sub(T), D_out)
|
||||
mask: Mask of output sequences. (B, 1, sub(T))
|
||||
"""
|
||||
if mask is not None:
|
||||
mask = self.create_new_mask(mask)
|
||||
olens = max(mask.eq(0).sum(1))
|
||||
|
||||
b, t, f = x.size()
|
||||
x = x.unsqueeze(1) # (b. 1. t. f)
|
||||
|
||||
if chunk_size is not None:
|
||||
max_input_length = int(
|
||||
chunk_size
|
||||
* self.subsampling_factor
|
||||
* (math.ceil(float(t) / (chunk_size * self.subsampling_factor)))
|
||||
)
|
||||
x = map(lambda inputs: pad_to_len(inputs, max_input_length, 1), x)
|
||||
x = list(x)
|
||||
x = torch.stack(x, dim=0)
|
||||
N_chunks = max_input_length // (chunk_size * self.subsampling_factor)
|
||||
x = x.view(b * N_chunks, 1, chunk_size * self.subsampling_factor, f)
|
||||
|
||||
x = self.conv(x)
|
||||
|
||||
_, c, _, f = x.size()
|
||||
if chunk_size is not None:
|
||||
x = x.transpose(1, 2).contiguous().view(b, -1, c * f)[:, :olens, :]
|
||||
else:
|
||||
x = x.transpose(1, 2).contiguous().view(b, -1, c * f)
|
||||
|
||||
if self.output is not None:
|
||||
x = self.output(x)
|
||||
|
||||
return x, mask[:, :olens][:, : x.size(1)]
|
||||
|
||||
def create_new_vgg_mask(self, mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Create a new mask for VGG output sequences.
|
||||
Args:
|
||||
mask: Mask of input sequences. (B, T)
|
||||
Returns:
|
||||
mask: Mask of output sequences. (B, sub(T))
|
||||
"""
|
||||
if self.subsampling_factor > 1:
|
||||
vgg1_t_len = mask.size(1) - (mask.size(1) % (self.subsampling_factor // 2))
|
||||
mask = mask[:, :vgg1_t_len][:, :: self.subsampling_factor // 2]
|
||||
|
||||
vgg2_t_len = mask.size(1) - (mask.size(1) % 2)
|
||||
mask = mask[:, :vgg2_t_len][:, ::2]
|
||||
else:
|
||||
mask = mask
|
||||
|
||||
return mask
|
||||
|
||||
def create_new_conv2d_mask(self, mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Create new conformer mask for Conv2d output sequences.
|
||||
Args:
|
||||
mask: Mask of input sequences. (B, T)
|
||||
Returns:
|
||||
mask: Mask of output sequences. (B, sub(T))
|
||||
"""
|
||||
if self.subsampling_factor > 1:
|
||||
return mask[:, ::2][:, :: self.stride_2]
|
||||
else:
|
||||
return mask
|
||||
|
||||
def get_size_before_subsampling(self, size: int) -> int:
|
||||
"""Return the original size before subsampling for a given size.
|
||||
Args:
|
||||
size: Number of frames after subsampling.
|
||||
Returns:
|
||||
: Number of frames before subsampling.
|
||||
"""
|
||||
return size * self.subsampling_factor
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copyright 2020 Emiru Tsunoo
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Subsampling layer definition."""
|
||||
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
class Conv2dSubsamplingWOPosEnc(torch.nn.Module):
|
||||
"""Convolutional 2D subsampling.
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
odim (int): Output dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
kernels (list): kernel sizes
|
||||
strides (list): stride sizes
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, odim, dropout_rate, kernels, strides):
|
||||
"""Construct an Conv2dSubsamplingWOPosEnc object."""
|
||||
assert len(kernels) == len(strides)
|
||||
super().__init__()
|
||||
conv = []
|
||||
olen = idim
|
||||
for i, (k, s) in enumerate(zip(kernels, strides)):
|
||||
conv += [
|
||||
torch.nn.Conv2d(1 if i == 0 else odim, odim, k, s),
|
||||
torch.nn.ReLU(),
|
||||
]
|
||||
olen = math.floor((olen - k) / s + 1)
|
||||
self.conv = torch.nn.Sequential(*conv)
|
||||
self.out = torch.nn.Linear(odim * olen, odim)
|
||||
self.strides = strides
|
||||
self.kernels = kernels
|
||||
|
||||
def forward(self, x, x_mask):
|
||||
"""Subsample x.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
x_mask (torch.Tensor): Input mask (#batch, 1, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Subsampled tensor (#batch, time', odim),
|
||||
where time' = time // 4.
|
||||
torch.Tensor: Subsampled mask (#batch, 1, time'),
|
||||
where time' = time // 4.
|
||||
|
||||
"""
|
||||
x = x.unsqueeze(1) # (b, c, t, f)
|
||||
x = self.conv(x)
|
||||
b, c, t, f = x.size()
|
||||
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
if x_mask is None:
|
||||
return x, None
|
||||
for k, s in zip(self.kernels, self.strides):
|
||||
x_mask = x_mask[:, :, : -k + 1 : s]
|
||||
return x, x_mask
|
||||
@@ -0,0 +1,88 @@
|
||||
"""VGG2L module definition for custom encoder."""
|
||||
|
||||
from typing import Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class VGG2L(torch.nn.Module):
|
||||
"""VGG2L module for custom encoder.
|
||||
|
||||
Args:
|
||||
idim: Input dimension.
|
||||
odim: Output dimension.
|
||||
pos_enc: Positional encoding class.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim: int, odim: int, pos_enc: torch.nn.Module = None):
|
||||
"""Construct a VGG2L object."""
|
||||
super().__init__()
|
||||
|
||||
self.vgg2l = torch.nn.Sequential(
|
||||
torch.nn.Conv2d(1, 64, 3, stride=1, padding=1),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(64, 64, 3, stride=1, padding=1),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((3, 2)),
|
||||
torch.nn.Conv2d(64, 128, 3, stride=1, padding=1),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(128, 128, 3, stride=1, padding=1),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.MaxPool2d((2, 2)),
|
||||
)
|
||||
|
||||
if pos_enc is not None:
|
||||
self.output = torch.nn.Sequential(
|
||||
torch.nn.Linear(128 * ((idim // 2) // 2), odim), pos_enc
|
||||
)
|
||||
else:
|
||||
self.output = torch.nn.Linear(128 * ((idim // 2) // 2), odim)
|
||||
|
||||
def forward(self, feats: torch.Tensor, feats_mask: torch.Tensor) -> Union[
|
||||
Tuple[torch.Tensor, torch.Tensor],
|
||||
Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor],
|
||||
]:
|
||||
"""Forward VGG2L bottleneck.
|
||||
|
||||
Args:
|
||||
feats: Feature sequences. (B, F, D_feats)
|
||||
feats_mask: Mask of feature sequences. (B, 1, F)
|
||||
|
||||
Returns:
|
||||
vgg_output: VGG output sequences.
|
||||
(B, sub(F), D_out) or ((B, sub(F), D_out), (B, sub(F), D_att))
|
||||
vgg_mask: Mask of VGG output sequences. (B, 1, sub(F))
|
||||
|
||||
"""
|
||||
feats = feats.unsqueeze(1)
|
||||
vgg_output = self.vgg2l(feats)
|
||||
|
||||
b, c, t, f = vgg_output.size()
|
||||
|
||||
vgg_output = self.output(vgg_output.transpose(1, 2).contiguous().view(b, t, c * f))
|
||||
|
||||
if feats_mask is not None:
|
||||
vgg_mask = self.create_new_mask(feats_mask)
|
||||
else:
|
||||
vgg_mask = feats_mask
|
||||
|
||||
return vgg_output, vgg_mask
|
||||
|
||||
def create_new_mask(self, feats_mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Create a subsampled mask of feature sequences.
|
||||
|
||||
Args:
|
||||
feats_mask: Mask of feature sequences. (B, 1, F)
|
||||
|
||||
Returns:
|
||||
vgg_mask: Mask of VGG2L output sequences. (B, 1, sub(F))
|
||||
|
||||
"""
|
||||
vgg1_t_len = feats_mask.size(2) - (feats_mask.size(2) % 3)
|
||||
vgg_mask = feats_mask[:, :, :vgg1_t_len][:, :, ::3]
|
||||
|
||||
vgg2_t_len = vgg_mask.size(2) - (vgg_mask.size(2) % 2)
|
||||
vgg_mask = vgg_mask[:, :, :vgg2_t_len][:, :, ::2]
|
||||
|
||||
return vgg_mask
|
||||
Reference in New Issue
Block a user