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,671 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Multi-Head Attention layer definition."""
|
||||
|
||||
import math
|
||||
|
||||
import numpy
|
||||
import torch
|
||||
from torch import nn
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch.nn.functional as F
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
import funasr.models.lora.layers as lora
|
||||
|
||||
|
||||
class MultiHeadedAttention(nn.Module):
|
||||
"""Multi-Head Attention layer.
|
||||
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate):
|
||||
"""Construct an MultiHeadedAttention object."""
|
||||
super(MultiHeadedAttention, self).__init__()
|
||||
assert n_feat % n_head == 0
|
||||
# We assume d_v always equals d_k
|
||||
self.d_k = n_feat // n_head
|
||||
self.h = n_head
|
||||
self.linear_q = nn.Linear(n_feat, n_feat)
|
||||
self.linear_k = nn.Linear(n_feat, n_feat)
|
||||
self.linear_v = nn.Linear(n_feat, n_feat)
|
||||
self.linear_out = nn.Linear(n_feat, n_feat)
|
||||
self.attn = None
|
||||
self.dropout = nn.Dropout(p=dropout_rate)
|
||||
|
||||
def forward_qkv(self, query, key, value):
|
||||
"""Transform query, key and value.
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Transformed query tensor (#batch, n_head, time1, d_k).
|
||||
torch.Tensor: Transformed key tensor (#batch, n_head, time2, d_k).
|
||||
torch.Tensor: Transformed value tensor (#batch, n_head, time2, d_k).
|
||||
|
||||
"""
|
||||
n_batch = query.size(0)
|
||||
q = self.linear_q(query).view(n_batch, -1, self.h, self.d_k)
|
||||
k = self.linear_k(key).view(n_batch, -1, self.h, self.d_k)
|
||||
v = self.linear_v(value).view(n_batch, -1, self.h, self.d_k)
|
||||
q = q.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
k = k.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
v = v.transpose(1, 2) # (batch, head, time2, d_k)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def forward_attention(self, value, scores, mask):
|
||||
"""Compute attention context vector.
|
||||
|
||||
Args:
|
||||
value (torch.Tensor): Transformed value (#batch, n_head, time2, d_k).
|
||||
scores (torch.Tensor): Attention score (#batch, n_head, time1, time2).
|
||||
mask (torch.Tensor): Mask (#batch, 1, time2) or (#batch, time1, time2).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Transformed value (#batch, time1, d_model)
|
||||
weighted by the attention score (#batch, time1, time2).
|
||||
|
||||
"""
|
||||
n_batch = value.size(0)
|
||||
if mask is not None:
|
||||
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
|
||||
|
||||
min_value = -float(
|
||||
"inf"
|
||||
) # min_value = float(np.finfo(torch.tensor(0, dtype=qk.dtype).numpy().dtype).min)
|
||||
scores = scores.masked_fill(mask, min_value)
|
||||
attn = torch.softmax(scores, dim=-1).masked_fill(
|
||||
mask, 0.0
|
||||
) # (batch, head, time1, time2)
|
||||
else:
|
||||
attn = torch.softmax(scores, dim=-1) # (batch, head, time1, time2)
|
||||
|
||||
p_attn = self.dropout(attn)
|
||||
x = torch.matmul(p_attn, value) # (batch, head, time1, d_k)
|
||||
x = (
|
||||
x.transpose(1, 2).contiguous().view(n_batch, -1, self.h * self.d_k)
|
||||
) # (batch, time1, d_model)
|
||||
|
||||
return self.linear_out(x) # (batch, time1, d_model)
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Compute scaled dot product attention.
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
|
||||
class MultiHeadedAttentionExport(nn.Module):
|
||||
def __init__(self, model):
|
||||
"""Initialize MultiHeadedAttentionExport.
|
||||
|
||||
Args:
|
||||
model: Model instance or model name.
|
||||
"""
|
||||
super().__init__()
|
||||
self.d_k = model.d_k
|
||||
self.h = model.h
|
||||
self.linear_q = model.linear_q
|
||||
self.linear_k = model.linear_k
|
||||
self.linear_v = model.linear_v
|
||||
self.linear_out = model.linear_out
|
||||
self.attn = None
|
||||
self.all_head_size = self.h * self.d_k
|
||||
|
||||
def forward(self, query, key, value, mask):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
query: TODO.
|
||||
key: Sample identifiers.
|
||||
value: TODO.
|
||||
mask: TODO.
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Transpose for scores.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
new_x_shape = x.size()[:-1] + (self.h, self.d_k)
|
||||
x = x.view(new_x_shape)
|
||||
return x.permute(0, 2, 1, 3)
|
||||
|
||||
def forward_qkv(self, query, key, value):
|
||||
"""Forward qkv.
|
||||
|
||||
Args:
|
||||
query: TODO.
|
||||
key: Sample identifiers.
|
||||
value: TODO.
|
||||
"""
|
||||
q = self.linear_q(query)
|
||||
k = self.linear_k(key)
|
||||
v = self.linear_v(value)
|
||||
q = self.transpose_for_scores(q)
|
||||
k = self.transpose_for_scores(k)
|
||||
v = self.transpose_for_scores(v)
|
||||
return q, k, v
|
||||
|
||||
def forward_attention(self, value, scores, mask):
|
||||
"""Forward attention.
|
||||
|
||||
Args:
|
||||
value: TODO.
|
||||
scores: TODO.
|
||||
mask: TODO.
|
||||
"""
|
||||
scores = scores + mask
|
||||
|
||||
attn = torch.softmax(scores, dim=-1)
|
||||
context_layer = torch.matmul(attn, value) # (batch, head, time1, d_k)
|
||||
|
||||
context_layer = context_layer.permute(0, 2, 1, 3).contiguous()
|
||||
new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)
|
||||
context_layer = context_layer.view(new_context_layer_shape)
|
||||
return self.linear_out(context_layer) # (batch, time1, d_model)
|
||||
|
||||
|
||||
class RelPosMultiHeadedAttentionExport(MultiHeadedAttentionExport):
|
||||
def __init__(self, model):
|
||||
"""Initialize RelPosMultiHeadedAttentionExport.
|
||||
|
||||
Args:
|
||||
model: Model instance or model name.
|
||||
"""
|
||||
super().__init__(model)
|
||||
self.linear_pos = model.linear_pos
|
||||
self.pos_bias_u = model.pos_bias_u
|
||||
self.pos_bias_v = model.pos_bias_v
|
||||
|
||||
def forward(self, query, key, value, pos_emb, mask):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
query: TODO.
|
||||
key: Sample identifiers.
|
||||
value: TODO.
|
||||
pos_emb: TODO.
|
||||
mask: TODO.
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
q = q.transpose(1, 2) # (batch, time1, head, d_k)
|
||||
|
||||
p = self.transpose_for_scores(self.linear_pos(pos_emb)) # (batch, head, time1, d_k)
|
||||
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
# compute attention score
|
||||
# first compute matrix a and matrix c
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
# (batch, head, time1, time2)
|
||||
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
|
||||
|
||||
# compute matrix b and matrix d
|
||||
# (batch, head, time1, time1)
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
|
||||
matrix_bd = self.rel_shift(matrix_bd)
|
||||
|
||||
scores = (matrix_ac + matrix_bd) / math.sqrt(self.d_k) # (batch, head, time1, time2)
|
||||
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
def rel_shift(self, x):
|
||||
"""Rel shift.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
|
||||
x_padded = torch.cat([zero_pad, x], dim=-1)
|
||||
|
||||
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
|
||||
x = x_padded[:, :, 1:].view_as(x)[
|
||||
:, :, :, : x.size(-1) // 2 + 1
|
||||
] # only keep the positions from 0 to time2
|
||||
return x
|
||||
|
||||
def forward_attention(self, value, scores, mask):
|
||||
"""Forward attention.
|
||||
|
||||
Args:
|
||||
value: TODO.
|
||||
scores: TODO.
|
||||
mask: TODO.
|
||||
"""
|
||||
scores = scores + mask
|
||||
|
||||
attn = torch.softmax(scores, dim=-1)
|
||||
context_layer = torch.matmul(attn, value) # (batch, head, time1, d_k)
|
||||
|
||||
context_layer = context_layer.permute(0, 2, 1, 3).contiguous()
|
||||
new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)
|
||||
context_layer = context_layer.view(new_context_layer_shape)
|
||||
return self.linear_out(context_layer) # (batch, time1, d_model)
|
||||
|
||||
|
||||
class LegacyRelPositionMultiHeadedAttention(MultiHeadedAttention):
|
||||
"""Multi-Head Attention layer with relative position encoding (old version).
|
||||
|
||||
Details can be found in https://github.com/espnet/espnet/pull/2816.
|
||||
|
||||
Paper: https://arxiv.org/abs/1901.02860
|
||||
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
zero_triu (bool): Whether to zero the upper triangular part of attention matrix.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate, zero_triu=False):
|
||||
"""Construct an RelPositionMultiHeadedAttention object."""
|
||||
super().__init__(n_head, n_feat, dropout_rate)
|
||||
self.zero_triu = zero_triu
|
||||
# linear transformation for positional encoding
|
||||
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
|
||||
# these two learnable bias are used in matrix c and matrix d
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_u)
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_v)
|
||||
|
||||
def rel_shift(self, x):
|
||||
"""Compute relative positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, head, time1, time2).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor.
|
||||
|
||||
"""
|
||||
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
|
||||
x_padded = torch.cat([zero_pad, x], dim=-1)
|
||||
|
||||
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
|
||||
x = x_padded[:, :, 1:].view_as(x)
|
||||
|
||||
if self.zero_triu:
|
||||
ones = torch.ones((x.size(2), x.size(3)))
|
||||
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
|
||||
|
||||
return x
|
||||
|
||||
def forward(self, query, key, value, pos_emb, mask):
|
||||
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
pos_emb (torch.Tensor): Positional embedding tensor (#batch, time1, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
q = q.transpose(1, 2) # (batch, time1, head, d_k)
|
||||
|
||||
n_batch_pos = pos_emb.size(0)
|
||||
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
|
||||
p = p.transpose(1, 2) # (batch, head, time1, d_k)
|
||||
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
# compute attention score
|
||||
# first compute matrix a and matrix c
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
# (batch, head, time1, time2)
|
||||
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
|
||||
|
||||
# compute matrix b and matrix d
|
||||
# (batch, head, time1, time1)
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
|
||||
matrix_bd = self.rel_shift(matrix_bd)
|
||||
|
||||
scores = (matrix_ac + matrix_bd) / math.sqrt(self.d_k) # (batch, head, time1, time2)
|
||||
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
|
||||
class RelPositionMultiHeadedAttention(MultiHeadedAttention):
|
||||
"""Multi-Head Attention layer with relative position encoding (new implementation).
|
||||
|
||||
Details can be found in https://github.com/espnet/espnet/pull/2816.
|
||||
|
||||
Paper: https://arxiv.org/abs/1901.02860
|
||||
|
||||
Args:
|
||||
n_head (int): The number of heads.
|
||||
n_feat (int): The number of features.
|
||||
dropout_rate (float): Dropout rate.
|
||||
zero_triu (bool): Whether to zero the upper triangular part of attention matrix.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, n_head, n_feat, dropout_rate, zero_triu=False):
|
||||
"""Construct an RelPositionMultiHeadedAttention object."""
|
||||
super().__init__(n_head, n_feat, dropout_rate)
|
||||
self.zero_triu = zero_triu
|
||||
# linear transformation for positional encoding
|
||||
self.linear_pos = nn.Linear(n_feat, n_feat, bias=False)
|
||||
# these two learnable bias are used in matrix c and matrix d
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
self.pos_bias_u = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
self.pos_bias_v = nn.Parameter(torch.Tensor(self.h, self.d_k))
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_u)
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_v)
|
||||
|
||||
def rel_shift(self, x):
|
||||
"""Compute relative positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, head, time1, 2*time1-1).
|
||||
time1 means the length of query vector.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor.
|
||||
|
||||
"""
|
||||
zero_pad = torch.zeros((*x.size()[:3], 1), device=x.device, dtype=x.dtype)
|
||||
x_padded = torch.cat([zero_pad, x], dim=-1)
|
||||
|
||||
x_padded = x_padded.view(*x.size()[:2], x.size(3) + 1, x.size(2))
|
||||
x = x_padded[:, :, 1:].view_as(x)[
|
||||
:, :, :, : x.size(-1) // 2 + 1
|
||||
] # only keep the positions from 0 to time2
|
||||
|
||||
if self.zero_triu:
|
||||
ones = torch.ones((x.size(2), x.size(3)), device=x.device)
|
||||
x = x * torch.tril(ones, x.size(3) - x.size(2))[None, None, :, :]
|
||||
|
||||
return x
|
||||
|
||||
def forward(self, query, key, value, pos_emb, mask):
|
||||
"""Compute 'Scaled Dot Product Attention' with rel. positional encoding.
|
||||
|
||||
Args:
|
||||
query (torch.Tensor): Query tensor (#batch, time1, size).
|
||||
key (torch.Tensor): Key tensor (#batch, time2, size).
|
||||
value (torch.Tensor): Value tensor (#batch, time2, size).
|
||||
pos_emb (torch.Tensor): Positional embedding tensor
|
||||
(#batch, 2*time1-1, size).
|
||||
mask (torch.Tensor): Mask tensor (#batch, 1, time2) or
|
||||
(#batch, time1, time2).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor (#batch, time1, d_model).
|
||||
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
q = q.transpose(1, 2) # (batch, time1, head, d_k)
|
||||
|
||||
n_batch_pos = pos_emb.size(0)
|
||||
p = self.linear_pos(pos_emb).view(n_batch_pos, -1, self.h, self.d_k)
|
||||
p = p.transpose(1, 2) # (batch, head, 2*time1-1, d_k)
|
||||
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_u = (q + self.pos_bias_u).transpose(1, 2)
|
||||
# (batch, head, time1, d_k)
|
||||
q_with_bias_v = (q + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
# compute attention score
|
||||
# first compute matrix a and matrix c
|
||||
# as described in https://arxiv.org/abs/1901.02860 Section 3.3
|
||||
# (batch, head, time1, time2)
|
||||
matrix_ac = torch.matmul(q_with_bias_u, k.transpose(-2, -1))
|
||||
|
||||
# compute matrix b and matrix d
|
||||
# (batch, head, time1, 2*time1-1)
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.transpose(-2, -1))
|
||||
matrix_bd = self.rel_shift(matrix_bd)
|
||||
|
||||
scores = (matrix_ac + matrix_bd) / math.sqrt(self.d_k) # (batch, head, time1, time2)
|
||||
|
||||
return self.forward_attention(v, scores, mask)
|
||||
|
||||
|
||||
class RelPositionMultiHeadedAttentionChunk(torch.nn.Module):
|
||||
"""RelPositionMultiHeadedAttention definition.
|
||||
Args:
|
||||
num_heads: Number of attention heads.
|
||||
embed_size: Embedding size.
|
||||
dropout_rate: Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
embed_size: int,
|
||||
dropout_rate: float = 0.0,
|
||||
simplified_attention_score: bool = False,
|
||||
) -> None:
|
||||
"""Construct an MultiHeadedAttention object."""
|
||||
super().__init__()
|
||||
|
||||
self.d_k = embed_size // num_heads
|
||||
self.num_heads = num_heads
|
||||
|
||||
assert self.d_k * num_heads == embed_size, (
|
||||
"embed_size (%d) must be divisible by num_heads (%d)",
|
||||
(embed_size, num_heads),
|
||||
)
|
||||
|
||||
self.linear_q = torch.nn.Linear(embed_size, embed_size)
|
||||
self.linear_k = torch.nn.Linear(embed_size, embed_size)
|
||||
self.linear_v = torch.nn.Linear(embed_size, embed_size)
|
||||
|
||||
self.linear_out = torch.nn.Linear(embed_size, embed_size)
|
||||
|
||||
if simplified_attention_score:
|
||||
self.linear_pos = torch.nn.Linear(embed_size, num_heads)
|
||||
|
||||
self.compute_att_score = self.compute_simplified_attention_score
|
||||
else:
|
||||
self.linear_pos = torch.nn.Linear(embed_size, embed_size, bias=False)
|
||||
|
||||
self.pos_bias_u = torch.nn.Parameter(torch.Tensor(num_heads, self.d_k))
|
||||
self.pos_bias_v = torch.nn.Parameter(torch.Tensor(num_heads, self.d_k))
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_u)
|
||||
torch.nn.init.xavier_uniform_(self.pos_bias_v)
|
||||
|
||||
self.compute_att_score = self.compute_attention_score
|
||||
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.attn = None
|
||||
|
||||
def rel_shift(self, x: torch.Tensor, left_context: int = 0) -> torch.Tensor:
|
||||
"""Compute relative positional encoding.
|
||||
Args:
|
||||
x: Input sequence. (B, H, T_1, 2 * T_1 - 1)
|
||||
left_context: Number of frames in left context.
|
||||
Returns:
|
||||
x: Output sequence. (B, H, T_1, T_2)
|
||||
"""
|
||||
batch_size, n_heads, time1, n = x.shape
|
||||
time2 = time1 + left_context
|
||||
|
||||
batch_stride, n_heads_stride, time1_stride, n_stride = x.stride()
|
||||
|
||||
return x.as_strided(
|
||||
(batch_size, n_heads, time1, time2),
|
||||
(batch_stride, n_heads_stride, time1_stride - n_stride, n_stride),
|
||||
storage_offset=(n_stride * (time1 - 1)),
|
||||
)
|
||||
|
||||
def compute_simplified_attention_score(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
pos_enc: torch.Tensor,
|
||||
left_context: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Simplified attention score computation.
|
||||
Reference: https://github.com/k2-fsa/icefall/pull/458
|
||||
Args:
|
||||
query: Transformed query tensor. (B, H, T_1, d_k)
|
||||
key: Transformed key tensor. (B, H, T_2, d_k)
|
||||
pos_enc: Positional embedding tensor. (B, 2 * T_1 - 1, size)
|
||||
left_context: Number of frames in left context.
|
||||
Returns:
|
||||
: Attention score. (B, H, T_1, T_2)
|
||||
"""
|
||||
pos_enc = self.linear_pos(pos_enc)
|
||||
|
||||
matrix_ac = torch.matmul(query, key.transpose(2, 3))
|
||||
|
||||
matrix_bd = self.rel_shift(
|
||||
pos_enc.transpose(1, 2).unsqueeze(2).repeat(1, 1, query.size(2), 1),
|
||||
left_context=left_context,
|
||||
)
|
||||
|
||||
return (matrix_ac + matrix_bd) / math.sqrt(self.d_k)
|
||||
|
||||
def compute_attention_score(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
pos_enc: torch.Tensor,
|
||||
left_context: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Attention score computation.
|
||||
Args:
|
||||
query: Transformed query tensor. (B, H, T_1, d_k)
|
||||
key: Transformed key tensor. (B, H, T_2, d_k)
|
||||
pos_enc: Positional embedding tensor. (B, 2 * T_1 - 1, size)
|
||||
left_context: Number of frames in left context.
|
||||
Returns:
|
||||
: Attention score. (B, H, T_1, T_2)
|
||||
"""
|
||||
p = self.linear_pos(pos_enc).view(pos_enc.size(0), -1, self.num_heads, self.d_k)
|
||||
|
||||
query = query.transpose(1, 2)
|
||||
q_with_bias_u = (query + self.pos_bias_u).transpose(1, 2)
|
||||
q_with_bias_v = (query + self.pos_bias_v).transpose(1, 2)
|
||||
|
||||
matrix_ac = torch.matmul(q_with_bias_u, key.transpose(-2, -1))
|
||||
|
||||
matrix_bd = torch.matmul(q_with_bias_v, p.permute(0, 2, 3, 1))
|
||||
matrix_bd = self.rel_shift(matrix_bd, left_context=left_context)
|
||||
|
||||
return (matrix_ac + matrix_bd) / math.sqrt(self.d_k)
|
||||
|
||||
def forward_qkv(
|
||||
self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Transform query, key and value.
|
||||
Args:
|
||||
query: Query tensor. (B, T_1, size)
|
||||
key: Key tensor. (B, T_2, size)
|
||||
v: Value tensor. (B, T_2, size)
|
||||
Returns:
|
||||
q: Transformed query tensor. (B, H, T_1, d_k)
|
||||
k: Transformed key tensor. (B, H, T_2, d_k)
|
||||
v: Transformed value tensor. (B, H, T_2, d_k)
|
||||
"""
|
||||
n_batch = query.size(0)
|
||||
|
||||
q = self.linear_q(query).view(n_batch, -1, self.num_heads, self.d_k).transpose(1, 2)
|
||||
k = self.linear_k(key).view(n_batch, -1, self.num_heads, self.d_k).transpose(1, 2)
|
||||
v = self.linear_v(value).view(n_batch, -1, self.num_heads, self.d_k).transpose(1, 2)
|
||||
|
||||
return q, k, v
|
||||
|
||||
def forward_attention(
|
||||
self,
|
||||
value: torch.Tensor,
|
||||
scores: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
chunk_mask: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Compute attention context vector.
|
||||
Args:
|
||||
value: Transformed value. (B, H, T_2, d_k)
|
||||
scores: Attention score. (B, H, T_1, T_2)
|
||||
mask: Source mask. (B, T_2)
|
||||
chunk_mask: Chunk mask. (T_1, T_1)
|
||||
Returns:
|
||||
attn_output: Transformed value weighted by attention score. (B, T_1, H * d_k)
|
||||
"""
|
||||
batch_size = scores.size(0)
|
||||
mask = mask.unsqueeze(1).unsqueeze(2)
|
||||
if chunk_mask is not None:
|
||||
mask = chunk_mask.unsqueeze(0).unsqueeze(1) | mask
|
||||
scores = scores.masked_fill(mask, float("-inf"))
|
||||
attn = torch.softmax(scores, dim=-1).masked_fill(mask, 0.0)
|
||||
|
||||
attn_output = self.dropout(attn)
|
||||
attn_output = torch.matmul(attn_output, value)
|
||||
|
||||
attn_output = self.linear_out(
|
||||
attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.num_heads * self.d_k)
|
||||
)
|
||||
|
||||
return attn_output
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
pos_enc: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
chunk_mask: Optional[torch.Tensor] = None,
|
||||
left_context: int = 0,
|
||||
) -> torch.Tensor:
|
||||
"""Compute scaled dot product attention with rel. positional encoding.
|
||||
Args:
|
||||
query: Query tensor. (B, T_1, size)
|
||||
key: Key tensor. (B, T_2, size)
|
||||
value: Value tensor. (B, T_2, size)
|
||||
pos_enc: Positional embedding tensor. (B, 2 * T_1 - 1, size)
|
||||
mask: Source mask. (B, T_2)
|
||||
chunk_mask: Chunk mask. (T_1, T_1)
|
||||
left_context: Number of frames in left context.
|
||||
Returns:
|
||||
: Output tensor. (B, T_1, H * d_k)
|
||||
"""
|
||||
q, k, v = self.forward_qkv(query, key, value)
|
||||
scores = self.compute_att_score(q, k, pos_enc, left_context=left_context)
|
||||
return self.forward_attention(v, scores, mask, chunk_mask=chunk_mask)
|
||||
@@ -0,0 +1,781 @@
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Decoder definition."""
|
||||
from typing import Any
|
||||
from typing import List
|
||||
from typing import Sequence
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
from funasr.models.transformer.attention import MultiHeadedAttention
|
||||
from funasr.models.transformer.utils.dynamic_conv import DynamicConvolution
|
||||
from funasr.models.transformer.utils.dynamic_conv2d import DynamicConvolution2D
|
||||
from funasr.models.transformer.embedding import PositionalEncoding
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
from funasr.models.transformer.utils.lightconv import LightweightConvolution
|
||||
from funasr.models.transformer.utils.lightconv2d import LightweightConvolution2D
|
||||
from funasr.models.transformer.utils.mask import subsequent_mask
|
||||
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.scorers.scorer_interface import BatchScorerInterface
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
class DecoderLayer(nn.Module):
|
||||
"""Single decoder layer module.
|
||||
|
||||
Args:
|
||||
size (int): Input dimension.
|
||||
self_attn (torch.nn.Module): Self-attention module instance.
|
||||
`MultiHeadedAttention` instance can be used as the argument.
|
||||
src_attn (torch.nn.Module): Self-attention module instance.
|
||||
`MultiHeadedAttention` 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)
|
||||
|
||||
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
self_attn,
|
||||
src_attn,
|
||||
feed_forward,
|
||||
dropout_rate,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
):
|
||||
"""Construct an DecoderLayer object."""
|
||||
super(DecoderLayer, self).__init__()
|
||||
self.size = size
|
||||
self.self_attn = self_attn
|
||||
self.src_attn = src_attn
|
||||
self.feed_forward = feed_forward
|
||||
self.norm1 = LayerNorm(size)
|
||||
self.norm2 = LayerNorm(size)
|
||||
self.norm3 = LayerNorm(size)
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.normalize_before = normalize_before
|
||||
self.concat_after = concat_after
|
||||
if self.concat_after:
|
||||
self.concat_linear1 = nn.Linear(size + size, size)
|
||||
self.concat_linear2 = nn.Linear(size + size, size)
|
||||
|
||||
def forward(self, tgt, tgt_mask, memory, memory_mask, cache=None):
|
||||
"""Compute decoded features.
|
||||
|
||||
Args:
|
||||
tgt (torch.Tensor): Input tensor (#batch, maxlen_out, size).
|
||||
tgt_mask (torch.Tensor): Mask for input tensor (#batch, maxlen_out).
|
||||
memory (torch.Tensor): Encoded memory, float32 (#batch, maxlen_in, size).
|
||||
memory_mask (torch.Tensor): Encoded memory mask (#batch, maxlen_in).
|
||||
cache (List[torch.Tensor]): List of cached tensors.
|
||||
Each tensor shape should be (#batch, maxlen_out - 1, size).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor(#batch, maxlen_out, size).
|
||||
torch.Tensor: Mask for output tensor (#batch, maxlen_out).
|
||||
torch.Tensor: Encoded memory (#batch, maxlen_in, size).
|
||||
torch.Tensor: Encoded memory mask (#batch, maxlen_in).
|
||||
|
||||
"""
|
||||
residual = tgt
|
||||
if self.normalize_before:
|
||||
tgt = self.norm1(tgt)
|
||||
|
||||
if cache is None:
|
||||
tgt_q = tgt
|
||||
tgt_q_mask = tgt_mask
|
||||
else:
|
||||
# compute only the last frame query keeping dim: max_time_out -> 1
|
||||
assert cache.shape == (
|
||||
tgt.shape[0],
|
||||
tgt.shape[1] - 1,
|
||||
self.size,
|
||||
), f"{cache.shape} == {(tgt.shape[0], tgt.shape[1] - 1, self.size)}"
|
||||
tgt_q = tgt[:, -1:, :]
|
||||
residual = residual[:, -1:, :]
|
||||
tgt_q_mask = None
|
||||
if tgt_mask is not None:
|
||||
tgt_q_mask = tgt_mask[:, -1:, :]
|
||||
|
||||
if self.concat_after:
|
||||
tgt_concat = torch.cat((tgt_q, self.self_attn(tgt_q, tgt, tgt, tgt_q_mask)), dim=-1)
|
||||
x = residual + self.concat_linear1(tgt_concat)
|
||||
else:
|
||||
x = residual + self.dropout(self.self_attn(tgt_q, tgt, tgt, tgt_q_mask))
|
||||
if not self.normalize_before:
|
||||
x = self.norm1(x)
|
||||
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm2(x)
|
||||
if self.concat_after:
|
||||
x_concat = torch.cat((x, self.src_attn(x, memory, memory, memory_mask)), dim=-1)
|
||||
x = residual + self.concat_linear2(x_concat)
|
||||
else:
|
||||
x = residual + self.dropout(self.src_attn(x, memory, memory, memory_mask))
|
||||
if not self.normalize_before:
|
||||
x = self.norm2(x)
|
||||
|
||||
residual = x
|
||||
if self.normalize_before:
|
||||
x = self.norm3(x)
|
||||
x = residual + self.dropout(self.feed_forward(x))
|
||||
if not self.normalize_before:
|
||||
x = self.norm3(x)
|
||||
|
||||
if cache is not None:
|
||||
x = torch.cat([cache, x], dim=1)
|
||||
|
||||
return x, tgt_mask, memory, memory_mask
|
||||
|
||||
|
||||
class DecoderLayerExport(nn.Module):
|
||||
def __init__(self, model):
|
||||
"""Initialize DecoderLayerExport.
|
||||
|
||||
Args:
|
||||
model: Model instance or model name.
|
||||
"""
|
||||
super().__init__()
|
||||
self.self_attn = model.self_attn
|
||||
self.src_attn = model.src_attn
|
||||
self.feed_forward = model.feed_forward
|
||||
self.norm1 = model.norm1
|
||||
self.norm2 = model.norm2
|
||||
self.norm3 = model.norm3
|
||||
|
||||
def forward(self, tgt, tgt_mask, memory, memory_mask, cache=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
tgt: TODO.
|
||||
tgt_mask: TODO.
|
||||
memory: TODO.
|
||||
memory_mask: TODO.
|
||||
cache: State cache dict for streaming inference.
|
||||
"""
|
||||
residual = tgt
|
||||
tgt = self.norm1(tgt)
|
||||
tgt_q = tgt
|
||||
tgt_q_mask = tgt_mask
|
||||
x = residual + self.self_attn(tgt_q, tgt, tgt, tgt_q_mask)
|
||||
|
||||
residual = x
|
||||
x = self.norm2(x)
|
||||
|
||||
x = residual + self.src_attn(x, memory, memory, memory_mask)
|
||||
|
||||
residual = x
|
||||
x = self.norm3(x)
|
||||
x = residual + self.feed_forward(x)
|
||||
|
||||
return x, tgt_mask, memory, memory_mask
|
||||
|
||||
|
||||
class BaseTransformerDecoder(nn.Module, BatchScorerInterface):
|
||||
"""Base class of Transfomer decoder module.
|
||||
|
||||
Args:
|
||||
vocab_size: output dim
|
||||
encoder_output_size: dimension of attention
|
||||
attention_heads: the number of heads of multi head attention
|
||||
linear_units: the number of units of position-wise feed forward
|
||||
num_blocks: the number of decoder blocks
|
||||
dropout_rate: dropout rate
|
||||
self_attention_dropout_rate: dropout rate for attention
|
||||
input_layer: input layer type
|
||||
use_output_layer: whether to use output layer
|
||||
pos_enc_class: PositionalEncoding or ScaledPositionalEncoding
|
||||
normalize_before: whether to use layer_norm before the first block
|
||||
concat_after: whether to concat attention layer's input and output
|
||||
if True, additional linear will be applied.
|
||||
i.e. x -> x + linear(concat(x, att(x)))
|
||||
if False, no additional linear will be applied.
|
||||
i.e. x -> x + att(x)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
):
|
||||
"""Initialize BaseTransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
attention_dim = encoder_output_size
|
||||
|
||||
if input_layer == "embed":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Embedding(vocab_size, attention_dim),
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
elif input_layer == "linear":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Linear(vocab_size, attention_dim),
|
||||
torch.nn.LayerNorm(attention_dim),
|
||||
torch.nn.Dropout(dropout_rate),
|
||||
torch.nn.ReLU(),
|
||||
pos_enc_class(attention_dim, positional_dropout_rate),
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"only 'embed' or 'linear' is supported: {input_layer}")
|
||||
|
||||
self.normalize_before = normalize_before
|
||||
if self.normalize_before:
|
||||
self.after_norm = LayerNorm(attention_dim)
|
||||
if use_output_layer:
|
||||
self.output_layer = torch.nn.Linear(attention_dim, vocab_size)
|
||||
else:
|
||||
self.output_layer = None
|
||||
|
||||
# Must set by the inheritance
|
||||
self.decoders = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hs_pad: torch.Tensor,
|
||||
hlens: torch.Tensor,
|
||||
ys_in_pad: torch.Tensor,
|
||||
ys_in_lens: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward decoder.
|
||||
|
||||
Args:
|
||||
hs_pad: encoded memory, float32 (batch, maxlen_in, feat)
|
||||
hlens: (batch)
|
||||
ys_in_pad:
|
||||
input token ids, int64 (batch, maxlen_out)
|
||||
if input_layer == "embed"
|
||||
input tensor (batch, maxlen_out, #mels) in the other cases
|
||||
ys_in_lens: (batch)
|
||||
Returns:
|
||||
(tuple): tuple containing:
|
||||
|
||||
x: decoded token score before softmax (batch, maxlen_out, token)
|
||||
if use_output_layer is True,
|
||||
olens: (batch, )
|
||||
"""
|
||||
tgt = ys_in_pad
|
||||
# tgt_mask: (B, 1, L)
|
||||
tgt_mask = (~make_pad_mask(ys_in_lens)[:, None, :]).to(tgt.device)
|
||||
# m: (1, L, L)
|
||||
m = subsequent_mask(tgt_mask.size(-1), device=tgt_mask.device).unsqueeze(0)
|
||||
# tgt_mask: (B, L, L)
|
||||
tgt_mask = tgt_mask & m
|
||||
|
||||
memory = hs_pad
|
||||
memory_mask = (~make_pad_mask(hlens, maxlen=memory.size(1)))[:, None, :].to(memory.device)
|
||||
# Padding for Longformer
|
||||
if memory_mask.shape[-1] != memory.shape[1]:
|
||||
padlen = memory.shape[1] - memory_mask.shape[-1]
|
||||
memory_mask = torch.nn.functional.pad(memory_mask, (0, padlen), "constant", False)
|
||||
|
||||
x = self.embed(tgt)
|
||||
x, tgt_mask, memory, memory_mask = self.decoders(x, tgt_mask, memory, memory_mask)
|
||||
if self.normalize_before:
|
||||
x = self.after_norm(x)
|
||||
if self.output_layer is not None:
|
||||
x = self.output_layer(x)
|
||||
|
||||
olens = tgt_mask.sum(1)
|
||||
return x, olens
|
||||
|
||||
def forward_one_step(
|
||||
self,
|
||||
tgt: torch.Tensor,
|
||||
tgt_mask: torch.Tensor,
|
||||
memory: torch.Tensor,
|
||||
cache: List[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, List[torch.Tensor]]:
|
||||
"""Forward one step.
|
||||
|
||||
Args:
|
||||
tgt: input token ids, int64 (batch, maxlen_out)
|
||||
tgt_mask: input token mask, (batch, maxlen_out)
|
||||
dtype=torch.uint8 in PyTorch 1.2-
|
||||
dtype=torch.bool in PyTorch 1.2+ (include 1.2)
|
||||
memory: encoded memory, float32 (batch, maxlen_in, feat)
|
||||
cache: cached output list of (batch, max_time_out-1, size)
|
||||
Returns:
|
||||
y, cache: NN output value and cache per `self.decoders`.
|
||||
y.shape` is (batch, maxlen_out, token)
|
||||
"""
|
||||
x = self.embed(tgt)
|
||||
if cache is None:
|
||||
cache = [None] * len(self.decoders)
|
||||
new_cache = []
|
||||
for c, decoder in zip(cache, self.decoders):
|
||||
x, tgt_mask, memory, memory_mask = decoder(x, tgt_mask, memory, None, cache=c)
|
||||
new_cache.append(x)
|
||||
|
||||
if self.normalize_before:
|
||||
y = self.after_norm(x[:, -1])
|
||||
else:
|
||||
y = x[:, -1]
|
||||
if self.output_layer is not None:
|
||||
y = torch.log_softmax(self.output_layer(y), dim=-1)
|
||||
|
||||
return y, new_cache
|
||||
|
||||
def score(self, ys, state, x):
|
||||
"""Score."""
|
||||
ys_mask = subsequent_mask(len(ys), device=x.device).unsqueeze(0)
|
||||
logp, state = self.forward_one_step(ys.unsqueeze(0), ys_mask, x.unsqueeze(0), cache=state)
|
||||
return logp.squeeze(0), state
|
||||
|
||||
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, n_vocab)`
|
||||
and next state list for ys.
|
||||
|
||||
"""
|
||||
# merge states
|
||||
n_batch = len(ys)
|
||||
n_layers = len(self.decoders)
|
||||
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
|
||||
ys_mask = subsequent_mask(ys.size(-1), device=xs.device).unsqueeze(0)
|
||||
logp, states = self.forward_one_step(ys, ys_mask, xs, cache=batch_state)
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
@tables.register("decoder_classes", "TransformerDecoder")
|
||||
class TransformerDecoder(BaseTransformerDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
self_attention_dropout_rate: float = 0.0,
|
||||
src_attention_dropout_rate: float = 0.0,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
):
|
||||
"""Initialize TransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
self_attention_dropout_rate: TODO.
|
||||
src_attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
"""
|
||||
super().__init__(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
dropout_rate=dropout_rate,
|
||||
positional_dropout_rate=positional_dropout_rate,
|
||||
input_layer=input_layer,
|
||||
use_output_layer=use_output_layer,
|
||||
pos_enc_class=pos_enc_class,
|
||||
normalize_before=normalize_before,
|
||||
)
|
||||
|
||||
attention_dim = encoder_output_size
|
||||
self.decoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: DecoderLayer(
|
||||
attention_dim,
|
||||
MultiHeadedAttention(attention_heads, attention_dim, self_attention_dropout_rate),
|
||||
MultiHeadedAttention(attention_heads, attention_dim, src_attention_dropout_rate),
|
||||
PositionwiseFeedForward(attention_dim, linear_units, dropout_rate),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@tables.register("decoder_classes", "LightweightConvolutionTransformerDecoder")
|
||||
class LightweightConvolutionTransformerDecoder(BaseTransformerDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
self_attention_dropout_rate: float = 0.0,
|
||||
src_attention_dropout_rate: float = 0.0,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
conv_wshare: int = 4,
|
||||
conv_kernel_length: Sequence[int] = (11, 11, 11, 11, 11, 11),
|
||||
conv_usebias: int = False,
|
||||
):
|
||||
"""Initialize LightweightConvolutionTransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
self_attention_dropout_rate: TODO.
|
||||
src_attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
conv_wshare: TODO.
|
||||
conv_kernel_length: TODO.
|
||||
conv_usebias: TODO.
|
||||
"""
|
||||
if len(conv_kernel_length) != num_blocks:
|
||||
raise ValueError(
|
||||
"conv_kernel_length must have equal number of values to num_blocks: "
|
||||
f"{len(conv_kernel_length)} != {num_blocks}"
|
||||
)
|
||||
super().__init__(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
dropout_rate=dropout_rate,
|
||||
positional_dropout_rate=positional_dropout_rate,
|
||||
input_layer=input_layer,
|
||||
use_output_layer=use_output_layer,
|
||||
pos_enc_class=pos_enc_class,
|
||||
normalize_before=normalize_before,
|
||||
)
|
||||
|
||||
attention_dim = encoder_output_size
|
||||
self.decoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: DecoderLayer(
|
||||
attention_dim,
|
||||
LightweightConvolution(
|
||||
wshare=conv_wshare,
|
||||
n_feat=attention_dim,
|
||||
dropout_rate=self_attention_dropout_rate,
|
||||
kernel_size=conv_kernel_length[lnum],
|
||||
use_kernel_mask=True,
|
||||
use_bias=conv_usebias,
|
||||
),
|
||||
MultiHeadedAttention(attention_heads, attention_dim, src_attention_dropout_rate),
|
||||
PositionwiseFeedForward(attention_dim, linear_units, dropout_rate),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@tables.register("decoder_classes", "LightweightConvolution2DTransformerDecoder")
|
||||
class LightweightConvolution2DTransformerDecoder(BaseTransformerDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
self_attention_dropout_rate: float = 0.0,
|
||||
src_attention_dropout_rate: float = 0.0,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
conv_wshare: int = 4,
|
||||
conv_kernel_length: Sequence[int] = (11, 11, 11, 11, 11, 11),
|
||||
conv_usebias: int = False,
|
||||
):
|
||||
"""Initialize LightweightConvolution2DTransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
self_attention_dropout_rate: TODO.
|
||||
src_attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
conv_wshare: TODO.
|
||||
conv_kernel_length: TODO.
|
||||
conv_usebias: TODO.
|
||||
"""
|
||||
if len(conv_kernel_length) != num_blocks:
|
||||
raise ValueError(
|
||||
"conv_kernel_length must have equal number of values to num_blocks: "
|
||||
f"{len(conv_kernel_length)} != {num_blocks}"
|
||||
)
|
||||
super().__init__(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
dropout_rate=dropout_rate,
|
||||
positional_dropout_rate=positional_dropout_rate,
|
||||
input_layer=input_layer,
|
||||
use_output_layer=use_output_layer,
|
||||
pos_enc_class=pos_enc_class,
|
||||
normalize_before=normalize_before,
|
||||
)
|
||||
|
||||
attention_dim = encoder_output_size
|
||||
self.decoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: DecoderLayer(
|
||||
attention_dim,
|
||||
LightweightConvolution2D(
|
||||
wshare=conv_wshare,
|
||||
n_feat=attention_dim,
|
||||
dropout_rate=self_attention_dropout_rate,
|
||||
kernel_size=conv_kernel_length[lnum],
|
||||
use_kernel_mask=True,
|
||||
use_bias=conv_usebias,
|
||||
),
|
||||
MultiHeadedAttention(attention_heads, attention_dim, src_attention_dropout_rate),
|
||||
PositionwiseFeedForward(attention_dim, linear_units, dropout_rate),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@tables.register("decoder_classes", "DynamicConvolutionTransformerDecoder")
|
||||
class DynamicConvolutionTransformerDecoder(BaseTransformerDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
self_attention_dropout_rate: float = 0.0,
|
||||
src_attention_dropout_rate: float = 0.0,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
conv_wshare: int = 4,
|
||||
conv_kernel_length: Sequence[int] = (11, 11, 11, 11, 11, 11),
|
||||
conv_usebias: int = False,
|
||||
):
|
||||
"""Initialize DynamicConvolutionTransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
self_attention_dropout_rate: TODO.
|
||||
src_attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
conv_wshare: TODO.
|
||||
conv_kernel_length: TODO.
|
||||
conv_usebias: TODO.
|
||||
"""
|
||||
if len(conv_kernel_length) != num_blocks:
|
||||
raise ValueError(
|
||||
"conv_kernel_length must have equal number of values to num_blocks: "
|
||||
f"{len(conv_kernel_length)} != {num_blocks}"
|
||||
)
|
||||
super().__init__(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
dropout_rate=dropout_rate,
|
||||
positional_dropout_rate=positional_dropout_rate,
|
||||
input_layer=input_layer,
|
||||
use_output_layer=use_output_layer,
|
||||
pos_enc_class=pos_enc_class,
|
||||
normalize_before=normalize_before,
|
||||
)
|
||||
attention_dim = encoder_output_size
|
||||
|
||||
self.decoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: DecoderLayer(
|
||||
attention_dim,
|
||||
DynamicConvolution(
|
||||
wshare=conv_wshare,
|
||||
n_feat=attention_dim,
|
||||
dropout_rate=self_attention_dropout_rate,
|
||||
kernel_size=conv_kernel_length[lnum],
|
||||
use_kernel_mask=True,
|
||||
use_bias=conv_usebias,
|
||||
),
|
||||
MultiHeadedAttention(attention_heads, attention_dim, src_attention_dropout_rate),
|
||||
PositionwiseFeedForward(attention_dim, linear_units, dropout_rate),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@tables.register("decoder_classes", "DynamicConvolution2DTransformerDecoder")
|
||||
class DynamicConvolution2DTransformerDecoder(BaseTransformerDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
encoder_output_size: int,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
self_attention_dropout_rate: float = 0.0,
|
||||
src_attention_dropout_rate: float = 0.0,
|
||||
input_layer: str = "embed",
|
||||
use_output_layer: bool = True,
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
conv_wshare: int = 4,
|
||||
conv_kernel_length: Sequence[int] = (11, 11, 11, 11, 11, 11),
|
||||
conv_usebias: int = False,
|
||||
):
|
||||
"""Initialize DynamicConvolution2DTransformerDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
self_attention_dropout_rate: TODO.
|
||||
src_attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
use_output_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
conv_wshare: TODO.
|
||||
conv_kernel_length: TODO.
|
||||
conv_usebias: TODO.
|
||||
"""
|
||||
if len(conv_kernel_length) != num_blocks:
|
||||
raise ValueError(
|
||||
"conv_kernel_length must have equal number of values to num_blocks: "
|
||||
f"{len(conv_kernel_length)} != {num_blocks}"
|
||||
)
|
||||
super().__init__(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
dropout_rate=dropout_rate,
|
||||
positional_dropout_rate=positional_dropout_rate,
|
||||
input_layer=input_layer,
|
||||
use_output_layer=use_output_layer,
|
||||
pos_enc_class=pos_enc_class,
|
||||
normalize_before=normalize_before,
|
||||
)
|
||||
attention_dim = encoder_output_size
|
||||
|
||||
self.decoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: DecoderLayer(
|
||||
attention_dim,
|
||||
DynamicConvolution2D(
|
||||
wshare=conv_wshare,
|
||||
n_feat=attention_dim,
|
||||
dropout_rate=self_attention_dropout_rate,
|
||||
kernel_size=conv_kernel_length[lnum],
|
||||
use_kernel_mask=True,
|
||||
use_bias=conv_usebias,
|
||||
),
|
||||
MultiHeadedAttention(attention_heads, attention_dim, src_attention_dropout_rate),
|
||||
PositionwiseFeedForward(attention_dim, linear_units, dropout_rate),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,581 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Positional Encoding Module."""
|
||||
|
||||
import math
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import einsum
|
||||
|
||||
|
||||
def _pre_hook(
|
||||
state_dict,
|
||||
prefix,
|
||||
local_metadata,
|
||||
strict,
|
||||
missing_keys,
|
||||
unexpected_keys,
|
||||
error_msgs,
|
||||
):
|
||||
"""Perform pre-hook in load_state_dict for backward compatibility.
|
||||
|
||||
Note:
|
||||
We saved self.pe until v.0.5.2 but we have omitted it later.
|
||||
Therefore, we remove the item "pe" from `state_dict` for backward compatibility.
|
||||
|
||||
"""
|
||||
k = prefix + "pe"
|
||||
if k in state_dict:
|
||||
state_dict.pop(k)
|
||||
|
||||
|
||||
class PositionalEncoding(torch.nn.Module):
|
||||
"""Positional encoding.
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
reverse (bool): Whether to reverse the input position. Only for
|
||||
the class LegacyRelPositionalEncoding. We remove it in the current
|
||||
class RelPositionalEncoding.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000, reverse=False):
|
||||
"""Construct an PositionalEncoding object."""
|
||||
super(PositionalEncoding, self).__init__()
|
||||
self.d_model = d_model
|
||||
self.reverse = reverse
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.pe = None
|
||||
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
|
||||
self._register_load_state_dict_pre_hook(_pre_hook)
|
||||
|
||||
def extend_pe(self, x):
|
||||
"""Reset the positional encodings."""
|
||||
if self.pe is not None:
|
||||
if self.pe.size(1) >= x.size(1):
|
||||
if self.pe.dtype != x.dtype or self.pe.device != x.device:
|
||||
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
|
||||
return
|
||||
pe = torch.zeros(x.size(1), self.d_model)
|
||||
if self.reverse:
|
||||
position = torch.arange(x.size(1) - 1, -1, -1.0, dtype=torch.float32).unsqueeze(1)
|
||||
else:
|
||||
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.d_model, 2, dtype=torch.float32)
|
||||
* -(math.log(10000.0) / self.d_model)
|
||||
)
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
pe = pe.unsqueeze(0)
|
||||
self.pe = pe.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
"""Add positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale + self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class ScaledPositionalEncoding(PositionalEncoding):
|
||||
"""Scaled positional encoding module.
|
||||
|
||||
See Sec. 3.2 https://arxiv.org/abs/1809.08895
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(d_model=d_model, dropout_rate=dropout_rate, max_len=max_len)
|
||||
self.alpha = torch.nn.Parameter(torch.tensor(1.0))
|
||||
|
||||
def reset_parameters(self):
|
||||
"""Reset parameters."""
|
||||
self.alpha.data = torch.tensor(1.0)
|
||||
|
||||
def forward(self, x):
|
||||
"""Add positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x + self.alpha * self.pe[:, : x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class LearnableFourierPosEnc(torch.nn.Module):
|
||||
"""Learnable Fourier Features for Positional Encoding.
|
||||
|
||||
See https://arxiv.org/pdf/2106.02795.pdf
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
gamma (float): init parameter for the positional kernel variance
|
||||
see https://arxiv.org/pdf/2106.02795.pdf.
|
||||
apply_scaling (bool): Whether to scale the input before adding the pos encoding.
|
||||
hidden_dim (int): if not None, we modulate the pos encodings with
|
||||
an MLP whose hidden layer has hidden_dim neurons.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model,
|
||||
dropout_rate=0.0,
|
||||
max_len=5000,
|
||||
gamma=1.0,
|
||||
apply_scaling=False,
|
||||
hidden_dim=None,
|
||||
):
|
||||
"""Initialize class."""
|
||||
super(LearnableFourierPosEnc, self).__init__()
|
||||
|
||||
self.d_model = d_model
|
||||
|
||||
if apply_scaling:
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
else:
|
||||
self.xscale = 1.0
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout_rate)
|
||||
self.max_len = max_len
|
||||
|
||||
self.gamma = gamma
|
||||
if self.gamma is None:
|
||||
self.gamma = self.d_model // 2
|
||||
|
||||
assert d_model % 2 == 0, "d_model should be divisible by two in order to use this layer."
|
||||
self.w_r = torch.nn.Parameter(torch.empty(1, d_model // 2))
|
||||
self._reset() # init the weights
|
||||
|
||||
self.hidden_dim = hidden_dim
|
||||
if self.hidden_dim is not None:
|
||||
self.mlp = torch.nn.Sequential(
|
||||
torch.nn.Linear(d_model, hidden_dim),
|
||||
torch.nn.GELU(),
|
||||
torch.nn.Linear(hidden_dim, d_model),
|
||||
)
|
||||
|
||||
def _reset(self):
|
||||
"""Internal: reset."""
|
||||
self.w_r.data = torch.normal(0, (1 / math.sqrt(self.gamma)), (1, self.d_model // 2))
|
||||
|
||||
def extend_pe(self, x):
|
||||
"""Reset the positional encodings."""
|
||||
position_v = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1).to(x)
|
||||
|
||||
cosine = torch.cos(torch.matmul(position_v, self.w_r))
|
||||
sine = torch.sin(torch.matmul(position_v, self.w_r))
|
||||
pos_enc = torch.cat((cosine, sine), -1)
|
||||
pos_enc /= math.sqrt(self.d_model)
|
||||
|
||||
if self.hidden_dim is None:
|
||||
return pos_enc.unsqueeze(0)
|
||||
else:
|
||||
return self.mlp(pos_enc.unsqueeze(0))
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
"""Add positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
"""
|
||||
pe = self.extend_pe(x)
|
||||
x = x * self.xscale + pe
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class LegacyRelPositionalEncoding(PositionalEncoding):
|
||||
"""Relative positional encoding module (old version).
|
||||
|
||||
Details can be found in https://github.com/espnet/espnet/pull/2816.
|
||||
|
||||
See : Appendix B in https://arxiv.org/abs/1901.02860
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Initialize class."""
|
||||
super().__init__(
|
||||
d_model=d_model,
|
||||
dropout_rate=dropout_rate,
|
||||
max_len=max_len,
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
"""Compute positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
torch.Tensor: Positional embedding tensor (1, time, `*`).
|
||||
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale
|
||||
pos_emb = self.pe[:, : x.size(1)]
|
||||
return self.dropout(x), self.dropout(pos_emb)
|
||||
|
||||
|
||||
class RelPositionalEncoding(torch.nn.Module):
|
||||
"""Relative positional encoding module (new implementation).
|
||||
|
||||
Details can be found in https://github.com/espnet/espnet/pull/2816.
|
||||
|
||||
See : Appendix B in https://arxiv.org/abs/1901.02860
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Construct an PositionalEncoding object."""
|
||||
super(RelPositionalEncoding, self).__init__()
|
||||
self.d_model = d_model
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.pe = None
|
||||
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
|
||||
|
||||
def extend_pe(self, x):
|
||||
"""Reset the positional encodings."""
|
||||
if self.pe is not None:
|
||||
# self.pe contains both positive and negative parts
|
||||
# the length of self.pe is 2 * input_len - 1
|
||||
if self.pe.size(1) >= x.size(1) * 2 - 1:
|
||||
if self.pe.dtype != x.dtype or self.pe.device != x.device:
|
||||
self.pe = self.pe.to(dtype=x.dtype, device=x.device)
|
||||
return
|
||||
# Suppose `i` means to the position of query vecotr and `j` means the
|
||||
# position of key vector. We use position relative positions when keys
|
||||
# are to the left (i>j) and negative relative positions otherwise (i<j).
|
||||
pe_positive = torch.zeros(x.size(1), self.d_model)
|
||||
pe_negative = torch.zeros(x.size(1), self.d_model)
|
||||
position = torch.arange(0, x.size(1), dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.d_model, 2, dtype=torch.float32)
|
||||
* -(math.log(10000.0) / self.d_model)
|
||||
)
|
||||
pe_positive[:, 0::2] = torch.sin(position * div_term)
|
||||
pe_positive[:, 1::2] = torch.cos(position * div_term)
|
||||
pe_negative[:, 0::2] = torch.sin(-1 * position * div_term)
|
||||
pe_negative[:, 1::2] = torch.cos(-1 * position * div_term)
|
||||
|
||||
# Reserve the order of positive indices and concat both positive and
|
||||
# negative indices. This is used to support the shifting trick
|
||||
# as in https://arxiv.org/abs/1901.02860
|
||||
pe_positive = torch.flip(pe_positive, [0]).unsqueeze(0)
|
||||
pe_negative = pe_negative[1:].unsqueeze(0)
|
||||
pe = torch.cat([pe_positive, pe_negative], dim=1)
|
||||
self.pe = pe.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
"""Add positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
|
||||
"""
|
||||
self.extend_pe(x)
|
||||
x = x * self.xscale
|
||||
pos_emb = self.pe[
|
||||
:,
|
||||
self.pe.size(1) // 2 - x.size(1) + 1 : self.pe.size(1) // 2 + x.size(1),
|
||||
]
|
||||
return self.dropout(x), self.dropout(pos_emb)
|
||||
|
||||
|
||||
class StreamPositionalEncoding(torch.nn.Module):
|
||||
"""Streaming Positional encoding.
|
||||
|
||||
Args:
|
||||
d_model (int): Embedding dimension.
|
||||
dropout_rate (float): Dropout rate.
|
||||
max_len (int): Maximum input length.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, d_model, dropout_rate, max_len=5000):
|
||||
"""Construct an PositionalEncoding object."""
|
||||
super(StreamPositionalEncoding, self).__init__()
|
||||
self.d_model = d_model
|
||||
self.xscale = math.sqrt(self.d_model)
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
self.pe = None
|
||||
self.tmp = torch.tensor(0.0).expand(1, max_len)
|
||||
self.extend_pe(self.tmp.size(1), self.tmp.device, self.tmp.dtype)
|
||||
self._register_load_state_dict_pre_hook(_pre_hook)
|
||||
|
||||
def extend_pe(self, length, device, dtype):
|
||||
"""Reset the positional encodings."""
|
||||
if self.pe is not None:
|
||||
if self.pe.size(1) >= length:
|
||||
if self.pe.dtype != dtype or self.pe.device != device:
|
||||
self.pe = self.pe.to(dtype=dtype, device=device)
|
||||
return
|
||||
pe = torch.zeros(length, self.d_model)
|
||||
position = torch.arange(0, length, dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.d_model, 2, dtype=torch.float32)
|
||||
* -(math.log(10000.0) / self.d_model)
|
||||
)
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
pe = pe.unsqueeze(0)
|
||||
self.pe = pe.to(device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor, start_idx: int = 0):
|
||||
"""Add positional encoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor (batch, time, `*`).
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Encoded tensor (batch, time, `*`).
|
||||
|
||||
"""
|
||||
self.extend_pe(x.size(1) + start_idx, x.device, x.dtype)
|
||||
x = x * self.xscale + self.pe[:, start_idx : start_idx + x.size(1)]
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class SinusoidalPositionEncoder(torch.nn.Module):
|
||||
""" """
|
||||
|
||||
def __init__(self, d_model=80, dropout_rate=0.1):
|
||||
"""Initialize SinusoidalPositionEncoder.
|
||||
|
||||
Args:
|
||||
d_model: D Model instance.
|
||||
dropout_rate: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
def encode(
|
||||
self, positions: torch.Tensor = None, depth: int = None, dtype: torch.dtype = torch.float32
|
||||
):
|
||||
"""Encode.
|
||||
|
||||
Args:
|
||||
positions: TODO.
|
||||
depth: TODO.
|
||||
dtype: TODO.
|
||||
"""
|
||||
batch_size = positions.size(0)
|
||||
positions = positions.type(dtype)
|
||||
device = positions.device
|
||||
log_timescale_increment = torch.log(torch.tensor([10000], dtype=dtype, device=device)) / (
|
||||
depth / 2 - 1
|
||||
)
|
||||
inv_timescales = torch.exp(
|
||||
torch.arange(depth / 2, device=device).type(dtype) * (-log_timescale_increment)
|
||||
)
|
||||
inv_timescales = torch.reshape(inv_timescales, [batch_size, -1])
|
||||
scaled_time = torch.reshape(positions, [1, -1, 1]) * torch.reshape(
|
||||
inv_timescales, [1, 1, -1]
|
||||
)
|
||||
encoding = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], dim=2)
|
||||
return encoding.type(dtype)
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
batch_size, timesteps, input_dim = x.size()
|
||||
positions = torch.arange(1, timesteps + 1, device=x.device)[None, :]
|
||||
position_encoding = self.encode(positions, input_dim, x.dtype).to(x.device)
|
||||
|
||||
return x + position_encoding
|
||||
|
||||
|
||||
class StreamSinusoidalPositionEncoder(torch.nn.Module):
|
||||
""" """
|
||||
|
||||
def __init__(self, d_model=80, dropout_rate=0.1):
|
||||
"""Initialize StreamSinusoidalPositionEncoder.
|
||||
|
||||
Args:
|
||||
d_model: D Model instance.
|
||||
dropout_rate: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
def encode(
|
||||
self, positions: torch.Tensor = None, depth: int = None, dtype: torch.dtype = torch.float32
|
||||
):
|
||||
"""Encode.
|
||||
|
||||
Args:
|
||||
positions: TODO.
|
||||
depth: TODO.
|
||||
dtype: TODO.
|
||||
"""
|
||||
batch_size = positions.size(0)
|
||||
positions = positions.type(dtype)
|
||||
log_timescale_increment = torch.log(torch.tensor([10000], dtype=dtype)) / (depth / 2 - 1)
|
||||
inv_timescales = torch.exp(torch.arange(depth / 2).type(dtype) * (-log_timescale_increment))
|
||||
inv_timescales = torch.reshape(inv_timescales, [batch_size, -1])
|
||||
scaled_time = torch.reshape(positions, [1, -1, 1]) * torch.reshape(
|
||||
inv_timescales, [1, 1, -1]
|
||||
)
|
||||
encoding = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], dim=2)
|
||||
return encoding.type(dtype)
|
||||
|
||||
def forward(self, x, cache=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
cache: State cache dict for streaming inference.
|
||||
"""
|
||||
batch_size, timesteps, input_dim = x.size()
|
||||
start_idx = 0
|
||||
if cache is not None:
|
||||
start_idx = cache["start_idx"]
|
||||
cache["start_idx"] += timesteps
|
||||
positions = torch.arange(1, timesteps + start_idx + 1)[None, :]
|
||||
position_encoding = self.encode(positions, input_dim, x.dtype).to(x.device)
|
||||
return x + position_encoding[:, start_idx : start_idx + timesteps]
|
||||
|
||||
|
||||
class StreamingRelPositionalEncoding(torch.nn.Module):
|
||||
"""Relative positional encoding.
|
||||
Args:
|
||||
size: Module size.
|
||||
max_len: Maximum input length.
|
||||
dropout_rate: Dropout rate.
|
||||
"""
|
||||
|
||||
def __init__(self, size: int, dropout_rate: float = 0.0, max_len: int = 5000) -> None:
|
||||
"""Construct a RelativePositionalEncoding object."""
|
||||
super().__init__()
|
||||
|
||||
self.size = size
|
||||
|
||||
self.pe = None
|
||||
self.dropout = torch.nn.Dropout(p=dropout_rate)
|
||||
|
||||
self.extend_pe(torch.tensor(0.0).expand(1, max_len))
|
||||
self._register_load_state_dict_pre_hook(_pre_hook)
|
||||
|
||||
def extend_pe(self, x: torch.Tensor, left_context: int = 0) -> None:
|
||||
"""Reset positional encoding.
|
||||
Args:
|
||||
x: Input sequences. (B, T, ?)
|
||||
left_context: Number of frames in left context.
|
||||
"""
|
||||
time1 = x.size(1) + left_context
|
||||
|
||||
if self.pe is not None:
|
||||
if self.pe.size(1) >= time1 * 2 - 1:
|
||||
if self.pe.dtype != x.dtype or self.pe.device != x.device:
|
||||
self.pe = self.pe.to(device=x.device, dtype=x.dtype)
|
||||
return
|
||||
|
||||
pe_positive = torch.zeros(time1, self.size)
|
||||
pe_negative = torch.zeros(time1, self.size)
|
||||
|
||||
position = torch.arange(0, time1, dtype=torch.float32).unsqueeze(1)
|
||||
div_term = torch.exp(
|
||||
torch.arange(0, self.size, 2, dtype=torch.float32) * -(math.log(10000.0) / self.size)
|
||||
)
|
||||
|
||||
pe_positive[:, 0::2] = torch.sin(position * div_term)
|
||||
pe_positive[:, 1::2] = torch.cos(position * div_term)
|
||||
pe_positive = torch.flip(pe_positive, [0]).unsqueeze(0)
|
||||
|
||||
pe_negative[:, 0::2] = torch.sin(-1 * position * div_term)
|
||||
pe_negative[:, 1::2] = torch.cos(-1 * position * div_term)
|
||||
pe_negative = pe_negative[1:].unsqueeze(0)
|
||||
|
||||
self.pe = torch.cat([pe_positive, pe_negative], dim=1).to(dtype=x.dtype, device=x.device)
|
||||
|
||||
def forward(self, x: torch.Tensor, left_context: int = 0) -> torch.Tensor:
|
||||
"""Compute positional encoding.
|
||||
Args:
|
||||
x: Input sequences. (B, T, ?)
|
||||
left_context: Number of frames in left context.
|
||||
Returns:
|
||||
pos_enc: Positional embedding sequences. (B, 2 * (T - 1), ?)
|
||||
"""
|
||||
self.extend_pe(x, left_context=left_context)
|
||||
|
||||
time1 = x.size(1) + left_context
|
||||
|
||||
pos_enc = self.pe[:, self.pe.size(1) // 2 - time1 + 1 : self.pe.size(1) // 2 + x.size(1)]
|
||||
pos_enc = self.dropout(pos_enc)
|
||||
|
||||
return pos_enc
|
||||
|
||||
|
||||
class ScaledSinuEmbedding(torch.nn.Module):
|
||||
def __init__(self, dim):
|
||||
"""Initialize ScaledSinuEmbedding.
|
||||
|
||||
Args:
|
||||
dim: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
1,
|
||||
)
|
||||
)
|
||||
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim))
|
||||
self.register_buffer("inv_freq", inv_freq)
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
n, device = x.shape[1], x.device
|
||||
t = torch.arange(n, device=device).type_as(self.inv_freq)
|
||||
sinu = einsum("i , j -> i j", t, self.inv_freq)
|
||||
emb = torch.cat((sinu.sin(), sinu.cos()), dim=-1)
|
||||
return emb * self.scale
|
||||
@@ -0,0 +1,351 @@
|
||||
# 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
|
||||
from funasr.models.transformer.utils.repeat import repeat
|
||||
from funasr.models.ctc.ctc import CTC
|
||||
|
||||
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
|
||||
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
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().__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
|
||||
|
||||
|
||||
@tables.register("encoder_classes", "TransformerEncoder")
|
||||
class TransformerEncoder(nn.Module):
|
||||
"""Transformer encoder module.
|
||||
|
||||
Args:
|
||||
input_size: input dim
|
||||
output_size: dimension of attention
|
||||
attention_heads: the number of heads of multi head attention
|
||||
linear_units: the number of units of position-wise feed forward
|
||||
num_blocks: the number of decoder blocks
|
||||
dropout_rate: dropout rate
|
||||
attention_dropout_rate: dropout rate in attention
|
||||
positional_dropout_rate: dropout rate after adding positional encoding
|
||||
input_layer: input layer type
|
||||
pos_enc_class: PositionalEncoding or ScaledPositionalEncoding
|
||||
normalize_before: whether to use layer_norm before the first block
|
||||
concat_after: 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: linear of conv1d
|
||||
positionwise_conv_kernel_size: kernel size of positionwise conv1d layer
|
||||
padding_idx: padding_idx for input_layer=embed
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int,
|
||||
output_size: int = 256,
|
||||
attention_heads: int = 4,
|
||||
linear_units: int = 2048,
|
||||
num_blocks: int = 6,
|
||||
dropout_rate: float = 0.1,
|
||||
positional_dropout_rate: float = 0.1,
|
||||
attention_dropout_rate: float = 0.0,
|
||||
input_layer: Optional[str] = "conv2d",
|
||||
pos_enc_class=PositionalEncoding,
|
||||
normalize_before: bool = True,
|
||||
concat_after: bool = False,
|
||||
positionwise_layer_type: str = "linear",
|
||||
positionwise_conv_kernel_size: int = 1,
|
||||
padding_idx: int = -1,
|
||||
interctc_layer_idx: List[int] = [],
|
||||
interctc_use_conditioning: bool = False,
|
||||
):
|
||||
"""Initialize TransformerEncoder.
|
||||
|
||||
Args:
|
||||
input_size: Size/dimension parameter.
|
||||
output_size: Size/dimension parameter.
|
||||
attention_heads: TODO.
|
||||
linear_units: TODO.
|
||||
num_blocks: TODO.
|
||||
dropout_rate: TODO.
|
||||
positional_dropout_rate: TODO.
|
||||
attention_dropout_rate: TODO.
|
||||
input_layer: TODO.
|
||||
pos_enc_class: TODO.
|
||||
normalize_before: TODO.
|
||||
concat_after: TODO.
|
||||
positionwise_layer_type: TODO.
|
||||
positionwise_conv_kernel_size: Size/dimension parameter.
|
||||
padding_idx: TODO.
|
||||
interctc_layer_idx: TODO.
|
||||
interctc_use_conditioning: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self._output_size = output_size
|
||||
|
||||
if input_layer == "linear":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Linear(input_size, output_size),
|
||||
torch.nn.LayerNorm(output_size),
|
||||
torch.nn.Dropout(dropout_rate),
|
||||
torch.nn.ReLU(),
|
||||
pos_enc_class(output_size, positional_dropout_rate),
|
||||
)
|
||||
elif input_layer == "conv2d":
|
||||
self.embed = Conv2dSubsampling(input_size, output_size, dropout_rate)
|
||||
elif input_layer == "conv2d2":
|
||||
self.embed = Conv2dSubsampling2(input_size, output_size, dropout_rate)
|
||||
elif input_layer == "conv2d6":
|
||||
self.embed = Conv2dSubsampling6(input_size, output_size, dropout_rate)
|
||||
elif input_layer == "conv2d8":
|
||||
self.embed = Conv2dSubsampling8(input_size, output_size, dropout_rate)
|
||||
elif input_layer == "embed":
|
||||
self.embed = torch.nn.Sequential(
|
||||
torch.nn.Embedding(input_size, output_size, padding_idx=padding_idx),
|
||||
pos_enc_class(output_size, positional_dropout_rate),
|
||||
)
|
||||
elif input_layer is None:
|
||||
if input_size == output_size:
|
||||
self.embed = None
|
||||
else:
|
||||
self.embed = torch.nn.Linear(input_size, output_size)
|
||||
else:
|
||||
raise ValueError("unknown input_layer: " + input_layer)
|
||||
self.normalize_before = normalize_before
|
||||
if positionwise_layer_type == "linear":
|
||||
positionwise_layer = PositionwiseFeedForward
|
||||
positionwise_layer_args = (
|
||||
output_size,
|
||||
linear_units,
|
||||
dropout_rate,
|
||||
)
|
||||
elif positionwise_layer_type == "conv1d":
|
||||
positionwise_layer = MultiLayeredConv1d
|
||||
positionwise_layer_args = (
|
||||
output_size,
|
||||
linear_units,
|
||||
positionwise_conv_kernel_size,
|
||||
dropout_rate,
|
||||
)
|
||||
elif positionwise_layer_type == "conv1d-linear":
|
||||
positionwise_layer = Conv1dLinear
|
||||
positionwise_layer_args = (
|
||||
output_size,
|
||||
linear_units,
|
||||
positionwise_conv_kernel_size,
|
||||
dropout_rate,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError("Support only linear or conv1d.")
|
||||
self.encoders = repeat(
|
||||
num_blocks,
|
||||
lambda lnum: EncoderLayer(
|
||||
output_size,
|
||||
MultiHeadedAttention(attention_heads, output_size, attention_dropout_rate),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
),
|
||||
)
|
||||
if self.normalize_before:
|
||||
self.after_norm = LayerNorm(output_size)
|
||||
|
||||
self.interctc_layer_idx = interctc_layer_idx
|
||||
if len(interctc_layer_idx) > 0:
|
||||
assert 0 < min(interctc_layer_idx) and max(interctc_layer_idx) < num_blocks
|
||||
self.interctc_use_conditioning = interctc_use_conditioning
|
||||
self.conditioning_layer = None
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self._output_size
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
ilens: torch.Tensor,
|
||||
prev_states: torch.Tensor = None,
|
||||
ctc: CTC = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Embed positions in tensor.
|
||||
|
||||
Args:
|
||||
xs_pad: input tensor (B, L, D)
|
||||
ilens: input length (B)
|
||||
prev_states: Not to be used now.
|
||||
Returns:
|
||||
position embedded tensor and mask
|
||||
"""
|
||||
masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
|
||||
|
||||
if self.embed is None:
|
||||
xs_pad = xs_pad
|
||||
elif (
|
||||
isinstance(self.embed, Conv2dSubsampling)
|
||||
or isinstance(self.embed, Conv2dSubsampling2)
|
||||
or isinstance(self.embed, Conv2dSubsampling6)
|
||||
or isinstance(self.embed, Conv2dSubsampling8)
|
||||
):
|
||||
short_status, limit_size = check_short_utt(self.embed, xs_pad.size(1))
|
||||
if short_status:
|
||||
raise TooShortUttError(
|
||||
f"has {xs_pad.size(1)} frames and is too short for subsampling "
|
||||
+ f"(it needs more than {limit_size} frames), return empty results",
|
||||
xs_pad.size(1),
|
||||
limit_size,
|
||||
)
|
||||
xs_pad, masks = self.embed(xs_pad, masks)
|
||||
else:
|
||||
xs_pad = self.embed(xs_pad)
|
||||
|
||||
intermediate_outs = []
|
||||
if len(self.interctc_layer_idx) == 0:
|
||||
xs_pad, masks = self.encoders(xs_pad, masks)
|
||||
else:
|
||||
for layer_idx, encoder_layer in enumerate(self.encoders):
|
||||
xs_pad, masks = encoder_layer(xs_pad, masks)
|
||||
|
||||
if layer_idx + 1 in self.interctc_layer_idx:
|
||||
encoder_out = xs_pad
|
||||
|
||||
# intermediate outputs are also normalized
|
||||
if self.normalize_before:
|
||||
encoder_out = self.after_norm(encoder_out)
|
||||
|
||||
intermediate_outs.append((layer_idx + 1, encoder_out))
|
||||
|
||||
if self.interctc_use_conditioning:
|
||||
ctc_out = ctc.softmax(encoder_out)
|
||||
xs_pad = xs_pad + self.conditioning_layer(ctc_out)
|
||||
|
||||
if self.normalize_before:
|
||||
xs_pad = self.after_norm(xs_pad)
|
||||
|
||||
olens = masks.squeeze(1).sum(1)
|
||||
if len(intermediate_outs) > 0:
|
||||
return (xs_pad, intermediate_outs), olens, None
|
||||
return xs_pad, olens, None
|
||||
@@ -0,0 +1,191 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Layer normalization module."""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class LayerNorm(torch.nn.LayerNorm):
|
||||
"""Layer normalization module.
|
||||
|
||||
Args:
|
||||
nout (int): Output dim size.
|
||||
dim (int): Dimension to be normalized.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, nout, dim=-1):
|
||||
"""Construct an LayerNorm object."""
|
||||
super(LayerNorm, self).__init__(nout, eps=1e-12)
|
||||
self.dim = dim
|
||||
|
||||
def forward(self, x):
|
||||
"""Apply layer normalization.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Normalized tensor.
|
||||
|
||||
"""
|
||||
if self.dim == -1:
|
||||
return super(LayerNorm, self).forward(x)
|
||||
return super(LayerNorm, self).forward(x.transpose(self.dim, -1)).transpose(self.dim, -1)
|
||||
|
||||
|
||||
class GlobalLayerNorm(nn.Module):
|
||||
"""Calculate Global Layer Normalization.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
dim : (int or list or torch.Size)
|
||||
Input shape from an expected input of size.
|
||||
eps : float
|
||||
A value added to the denominator for numerical stability.
|
||||
elementwise_affine : bool
|
||||
A boolean value that when set to True,
|
||||
this module has learnable per-element affine parameters
|
||||
initialized to ones (for weights) and zeros (for biases).
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> x = torch.randn(5, 10, 20)
|
||||
>>> GLN = GlobalLayerNorm(10, 3)
|
||||
>>> x_norm = GLN(x)
|
||||
"""
|
||||
|
||||
def __init__(self, dim, shape, eps=1e-8, elementwise_affine=True):
|
||||
"""Initialize GlobalLayerNorm.
|
||||
|
||||
Args:
|
||||
dim: TODO.
|
||||
shape: TODO.
|
||||
eps: TODO.
|
||||
elementwise_affine: TODO.
|
||||
"""
|
||||
super(GlobalLayerNorm, self).__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.elementwise_affine = elementwise_affine
|
||||
|
||||
if self.elementwise_affine:
|
||||
if shape == 3:
|
||||
self.weight = nn.Parameter(torch.ones(self.dim, 1))
|
||||
self.bias = nn.Parameter(torch.zeros(self.dim, 1))
|
||||
if shape == 4:
|
||||
self.weight = nn.Parameter(torch.ones(self.dim, 1, 1))
|
||||
self.bias = nn.Parameter(torch.zeros(self.dim, 1, 1))
|
||||
else:
|
||||
self.register_parameter("weight", None)
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def forward(self, x):
|
||||
"""Returns the normalized tensor.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
x : torch.Tensor
|
||||
Tensor of size [N, C, K, S] or [N, C, L].
|
||||
"""
|
||||
# x = N x C x K x S or N x C x L
|
||||
# N x 1 x 1
|
||||
# cln: mean,var N x 1 x K x S
|
||||
# gln: mean,var N x 1 x 1
|
||||
if x.dim() == 3:
|
||||
mean = torch.mean(x, (1, 2), keepdim=True)
|
||||
var = torch.mean((x - mean) ** 2, (1, 2), keepdim=True)
|
||||
if self.elementwise_affine:
|
||||
x = self.weight * (x - mean) / torch.sqrt(var + self.eps) + self.bias
|
||||
else:
|
||||
x = (x - mean) / torch.sqrt(var + self.eps)
|
||||
|
||||
if x.dim() == 4:
|
||||
mean = torch.mean(x, (1, 2, 3), keepdim=True)
|
||||
var = torch.mean((x - mean) ** 2, (1, 2, 3), keepdim=True)
|
||||
if self.elementwise_affine:
|
||||
x = self.weight * (x - mean) / torch.sqrt(var + self.eps) + self.bias
|
||||
else:
|
||||
x = (x - mean) / torch.sqrt(var + self.eps)
|
||||
return x
|
||||
|
||||
|
||||
class CumulativeLayerNorm(nn.LayerNorm):
|
||||
"""Calculate Cumulative Layer Normalization.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
dim : int
|
||||
Dimension that you want to normalize.
|
||||
elementwise_affine : True
|
||||
Learnable per-element affine parameters.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> x = torch.randn(5, 10, 20)
|
||||
>>> CLN = CumulativeLayerNorm(10)
|
||||
>>> x_norm = CLN(x)
|
||||
"""
|
||||
|
||||
def __init__(self, dim, elementwise_affine=True):
|
||||
"""Initialize CumulativeLayerNorm.
|
||||
|
||||
Args:
|
||||
dim: TODO.
|
||||
elementwise_affine: TODO.
|
||||
"""
|
||||
super(CumulativeLayerNorm, self).__init__(
|
||||
dim, elementwise_affine=elementwise_affine, eps=1e-8
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
"""Returns the normalized tensor.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
x : torch.Tensor
|
||||
Tensor size [N, C, K, S] or [N, C, L]
|
||||
"""
|
||||
# x: N x C x K x S or N x C x L
|
||||
# N x K x S x C
|
||||
if x.dim() == 4:
|
||||
x = x.permute(0, 2, 3, 1).contiguous()
|
||||
# N x K x S x C == only channel norm
|
||||
x = super().forward(x)
|
||||
# N x C x K x S
|
||||
x = x.permute(0, 3, 1, 2).contiguous()
|
||||
if x.dim() == 3:
|
||||
x = torch.transpose(x, 1, 2)
|
||||
# N x L x C == only channel norm
|
||||
x = super().forward(x)
|
||||
# N x C x L
|
||||
x = torch.transpose(x, 1, 2)
|
||||
return x
|
||||
|
||||
|
||||
class ScaleNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
"""Initialize ScaleNorm.
|
||||
|
||||
Args:
|
||||
dim: TODO.
|
||||
eps: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.scale = dim**-0.5
|
||||
self.eps = eps
|
||||
self.g = nn.Parameter(torch.ones(1))
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
norm = torch.norm(x, dim=-1, keepdim=True) * self.scale
|
||||
return x / norm.clamp(min=self.eps) * self.g
|
||||
@@ -0,0 +1,629 @@
|
||||
import logging
|
||||
from typing import Union, Dict, List, Tuple, Optional
|
||||
|
||||
import time
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.cuda.amp import autocast
|
||||
|
||||
from funasr.losses.label_smoothing_loss import LabelSmoothingLoss
|
||||
from funasr.models.ctc.ctc import CTC
|
||||
from funasr.models.transformer.utils.add_sos_eos import add_sos_eos
|
||||
from funasr.metrics.compute_acc import th_accuracy
|
||||
|
||||
# from funasr.models.e2e_asr_common import ErrorCalculator
|
||||
from funasr.train_utils.device_funcs import force_gatherable
|
||||
from funasr.losses.cr_ctc import cr_ctc_loss
|
||||
from funasr.utils.load_utils import load_audio_text_image_video, extract_fbank
|
||||
from funasr.utils import postprocess_utils
|
||||
from funasr.utils.datadir_writer import DatadirWriter
|
||||
from funasr.register import tables
|
||||
|
||||
|
||||
@tables.register("model_classes", "Transformer")
|
||||
class Transformer(nn.Module):
|
||||
"""Transformer: Base encoder-decoder ASR model.
|
||||
|
||||
Standard CTC-attention hybrid architecture with:
|
||||
- Encoder (self-attention + position encoding)
|
||||
- CTC branch for auxiliary loss
|
||||
- Attention decoder for sequence generation
|
||||
- Beam search with LM fusion
|
||||
|
||||
Base class for Conformer, Branchformer, etc.
|
||||
Output: {"key": str, "text": str}
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
specaug: str = None,
|
||||
specaug_conf: dict = None,
|
||||
normalize: str = None,
|
||||
normalize_conf: dict = None,
|
||||
encoder: str = None,
|
||||
encoder_conf: dict = None,
|
||||
decoder: str = None,
|
||||
decoder_conf: dict = None,
|
||||
ctc: str = None,
|
||||
ctc_conf: dict = None,
|
||||
ctc_weight: float = 0.5,
|
||||
interctc_weight: float = 0.0,
|
||||
input_size: int = 80,
|
||||
vocab_size: int = -1,
|
||||
ignore_id: int = -1,
|
||||
blank_id: int = 0,
|
||||
sos: int = 1,
|
||||
eos: int = 2,
|
||||
lsm_weight: float = 0.0,
|
||||
length_normalized_loss: bool = False,
|
||||
report_cer: bool = True,
|
||||
report_wer: bool = True,
|
||||
sym_space: str = "<space>",
|
||||
sym_blank: str = "<blank>",
|
||||
# extract_feats_in_collect_stats: bool = True,
|
||||
share_embedding: bool = False,
|
||||
# preencoder: Optional[AbsPreEncoder] = None,
|
||||
# postencoder: Optional[AbsPostEncoder] = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Initialize Transformer.
|
||||
|
||||
Args:
|
||||
specaug: TODO.
|
||||
specaug_conf: Configuration dict for specaug.
|
||||
normalize: TODO.
|
||||
normalize_conf: Configuration dict for normalize.
|
||||
encoder: TODO.
|
||||
encoder_conf: Configuration dict for encoder.
|
||||
decoder: TODO.
|
||||
decoder_conf: Configuration dict for decoder.
|
||||
ctc: TODO.
|
||||
ctc_conf: Configuration dict for ctc.
|
||||
ctc_weight: TODO.
|
||||
interctc_weight: TODO.
|
||||
input_size: Size/dimension parameter.
|
||||
vocab_size: Size/dimension parameter.
|
||||
ignore_id: TODO.
|
||||
blank_id: TODO.
|
||||
sos: TODO.
|
||||
eos: TODO.
|
||||
lsm_weight: TODO.
|
||||
length_normalized_loss: TODO.
|
||||
report_cer: TODO.
|
||||
report_wer: TODO.
|
||||
sym_space: TODO.
|
||||
sym_blank: TODO.
|
||||
share_embedding: TODO.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
if specaug is not None:
|
||||
specaug_class = tables.specaug_classes.get(specaug)
|
||||
specaug = specaug_class(**specaug_conf)
|
||||
if normalize is not None:
|
||||
normalize_class = tables.normalize_classes.get(normalize)
|
||||
normalize = normalize_class(**normalize_conf)
|
||||
encoder_class = tables.encoder_classes.get(encoder)
|
||||
encoder = encoder_class(input_size=input_size, **encoder_conf)
|
||||
encoder_output_size = encoder.output_size()
|
||||
if decoder is not None:
|
||||
decoder_class = tables.decoder_classes.get(decoder)
|
||||
decoder = decoder_class(
|
||||
vocab_size=vocab_size,
|
||||
encoder_output_size=encoder_output_size,
|
||||
**decoder_conf,
|
||||
)
|
||||
if ctc_weight > 0.0:
|
||||
|
||||
if ctc_conf is None:
|
||||
ctc_conf = {}
|
||||
|
||||
ctc = CTC(odim=vocab_size, encoder_output_size=encoder_output_size, **ctc_conf)
|
||||
|
||||
self.blank_id = blank_id
|
||||
self.sos = sos if sos is not None else vocab_size - 1
|
||||
self.eos = eos if eos is not None else vocab_size - 1
|
||||
self.vocab_size = vocab_size
|
||||
self.ignore_id = ignore_id
|
||||
self.ctc_weight = ctc_weight
|
||||
self.cr_ctc_weight = kwargs.get("cr_ctc_weight", 0.0)
|
||||
self.specaug = specaug
|
||||
self.normalize = normalize
|
||||
self.encoder = encoder
|
||||
|
||||
if not hasattr(self.encoder, "interctc_use_conditioning"):
|
||||
self.encoder.interctc_use_conditioning = False
|
||||
if self.encoder.interctc_use_conditioning:
|
||||
self.encoder.conditioning_layer = torch.nn.Linear(
|
||||
vocab_size, self.encoder.output_size()
|
||||
)
|
||||
self.interctc_weight = interctc_weight
|
||||
|
||||
# self.error_calculator = None
|
||||
if ctc_weight == 1.0:
|
||||
self.decoder = None
|
||||
else:
|
||||
self.decoder = decoder
|
||||
|
||||
self.criterion_att = LabelSmoothingLoss(
|
||||
size=vocab_size,
|
||||
padding_idx=ignore_id,
|
||||
smoothing=lsm_weight,
|
||||
normalize_length=length_normalized_loss,
|
||||
)
|
||||
#
|
||||
# if report_cer or report_wer:
|
||||
# self.error_calculator = ErrorCalculator(
|
||||
# token_list, sym_space, sym_blank, report_cer, report_wer
|
||||
# )
|
||||
#
|
||||
self.error_calculator = None
|
||||
if ctc_weight == 0.0:
|
||||
self.ctc = None
|
||||
else:
|
||||
self.ctc = ctc
|
||||
|
||||
self.share_embedding = share_embedding
|
||||
if self.share_embedding:
|
||||
self.decoder.embed = None
|
||||
|
||||
self.length_normalized_loss = length_normalized_loss
|
||||
self.beam_search = None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor,
|
||||
text: torch.Tensor,
|
||||
text_lengths: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor], torch.Tensor]:
|
||||
"""Encoder + Decoder + Calc loss
|
||||
Args:
|
||||
speech: (Batch, Length, ...)
|
||||
speech_lengths: (Batch, )
|
||||
text: (Batch, Length)
|
||||
text_lengths: (Batch,)
|
||||
"""
|
||||
if len(text_lengths.size()) > 1:
|
||||
text_lengths = text_lengths[:, 0]
|
||||
if len(speech_lengths.size()) > 1:
|
||||
speech_lengths = speech_lengths[:, 0]
|
||||
|
||||
batch_size = speech.shape[0]
|
||||
|
||||
# 1. Encoder
|
||||
encoder_out, encoder_out_lens = self.encode(speech, speech_lengths)
|
||||
intermediate_outs = None
|
||||
if isinstance(encoder_out, tuple):
|
||||
intermediate_outs = encoder_out[1]
|
||||
encoder_out = encoder_out[0]
|
||||
|
||||
loss_att, acc_att, cer_att, wer_att = None, None, None, None
|
||||
loss_ctc, cer_ctc = None, None
|
||||
stats = dict()
|
||||
|
||||
# decoder: CTC branch
|
||||
if self.ctc_weight != 0.0:
|
||||
loss_ctc, cer_ctc = self._calc_ctc_loss(
|
||||
encoder_out, encoder_out_lens, text, text_lengths
|
||||
)
|
||||
|
||||
# Collect CTC branch stats
|
||||
stats["loss_ctc"] = loss_ctc.detach() if loss_ctc is not None else None
|
||||
stats["cer_ctc"] = cer_ctc
|
||||
|
||||
# Intermediate CTC (optional)
|
||||
loss_interctc = 0.0
|
||||
if self.interctc_weight != 0.0 and intermediate_outs is not None:
|
||||
for layer_idx, intermediate_out in intermediate_outs:
|
||||
# we assume intermediate_out has the same length & padding
|
||||
# as those of encoder_out
|
||||
loss_ic, cer_ic = self._calc_ctc_loss(
|
||||
intermediate_out, encoder_out_lens, text, text_lengths
|
||||
)
|
||||
loss_interctc = loss_interctc + loss_ic
|
||||
|
||||
# Collect Intermedaite CTC stats
|
||||
stats["loss_interctc_layer{}".format(layer_idx)] = (
|
||||
loss_ic.detach() if loss_ic is not None else None
|
||||
)
|
||||
stats["cer_interctc_layer{}".format(layer_idx)] = cer_ic
|
||||
|
||||
loss_interctc = loss_interctc / len(intermediate_outs)
|
||||
|
||||
# calculate whole encoder loss
|
||||
loss_ctc = (1 - self.interctc_weight) * loss_ctc + self.interctc_weight * loss_interctc
|
||||
|
||||
# CR-CTC: consistency regularization
|
||||
loss_cr_ctc = None
|
||||
if self.cr_ctc_weight > 0.0 and self.training and self.ctc_weight != 0.0:
|
||||
# Second forward pass WITHOUT SpecAug
|
||||
specaug_backup = self.specaug
|
||||
self.specaug = None
|
||||
encoder_out_clean, encoder_out_lens_clean = self.encode(speech, speech_lengths)
|
||||
self.specaug = specaug_backup
|
||||
if isinstance(encoder_out_clean, tuple):
|
||||
encoder_out_clean = encoder_out_clean[0]
|
||||
# Compute CTC log probs for both augmented and clean
|
||||
ctc_logprobs_aug = self.ctc.log_softmax(encoder_out)
|
||||
ctc_logprobs_clean = self.ctc.log_softmax(encoder_out_clean).detach()
|
||||
loss_cr_ctc = cr_ctc_loss(ctc_logprobs_aug, ctc_logprobs_clean, encoder_out_lens)
|
||||
stats["loss_cr_ctc"] = loss_cr_ctc.detach()
|
||||
|
||||
# decoder: Attention decoder branch
|
||||
loss_att, acc_att, cer_att, wer_att = self._calc_att_loss(
|
||||
encoder_out, encoder_out_lens, text, text_lengths
|
||||
)
|
||||
|
||||
# 3. CTC-Att loss definition
|
||||
if self.ctc_weight == 0.0:
|
||||
loss = loss_att
|
||||
elif self.ctc_weight == 1.0:
|
||||
loss = loss_ctc
|
||||
else:
|
||||
loss = self.ctc_weight * loss_ctc + (1 - self.ctc_weight) * loss_att
|
||||
|
||||
# Add CR-CTC loss
|
||||
if loss_cr_ctc is not None:
|
||||
loss = loss + self.cr_ctc_weight * loss_cr_ctc
|
||||
|
||||
# Collect Attn branch stats
|
||||
stats["loss_att"] = loss_att.detach() if loss_att is not None else None
|
||||
stats["acc"] = acc_att
|
||||
stats["cer"] = cer_att
|
||||
stats["wer"] = wer_att
|
||||
|
||||
# Collect total loss stats
|
||||
stats["loss"] = torch.clone(loss.detach())
|
||||
|
||||
# force_gatherable: to-device and to-tensor if scalar for DataParallel
|
||||
if self.length_normalized_loss:
|
||||
batch_size = int((text_lengths + 1).sum())
|
||||
loss, stats, weight = force_gatherable((loss, stats, batch_size), loss.device)
|
||||
return loss, stats, weight
|
||||
|
||||
def encode(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Frontend + Encoder. Note that this method is used by asr_inference.py
|
||||
Args:
|
||||
speech: (Batch, Length, ...)
|
||||
speech_lengths: (Batch, )
|
||||
ind: int
|
||||
"""
|
||||
with autocast(False):
|
||||
|
||||
# Data augmentation
|
||||
if self.specaug is not None and self.training:
|
||||
speech, speech_lengths = self.specaug(speech, speech_lengths)
|
||||
|
||||
# Normalization for feature: e.g. Global-CMVN, Utterance-CMVN
|
||||
if self.normalize is not None:
|
||||
speech, speech_lengths = self.normalize(speech, speech_lengths)
|
||||
|
||||
# Forward encoder
|
||||
# feats: (Batch, Length, Dim)
|
||||
# -> encoder_out: (Batch, Length2, Dim2)
|
||||
if self.encoder.interctc_use_conditioning:
|
||||
encoder_out, encoder_out_lens, _ = self.encoder(speech, speech_lengths, ctc=self.ctc)
|
||||
else:
|
||||
encoder_out, encoder_out_lens, _ = self.encoder(speech, speech_lengths)
|
||||
intermediate_outs = None
|
||||
if isinstance(encoder_out, tuple):
|
||||
intermediate_outs = encoder_out[1]
|
||||
encoder_out = encoder_out[0]
|
||||
|
||||
if intermediate_outs is not None:
|
||||
return (encoder_out, intermediate_outs), encoder_out_lens
|
||||
|
||||
return encoder_out, encoder_out_lens
|
||||
|
||||
def _calc_att_loss(
|
||||
self,
|
||||
encoder_out: torch.Tensor,
|
||||
encoder_out_lens: torch.Tensor,
|
||||
ys_pad: torch.Tensor,
|
||||
ys_pad_lens: torch.Tensor,
|
||||
):
|
||||
"""Internal: calc att loss.
|
||||
|
||||
Args:
|
||||
encoder_out: Encoder output tensor.
|
||||
encoder_out_lens: Encoder output lengths.
|
||||
ys_pad: TODO.
|
||||
ys_pad_lens: Lengths of ys_pad.
|
||||
"""
|
||||
ys_in_pad, ys_out_pad = add_sos_eos(ys_pad, self.sos, self.eos, self.ignore_id)
|
||||
ys_in_lens = ys_pad_lens + 1
|
||||
|
||||
# 1. Forward decoder
|
||||
decoder_out, _ = self.decoder(encoder_out, encoder_out_lens, ys_in_pad, ys_in_lens)
|
||||
|
||||
# 2. Compute attention loss
|
||||
loss_att = self.criterion_att(decoder_out, ys_out_pad)
|
||||
acc_att = th_accuracy(
|
||||
decoder_out.view(-1, self.vocab_size),
|
||||
ys_out_pad,
|
||||
ignore_label=self.ignore_id,
|
||||
)
|
||||
|
||||
# Compute cer/wer using attention-decoder
|
||||
if self.training or self.error_calculator is None:
|
||||
cer_att, wer_att = None, None
|
||||
else:
|
||||
ys_hat = decoder_out.argmax(dim=-1)
|
||||
cer_att, wer_att = self.error_calculator(ys_hat.cpu(), ys_pad.cpu())
|
||||
|
||||
return loss_att, acc_att, cer_att, wer_att
|
||||
|
||||
def _calc_ctc_loss(
|
||||
self,
|
||||
encoder_out: torch.Tensor,
|
||||
encoder_out_lens: torch.Tensor,
|
||||
ys_pad: torch.Tensor,
|
||||
ys_pad_lens: torch.Tensor,
|
||||
):
|
||||
# Calc CTC loss
|
||||
"""Internal: calc ctc loss.
|
||||
|
||||
Args:
|
||||
encoder_out: Encoder output tensor.
|
||||
encoder_out_lens: Encoder output lengths.
|
||||
ys_pad: TODO.
|
||||
ys_pad_lens: Lengths of ys_pad.
|
||||
"""
|
||||
loss_ctc = self.ctc(encoder_out, encoder_out_lens, ys_pad, ys_pad_lens)
|
||||
|
||||
# Calc CER using CTC
|
||||
cer_ctc = None
|
||||
if not self.training and self.error_calculator is not None:
|
||||
ys_hat = self.ctc.argmax(encoder_out).data
|
||||
cer_ctc = self.error_calculator(ys_hat.cpu(), ys_pad.cpu(), is_ctc=True)
|
||||
return loss_ctc, cer_ctc
|
||||
|
||||
def inference_batch_ctc(
|
||||
self,
|
||||
data_in,
|
||||
data_lengths=None,
|
||||
key: list = None,
|
||||
tokenizer=None,
|
||||
frontend=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Batch CTC greedy decoding for fast inference.
|
||||
|
||||
Uses CTC output with greedy decoding (argmax + collapse repeats + remove blanks).
|
||||
Much faster than autoregressive beam search, with comparable accuracy.
|
||||
"""
|
||||
meta_data = {}
|
||||
|
||||
# extract fbank feats
|
||||
time1 = time.perf_counter()
|
||||
audio_sample_list = load_audio_text_image_video(
|
||||
data_in,
|
||||
fs=frontend.fs,
|
||||
audio_fs=kwargs.get("fs", 16000),
|
||||
data_type=kwargs.get("data_type", "sound"),
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
time2 = time.perf_counter()
|
||||
meta_data["load_data"] = f"{time2 - time1:0.3f}"
|
||||
speech, speech_lengths = extract_fbank(
|
||||
audio_sample_list, data_type=kwargs.get("data_type", "sound"), frontend=frontend
|
||||
)
|
||||
time3 = time.perf_counter()
|
||||
meta_data["extract_feat"] = f"{time3 - time2:0.3f}"
|
||||
meta_data["batch_data_time"] = (
|
||||
speech_lengths.sum().item() * frontend.frame_shift * frontend.lfr_n / 1000
|
||||
)
|
||||
|
||||
speech = speech.to(device=kwargs["device"])
|
||||
speech_lengths = speech_lengths.to(device=kwargs["device"])
|
||||
|
||||
# Encoder
|
||||
encoder_out, encoder_out_lens = self.encode(speech, speech_lengths)
|
||||
if isinstance(encoder_out, tuple):
|
||||
encoder_out = encoder_out[0]
|
||||
|
||||
# CTC log probs
|
||||
ctc_logprobs = self.ctc.log_softmax(encoder_out)
|
||||
|
||||
results = []
|
||||
b = encoder_out.size(0)
|
||||
if key is None:
|
||||
key = [f"utt_{i}" for i in range(b)]
|
||||
|
||||
for i in range(b):
|
||||
x = ctc_logprobs[i, :encoder_out_lens[i].item(), :]
|
||||
yseq = x.argmax(dim=-1)
|
||||
yseq = torch.unique_consecutive(yseq, dim=-1)
|
||||
mask = yseq != self.blank_id
|
||||
token_int = yseq[mask].tolist()
|
||||
|
||||
token = tokenizer.ids2tokens(token_int)
|
||||
text_postprocessed, _ = postprocess_utils.sentence_postprocess(token)
|
||||
|
||||
result_i = {"key": key[i], "text": text_postprocessed}
|
||||
results.append(result_i)
|
||||
|
||||
if kwargs.get("output_dir") is not None:
|
||||
if not hasattr(self, "writer"):
|
||||
self.writer = DatadirWriter(kwargs.get("output_dir"))
|
||||
ibest_writer = self.writer["1best_recog"]
|
||||
ibest_writer["token"][key[i]] = " ".join(token)
|
||||
ibest_writer["text"][key[i]] = text_postprocessed
|
||||
|
||||
return results, meta_data
|
||||
|
||||
def init_beam_search(
|
||||
self,
|
||||
**kwargs,
|
||||
):
|
||||
"""Init beam search.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
from funasr.models.transformer.search import BeamSearch
|
||||
from funasr.models.transformer.scorers.ctc import CTCPrefixScorer
|
||||
from funasr.models.transformer.scorers.length_bonus import LengthBonus
|
||||
|
||||
# 1. Build ASR model
|
||||
scorers = {}
|
||||
|
||||
if self.ctc != None:
|
||||
ctc = CTCPrefixScorer(ctc=self.ctc, eos=self.eos)
|
||||
scorers.update(ctc=ctc)
|
||||
token_list = kwargs.get("token_list")
|
||||
scorers.update(
|
||||
decoder=self.decoder,
|
||||
length_bonus=LengthBonus(len(token_list)),
|
||||
)
|
||||
|
||||
# 3. Build ngram model
|
||||
# ngram is not supported now
|
||||
ngram = None
|
||||
scorers["ngram"] = ngram
|
||||
|
||||
weights = dict(
|
||||
decoder=1.0 - kwargs.get("decoding_ctc_weight", 0.5),
|
||||
ctc=kwargs.get("decoding_ctc_weight", 0.5),
|
||||
lm=kwargs.get("lm_weight", 0.0),
|
||||
ngram=kwargs.get("ngram_weight", 0.0),
|
||||
length_bonus=kwargs.get("penalty", 0.0),
|
||||
)
|
||||
beam_search = BeamSearch(
|
||||
beam_size=kwargs.get("beam_size", 10),
|
||||
weights=weights,
|
||||
scorers=scorers,
|
||||
sos=self.sos,
|
||||
eos=self.eos,
|
||||
vocab_size=len(token_list),
|
||||
token_list=token_list,
|
||||
pre_beam_score_key=None if self.ctc_weight == 1.0 else "full",
|
||||
)
|
||||
|
||||
self.beam_search = beam_search
|
||||
|
||||
def inference(
|
||||
self,
|
||||
data_in,
|
||||
data_lengths=None,
|
||||
key: list = None,
|
||||
tokenizer=None,
|
||||
frontend=None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
"""Run inference on input data.
|
||||
|
||||
Args:
|
||||
data_in: Input data (audio samples, file paths, or text).
|
||||
data_lengths: Lengths of each input sample in the batch.
|
||||
key: Sample identifiers.
|
||||
tokenizer: Tokenizer instance for text encoding/decoding.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
if kwargs.get("batch_size", 1) > 1:
|
||||
return self.inference_batch_ctc(
|
||||
data_in, data_lengths=data_lengths, key=key,
|
||||
tokenizer=tokenizer, frontend=frontend, **kwargs
|
||||
)
|
||||
|
||||
# init beamsearch
|
||||
if self.beam_search is None:
|
||||
logging.info("enable beam_search")
|
||||
self.init_beam_search(**kwargs)
|
||||
self.nbest = kwargs.get("nbest", 1)
|
||||
|
||||
meta_data = {}
|
||||
if (
|
||||
isinstance(data_in, torch.Tensor) and kwargs.get("data_type", "sound") == "fbank"
|
||||
): # fbank
|
||||
speech, speech_lengths = data_in, data_lengths
|
||||
if len(speech.shape) < 3:
|
||||
speech = speech[None, :, :]
|
||||
if speech_lengths is None:
|
||||
speech_lengths = speech.shape[1]
|
||||
else:
|
||||
# extract fbank feats
|
||||
time1 = time.perf_counter()
|
||||
audio_sample_list = load_audio_text_image_video(
|
||||
data_in,
|
||||
fs=frontend.fs,
|
||||
audio_fs=kwargs.get("fs", 16000),
|
||||
data_type=kwargs.get("data_type", "sound"),
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
time2 = time.perf_counter()
|
||||
meta_data["load_data"] = f"{time2 - time1:0.3f}"
|
||||
speech, speech_lengths = extract_fbank(
|
||||
audio_sample_list, data_type=kwargs.get("data_type", "sound"), frontend=frontend
|
||||
)
|
||||
time3 = time.perf_counter()
|
||||
meta_data["extract_feat"] = f"{time3 - time2:0.3f}"
|
||||
meta_data["batch_data_time"] = (
|
||||
speech_lengths.sum().item() * frontend.frame_shift * frontend.lfr_n / 1000
|
||||
)
|
||||
|
||||
speech = speech.to(device=kwargs["device"])
|
||||
speech_lengths = speech_lengths.to(device=kwargs["device"])
|
||||
# Encoder
|
||||
encoder_out, encoder_out_lens = self.encode(speech, speech_lengths)
|
||||
if isinstance(encoder_out, tuple):
|
||||
encoder_out = encoder_out[0]
|
||||
|
||||
# c. Passed the encoder result and the beam search
|
||||
nbest_hyps = self.beam_search(
|
||||
x=encoder_out[0],
|
||||
maxlenratio=kwargs.get("maxlenratio", 0.0),
|
||||
minlenratio=kwargs.get("minlenratio", 0.0),
|
||||
)
|
||||
|
||||
nbest_hyps = nbest_hyps[: self.nbest]
|
||||
|
||||
results = []
|
||||
b, n, d = encoder_out.size()
|
||||
for i in range(b):
|
||||
|
||||
for nbest_idx, hyp in enumerate(nbest_hyps):
|
||||
ibest_writer = None
|
||||
if kwargs.get("output_dir") is not None:
|
||||
if not hasattr(self, "writer"):
|
||||
self.writer = DatadirWriter(kwargs.get("output_dir"))
|
||||
ibest_writer = self.writer[f"{nbest_idx + 1}best_recog"]
|
||||
|
||||
# remove sos/eos and get results
|
||||
last_pos = -1
|
||||
if isinstance(hyp.yseq, list):
|
||||
token_int = hyp.yseq[1:last_pos]
|
||||
else:
|
||||
token_int = hyp.yseq[1:last_pos].tolist()
|
||||
|
||||
# remove blank symbol id, which is assumed to be 0
|
||||
token_int = list(
|
||||
filter(
|
||||
lambda x: x != self.eos and x != self.sos and x != self.blank_id, token_int
|
||||
)
|
||||
)
|
||||
|
||||
# Change integer-ids to tokens
|
||||
token = tokenizer.ids2tokens(token_int)
|
||||
text = tokenizer.tokens2text(token)
|
||||
|
||||
text_postprocessed, _ = postprocess_utils.sentence_postprocess(token)
|
||||
result_i = {"key": key[i], "token": token, "text": text_postprocessed}
|
||||
results.append(result_i)
|
||||
|
||||
if ibest_writer is not None:
|
||||
ibest_writer["token"][key[i]] = " ".join(token)
|
||||
ibest_writer["text"][key[i]] = text_postprocessed
|
||||
|
||||
return results, meta_data
|
||||
@@ -0,0 +1,58 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2019 Shigeki Karita
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
"""Positionwise feed forward layer definition."""
|
||||
|
||||
import torch
|
||||
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
|
||||
|
||||
class PositionwiseFeedForward(torch.nn.Module):
|
||||
"""Positionwise feed forward layer.
|
||||
|
||||
Args:
|
||||
idim (int): Input dimenstion.
|
||||
hidden_units (int): The number of hidden units.
|
||||
dropout_rate (float): Dropout rate.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, idim, hidden_units, dropout_rate, activation=torch.nn.ReLU()):
|
||||
"""Construct an PositionwiseFeedForward object."""
|
||||
super(PositionwiseFeedForward, self).__init__()
|
||||
self.w_1 = torch.nn.Linear(idim, hidden_units)
|
||||
self.w_2 = torch.nn.Linear(hidden_units, idim)
|
||||
self.dropout = torch.nn.Dropout(dropout_rate)
|
||||
self.activation = activation
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
return self.w_2(self.dropout(self.activation(self.w_1(x))))
|
||||
|
||||
|
||||
class PositionwiseFeedForwardDecoderSANMExport(torch.nn.Module):
|
||||
def __init__(self, model):
|
||||
"""Initialize PositionwiseFeedForwardDecoderSANMExport.
|
||||
|
||||
Args:
|
||||
model: Model instance or model name.
|
||||
"""
|
||||
super().__init__()
|
||||
self.w_1 = model.w_1
|
||||
self.w_2 = model.w_2
|
||||
self.activation = model.activation
|
||||
self.norm = model.norm
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
x = self.activation(self.w_1(x))
|
||||
x = self.w_2(self.norm(x))
|
||||
return x
|
||||
@@ -0,0 +1 @@
|
||||
"""Initialize sub package."""
|
||||
@@ -0,0 +1,155 @@
|
||||
"""ScorerInterface implementation for CTC."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from funasr.models.transformer.scorers.ctc_prefix_score import CTCPrefixScore
|
||||
from funasr.models.transformer.scorers.ctc_prefix_score import CTCPrefixScoreTH
|
||||
from funasr.models.transformer.scorers.scorer_interface import BatchPartialScorerInterface
|
||||
|
||||
class CTCPrefixScorer(BatchPartialScorerInterface):
|
||||
"""Decoder interface wrapper for CTCPrefixScore."""
|
||||
|
||||
def __init__(self, ctc: torch.nn.Module, eos: int):
|
||||
"""Initialize class.
|
||||
|
||||
Args:
|
||||
ctc (torch.nn.Module): The CTC implementation.
|
||||
For example, :class:`espnet.nets.pytorch_backend.ctc.CTC`
|
||||
eos (int): The end-of-sequence id.
|
||||
|
||||
"""
|
||||
self.ctc = ctc
|
||||
self.eos = eos
|
||||
self.impl = None
|
||||
|
||||
def init_state(self, x: torch.Tensor):
|
||||
"""Get an initial state for decoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoded feature tensor
|
||||
|
||||
Returns: initial state
|
||||
|
||||
"""
|
||||
logp = self.ctc.log_softmax(x.unsqueeze(0)).detach().squeeze(0).cpu().numpy()
|
||||
# TODO(karita): use CTCPrefixScoreTH
|
||||
self.impl = CTCPrefixScore(logp, 0, self.eos, np)
|
||||
return 0, self.impl.initial_state()
|
||||
|
||||
def select_state(self, state, i, new_id=None):
|
||||
"""Select state with relative ids in the main beam search.
|
||||
|
||||
Args:
|
||||
state: Decoder state for prefix tokens
|
||||
i (int): Index to select a state in the main beam search
|
||||
new_id (int): New label id to select a state if necessary
|
||||
|
||||
Returns:
|
||||
state: pruned state
|
||||
|
||||
"""
|
||||
if type(state) == tuple:
|
||||
if len(state) == 2: # for CTCPrefixScore
|
||||
sc, st = state
|
||||
return sc[i], st[i]
|
||||
else: # for CTCPrefixScoreTH (need new_id > 0)
|
||||
r, log_psi, f_min, f_max, scoring_idmap = state
|
||||
s = log_psi[i, new_id].expand(log_psi.size(1))
|
||||
if scoring_idmap is not None:
|
||||
return r[:, :, i, scoring_idmap[i, new_id]], s, f_min, f_max
|
||||
else:
|
||||
return r[:, :, i, new_id], s, f_min, f_max
|
||||
return None if state is None else state[i]
|
||||
|
||||
def score_partial(self, y, ids, state, x):
|
||||
"""Score new token.
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D prefix token
|
||||
next_tokens (torch.Tensor): torch.int64 next token to score
|
||||
state: decoder state for prefix tokens
|
||||
x (torch.Tensor): 2D encoder feature that generates ys
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]:
|
||||
Tuple of a score tensor for y that has a shape `(len(next_tokens),)`
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
prev_score, state = state
|
||||
presub_score, new_st = self.impl(y.cpu(), ids.cpu(), state)
|
||||
tscore = torch.as_tensor(presub_score - prev_score, device=x.device, dtype=x.dtype)
|
||||
return tscore, (presub_score, new_st)
|
||||
|
||||
def batch_init_state(self, x: torch.Tensor):
|
||||
"""Get an initial state for decoding.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoded feature tensor
|
||||
|
||||
Returns: initial state
|
||||
|
||||
"""
|
||||
logp = self.ctc.log_softmax(x.unsqueeze(0)) # assuming batch_size = 1
|
||||
xlen = torch.tensor([logp.size(1)])
|
||||
self.impl = CTCPrefixScoreTH(logp, xlen, 0, self.eos)
|
||||
return None
|
||||
|
||||
def batch_score_partial(self, y, ids, state, x):
|
||||
"""Score new token.
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D prefix token
|
||||
ids (torch.Tensor): torch.int64 next token to score
|
||||
state: decoder state for prefix tokens
|
||||
x (torch.Tensor): 2D encoder feature that generates ys
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]:
|
||||
Tuple of a score tensor for y that has a shape `(len(next_tokens),)`
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
batch_state = (
|
||||
(
|
||||
torch.stack([s[0] for s in state], dim=2),
|
||||
torch.stack([s[1] for s in state]),
|
||||
state[0][2],
|
||||
state[0][3],
|
||||
)
|
||||
if state[0] is not None
|
||||
else None
|
||||
)
|
||||
return self.impl(y, batch_state, ids)
|
||||
|
||||
def extend_prob(self, x: torch.Tensor):
|
||||
"""Extend probs for decoding.
|
||||
|
||||
This extension is for streaming decoding
|
||||
as in Eq (14) in https://arxiv.org/abs/2006.14941
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoded feature tensor
|
||||
|
||||
"""
|
||||
logp = self.ctc.log_softmax(x.unsqueeze(0))
|
||||
self.impl.extend_prob(logp)
|
||||
|
||||
def extend_state(self, state):
|
||||
"""Extend state for decoding.
|
||||
|
||||
This extension is for streaming decoding
|
||||
as in Eq (14) in https://arxiv.org/abs/2006.14941
|
||||
|
||||
Args:
|
||||
state: The states of hyps
|
||||
|
||||
Returns: exteded state
|
||||
|
||||
"""
|
||||
new_state = []
|
||||
for s in state:
|
||||
new_state.append(self.impl.extend_state(s))
|
||||
|
||||
return new_state
|
||||
@@ -0,0 +1,345 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
# Copyright 2018 Mitsubishi Electric Research Labs (Takaaki Hori)
|
||||
# Apache 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
import torch
|
||||
|
||||
import numpy as np
|
||||
import six
|
||||
|
||||
|
||||
class CTCPrefixScoreTH(object):
|
||||
"""Batch processing of CTCPrefixScore
|
||||
|
||||
which is based on Algorithm 2 in WATANABE et al.
|
||||
"HYBRID CTC/ATTENTION ARCHITECTURE FOR END-TO-END SPEECH RECOGNITION,"
|
||||
but extended to efficiently compute the label probablities for multiple
|
||||
hypotheses simultaneously
|
||||
See also Seki et al. "Vectorized Beam Search for CTC-Attention-Based
|
||||
Speech Recognition," In INTERSPEECH (pp. 3825-3829), 2019.
|
||||
"""
|
||||
|
||||
def __init__(self, x, xlens, blank, eos, margin=0):
|
||||
"""Construct CTC prefix scorer
|
||||
|
||||
:param torch.Tensor x: input label posterior sequences (B, T, O)
|
||||
:param torch.Tensor xlens: input lengths (B,)
|
||||
:param int blank: blank label id
|
||||
:param int eos: end-of-sequence id
|
||||
:param int margin: margin parameter for windowing (0 means no windowing)
|
||||
"""
|
||||
# In the comment lines,
|
||||
# we assume T: input_length, B: batch size, W: beam width, O: output dim.
|
||||
self.logzero = -10000000000.0
|
||||
self.blank = blank
|
||||
self.eos = eos
|
||||
self.batch = x.size(0)
|
||||
self.input_length = x.size(1)
|
||||
self.odim = x.size(2)
|
||||
self.dtype = x.dtype
|
||||
self.device = torch.device("cuda:%d" % x.get_device()) if x.is_cuda else torch.device("cpu")
|
||||
# Pad the rest of posteriors in the batch
|
||||
# TODO(takaaki-hori): need a better way without for-loops
|
||||
for i, l in enumerate(xlens):
|
||||
if l < self.input_length:
|
||||
x[i, l:, :] = self.logzero
|
||||
x[i, l:, blank] = 0
|
||||
# Reshape input x
|
||||
xn = x.transpose(0, 1) # (B, T, O) -> (T, B, O)
|
||||
xb = xn[:, :, self.blank].unsqueeze(2).expand(-1, -1, self.odim)
|
||||
self.x = torch.stack([xn, xb]) # (2, T, B, O)
|
||||
self.end_frames = torch.as_tensor(xlens) - 1
|
||||
|
||||
# Setup CTC windowing
|
||||
self.margin = margin
|
||||
if margin > 0:
|
||||
self.frame_ids = torch.arange(self.input_length, dtype=self.dtype, device=self.device)
|
||||
# Base indices for index conversion
|
||||
self.idx_bh = None
|
||||
self.idx_b = torch.arange(self.batch, device=self.device)
|
||||
self.idx_bo = (self.idx_b * self.odim).unsqueeze(1)
|
||||
|
||||
def __call__(self, y, state, scoring_ids=None, att_w=None):
|
||||
"""Compute CTC prefix scores for next labels
|
||||
|
||||
:param list y: prefix label sequences
|
||||
:param tuple state: previous CTC state
|
||||
:param torch.Tensor pre_scores: scores for pre-selection of hypotheses (BW, O)
|
||||
:param torch.Tensor att_w: attention weights to decide CTC window
|
||||
:return new_state, ctc_local_scores (BW, O)
|
||||
"""
|
||||
output_length = len(y[0]) - 1 # ignore sos
|
||||
last_ids = [yi[-1] for yi in y] # last output label ids
|
||||
n_bh = len(last_ids) # batch * hyps
|
||||
n_hyps = n_bh // self.batch # assuming each utterance has the same # of hyps
|
||||
self.scoring_num = scoring_ids.size(-1) if scoring_ids is not None else 0
|
||||
# prepare state info
|
||||
if state is None:
|
||||
r_prev = torch.full(
|
||||
(self.input_length, 2, self.batch, n_hyps),
|
||||
self.logzero,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
r_prev[:, 1] = torch.cumsum(self.x[0, :, :, self.blank], 0).unsqueeze(2)
|
||||
r_prev = r_prev.view(-1, 2, n_bh)
|
||||
s_prev = 0.0
|
||||
f_min_prev = 0
|
||||
f_max_prev = 1
|
||||
else:
|
||||
r_prev, s_prev, f_min_prev, f_max_prev = state
|
||||
|
||||
# select input dimensions for scoring
|
||||
if self.scoring_num > 0:
|
||||
scoring_idmap = torch.full((n_bh, self.odim), -1, dtype=torch.long, device=self.device)
|
||||
snum = self.scoring_num
|
||||
if self.idx_bh is None or n_bh > len(self.idx_bh):
|
||||
self.idx_bh = torch.arange(n_bh, device=self.device).view(-1, 1)
|
||||
scoring_idmap[self.idx_bh[:n_bh], scoring_ids] = torch.arange(snum, device=self.device)
|
||||
scoring_idx = (scoring_ids + self.idx_bo.repeat(1, n_hyps).view(-1, 1)).view(-1)
|
||||
x_ = torch.index_select(
|
||||
self.x.view(2, -1, self.batch * self.odim), 2, scoring_idx
|
||||
).view(2, -1, n_bh, snum)
|
||||
else:
|
||||
scoring_ids = None
|
||||
scoring_idmap = None
|
||||
snum = self.odim
|
||||
x_ = self.x.unsqueeze(3).repeat(1, 1, 1, n_hyps, 1).view(2, -1, n_bh, snum)
|
||||
|
||||
# new CTC forward probs are prepared as a (T x 2 x BW x S) tensor
|
||||
# that corresponds to r_t^n(h) and r_t^b(h) in a batch.
|
||||
r = torch.full(
|
||||
(self.input_length, 2, n_bh, snum),
|
||||
self.logzero,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
if output_length == 0:
|
||||
r[0, 0] = x_[0, 0]
|
||||
|
||||
r_sum = torch.logsumexp(r_prev, 1)
|
||||
log_phi = r_sum.unsqueeze(2).repeat(1, 1, snum)
|
||||
if scoring_ids is not None:
|
||||
for idx in range(n_bh):
|
||||
pos = scoring_idmap[idx, last_ids[idx]]
|
||||
if pos >= 0:
|
||||
log_phi[:, idx, pos] = r_prev[:, 1, idx]
|
||||
else:
|
||||
for idx in range(n_bh):
|
||||
log_phi[:, idx, last_ids[idx]] = r_prev[:, 1, idx]
|
||||
|
||||
# decide start and end frames based on attention weights
|
||||
if att_w is not None and self.margin > 0:
|
||||
f_arg = torch.matmul(att_w, self.frame_ids)
|
||||
f_min = max(int(f_arg.min().cpu()), f_min_prev)
|
||||
f_max = max(int(f_arg.max().cpu()), f_max_prev)
|
||||
start = min(f_max_prev, max(f_min - self.margin, output_length, 1))
|
||||
end = min(f_max + self.margin, self.input_length)
|
||||
else:
|
||||
f_min = f_max = 0
|
||||
start = max(output_length, 1)
|
||||
end = self.input_length
|
||||
|
||||
# compute forward probabilities log(r_t^n(h)) and log(r_t^b(h))
|
||||
for t in range(start, end):
|
||||
rp = r[t - 1]
|
||||
rr = torch.stack([rp[0], log_phi[t - 1], rp[0], rp[1]]).view(2, 2, n_bh, snum)
|
||||
r[t] = torch.logsumexp(rr, 1) + x_[:, t]
|
||||
|
||||
# compute log prefix probabilities log(psi)
|
||||
log_phi_x = torch.cat((log_phi[0].unsqueeze(0), log_phi[:-1]), dim=0) + x_[0]
|
||||
if scoring_ids is not None:
|
||||
log_psi = torch.full(
|
||||
(n_bh, self.odim), self.logzero, dtype=self.dtype, device=self.device
|
||||
)
|
||||
log_psi_ = torch.logsumexp(
|
||||
torch.cat((log_phi_x[start:end], r[start - 1, 0].unsqueeze(0)), dim=0),
|
||||
dim=0,
|
||||
)
|
||||
for si in range(n_bh):
|
||||
log_psi[si, scoring_ids[si]] = log_psi_[si]
|
||||
else:
|
||||
log_psi = torch.logsumexp(
|
||||
torch.cat((log_phi_x[start:end], r[start - 1, 0].unsqueeze(0)), dim=0),
|
||||
dim=0,
|
||||
)
|
||||
|
||||
for si in range(n_bh):
|
||||
log_psi[si, self.eos] = r_sum[self.end_frames[si // n_hyps], si]
|
||||
|
||||
# exclude blank probs
|
||||
log_psi[:, self.blank] = self.logzero
|
||||
|
||||
return (log_psi - s_prev), (r, log_psi, f_min, f_max, scoring_idmap)
|
||||
|
||||
def index_select_state(self, state, best_ids):
|
||||
"""Select CTC states according to best ids
|
||||
|
||||
:param state : CTC state
|
||||
:param best_ids : index numbers selected by beam pruning (B, W)
|
||||
:return selected_state
|
||||
"""
|
||||
r, s, f_min, f_max, scoring_idmap = state
|
||||
# convert ids to BHO space
|
||||
n_bh = len(s)
|
||||
n_hyps = n_bh // self.batch
|
||||
vidx = (best_ids + (self.idx_b * (n_hyps * self.odim)).view(-1, 1)).view(-1)
|
||||
# select hypothesis scores
|
||||
s_new = torch.index_select(s.view(-1), 0, vidx)
|
||||
s_new = s_new.view(-1, 1).repeat(1, self.odim).view(n_bh, self.odim)
|
||||
# convert ids to BHS space (S: scoring_num)
|
||||
if scoring_idmap is not None:
|
||||
snum = self.scoring_num
|
||||
hyp_idx = (best_ids // self.odim + (self.idx_b * n_hyps).view(-1, 1)).view(-1)
|
||||
label_ids = torch.fmod(best_ids, self.odim).view(-1)
|
||||
score_idx = scoring_idmap[hyp_idx, label_ids]
|
||||
score_idx[score_idx == -1] = 0
|
||||
vidx = score_idx + hyp_idx * snum
|
||||
else:
|
||||
snum = self.odim
|
||||
# select forward probabilities
|
||||
r_new = torch.index_select(r.view(-1, 2, n_bh * snum), 2, vidx).view(-1, 2, n_bh)
|
||||
return r_new, s_new, f_min, f_max
|
||||
|
||||
def extend_prob(self, x):
|
||||
"""Extend CTC prob.
|
||||
|
||||
:param torch.Tensor x: input label posterior sequences (B, T, O)
|
||||
"""
|
||||
|
||||
if self.x.shape[1] < x.shape[1]: # self.x (2,T,B,O); x (B,T,O)
|
||||
# Pad the rest of posteriors in the batch
|
||||
# TODO(takaaki-hori): need a better way without for-loops
|
||||
xlens = [x.size(1)]
|
||||
for i, l in enumerate(xlens):
|
||||
if l < self.input_length:
|
||||
x[i, l:, :] = self.logzero
|
||||
x[i, l:, self.blank] = 0
|
||||
tmp_x = self.x
|
||||
xn = x.transpose(0, 1) # (B, T, O) -> (T, B, O)
|
||||
xb = xn[:, :, self.blank].unsqueeze(2).expand(-1, -1, self.odim)
|
||||
self.x = torch.stack([xn, xb]) # (2, T, B, O)
|
||||
self.x[:, : tmp_x.shape[1], :, :] = tmp_x
|
||||
self.input_length = x.size(1)
|
||||
self.end_frames = torch.as_tensor(xlens) - 1
|
||||
|
||||
def extend_state(self, state):
|
||||
"""Compute CTC prefix state.
|
||||
|
||||
|
||||
:param state : CTC state
|
||||
:return ctc_state
|
||||
"""
|
||||
|
||||
if state is None:
|
||||
# nothing to do
|
||||
return state
|
||||
else:
|
||||
r_prev, s_prev, f_min_prev, f_max_prev = state
|
||||
|
||||
r_prev_new = torch.full(
|
||||
(self.input_length, 2),
|
||||
self.logzero,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
start = max(r_prev.shape[0], 1)
|
||||
r_prev_new[0:start] = r_prev
|
||||
for t in six.moves.range(start, self.input_length):
|
||||
r_prev_new[t, 1] = r_prev_new[t - 1, 1] + self.x[0, t, :, self.blank]
|
||||
|
||||
return (r_prev_new, s_prev, f_min_prev, f_max_prev)
|
||||
|
||||
|
||||
class CTCPrefixScore(object):
|
||||
"""Compute CTC label sequence scores
|
||||
|
||||
which is based on Algorithm 2 in WATANABE et al.
|
||||
"HYBRID CTC/ATTENTION ARCHITECTURE FOR END-TO-END SPEECH RECOGNITION,"
|
||||
but extended to efficiently compute the probablities of multiple labels
|
||||
simultaneously
|
||||
"""
|
||||
|
||||
def __init__(self, x, blank, eos, xp):
|
||||
"""Initialize CTCPrefixScore.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
blank: TODO.
|
||||
eos: TODO.
|
||||
xp: TODO.
|
||||
"""
|
||||
self.xp = xp
|
||||
self.logzero = -10000000000.0
|
||||
self.blank = blank
|
||||
self.eos = eos
|
||||
self.input_length = len(x)
|
||||
self.x = x
|
||||
|
||||
def initial_state(self):
|
||||
"""Obtain an initial CTC state
|
||||
|
||||
:return: CTC state
|
||||
"""
|
||||
# initial CTC state is made of a frame x 2 tensor that corresponds to
|
||||
# r_t^n(<sos>) and r_t^b(<sos>), where 0 and 1 of axis=1 represent
|
||||
# superscripts n and b (non-blank and blank), respectively.
|
||||
r = self.xp.full((self.input_length, 2), self.logzero, dtype=np.float32)
|
||||
r[0, 1] = self.x[0, self.blank]
|
||||
for i in six.moves.range(1, self.input_length):
|
||||
r[i, 1] = r[i - 1, 1] + self.x[i, self.blank]
|
||||
return r
|
||||
|
||||
def __call__(self, y, cs, r_prev):
|
||||
"""Compute CTC prefix scores for next labels
|
||||
|
||||
:param y : prefix label sequence
|
||||
:param cs : array of next labels
|
||||
:param r_prev: previous CTC state
|
||||
:return ctc_scores, ctc_states
|
||||
"""
|
||||
# initialize CTC states
|
||||
output_length = len(y) - 1 # ignore sos
|
||||
# new CTC states are prepared as a frame x (n or b) x n_labels tensor
|
||||
# that corresponds to r_t^n(h) and r_t^b(h).
|
||||
r = self.xp.ndarray((self.input_length, 2, len(cs)), dtype=np.float32)
|
||||
xs = self.x[:, cs]
|
||||
if output_length == 0:
|
||||
r[0, 0] = xs[0]
|
||||
r[0, 1] = self.logzero
|
||||
else:
|
||||
r[output_length - 1] = self.logzero
|
||||
|
||||
# prepare forward probabilities for the last label
|
||||
r_sum = self.xp.logaddexp(r_prev[:, 0], r_prev[:, 1]) # log(r_t^n(g) + r_t^b(g))
|
||||
last = y[-1]
|
||||
if output_length > 0 and last in cs:
|
||||
log_phi = self.xp.ndarray((self.input_length, len(cs)), dtype=np.float32)
|
||||
for i in six.moves.range(len(cs)):
|
||||
log_phi[:, i] = r_sum if cs[i] != last else r_prev[:, 1]
|
||||
else:
|
||||
log_phi = r_sum
|
||||
|
||||
# compute forward probabilities log(r_t^n(h)), log(r_t^b(h)),
|
||||
# and log prefix probabilities log(psi)
|
||||
start = max(output_length, 1)
|
||||
log_psi = r[start - 1, 0]
|
||||
for t in six.moves.range(start, self.input_length):
|
||||
r[t, 0] = self.xp.logaddexp(r[t - 1, 0], log_phi[t - 1]) + xs[t]
|
||||
r[t, 1] = self.xp.logaddexp(r[t - 1, 0], r[t - 1, 1]) + self.x[t, self.blank]
|
||||
log_psi = self.xp.logaddexp(log_psi, log_phi[t - 1] + xs[t])
|
||||
|
||||
# get P(...eos|X) that ends with the prefix itself
|
||||
eos_pos = self.xp.where(cs == self.eos)[0]
|
||||
if len(eos_pos) > 0:
|
||||
log_psi[eos_pos] = r_sum[-1] # log(r_T^n(g) + r_T^b(g))
|
||||
|
||||
# exclude blank probs
|
||||
blank_pos = self.xp.where(cs == self.blank)[0]
|
||||
if len(blank_pos) > 0:
|
||||
log_psi[blank_pos] = self.logzero
|
||||
|
||||
# return the log prefix probability and CTC states, where the label axis
|
||||
# of the CTC states is moved to the first axis to slice it easily
|
||||
return log_psi, self.xp.rollaxis(r, 2)
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Length bonus module."""
|
||||
|
||||
from typing import Any
|
||||
from typing import List
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from funasr.models.transformer.scorers.scorer_interface import BatchScorerInterface
|
||||
|
||||
|
||||
class LengthBonus(BatchScorerInterface):
|
||||
"""Length bonus in beam search."""
|
||||
|
||||
def __init__(self, n_vocab: int):
|
||||
"""Initialize class.
|
||||
|
||||
Args:
|
||||
n_vocab (int): The number of tokens in vocabulary for beam search
|
||||
|
||||
"""
|
||||
self.n = n_vocab
|
||||
|
||||
def score(self, y, state, x):
|
||||
"""Score new token.
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D torch.int64 prefix tokens.
|
||||
state: Scorer state for prefix tokens
|
||||
x (torch.Tensor): 2D encoder feature that generates ys.
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]: Tuple of
|
||||
torch.float32 scores for next token (n_vocab)
|
||||
and None
|
||||
|
||||
"""
|
||||
return torch.tensor([1.0], device=x.device, dtype=x.dtype).expand(self.n), None
|
||||
|
||||
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, n_vocab)`
|
||||
and next state list for ys.
|
||||
|
||||
"""
|
||||
return (
|
||||
torch.tensor([1.0], device=xs.device, dtype=xs.dtype).expand(ys.shape[0], self.n),
|
||||
None,
|
||||
)
|
||||
@@ -0,0 +1,186 @@
|
||||
"""Scorer interface module."""
|
||||
|
||||
from typing import Any
|
||||
from typing import List
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import warnings
|
||||
|
||||
|
||||
class ScorerInterface:
|
||||
"""Scorer interface for beam search.
|
||||
|
||||
The scorer performs scoring of the all tokens in vocabulary.
|
||||
|
||||
Examples:
|
||||
* Search heuristics
|
||||
* :class:`espnet.nets.scorers.length_bonus.LengthBonus`
|
||||
* Decoder networks of the sequence-to-sequence models
|
||||
* :class:`espnet.nets.pytorch_backend.nets.transformer.decoder.Decoder`
|
||||
* :class:`espnet.nets.pytorch_backend.nets.rnn.decoders.Decoder`
|
||||
* Neural language models
|
||||
* :class:`espnet.nets.pytorch_backend.lm.transformer.TransformerLM`
|
||||
* :class:`espnet.nets.pytorch_backend.lm.default.DefaultRNNLM`
|
||||
* :class:`espnet.nets.pytorch_backend.lm.seq_rnn.SequentialRNNLM`
|
||||
|
||||
"""
|
||||
|
||||
def init_state(self, x: torch.Tensor) -> Any:
|
||||
"""Get an initial state for decoding (optional).
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoded feature tensor
|
||||
|
||||
Returns: initial state
|
||||
|
||||
"""
|
||||
return None
|
||||
|
||||
def select_state(self, state: Any, i: int, new_id: int = None) -> Any:
|
||||
"""Select state with relative ids in the main beam search.
|
||||
|
||||
Args:
|
||||
state: Decoder state for prefix tokens
|
||||
i (int): Index to select a state in the main beam search
|
||||
new_id (int): New label index to select a state if necessary
|
||||
|
||||
Returns:
|
||||
state: pruned state
|
||||
|
||||
"""
|
||||
return None if state is None else state[i]
|
||||
|
||||
def score(self, y: torch.Tensor, state: Any, x: torch.Tensor) -> Tuple[torch.Tensor, Any]:
|
||||
"""Score new token (required).
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D torch.int64 prefix tokens.
|
||||
state: Scorer state for prefix tokens
|
||||
x (torch.Tensor): The encoder feature that generates ys.
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]: Tuple of
|
||||
scores for next token that has a shape of `(n_vocab)`
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def final_score(self, state: Any) -> float:
|
||||
"""Score eos (optional).
|
||||
|
||||
Args:
|
||||
state: Scorer state for prefix tokens
|
||||
|
||||
Returns:
|
||||
float: final score
|
||||
|
||||
"""
|
||||
return 0.0
|
||||
|
||||
|
||||
class BatchScorerInterface(ScorerInterface):
|
||||
"""Batch scorer interface."""
|
||||
|
||||
def batch_init_state(self, x: torch.Tensor) -> Any:
|
||||
"""Get an initial state for decoding (optional).
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoded feature tensor
|
||||
|
||||
Returns: initial state
|
||||
|
||||
"""
|
||||
return self.init_state(x)
|
||||
|
||||
def batch_score(
|
||||
self, ys: torch.Tensor, states: List[Any], xs: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, List[Any]]:
|
||||
"""Score new token batch (required).
|
||||
|
||||
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.
|
||||
|
||||
"""
|
||||
warnings.warn(
|
||||
"{} batch score is implemented through for loop not parallelized".format(
|
||||
self.__class__.__name__
|
||||
)
|
||||
)
|
||||
scores = list()
|
||||
outstates = list()
|
||||
for i, (y, state, x) in enumerate(zip(ys, states, xs)):
|
||||
score, outstate = self.score(y, state, x)
|
||||
outstates.append(outstate)
|
||||
scores.append(score)
|
||||
scores = torch.cat(scores, 0).view(ys.shape[0], -1)
|
||||
return scores, outstates
|
||||
|
||||
|
||||
class PartialScorerInterface(ScorerInterface):
|
||||
"""Partial scorer interface for beam search.
|
||||
|
||||
The partial scorer performs scoring when non-partial scorer finished scoring,
|
||||
and receives pre-pruned next tokens to score because it is too heavy to score
|
||||
all the tokens.
|
||||
|
||||
Examples:
|
||||
* Prefix search for connectionist-temporal-classification models
|
||||
* :class:`espnet.nets.scorers.ctc.CTCPrefixScorer`
|
||||
|
||||
"""
|
||||
|
||||
def score_partial(
|
||||
self, y: torch.Tensor, next_tokens: torch.Tensor, state: Any, x: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, Any]:
|
||||
"""Score new token (required).
|
||||
|
||||
Args:
|
||||
y (torch.Tensor): 1D prefix token
|
||||
next_tokens (torch.Tensor): torch.int64 next token to score
|
||||
state: decoder state for prefix tokens
|
||||
x (torch.Tensor): The encoder feature that generates ys
|
||||
|
||||
Returns:
|
||||
tuple[torch.Tensor, Any]:
|
||||
Tuple of a score tensor for y that has a shape `(len(next_tokens),)`
|
||||
and next state for ys
|
||||
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class BatchPartialScorerInterface(BatchScorerInterface, PartialScorerInterface):
|
||||
"""Batch partial scorer interface for beam search."""
|
||||
|
||||
def batch_score_partial(
|
||||
self,
|
||||
ys: torch.Tensor,
|
||||
next_tokens: torch.Tensor,
|
||||
states: List[Any],
|
||||
xs: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, Any]:
|
||||
"""Score new token (required).
|
||||
|
||||
Args:
|
||||
ys (torch.Tensor): torch.int64 prefix tokens (n_batch, ylen).
|
||||
next_tokens (torch.Tensor): torch.int64 tokens to score (n_batch, n_token).
|
||||
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, Any]:
|
||||
Tuple of a score tensor for ys that has a shape `(n_batch, n_vocab)`
|
||||
and next states for ys
|
||||
"""
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,449 @@
|
||||
from itertools import chain
|
||||
import logging
|
||||
from typing import Any
|
||||
from typing import Dict
|
||||
from typing import List
|
||||
from typing import NamedTuple
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from funasr.metrics.common import end_detect
|
||||
from funasr.models.transformer.scorers.scorer_interface import PartialScorerInterface
|
||||
from funasr.models.transformer.scorers.scorer_interface import ScorerInterface
|
||||
|
||||
|
||||
class Hypothesis(NamedTuple):
|
||||
"""Hypothesis data type."""
|
||||
|
||||
yseq: torch.Tensor
|
||||
score: Union[float, torch.Tensor] = 0
|
||||
scores: Dict[str, Union[float, torch.Tensor]] = dict()
|
||||
states: Dict[str, Any] = dict()
|
||||
|
||||
def asdict(self) -> dict:
|
||||
"""Convert data to JSON-friendly dict."""
|
||||
return self._replace(
|
||||
yseq=self.yseq.tolist(),
|
||||
score=float(self.score),
|
||||
scores={k: float(v) for k, v in self.scores.items()},
|
||||
)._asdict()
|
||||
|
||||
|
||||
class BeamSearch(torch.nn.Module):
|
||||
"""Beam search implementation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scorers: Dict[str, ScorerInterface],
|
||||
weights: Dict[str, float],
|
||||
beam_size: int,
|
||||
vocab_size: int,
|
||||
sos: int,
|
||||
eos: int,
|
||||
token_list: List[str] = None,
|
||||
pre_beam_ratio: float = 1.5,
|
||||
pre_beam_score_key: str = None,
|
||||
):
|
||||
"""Initialize beam search.
|
||||
|
||||
Args:
|
||||
scorers (dict[str, ScorerInterface]): Dict of decoder modules
|
||||
e.g., Decoder, CTCPrefixScorer, LM
|
||||
The scorer will be ignored if it is `None`
|
||||
weights (dict[str, float]): Dict of weights for each scorers
|
||||
The scorer will be ignored if its weight is 0
|
||||
beam_size (int): The number of hypotheses kept during search
|
||||
vocab_size (int): The number of vocabulary
|
||||
sos (int): Start of sequence id
|
||||
eos (int): End of sequence id
|
||||
token_list (list[str]): List of tokens for debug log
|
||||
pre_beam_score_key (str): key of scores to perform pre-beam search
|
||||
pre_beam_ratio (float): beam size in the pre-beam search
|
||||
will be `int(pre_beam_ratio * beam_size)`
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
# set scorers
|
||||
self.weights = weights
|
||||
self.scorers = dict()
|
||||
self.full_scorers = dict()
|
||||
self.part_scorers = dict()
|
||||
# this module dict is required for recursive cast
|
||||
# `self.to(device, dtype)` in `recog.py`
|
||||
self.nn_dict = torch.nn.ModuleDict()
|
||||
for k, v in scorers.items():
|
||||
w = weights.get(k, 0)
|
||||
if w == 0 or v is None:
|
||||
continue
|
||||
assert isinstance(
|
||||
v, ScorerInterface
|
||||
), f"{k} ({type(v)}) does not implement ScorerInterface"
|
||||
self.scorers[k] = v
|
||||
if isinstance(v, PartialScorerInterface):
|
||||
self.part_scorers[k] = v
|
||||
else:
|
||||
self.full_scorers[k] = v
|
||||
if isinstance(v, torch.nn.Module):
|
||||
self.nn_dict[k] = v
|
||||
|
||||
# set configurations
|
||||
self.sos = sos
|
||||
self.eos = eos
|
||||
self.token_list = token_list
|
||||
self.pre_beam_size = int(pre_beam_ratio * beam_size)
|
||||
self.beam_size = beam_size
|
||||
self.n_vocab = vocab_size
|
||||
if (
|
||||
pre_beam_score_key is not None
|
||||
and pre_beam_score_key != "full"
|
||||
and pre_beam_score_key not in self.full_scorers
|
||||
):
|
||||
raise KeyError(f"{pre_beam_score_key} is not found in {self.full_scorers}")
|
||||
self.pre_beam_score_key = pre_beam_score_key
|
||||
self.do_pre_beam = (
|
||||
self.pre_beam_score_key is not None
|
||||
and self.pre_beam_size < self.n_vocab
|
||||
and len(self.part_scorers) > 0
|
||||
)
|
||||
|
||||
def init_hyp(self, x: torch.Tensor) -> List[Hypothesis]:
|
||||
"""Get an initial hypothesis data.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The encoder output feature
|
||||
|
||||
Returns:
|
||||
Hypothesis: The initial hypothesis.
|
||||
|
||||
"""
|
||||
init_states = dict()
|
||||
init_scores = dict()
|
||||
for k, d in self.scorers.items():
|
||||
init_states[k] = d.init_state(x)
|
||||
init_scores[k] = 0.0
|
||||
return [
|
||||
Hypothesis(
|
||||
score=0.0,
|
||||
scores=init_scores,
|
||||
states=init_states,
|
||||
yseq=torch.tensor([self.sos], device=x.device),
|
||||
)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def append_token(xs: torch.Tensor, x: int) -> torch.Tensor:
|
||||
"""Append new token to prefix tokens.
|
||||
|
||||
Args:
|
||||
xs (torch.Tensor): The prefix token
|
||||
x (int): The new token to append
|
||||
|
||||
Returns:
|
||||
torch.Tensor: New tensor contains: xs + [x] with xs.dtype and xs.device
|
||||
|
||||
"""
|
||||
x = torch.tensor([x], dtype=xs.dtype, device=xs.device)
|
||||
return torch.cat((xs, x))
|
||||
|
||||
def score_full(
|
||||
self, hyp: Hypothesis, x: torch.Tensor
|
||||
) -> Tuple[Dict[str, torch.Tensor], Dict[str, Any]]:
|
||||
"""Score new hypothesis by `self.full_scorers`.
|
||||
|
||||
Args:
|
||||
hyp (Hypothesis): Hypothesis with prefix tokens to score
|
||||
x (torch.Tensor): Corresponding input feature
|
||||
|
||||
Returns:
|
||||
Tuple[Dict[str, torch.Tensor], Dict[str, Any]]: Tuple of
|
||||
score dict of `hyp` that has string keys of `self.full_scorers`
|
||||
and tensor score values of shape: `(self.n_vocab,)`,
|
||||
and state dict that has string keys
|
||||
and state values of `self.full_scorers`
|
||||
|
||||
"""
|
||||
scores = dict()
|
||||
states = dict()
|
||||
for k, d in self.full_scorers.items():
|
||||
scores[k], states[k] = d.score(hyp.yseq, hyp.states[k], x)
|
||||
return scores, states
|
||||
|
||||
def score_partial(
|
||||
self, hyp: Hypothesis, ids: torch.Tensor, x: torch.Tensor
|
||||
) -> Tuple[Dict[str, torch.Tensor], Dict[str, Any]]:
|
||||
"""Score new hypothesis by `self.part_scorers`.
|
||||
|
||||
Args:
|
||||
hyp (Hypothesis): Hypothesis with prefix tokens to score
|
||||
ids (torch.Tensor): 1D tensor of new partial tokens to score
|
||||
x (torch.Tensor): Corresponding input feature
|
||||
|
||||
Returns:
|
||||
Tuple[Dict[str, torch.Tensor], Dict[str, Any]]: Tuple of
|
||||
score dict of `hyp` that has string keys of `self.part_scorers`
|
||||
and tensor score values of shape: `(len(ids),)`,
|
||||
and state dict that has string keys
|
||||
and state values of `self.part_scorers`
|
||||
|
||||
"""
|
||||
scores = dict()
|
||||
states = dict()
|
||||
for k, d in self.part_scorers.items():
|
||||
scores[k], states[k] = d.score_partial(hyp.yseq, ids, hyp.states[k], x)
|
||||
return scores, states
|
||||
|
||||
def beam(
|
||||
self, weighted_scores: torch.Tensor, ids: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compute topk full token ids and partial token ids.
|
||||
|
||||
Args:
|
||||
weighted_scores (torch.Tensor): The weighted sum scores for each tokens.
|
||||
Its shape is `(self.n_vocab,)`.
|
||||
ids (torch.Tensor): The partial token ids to compute topk
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]:
|
||||
The topk full token ids and partial token ids.
|
||||
Their shapes are `(self.beam_size,)`
|
||||
|
||||
"""
|
||||
# no pre beam performed
|
||||
if weighted_scores.size(0) == ids.size(0):
|
||||
top_ids = weighted_scores.topk(self.beam_size)[1]
|
||||
return top_ids, top_ids
|
||||
|
||||
# mask pruned in pre-beam not to select in topk
|
||||
tmp = weighted_scores[ids]
|
||||
weighted_scores[:] = -float("inf")
|
||||
weighted_scores[ids] = tmp
|
||||
top_ids = weighted_scores.topk(self.beam_size)[1]
|
||||
local_ids = weighted_scores[ids].topk(self.beam_size)[1]
|
||||
return top_ids, local_ids
|
||||
|
||||
@staticmethod
|
||||
def merge_scores(
|
||||
prev_scores: Dict[str, float],
|
||||
next_full_scores: Dict[str, torch.Tensor],
|
||||
full_idx: int,
|
||||
next_part_scores: Dict[str, torch.Tensor],
|
||||
part_idx: int,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Merge scores for new hypothesis.
|
||||
|
||||
Args:
|
||||
prev_scores (Dict[str, float]):
|
||||
The previous hypothesis scores by `self.scorers`
|
||||
next_full_scores (Dict[str, torch.Tensor]): scores by `self.full_scorers`
|
||||
full_idx (int): The next token id for `next_full_scores`
|
||||
next_part_scores (Dict[str, torch.Tensor]):
|
||||
scores of partial tokens by `self.part_scorers`
|
||||
part_idx (int): The new token id for `next_part_scores`
|
||||
|
||||
Returns:
|
||||
Dict[str, torch.Tensor]: The new score dict.
|
||||
Its keys are names of `self.full_scorers` and `self.part_scorers`.
|
||||
Its values are scalar tensors by the scorers.
|
||||
|
||||
"""
|
||||
new_scores = dict()
|
||||
for k, v in next_full_scores.items():
|
||||
new_scores[k] = prev_scores[k] + v[full_idx]
|
||||
for k, v in next_part_scores.items():
|
||||
new_scores[k] = prev_scores[k] + v[part_idx]
|
||||
return new_scores
|
||||
|
||||
def merge_states(self, states: Any, part_states: Any, part_idx: int) -> Any:
|
||||
"""Merge states for new hypothesis.
|
||||
|
||||
Args:
|
||||
states: states of `self.full_scorers`
|
||||
part_states: states of `self.part_scorers`
|
||||
part_idx (int): The new token id for `part_scores`
|
||||
|
||||
Returns:
|
||||
Dict[str, torch.Tensor]: The new score dict.
|
||||
Its keys are names of `self.full_scorers` and `self.part_scorers`.
|
||||
Its values are states of the scorers.
|
||||
|
||||
"""
|
||||
new_states = dict()
|
||||
for k, v in states.items():
|
||||
new_states[k] = v
|
||||
for k, d in self.part_scorers.items():
|
||||
new_states[k] = d.select_state(part_states[k], part_idx)
|
||||
return new_states
|
||||
|
||||
def search(self, running_hyps: List[Hypothesis], x: torch.Tensor) -> List[Hypothesis]:
|
||||
"""Search new tokens for running hypotheses and encoded speech x.
|
||||
|
||||
Args:
|
||||
running_hyps (List[Hypothesis]): Running hypotheses on beam
|
||||
x (torch.Tensor): Encoded speech feature (T, D)
|
||||
|
||||
Returns:
|
||||
List[Hypotheses]: Best sorted hypotheses
|
||||
|
||||
"""
|
||||
best_hyps = []
|
||||
part_ids = torch.arange(self.n_vocab, device=x.device) # no pre-beam
|
||||
for hyp in running_hyps:
|
||||
# scoring
|
||||
weighted_scores = torch.zeros(self.n_vocab, dtype=x.dtype, device=x.device)
|
||||
scores, states = self.score_full(hyp, x)
|
||||
for k in self.full_scorers:
|
||||
weighted_scores += self.weights[k] * scores[k]
|
||||
# partial scoring
|
||||
if self.do_pre_beam:
|
||||
pre_beam_scores = (
|
||||
weighted_scores
|
||||
if self.pre_beam_score_key == "full"
|
||||
else scores[self.pre_beam_score_key]
|
||||
)
|
||||
part_ids = torch.topk(pre_beam_scores, self.pre_beam_size)[1]
|
||||
part_scores, part_states = self.score_partial(hyp, part_ids, x)
|
||||
for k in self.part_scorers:
|
||||
weighted_scores[part_ids] += self.weights[k] * part_scores[k]
|
||||
# add previous hyp score
|
||||
weighted_scores += hyp.score
|
||||
|
||||
# update hyps
|
||||
for j, part_j in zip(*self.beam(weighted_scores, part_ids)):
|
||||
# will be (2 x beam at most)
|
||||
best_hyps.append(
|
||||
Hypothesis(
|
||||
score=weighted_scores[j],
|
||||
yseq=self.append_token(hyp.yseq, j),
|
||||
scores=self.merge_scores(hyp.scores, scores, j, part_scores, part_j),
|
||||
states=self.merge_states(states, part_states, part_j),
|
||||
)
|
||||
)
|
||||
|
||||
# sort and prune 2 x beam -> beam
|
||||
best_hyps = sorted(best_hyps, key=lambda x: x.score, reverse=True)[
|
||||
: min(len(best_hyps), self.beam_size)
|
||||
]
|
||||
return best_hyps
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, maxlenratio: float = 0.0, minlenratio: float = 0.0
|
||||
) -> List[Hypothesis]:
|
||||
"""Perform beam search.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Encoded speech feature (T, D)
|
||||
maxlenratio (float): Input length ratio to obtain max output length.
|
||||
If maxlenratio=0.0 (default), it uses a end-detect function
|
||||
to automatically find maximum hypothesis lengths
|
||||
If maxlenratio<0.0, its absolute value is interpreted
|
||||
as a constant max output length.
|
||||
minlenratio (float): Input length ratio to obtain min output length.
|
||||
|
||||
Returns:
|
||||
list[Hypothesis]: N-best decoding results
|
||||
|
||||
"""
|
||||
# set length bounds
|
||||
if maxlenratio == 0:
|
||||
maxlen = x.shape[0]
|
||||
elif maxlenratio < 0:
|
||||
maxlen = -1 * int(maxlenratio)
|
||||
else:
|
||||
maxlen = max(1, int(maxlenratio * x.size(0)))
|
||||
minlen = int(minlenratio * x.size(0))
|
||||
logging.info("decoder input length: " + str(x.shape[0]))
|
||||
logging.info("max output length: " + str(maxlen))
|
||||
logging.info("min output length: " + str(minlen))
|
||||
|
||||
# main loop of prefix search
|
||||
running_hyps = self.init_hyp(x)
|
||||
ended_hyps = []
|
||||
for i in range(maxlen):
|
||||
logging.debug("position " + str(i))
|
||||
best = self.search(running_hyps, x)
|
||||
# post process of one iteration
|
||||
running_hyps = self.post_process(i, maxlen, maxlenratio, best, ended_hyps)
|
||||
# end detection
|
||||
if maxlenratio == 0.0 and end_detect([h.asdict() for h in ended_hyps], i):
|
||||
logging.info(f"end detected at {i}")
|
||||
break
|
||||
if len(running_hyps) == 0:
|
||||
logging.info("no hypothesis. Finish decoding.")
|
||||
break
|
||||
else:
|
||||
logging.debug(f"remained hypotheses: {len(running_hyps)}")
|
||||
|
||||
nbest_hyps = sorted(ended_hyps, key=lambda x: x.score, reverse=True)
|
||||
# check the number of hypotheses reaching to eos
|
||||
if len(nbest_hyps) == 0:
|
||||
logging.warning(
|
||||
"there is no N-best results, perform recognition " "again with smaller minlenratio."
|
||||
)
|
||||
return (
|
||||
[]
|
||||
if minlenratio < 0.1
|
||||
else self.forward(x, maxlenratio, max(0.0, minlenratio - 0.1))
|
||||
)
|
||||
|
||||
# report the best result
|
||||
best = nbest_hyps[0]
|
||||
for k, v in best.scores.items():
|
||||
logging.info(f"{v:6.2f} * {self.weights[k]:3} = {v * self.weights[k]:6.2f} for {k}")
|
||||
logging.info(f"total log probability: {best.score:.2f}")
|
||||
logging.info(f"normalized log probability: {best.score / len(best.yseq):.2f}")
|
||||
logging.info(f"total number of ended hypotheses: {len(nbest_hyps)}")
|
||||
if self.token_list is not None:
|
||||
logging.info(
|
||||
"best hypo: " + "".join([self.token_list[x] for x in best.yseq[1:-1]]) + "\n"
|
||||
)
|
||||
return nbest_hyps
|
||||
|
||||
def post_process(
|
||||
self,
|
||||
i: int,
|
||||
maxlen: int,
|
||||
maxlenratio: float,
|
||||
running_hyps: List[Hypothesis],
|
||||
ended_hyps: List[Hypothesis],
|
||||
) -> List[Hypothesis]:
|
||||
"""Perform post-processing of beam search iterations.
|
||||
|
||||
Args:
|
||||
i (int): The length of hypothesis tokens.
|
||||
maxlen (int): The maximum length of tokens in beam search.
|
||||
maxlenratio (int): The maximum length ratio in beam search.
|
||||
running_hyps (List[Hypothesis]): The running hypotheses in beam search.
|
||||
ended_hyps (List[Hypothesis]): The ended hypotheses in beam search.
|
||||
|
||||
Returns:
|
||||
List[Hypothesis]: The new running hypotheses.
|
||||
|
||||
"""
|
||||
logging.debug(f"the number of running hypotheses: {len(running_hyps)}")
|
||||
if self.token_list is not None:
|
||||
logging.debug(
|
||||
"best hypo: " + "".join([self.token_list[x] for x in running_hyps[0].yseq[1:]])
|
||||
)
|
||||
# add eos in the final loop to avoid that there are no ended hyps
|
||||
if i == maxlen - 1:
|
||||
logging.info("adding <eos> in the last position in the loop")
|
||||
running_hyps = [
|
||||
h._replace(yseq=self.append_token(h.yseq, self.eos)) for h in running_hyps
|
||||
]
|
||||
|
||||
# add ended hypotheses to a final list, and removed them from current hypotheses
|
||||
# (this will be a problem, number of hyps < beam)
|
||||
remained_hyps = []
|
||||
for hyp in running_hyps:
|
||||
if hyp.yseq[-1] == self.eos:
|
||||
# e.g., Word LM needs to add final <eos> score
|
||||
for k, d in chain(self.full_scorers.items(), self.part_scorers.items()):
|
||||
s = d.final_score(hyp.states[k])
|
||||
hyp.scores[k] += s
|
||||
hyp = hyp._replace(score=hyp.score + self.weights[k] * s)
|
||||
ended_hyps.append(hyp)
|
||||
else:
|
||||
remained_hyps.append(hyp)
|
||||
return remained_hyps
|
||||
@@ -0,0 +1,110 @@
|
||||
# This is an example that demonstrates how to configure a model file.
|
||||
# You can modify the configuration according to your own requirements.
|
||||
|
||||
# to print the register_table:
|
||||
# from funasr.register import tables
|
||||
# tables.print()
|
||||
|
||||
# network architecture
|
||||
model: Transformer
|
||||
model_conf:
|
||||
ctc_weight: 0.3
|
||||
lsm_weight: 0.1 # label smoothing option
|
||||
length_normalized_loss: false
|
||||
|
||||
# encoder
|
||||
encoder: TransformerEncoder
|
||||
encoder_conf:
|
||||
output_size: 256 # dimension of attention
|
||||
attention_heads: 4
|
||||
linear_units: 2048 # the number of units of position-wise feed forward
|
||||
num_blocks: 12 # the number of encoder blocks
|
||||
dropout_rate: 0.1
|
||||
positional_dropout_rate: 0.1
|
||||
attention_dropout_rate: 0.0
|
||||
input_layer: conv2d # encoder architecture type
|
||||
normalize_before: true
|
||||
|
||||
# decoder
|
||||
decoder: TransformerDecoder
|
||||
decoder_conf:
|
||||
attention_heads: 4
|
||||
linear_units: 2048
|
||||
num_blocks: 6
|
||||
dropout_rate: 0.1
|
||||
positional_dropout_rate: 0.1
|
||||
self_attention_dropout_rate: 0.0
|
||||
src_attention_dropout_rate: 0.0
|
||||
|
||||
|
||||
# frontend related
|
||||
frontend: WavFrontend
|
||||
frontend_conf:
|
||||
fs: 16000
|
||||
window: hamming
|
||||
n_mels: 80
|
||||
frame_length: 25
|
||||
frame_shift: 10
|
||||
lfr_m: 1
|
||||
lfr_n: 1
|
||||
|
||||
specaug: SpecAug
|
||||
specaug_conf:
|
||||
apply_time_warp: true
|
||||
time_warp_window: 5
|
||||
time_warp_mode: bicubic
|
||||
apply_freq_mask: true
|
||||
freq_mask_width_range:
|
||||
- 0
|
||||
- 30
|
||||
num_freq_mask: 2
|
||||
apply_time_mask: true
|
||||
time_mask_width_range:
|
||||
- 0
|
||||
- 40
|
||||
num_time_mask: 2
|
||||
|
||||
train_conf:
|
||||
accum_grad: 1
|
||||
grad_clip: 5
|
||||
max_epoch: 150
|
||||
val_scheduler_criterion:
|
||||
- valid
|
||||
- acc
|
||||
best_model_criterion:
|
||||
- - valid
|
||||
- acc
|
||||
- max
|
||||
keep_nbest_models: 10
|
||||
log_interval: 50
|
||||
|
||||
optim: adam
|
||||
optim_conf:
|
||||
lr: 0.002
|
||||
scheduler: warmuplr
|
||||
scheduler_conf:
|
||||
warmup_steps: 30000
|
||||
|
||||
dataset: AudioDataset
|
||||
dataset_conf:
|
||||
index_ds: IndexDSJsonl
|
||||
batch_sampler: BatchSampler
|
||||
batch_type: example # example or length
|
||||
batch_size: 1 # if batch_type is example, batch_size is the numbers of samples; if length, batch_size is source_token_len+target_token_len;
|
||||
max_token_length: 2048 # filter samples if source_token_len+target_token_len > max_token_length,
|
||||
buffer_size: 500
|
||||
shuffle: True
|
||||
num_workers: 0
|
||||
|
||||
tokenizer: CharTokenizer
|
||||
tokenizer_conf:
|
||||
unk_symbol: <unk>
|
||||
split_with_space: true
|
||||
|
||||
|
||||
ctc_conf:
|
||||
dropout_rate: 0.0
|
||||
ctc_type: builtin
|
||||
reduce: true
|
||||
ignore_nan_grad: true
|
||||
normalize: null
|
||||
@@ -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