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 @@
|
||||
"""Initialize sub package."""
|
||||
@@ -0,0 +1,148 @@
|
||||
# Copyright 2020 Hirofumi Inaguma
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Conformer common arguments."""
|
||||
|
||||
|
||||
def add_arguments_rnn_encoder_common(group):
|
||||
"""Define common arguments for RNN encoder."""
|
||||
group.add_argument(
|
||||
"--etype",
|
||||
default="blstmp",
|
||||
type=str,
|
||||
choices=[
|
||||
"lstm",
|
||||
"blstm",
|
||||
"lstmp",
|
||||
"blstmp",
|
||||
"vgglstmp",
|
||||
"vggblstmp",
|
||||
"vgglstm",
|
||||
"vggblstm",
|
||||
"gru",
|
||||
"bgru",
|
||||
"grup",
|
||||
"bgrup",
|
||||
"vgggrup",
|
||||
"vggbgrup",
|
||||
"vgggru",
|
||||
"vggbgru",
|
||||
],
|
||||
help="Type of encoder network architecture",
|
||||
)
|
||||
group.add_argument(
|
||||
"--elayers",
|
||||
default=4,
|
||||
type=int,
|
||||
help="Number of encoder layers",
|
||||
)
|
||||
group.add_argument(
|
||||
"--eunits",
|
||||
"-u",
|
||||
default=300,
|
||||
type=int,
|
||||
help="Number of encoder hidden units",
|
||||
)
|
||||
group.add_argument("--eprojs", default=320, type=int, help="Number of encoder projection units")
|
||||
group.add_argument(
|
||||
"--subsample",
|
||||
default="1",
|
||||
type=str,
|
||||
help="Subsample input frames x_y_z means "
|
||||
"subsample every x frame at 1st layer, "
|
||||
"every y frame at 2nd layer etc.",
|
||||
)
|
||||
return group
|
||||
|
||||
|
||||
def add_arguments_rnn_decoder_common(group):
|
||||
"""Define common arguments for RNN decoder."""
|
||||
group.add_argument(
|
||||
"--dtype",
|
||||
default="lstm",
|
||||
type=str,
|
||||
choices=["lstm", "gru"],
|
||||
help="Type of decoder network architecture",
|
||||
)
|
||||
group.add_argument("--dlayers", default=1, type=int, help="Number of decoder layers")
|
||||
group.add_argument("--dunits", default=320, type=int, help="Number of decoder hidden units")
|
||||
group.add_argument(
|
||||
"--dropout-rate-decoder",
|
||||
default=0.0,
|
||||
type=float,
|
||||
help="Dropout rate for the decoder",
|
||||
)
|
||||
group.add_argument(
|
||||
"--sampling-probability",
|
||||
default=0.0,
|
||||
type=float,
|
||||
help="Ratio of predicted labels fed back to decoder",
|
||||
)
|
||||
group.add_argument(
|
||||
"--lsm-type",
|
||||
const="",
|
||||
default="",
|
||||
type=str,
|
||||
nargs="?",
|
||||
choices=["", "unigram"],
|
||||
help="Apply label smoothing with a specified distribution type",
|
||||
)
|
||||
return group
|
||||
|
||||
|
||||
def add_arguments_rnn_attention_common(group):
|
||||
"""Define common arguments for RNN attention."""
|
||||
group.add_argument(
|
||||
"--atype",
|
||||
default="dot",
|
||||
type=str,
|
||||
choices=[
|
||||
"noatt",
|
||||
"dot",
|
||||
"add",
|
||||
"location",
|
||||
"coverage",
|
||||
"coverage_location",
|
||||
"location2d",
|
||||
"location_recurrent",
|
||||
"multi_head_dot",
|
||||
"multi_head_add",
|
||||
"multi_head_loc",
|
||||
"multi_head_multi_res_loc",
|
||||
],
|
||||
help="Type of attention architecture",
|
||||
)
|
||||
group.add_argument(
|
||||
"--adim",
|
||||
default=320,
|
||||
type=int,
|
||||
help="Number of attention transformation dimensions",
|
||||
)
|
||||
group.add_argument("--awin", default=5, type=int, help="Window size for location2d attention")
|
||||
group.add_argument(
|
||||
"--aheads",
|
||||
default=4,
|
||||
type=int,
|
||||
help="Number of heads for multi head attention",
|
||||
)
|
||||
group.add_argument(
|
||||
"--aconv-chans",
|
||||
default=-1,
|
||||
type=int,
|
||||
help="Number of attention convolution channels \
|
||||
(negative value indicates no location-aware attention)",
|
||||
)
|
||||
group.add_argument(
|
||||
"--aconv-filts",
|
||||
default=100,
|
||||
type=int,
|
||||
help="Number of attention convolution filters \
|
||||
(negative value indicates no location-aware attention)",
|
||||
)
|
||||
group.add_argument(
|
||||
"--dropout-rate",
|
||||
default=0.0,
|
||||
type=float,
|
||||
help="Dropout rate for the encoder",
|
||||
)
|
||||
return group
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,398 @@
|
||||
import logging
|
||||
|
||||
import numpy as np
|
||||
import six
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.utils.rnn import pack_padded_sequence
|
||||
from torch.nn.utils.rnn import pad_packed_sequence
|
||||
|
||||
from funasr.metrics.common import get_vgg2l_odim
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.transformer.utils.nets_utils import to_device
|
||||
|
||||
|
||||
class RNNP(torch.nn.Module):
|
||||
"""RNN with projection layer module
|
||||
|
||||
:param int idim: dimension of inputs
|
||||
:param int elayers: number of encoder layers
|
||||
:param int cdim: number of rnn units (resulted in cdim * 2 if bidirectional)
|
||||
:param int hdim: number of projection units
|
||||
:param np.ndarray subsample: list of subsampling numbers
|
||||
:param float dropout: dropout rate
|
||||
:param str typ: The RNN type
|
||||
"""
|
||||
|
||||
def __init__(self, idim, elayers, cdim, hdim, subsample, dropout, typ="blstm"):
|
||||
"""Initialize RNNP.
|
||||
|
||||
Args:
|
||||
idim: TODO.
|
||||
elayers: TODO.
|
||||
cdim: TODO.
|
||||
hdim: TODO.
|
||||
subsample: TODO.
|
||||
dropout: TODO.
|
||||
typ: TODO.
|
||||
"""
|
||||
super(RNNP, self).__init__()
|
||||
bidir = typ[0] == "b"
|
||||
for i in six.moves.range(elayers):
|
||||
if i == 0:
|
||||
inputdim = idim
|
||||
else:
|
||||
inputdim = hdim
|
||||
|
||||
RNN = torch.nn.LSTM if "lstm" in typ else torch.nn.GRU
|
||||
rnn = RNN(inputdim, cdim, num_layers=1, bidirectional=bidir, batch_first=True)
|
||||
|
||||
setattr(self, "%s%d" % ("birnn" if bidir else "rnn", i), rnn)
|
||||
|
||||
# bottleneck layer to merge
|
||||
if bidir:
|
||||
setattr(self, "bt%d" % i, torch.nn.Linear(2 * cdim, hdim))
|
||||
else:
|
||||
setattr(self, "bt%d" % i, torch.nn.Linear(cdim, hdim))
|
||||
|
||||
self.elayers = elayers
|
||||
self.cdim = cdim
|
||||
self.subsample = subsample
|
||||
self.typ = typ
|
||||
self.bidir = bidir
|
||||
self.dropout = dropout
|
||||
|
||||
def forward(self, xs_pad, ilens, prev_state=None):
|
||||
"""RNNP forward
|
||||
|
||||
:param torch.Tensor xs_pad: batch of padded input sequences (B, Tmax, idim)
|
||||
:param torch.Tensor ilens: batch of lengths of input sequences (B)
|
||||
:param torch.Tensor prev_state: batch of previous RNN states
|
||||
:return: batch of hidden state sequences (B, Tmax, hdim)
|
||||
:rtype: torch.Tensor
|
||||
"""
|
||||
logging.debug(self.__class__.__name__ + " input lengths: " + str(ilens))
|
||||
elayer_states = []
|
||||
for layer in six.moves.range(self.elayers):
|
||||
if not isinstance(ilens, torch.Tensor):
|
||||
ilens = torch.tensor(ilens)
|
||||
xs_pack = pack_padded_sequence(xs_pad, ilens.cpu(), batch_first=True)
|
||||
rnn = getattr(self, ("birnn" if self.bidir else "rnn") + str(layer))
|
||||
rnn.flatten_parameters()
|
||||
if prev_state is not None and rnn.bidirectional:
|
||||
prev_state = reset_backward_rnn_state(prev_state)
|
||||
ys, states = rnn(xs_pack, hx=None if prev_state is None else prev_state[layer])
|
||||
elayer_states.append(states)
|
||||
# ys: utt list of frame x cdim x 2 (2: means bidirectional)
|
||||
ys_pad, ilens = pad_packed_sequence(ys, batch_first=True)
|
||||
sub = self.subsample[layer + 1]
|
||||
if sub > 1:
|
||||
ys_pad = ys_pad[:, ::sub]
|
||||
ilens = torch.tensor([int(i + 1) // sub for i in ilens])
|
||||
# (sum _utt frame_utt) x dim
|
||||
projection_layer = getattr(self, "bt%d" % layer)
|
||||
projected = projection_layer(ys_pad.contiguous().view(-1, ys_pad.size(2)))
|
||||
xs_pad = projected.view(ys_pad.size(0), ys_pad.size(1), -1)
|
||||
if layer < self.elayers - 1:
|
||||
xs_pad = torch.tanh(F.dropout(xs_pad, p=self.dropout))
|
||||
|
||||
return xs_pad, ilens, elayer_states # x: utt list of frame x dim
|
||||
|
||||
|
||||
class RNN(torch.nn.Module):
|
||||
"""RNN module
|
||||
|
||||
:param int idim: dimension of inputs
|
||||
:param int elayers: number of encoder layers
|
||||
:param int cdim: number of rnn units (resulted in cdim * 2 if bidirectional)
|
||||
:param int hdim: number of final projection units
|
||||
:param float dropout: dropout rate
|
||||
:param str typ: The RNN type
|
||||
"""
|
||||
|
||||
def __init__(self, idim, elayers, cdim, hdim, dropout, typ="blstm"):
|
||||
"""Initialize RNN.
|
||||
|
||||
Args:
|
||||
idim: TODO.
|
||||
elayers: TODO.
|
||||
cdim: TODO.
|
||||
hdim: TODO.
|
||||
dropout: TODO.
|
||||
typ: TODO.
|
||||
"""
|
||||
super(RNN, self).__init__()
|
||||
bidir = typ[0] == "b"
|
||||
self.nbrnn = (
|
||||
torch.nn.LSTM(
|
||||
idim,
|
||||
cdim,
|
||||
elayers,
|
||||
batch_first=True,
|
||||
dropout=dropout,
|
||||
bidirectional=bidir,
|
||||
)
|
||||
if "lstm" in typ
|
||||
else torch.nn.GRU(
|
||||
idim,
|
||||
cdim,
|
||||
elayers,
|
||||
batch_first=True,
|
||||
dropout=dropout,
|
||||
bidirectional=bidir,
|
||||
)
|
||||
)
|
||||
if bidir:
|
||||
self.l_last = torch.nn.Linear(cdim * 2, hdim)
|
||||
else:
|
||||
self.l_last = torch.nn.Linear(cdim, hdim)
|
||||
self.typ = typ
|
||||
|
||||
def forward(self, xs_pad, ilens, prev_state=None):
|
||||
"""RNN forward
|
||||
|
||||
:param torch.Tensor xs_pad: batch of padded input sequences (B, Tmax, D)
|
||||
:param torch.Tensor ilens: batch of lengths of input sequences (B)
|
||||
:param torch.Tensor prev_state: batch of previous RNN states
|
||||
:return: batch of hidden state sequences (B, Tmax, eprojs)
|
||||
:rtype: torch.Tensor
|
||||
"""
|
||||
logging.debug(self.__class__.__name__ + " input lengths: " + str(ilens))
|
||||
if not isinstance(ilens, torch.Tensor):
|
||||
ilens = torch.tensor(ilens)
|
||||
xs_pack = pack_padded_sequence(xs_pad, ilens.cpu(), batch_first=True)
|
||||
self.nbrnn.flatten_parameters()
|
||||
if prev_state is not None and self.nbrnn.bidirectional:
|
||||
# We assume that when previous state is passed,
|
||||
# it means that we're streaming the input
|
||||
# and therefore cannot propagate backward BRNN state
|
||||
# (otherwise it goes in the wrong direction)
|
||||
prev_state = reset_backward_rnn_state(prev_state)
|
||||
ys, states = self.nbrnn(xs_pack, hx=prev_state)
|
||||
# ys: utt list of frame x cdim x 2 (2: means bidirectional)
|
||||
ys_pad, ilens = pad_packed_sequence(ys, batch_first=True)
|
||||
# (sum _utt frame_utt) x dim
|
||||
projected = torch.tanh(self.l_last(ys_pad.contiguous().view(-1, ys_pad.size(2))))
|
||||
xs_pad = projected.view(ys_pad.size(0), ys_pad.size(1), -1)
|
||||
return xs_pad, ilens, states # x: utt list of frame x dim
|
||||
|
||||
|
||||
def reset_backward_rnn_state(states):
|
||||
"""Sets backward BRNN states to zeroes
|
||||
|
||||
Useful in processing of sliding windows over the inputs
|
||||
"""
|
||||
if isinstance(states, (list, tuple)):
|
||||
for state in states:
|
||||
state[1::2] = 0.0
|
||||
else:
|
||||
states[1::2] = 0.0
|
||||
return states
|
||||
|
||||
|
||||
class VGG2L(torch.nn.Module):
|
||||
"""VGG-like module
|
||||
|
||||
:param int in_channel: number of input channels
|
||||
"""
|
||||
|
||||
def __init__(self, in_channel=1):
|
||||
"""Initialize VGG2L.
|
||||
|
||||
Args:
|
||||
in_channel: TODO.
|
||||
"""
|
||||
super(VGG2L, self).__init__()
|
||||
# CNN layer (VGG motivated)
|
||||
self.conv1_1 = torch.nn.Conv2d(in_channel, 64, 3, stride=1, padding=1)
|
||||
self.conv1_2 = torch.nn.Conv2d(64, 64, 3, stride=1, padding=1)
|
||||
self.conv2_1 = torch.nn.Conv2d(64, 128, 3, stride=1, padding=1)
|
||||
self.conv2_2 = torch.nn.Conv2d(128, 128, 3, stride=1, padding=1)
|
||||
|
||||
self.in_channel = in_channel
|
||||
|
||||
def forward(self, xs_pad, ilens, **kwargs):
|
||||
"""VGG2L forward
|
||||
|
||||
:param torch.Tensor xs_pad: batch of padded input sequences (B, Tmax, D)
|
||||
:param torch.Tensor ilens: batch of lengths of input sequences (B)
|
||||
:return: batch of padded hidden state sequences (B, Tmax // 4, 128 * D // 4)
|
||||
:rtype: torch.Tensor
|
||||
"""
|
||||
logging.debug(self.__class__.__name__ + " input lengths: " + str(ilens))
|
||||
|
||||
# x: utt x frame x dim
|
||||
# xs_pad = F.pad_sequence(xs_pad)
|
||||
|
||||
# x: utt x 1 (input channel num) x frame x dim
|
||||
xs_pad = xs_pad.view(
|
||||
xs_pad.size(0),
|
||||
xs_pad.size(1),
|
||||
self.in_channel,
|
||||
xs_pad.size(2) // self.in_channel,
|
||||
).transpose(1, 2)
|
||||
|
||||
# NOTE: max_pool1d ?
|
||||
xs_pad = F.relu(self.conv1_1(xs_pad))
|
||||
xs_pad = F.relu(self.conv1_2(xs_pad))
|
||||
xs_pad = F.max_pool2d(xs_pad, 2, stride=2, ceil_mode=True)
|
||||
|
||||
xs_pad = F.relu(self.conv2_1(xs_pad))
|
||||
xs_pad = F.relu(self.conv2_2(xs_pad))
|
||||
xs_pad = F.max_pool2d(xs_pad, 2, stride=2, ceil_mode=True)
|
||||
if torch.is_tensor(ilens):
|
||||
ilens = ilens.cpu().numpy()
|
||||
else:
|
||||
ilens = np.array(ilens, dtype=np.float32)
|
||||
ilens = np.array(np.ceil(ilens / 2), dtype=np.int64)
|
||||
ilens = np.array(np.ceil(np.array(ilens, dtype=np.float32) / 2), dtype=np.int64).tolist()
|
||||
|
||||
# x: utt_list of frame (remove zeropaded frames) x (input channel num x dim)
|
||||
xs_pad = xs_pad.transpose(1, 2)
|
||||
xs_pad = xs_pad.contiguous().view(
|
||||
xs_pad.size(0), xs_pad.size(1), xs_pad.size(2) * xs_pad.size(3)
|
||||
)
|
||||
return xs_pad, ilens, None # no state in this layer
|
||||
|
||||
|
||||
class Encoder(torch.nn.Module):
|
||||
"""Encoder module
|
||||
|
||||
:param str etype: type of encoder network
|
||||
:param int idim: number of dimensions of encoder network
|
||||
:param int elayers: number of layers of encoder network
|
||||
:param int eunits: number of lstm units of encoder network
|
||||
:param int eprojs: number of projection units of encoder network
|
||||
:param np.ndarray subsample: list of subsampling numbers
|
||||
:param float dropout: dropout rate
|
||||
:param int in_channel: number of input channels
|
||||
"""
|
||||
|
||||
def __init__(self, etype, idim, elayers, eunits, eprojs, subsample, dropout, in_channel=1):
|
||||
"""Initialize Encoder.
|
||||
|
||||
Args:
|
||||
etype: TODO.
|
||||
idim: TODO.
|
||||
elayers: TODO.
|
||||
eunits: TODO.
|
||||
eprojs: TODO.
|
||||
subsample: TODO.
|
||||
dropout: TODO.
|
||||
in_channel: TODO.
|
||||
"""
|
||||
super(Encoder, self).__init__()
|
||||
typ = etype.lstrip("vgg").rstrip("p")
|
||||
if typ not in ["lstm", "gru", "blstm", "bgru"]:
|
||||
logging.error("Error: need to specify an appropriate encoder architecture")
|
||||
|
||||
if etype.startswith("vgg"):
|
||||
if etype[-1] == "p":
|
||||
self.enc = torch.nn.ModuleList(
|
||||
[
|
||||
VGG2L(in_channel),
|
||||
RNNP(
|
||||
get_vgg2l_odim(idim, in_channel=in_channel),
|
||||
elayers,
|
||||
eunits,
|
||||
eprojs,
|
||||
subsample,
|
||||
dropout,
|
||||
typ=typ,
|
||||
),
|
||||
]
|
||||
)
|
||||
logging.info("Use CNN-VGG + " + typ.upper() + "P for encoder")
|
||||
else:
|
||||
self.enc = torch.nn.ModuleList(
|
||||
[
|
||||
VGG2L(in_channel),
|
||||
RNN(
|
||||
get_vgg2l_odim(idim, in_channel=in_channel),
|
||||
elayers,
|
||||
eunits,
|
||||
eprojs,
|
||||
dropout,
|
||||
typ=typ,
|
||||
),
|
||||
]
|
||||
)
|
||||
logging.info("Use CNN-VGG + " + typ.upper() + " for encoder")
|
||||
self.conv_subsampling_factor = 4
|
||||
else:
|
||||
if etype[-1] == "p":
|
||||
self.enc = torch.nn.ModuleList(
|
||||
[RNNP(idim, elayers, eunits, eprojs, subsample, dropout, typ=typ)]
|
||||
)
|
||||
logging.info(typ.upper() + " with every-layer projection for encoder")
|
||||
else:
|
||||
self.enc = torch.nn.ModuleList(
|
||||
[RNN(idim, elayers, eunits, eprojs, dropout, typ=typ)]
|
||||
)
|
||||
logging.info(typ.upper() + " without projection for encoder")
|
||||
self.conv_subsampling_factor = 1
|
||||
|
||||
def forward(self, xs_pad, ilens, prev_states=None):
|
||||
"""Encoder forward
|
||||
|
||||
:param torch.Tensor xs_pad: batch of padded input sequences (B, Tmax, D)
|
||||
:param torch.Tensor ilens: batch of lengths of input sequences (B)
|
||||
:param torch.Tensor prev_state: batch of previous encoder hidden states (?, ...)
|
||||
:return: batch of hidden state sequences (B, Tmax, eprojs)
|
||||
:rtype: torch.Tensor
|
||||
"""
|
||||
if prev_states is None:
|
||||
prev_states = [None] * len(self.enc)
|
||||
assert len(prev_states) == len(self.enc)
|
||||
|
||||
current_states = []
|
||||
for module, prev_state in zip(self.enc, prev_states):
|
||||
xs_pad, ilens, states = module(xs_pad, ilens, prev_state=prev_state)
|
||||
current_states.append(states)
|
||||
|
||||
# make mask to remove bias value in padded part
|
||||
mask = to_device(xs_pad, make_pad_mask(ilens).unsqueeze(-1))
|
||||
|
||||
return xs_pad.masked_fill(mask, 0.0), ilens, current_states
|
||||
|
||||
|
||||
def encoder_for(args, idim, subsample):
|
||||
"""Instantiates an encoder module given the program arguments
|
||||
|
||||
:param Namespace args: The arguments
|
||||
:param int or List of integer idim: dimension of input, e.g. 83, or
|
||||
List of dimensions of inputs, e.g. [83,83]
|
||||
:param List or List of List subsample: subsample factors, e.g. [1,2,2,1,1], or
|
||||
List of subsample factors of each encoder.
|
||||
e.g. [[1,2,2,1,1], [1,2,2,1,1]]
|
||||
:rtype torch.nn.Module
|
||||
:return: The encoder module
|
||||
"""
|
||||
num_encs = getattr(args, "num_encs", 1) # use getattr to keep compatibility
|
||||
if num_encs == 1:
|
||||
# compatible with single encoder asr mode
|
||||
return Encoder(
|
||||
args.etype,
|
||||
idim,
|
||||
args.elayers,
|
||||
args.eunits,
|
||||
args.eprojs,
|
||||
subsample,
|
||||
args.dropout_rate,
|
||||
)
|
||||
elif num_encs >= 1:
|
||||
enc_list = torch.nn.ModuleList()
|
||||
for idx in range(num_encs):
|
||||
enc = Encoder(
|
||||
args.etype[idx],
|
||||
idim[idx],
|
||||
args.elayers[idx],
|
||||
args.eunits[idx],
|
||||
args.eprojs,
|
||||
subsample[idx],
|
||||
args.dropout_rate[idx],
|
||||
)
|
||||
enc_list.append(enc)
|
||||
return enc_list
|
||||
else:
|
||||
raise ValueError("Number of encoders needs to be more than one. {}".format(num_encs))
|
||||
@@ -0,0 +1,186 @@
|
||||
"""Sequential implementation of Recurrent Neural Network Language Model."""
|
||||
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from funasr.train.abs_model import AbsLM
|
||||
|
||||
|
||||
class SequentialRNNLM(AbsLM):
|
||||
"""Sequential RNNLM.
|
||||
|
||||
See also:
|
||||
https://github.com/pytorch/examples/blob/4581968193699de14b56527296262dd76ab43557/word_language_model/model.py
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
unit: int = 650,
|
||||
nhid: int = None,
|
||||
nlayers: int = 2,
|
||||
dropout_rate: float = 0.0,
|
||||
tie_weights: bool = False,
|
||||
rnn_type: str = "lstm",
|
||||
ignore_id: int = 0,
|
||||
):
|
||||
"""Initialize SequentialRNNLM.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
unit: TODO.
|
||||
nhid: TODO.
|
||||
nlayers: TODO.
|
||||
dropout_rate: TODO.
|
||||
tie_weights: TODO.
|
||||
rnn_type: TODO.
|
||||
ignore_id: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
ninp = unit
|
||||
if nhid is None:
|
||||
nhid = unit
|
||||
rnn_type = rnn_type.upper()
|
||||
|
||||
self.drop = nn.Dropout(dropout_rate)
|
||||
self.encoder = nn.Embedding(vocab_size, ninp, padding_idx=ignore_id)
|
||||
if rnn_type in ["LSTM", "GRU"]:
|
||||
rnn_class = getattr(nn, rnn_type)
|
||||
self.rnn = rnn_class(ninp, nhid, nlayers, dropout=dropout_rate, batch_first=True)
|
||||
else:
|
||||
try:
|
||||
nonlinearity = {"RNN_TANH": "tanh", "RNN_RELU": "relu"}[rnn_type]
|
||||
except KeyError:
|
||||
raise ValueError(
|
||||
"""An invalid option for `--model` was supplied,
|
||||
options are ['LSTM', 'GRU', 'RNN_TANH' or 'RNN_RELU']"""
|
||||
)
|
||||
self.rnn = nn.RNN(
|
||||
ninp,
|
||||
nhid,
|
||||
nlayers,
|
||||
nonlinearity=nonlinearity,
|
||||
dropout=dropout_rate,
|
||||
batch_first=True,
|
||||
)
|
||||
self.decoder = nn.Linear(nhid, vocab_size)
|
||||
|
||||
# Optionally tie weights as in:
|
||||
# "Using the Output Embedding to Improve Language Models"
|
||||
# (Press & Wolf 2016) https://arxiv.org/abs/1608.05859
|
||||
# and
|
||||
# "Tying Word Vectors and Word Classifiers:
|
||||
# A Loss Framework for Language Modeling" (Inan et al. 2016)
|
||||
# https://arxiv.org/abs/1611.01462
|
||||
if tie_weights:
|
||||
if nhid != ninp:
|
||||
raise ValueError("When using the tied flag, nhid must be equal to emsize")
|
||||
self.decoder.weight = self.encoder.weight
|
||||
|
||||
self.rnn_type = rnn_type
|
||||
self.nhid = nhid
|
||||
self.nlayers = nlayers
|
||||
|
||||
def zero_state(self):
|
||||
"""Initialize LM state filled with zero values."""
|
||||
if isinstance(self.rnn, torch.nn.LSTM):
|
||||
h = torch.zeros((self.nlayers, self.nhid), dtype=torch.float)
|
||||
c = torch.zeros((self.nlayers, self.nhid), dtype=torch.float)
|
||||
state = h, c
|
||||
else:
|
||||
state = torch.zeros((self.nlayers, self.nhid), dtype=torch.float)
|
||||
|
||||
return state
|
||||
|
||||
def forward(
|
||||
self, input: torch.Tensor, hidden: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
input: Input audio/text data.
|
||||
hidden: TODO.
|
||||
"""
|
||||
emb = self.drop(self.encoder(input))
|
||||
output, hidden = self.rnn(emb, hidden)
|
||||
output = self.drop(output)
|
||||
decoded = self.decoder(
|
||||
output.contiguous().view(output.size(0) * output.size(1), output.size(2))
|
||||
)
|
||||
return (
|
||||
decoded.view(output.size(0), output.size(1), decoded.size(1)),
|
||||
hidden,
|
||||
)
|
||||
|
||||
def score(
|
||||
self,
|
||||
y: torch.Tensor,
|
||||
state: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
|
||||
x: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]]:
|
||||
"""Score new token.
|
||||
|
||||
Args:
|
||||
y: 1D torch.int64 prefix tokens.
|
||||
state: Scorer state for prefix tokens
|
||||
x: 2D encoder feature that generates ys.
|
||||
|
||||
Returns:
|
||||
Tuple of
|
||||
torch.float32 scores for next token (n_vocab)
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
y, new_state = self(y[-1].view(1, 1), state)
|
||||
logp = y.log_softmax(dim=-1).view(-1)
|
||||
return logp, new_state
|
||||
|
||||
def batch_score(
|
||||
self, ys: torch.Tensor, states: torch.Tensor, xs: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Score new token batch.
|
||||
|
||||
Args:
|
||||
ys (torch.Tensor): torch.int64 prefix tokens (n_batch, ylen).
|
||||
states (List[Any]): Scorer states for prefix tokens.
|
||||
xs (torch.Tensor):
|
||||
The encoder feature that generates ys (n_batch, xlen, n_feat).
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, List[Any]]: Tuple of
|
||||
batchfied scores for next token with shape of `(n_batch, n_vocab)`
|
||||
and next state list for ys.
|
||||
|
||||
"""
|
||||
if states[0] is None:
|
||||
states = None
|
||||
elif isinstance(self.rnn, torch.nn.LSTM):
|
||||
# states: Batch x 2 x (Nlayers, Dim) -> 2 x (Nlayers, Batch, Dim)
|
||||
h = torch.stack([h for h, c in states], dim=1)
|
||||
c = torch.stack([c for h, c in states], dim=1)
|
||||
states = h, c
|
||||
else:
|
||||
# states: Batch x (Nlayers, Dim) -> (Nlayers, Batch, Dim)
|
||||
states = torch.stack(states, dim=1)
|
||||
|
||||
ys, states = self(ys[:, -1:], states)
|
||||
# ys: (Batch, 1, Nvocab) -> (Batch, NVocab)
|
||||
assert ys.size(1) == 1, ys.shape
|
||||
ys = ys.squeeze(1)
|
||||
logp = ys.log_softmax(dim=-1)
|
||||
|
||||
# state: Change to batch first
|
||||
if isinstance(self.rnn, torch.nn.LSTM):
|
||||
# h, c: (Nlayers, Batch, Dim)
|
||||
h, c = states
|
||||
# states: Batch x 2 x (Nlayers, Dim)
|
||||
states = [(h[:, i], c[:, i]) for i in range(h.size(1))]
|
||||
else:
|
||||
# states: (Nlayers, Batch, Dim) -> Batch x (Nlayers, Dim)
|
||||
states = [states[:, i] for i in range(states.size(1))]
|
||||
|
||||
return logp, states
|
||||
@@ -0,0 +1,459 @@
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Transformer encoder definition."""
|
||||
|
||||
from typing import List
|
||||
from typing import Optional
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import logging
|
||||
|
||||
from funasr.models.transformer.attention import MultiHeadedAttention
|
||||
from funasr.models.transformer.embedding import PositionalEncoding
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
from funasr.models.transformer.utils.multi_layer_conv import Conv1dLinear
|
||||
from funasr.models.transformer.utils.multi_layer_conv import MultiLayeredConv1d
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.transformer.positionwise_feed_forward import (
|
||||
PositionwiseFeedForward, # noqa: H301
|
||||
)
|
||||
from funasr.models.transformer.utils.repeat import repeat
|
||||
from funasr.models.transformer.utils.dynamic_conv import DynamicConvolution
|
||||
from funasr.models.transformer.utils.dynamic_conv2d import DynamicConvolution2D
|
||||
from funasr.models.transformer.utils.lightconv import LightweightConvolution
|
||||
from funasr.models.transformer.utils.lightconv2d import LightweightConvolution2D
|
||||
from funasr.models.transformer.utils.subsampling import Conv2dSubsampling
|
||||
from funasr.models.transformer.utils.subsampling import Conv2dSubsampling2
|
||||
from funasr.models.transformer.utils.subsampling import Conv2dSubsampling6
|
||||
from funasr.models.transformer.utils.subsampling import Conv2dSubsampling8
|
||||
from funasr.models.transformer.utils.subsampling import TooShortUttError
|
||||
from funasr.models.transformer.utils.subsampling import check_short_utt
|
||||
|
||||
|
||||
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.
|
||||
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)
|
||||
stochastic_depth_rate (float): Proability to skip this layer.
|
||||
During training, the layer may skip residual computation and return input
|
||||
as-is with given probability.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
self_attn,
|
||||
feed_forward,
|
||||
dropout_rate,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
stochastic_depth_rate=0.0,
|
||||
):
|
||||
"""Construct an EncoderLayer object."""
|
||||
super(EncoderLayer, self).__init__()
|
||||
self.self_attn = self_attn
|
||||
self.feed_forward = feed_forward
|
||||
self.norm1 = LayerNorm(size)
|
||||
self.norm2 = LayerNorm(size)
|
||||
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)
|
||||
self.stochastic_depth_rate = stochastic_depth_rate
|
||||
|
||||
def forward(self, x, mask, cache=None):
|
||||
"""Compute encoded features.
|
||||
|
||||
Args:
|
||||
x_input (torch.Tensor): Input 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).
|
||||
|
||||
"""
|
||||
skip_layer = False
|
||||
# with stochastic depth, residual connection `x + f(x)` becomes
|
||||
# `x <- x + 1 / (1 - p) * f(x)` at training time.
|
||||
stoch_layer_coeff = 1.0
|
||||
if self.training and self.stochastic_depth_rate > 0:
|
||||
skip_layer = torch.rand(1).item() < self.stochastic_depth_rate
|
||||
stoch_layer_coeff = 1.0 / (1 - self.stochastic_depth_rate)
|
||||
|
||||
if skip_layer:
|
||||
if cache is not None:
|
||||
x = torch.cat([cache, x], dim=1)
|
||||
return x, mask
|
||||
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm1(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 self.concat_after:
|
||||
x_concat = torch.cat((x, self.self_attn(x_q, x, x, mask)), dim=-1)
|
||||
x = residual + stoch_layer_coeff * self.concat_linear(x_concat)
|
||||
else:
|
||||
x = residual + stoch_layer_coeff * self.dropout(self.self_attn(x_q, x, x, mask))
|
||||
if not self.normalize_before:
|
||||
x = self.norm1(x)
|
||||
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm2(x)
|
||||
x = residual + stoch_layer_coeff * self.dropout(self.feed_forward(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm2(x)
|
||||
|
||||
if cache is not None:
|
||||
x = torch.cat([cache, x], dim=1)
|
||||
|
||||
return x, mask
|
||||
|
||||
|
||||
class TransformerEncoder_lm(nn.Module):
|
||||
"""Transformer encoder module.
|
||||
|
||||
Args:
|
||||
idim (int): Input dimension.
|
||||
attention_dim (int): Dimension of attention.
|
||||
attention_heads (int): The number of heads of multi head attention.
|
||||
conv_wshare (int): The number of kernel of convolution. Only used in
|
||||
selfattention_layer_type == "lightconv*" or "dynamiconv*".
|
||||
conv_kernel_length (Union[int, str]): Kernel size str of convolution
|
||||
(e.g. 71_71_71_71_71_71). Only used in selfattention_layer_type
|
||||
== "lightconv*" or "dynamiconv*".
|
||||
conv_usebias (bool): Whether to use bias in convolution. Only used in
|
||||
selfattention_layer_type == "lightconv*" or "dynamiconv*".
|
||||
linear_units (int): The number of units of position-wise feed forward.
|
||||
num_blocks (int): The number of decoder blocks.
|
||||
dropout_rate (float): Dropout rate.
|
||||
positional_dropout_rate (float): Dropout rate after adding positional encoding.
|
||||
attention_dropout_rate (float): Dropout rate in attention.
|
||||
input_layer (Union[str, torch.nn.Module]): Input layer type.
|
||||
pos_enc_class (torch.nn.Module): Positional encoding module class.
|
||||
`PositionalEncoding `or `ScaledPositionalEncoding`
|
||||
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)
|
||||
positionwise_layer_type (str): "linear", "conv1d", or "conv1d-linear".
|
||||
positionwise_conv_kernel_size (int): Kernel size of positionwise conv1d layer.
|
||||
selfattention_layer_type (str): Encoder attention layer type.
|
||||
padding_idx (int): Padding idx for input_layer=embed.
|
||||
stochastic_depth_rate (float): Maximum probability to skip the encoder layer.
|
||||
intermediate_layers (Union[List[int], None]): indices of intermediate CTC layer.
|
||||
indices start from 1.
|
||||
if not None, intermediate outputs are returned (which changes return type
|
||||
signature.)
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
idim,
|
||||
attention_dim=256,
|
||||
attention_heads=4,
|
||||
conv_wshare=4,
|
||||
conv_kernel_length="11",
|
||||
conv_usebias=False,
|
||||
linear_units=2048,
|
||||
num_blocks=6,
|
||||
dropout_rate=0.1,
|
||||
positional_dropout_rate=0.1,
|
||||
attention_dropout_rate=0.0,
|
||||
input_layer="conv2d",
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
positionwise_layer_type="linear",
|
||||
positionwise_conv_kernel_size=1,
|
||||
selfattention_layer_type="selfattn",
|
||||
padding_idx=-1,
|
||||
stochastic_depth_rate=0.0,
|
||||
intermediate_layers=None,
|
||||
ctc_softmax=None,
|
||||
conditioning_layer_dim=None,
|
||||
):
|
||||
"""Construct an Encoder object."""
|
||||
super().__init__()
|
||||
|
||||
self.conv_subsampling_factor = 1
|
||||
if input_layer == "linear":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Linear(idim, attention_dim),
|
||||
torch.nn.LayerNorm(attention_dim),
|
||||
torch.nn.Dropout(dropout_rate),
|
||||
torch.nn.ReLU(),
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
elif input_layer == "conv2d":
|
||||
self.embed = Conv2dSubsampling(idim, attention_dim, dropout_rate)
|
||||
self.conv_subsampling_factor = 4
|
||||
elif input_layer == "conv2d-scaled-pos-enc":
|
||||
self.embed = Conv2dSubsampling(
|
||||
idim,
|
||||
attention_dim,
|
||||
dropout_rate,
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
self.conv_subsampling_factor = 4
|
||||
elif input_layer == "conv2d6":
|
||||
self.embed = Conv2dSubsampling6(idim, attention_dim, dropout_rate)
|
||||
self.conv_subsampling_factor = 6
|
||||
elif input_layer == "conv2d8":
|
||||
self.embed = Conv2dSubsampling8(idim, attention_dim, dropout_rate)
|
||||
self.conv_subsampling_factor = 8
|
||||
elif input_layer == "embed":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Embedding(idim, attention_dim, padding_idx=padding_idx),
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
elif isinstance(input_layer, torch.nn.Module):
|
||||
self.embed = torch.nn.Sequential(
|
||||
input_layer,
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
elif input_layer is None:
|
||||
self.embed = torch.nn.Sequential(pos_enc_class(attention_dim, positional_dropout_rate))
|
||||
else:
|
||||
raise ValueError("unknown input_layer: " + input_layer)
|
||||
self.normalize_before = normalize_before
|
||||
positionwise_layer, positionwise_layer_args = self.get_positionwise_layer(
|
||||
positionwise_layer_type,
|
||||
attention_dim,
|
||||
linear_units,
|
||||
dropout_rate,
|
||||
positionwise_conv_kernel_size,
|
||||
)
|
||||
if selfattention_layer_type in [
|
||||
"selfattn",
|
||||
"rel_selfattn",
|
||||
"legacy_rel_selfattn",
|
||||
]:
|
||||
logging.info("encoder self-attention layer type = self-attention")
|
||||
encoder_selfattn_layer = MultiHeadedAttention
|
||||
encoder_selfattn_layer_args = [
|
||||
(
|
||||
attention_heads,
|
||||
attention_dim,
|
||||
attention_dropout_rate,
|
||||
)
|
||||
] * num_blocks
|
||||
elif selfattention_layer_type == "lightconv":
|
||||
logging.info("encoder self-attention layer type = lightweight convolution")
|
||||
encoder_selfattn_layer = LightweightConvolution
|
||||
encoder_selfattn_layer_args = [
|
||||
(
|
||||
conv_wshare,
|
||||
attention_dim,
|
||||
attention_dropout_rate,
|
||||
int(conv_kernel_length.split("_")[lnum]),
|
||||
False,
|
||||
conv_usebias,
|
||||
)
|
||||
for lnum in range(num_blocks)
|
||||
]
|
||||
elif selfattention_layer_type == "lightconv2d":
|
||||
logging.info(
|
||||
"encoder self-attention layer " "type = lightweight convolution 2-dimensional"
|
||||
)
|
||||
encoder_selfattn_layer = LightweightConvolution2D
|
||||
encoder_selfattn_layer_args = [
|
||||
(
|
||||
conv_wshare,
|
||||
attention_dim,
|
||||
attention_dropout_rate,
|
||||
int(conv_kernel_length.split("_")[lnum]),
|
||||
False,
|
||||
conv_usebias,
|
||||
)
|
||||
for lnum in range(num_blocks)
|
||||
]
|
||||
elif selfattention_layer_type == "dynamicconv":
|
||||
logging.info("encoder self-attention layer type = dynamic convolution")
|
||||
encoder_selfattn_layer = DynamicConvolution
|
||||
encoder_selfattn_layer_args = [
|
||||
(
|
||||
conv_wshare,
|
||||
attention_dim,
|
||||
attention_dropout_rate,
|
||||
int(conv_kernel_length.split("_")[lnum]),
|
||||
False,
|
||||
conv_usebias,
|
||||
)
|
||||
for lnum in range(num_blocks)
|
||||
]
|
||||
elif selfattention_layer_type == "dynamicconv2d":
|
||||
logging.info("encoder self-attention layer type = dynamic convolution 2-dimensional")
|
||||
encoder_selfattn_layer = DynamicConvolution2D
|
||||
encoder_selfattn_layer_args = [
|
||||
(
|
||||
conv_wshare,
|
||||
attention_dim,
|
||||
attention_dropout_rate,
|
||||
int(conv_kernel_length.split("_")[lnum]),
|
||||
False,
|
||||
conv_usebias,
|
||||
)
|
||||
for lnum in range(num_blocks)
|
||||
]
|
||||
else:
|
||||
raise NotImplementedError(selfattention_layer_type)
|
||||
|
||||
self.encoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: EncoderLayer(
|
||||
attention_dim,
|
||||
encoder_selfattn_layer(*encoder_selfattn_layer_args[lnum]),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
stochastic_depth_rate * float(1 + lnum) / num_blocks,
|
||||
),
|
||||
)
|
||||
if self.normalize_before:
|
||||
self.after_norm = LayerNorm(attention_dim)
|
||||
|
||||
self.intermediate_layers = intermediate_layers
|
||||
self.use_conditioning = True if ctc_softmax is not None else False
|
||||
if self.use_conditioning:
|
||||
self.ctc_softmax = ctc_softmax
|
||||
self.conditioning_layer = torch.nn.Linear(conditioning_layer_dim, attention_dim)
|
||||
|
||||
def get_positionwise_layer(
|
||||
self,
|
||||
positionwise_layer_type="linear",
|
||||
attention_dim=256,
|
||||
linear_units=2048,
|
||||
dropout_rate=0.1,
|
||||
positionwise_conv_kernel_size=1,
|
||||
):
|
||||
"""Define positionwise layer."""
|
||||
if positionwise_layer_type == "linear":
|
||||
positionwise_layer = PositionwiseFeedForward
|
||||
positionwise_layer_args = (attention_dim, linear_units, dropout_rate)
|
||||
elif positionwise_layer_type == "conv1d":
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (
|
||||
attention_dim,
|
||||
linear_units,
|
||||
positionwise_conv_kernel_size,
|
||||
dropout_rate,
|
||||
)
|
||||
elif positionwise_layer_type == "conv1d-linear":
|
||||
positionwise_layer = Conv1dLinear
|
||||
positionwise_layer_args = (
|
||||
attention_dim,
|
||||
linear_units,
|
||||
positionwise_conv_kernel_size,
|
||||
dropout_rate,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError("Support only linear or conv1d.")
|
||||
return positionwise_layer, positionwise_layer_args
|
||||
|
||||
def forward(self, xs, masks):
|
||||
"""Encode input sequence.
|
||||
|
||||
Args:
|
||||
xs (torch.Tensor): Input tensor (#batch, time, idim).
|
||||
masks (torch.Tensor): Mask tensor (#batch, time).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time, attention_dim).
|
||||
torch.Tensor: Mask tensor (#batch, time).
|
||||
|
||||
"""
|
||||
if isinstance(
|
||||
self.embed,
|
||||
(Conv2dSubsampling, Conv2dSubsampling6, Conv2dSubsampling8),
|
||||
):
|
||||
xs, masks = self.embed(xs, masks)
|
||||
else:
|
||||
xs = self.embed(xs)
|
||||
|
||||
if self.intermediate_layers is None:
|
||||
xs, masks = self.encoders(xs, masks)
|
||||
else:
|
||||
intermediate_outputs = []
|
||||
for layer_idx, encoder_layer in enumerate(self.encoders):
|
||||
xs, masks = encoder_layer(xs, masks)
|
||||
|
||||
if (
|
||||
self.intermediate_layers is not None
|
||||
and layer_idx + 1 in self.intermediate_layers
|
||||
):
|
||||
encoder_output = xs
|
||||
# intermediate branches also require normalization.
|
||||
if self.normalize_before:
|
||||
encoder_output = self.after_norm(encoder_output)
|
||||
intermediate_outputs.append(encoder_output)
|
||||
|
||||
if self.use_conditioning:
|
||||
intermediate_result = self.ctc_softmax(encoder_output)
|
||||
xs = xs + self.conditioning_layer(intermediate_result)
|
||||
|
||||
if self.normalize_before:
|
||||
xs = self.after_norm(xs)
|
||||
|
||||
if self.intermediate_layers is not None:
|
||||
return xs, masks, intermediate_outputs
|
||||
return xs, masks
|
||||
|
||||
def forward_one_step(self, xs, masks, cache=None):
|
||||
"""Encode input frame.
|
||||
|
||||
Args:
|
||||
xs (torch.Tensor): Input tensor.
|
||||
masks (torch.Tensor): Mask tensor.
|
||||
cache (List[torch.Tensor]): List of cache tensors.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor.
|
||||
torch.Tensor: Mask tensor.
|
||||
List[torch.Tensor]: List of new cache tensors.
|
||||
|
||||
"""
|
||||
if isinstance(self.embed, Conv2dSubsampling):
|
||||
xs, masks = self.embed(xs, masks)
|
||||
else:
|
||||
xs = self.embed(xs)
|
||||
if cache is None:
|
||||
cache = [None for _ in range(len(self.encoders))]
|
||||
new_cache = []
|
||||
for c, e in zip(cache, self.encoders):
|
||||
xs, masks = e(xs, masks, cache=c)
|
||||
new_cache.append(xs)
|
||||
if self.normalize_before:
|
||||
xs = self.after_norm(xs)
|
||||
return xs, masks, new_cache
|
||||
@@ -0,0 +1,151 @@
|
||||
from typing import Any
|
||||
from typing import List
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from funasr.models.transformer.embedding import PositionalEncoding
|
||||
from funasr.models.encoder.transformer_encoder import TransformerEncoder_s0 as Encoder
|
||||
from funasr.models.transformer.utils.mask import subsequent_mask
|
||||
from funasr.train.abs_model import AbsLM
|
||||
|
||||
|
||||
class TransformerLM(AbsLM):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
pos_enc: str = None,
|
||||
embed_unit: int = 128,
|
||||
att_unit: int = 256,
|
||||
head: int = 2,
|
||||
unit: int = 1024,
|
||||
layer: int = 4,
|
||||
dropout_rate: float = 0.5,
|
||||
):
|
||||
"""Initialize TransformerLM.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
pos_enc: TODO.
|
||||
embed_unit: TODO.
|
||||
att_unit: TODO.
|
||||
head: TODO.
|
||||
unit: TODO.
|
||||
layer: TODO.
|
||||
dropout_rate: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
if pos_enc == "sinusoidal":
|
||||
pos_enc_class = PositionalEncoding
|
||||
elif pos_enc is None:
|
||||
|
||||
def pos_enc_class(*args, **kwargs):
|
||||
"""Pos enc class.
|
||||
|
||||
Args:
|
||||
*args: Variable positional arguments.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
return nn.Sequential() # indentity
|
||||
|
||||
else:
|
||||
raise ValueError(f"unknown pos-enc option: {pos_enc}")
|
||||
|
||||
self.embed = nn.Embedding(vocab_size, embed_unit)
|
||||
self.encoder = Encoder(
|
||||
idim=embed_unit,
|
||||
attention_dim=att_unit,
|
||||
attention_heads=head,
|
||||
linear_units=unit,
|
||||
num_blocks=layer,
|
||||
dropout_rate=dropout_rate,
|
||||
input_layer="linear",
|
||||
pos_enc_class=pos_enc_class,
|
||||
)
|
||||
self.decoder = nn.Linear(att_unit, vocab_size)
|
||||
|
||||
def _target_mask(self, ys_in_pad):
|
||||
"""Internal: target mask.
|
||||
|
||||
Args:
|
||||
ys_in_pad: TODO.
|
||||
"""
|
||||
ys_mask = ys_in_pad != 0
|
||||
m = subsequent_mask(ys_mask.size(-1), device=ys_mask.device).unsqueeze(0)
|
||||
return ys_mask.unsqueeze(-2) & m
|
||||
|
||||
def forward(self, input: torch.Tensor, hidden: None) -> Tuple[torch.Tensor, None]:
|
||||
"""Compute LM loss value from buffer sequences.
|
||||
|
||||
Args:
|
||||
input (torch.Tensor): Input ids. (batch, len)
|
||||
hidden (torch.Tensor): Target ids. (batch, len)
|
||||
|
||||
"""
|
||||
x = self.embed(input)
|
||||
mask = self._target_mask(input)
|
||||
h, _ = self.encoder(x, mask)
|
||||
y = self.decoder(h)
|
||||
return y, None
|
||||
|
||||
def score(self, y: torch.Tensor, state: Any, x: torch.Tensor) -> Tuple[torch.Tensor, Any]:
|
||||
"""Score new token.
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D torch.int64 prefix tokens.
|
||||
state: Scorer state for prefix tokens
|
||||
x (torch.Tensor): encoder feature that generates ys.
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]: Tuple of
|
||||
torch.float32 scores for next token (vocab_size)
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
y = y.unsqueeze(0)
|
||||
h, _, cache = self.encoder.forward_one_step(
|
||||
self.embed(y), self._target_mask(y), cache=state
|
||||
)
|
||||
h = self.decoder(h[:, -1])
|
||||
logp = h.log_softmax(dim=-1).squeeze(0)
|
||||
return logp, cache
|
||||
|
||||
def batch_score(
|
||||
self, ys: torch.Tensor, states: List[Any], xs: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, List[Any]]:
|
||||
"""Score new token batch.
|
||||
|
||||
Args:
|
||||
ys (torch.Tensor): torch.int64 prefix tokens (n_batch, ylen).
|
||||
states (List[Any]): Scorer states for prefix tokens.
|
||||
xs (torch.Tensor):
|
||||
The encoder feature that generates ys (n_batch, xlen, n_feat).
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, List[Any]]: Tuple of
|
||||
batchfied scores for next token with shape of `(n_batch, vocab_size)`
|
||||
and next state list for ys.
|
||||
|
||||
"""
|
||||
# merge states
|
||||
n_batch = len(ys)
|
||||
n_layers = len(self.encoder.encoders)
|
||||
if states[0] is None:
|
||||
batch_state = None
|
||||
else:
|
||||
# transpose state of [batch, layer] into [layer, batch]
|
||||
batch_state = [
|
||||
torch.stack([states[b][i] for b in range(n_batch)]) for i in range(n_layers)
|
||||
]
|
||||
|
||||
# batch decoding
|
||||
h, _, states = self.encoder.forward_one_step(
|
||||
self.embed(ys), self._target_mask(ys), cache=batch_state
|
||||
)
|
||||
h = self.decoder(h[:, -1])
|
||||
logp = h.log_softmax(dim=-1)
|
||||
|
||||
# transpose state of [layer, batch] into [batch, layer]
|
||||
state_list = [[states[i][b] for i in range(n_layers)] for b in range(n_batch)]
|
||||
return logp, state_list
|
||||
Reference in New Issue
Block a user