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,322 @@
|
||||
#!/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(numpy.finfo(torch.tensor(0, dtype=scores.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 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 MultiHeadSelfAttention(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, in_feat, n_feat, dropout_rate):
|
||||
"""Construct an MultiHeadedAttention object."""
|
||||
super(MultiHeadSelfAttention, 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_out = nn.Linear(n_feat, n_feat)
|
||||
self.linear_q_k_v = nn.Linear(in_feat, n_feat * 3)
|
||||
self.attn = None
|
||||
self.dropout = nn.Dropout(p=dropout_rate)
|
||||
|
||||
def forward_qkv(self, x):
|
||||
"""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).
|
||||
|
||||
"""
|
||||
b, t, d = x.size()
|
||||
q_k_v = self.linear_q_k_v(x)
|
||||
q, k, v = torch.split(q_k_v, int(self.h * self.d_k), dim=-1)
|
||||
q_h = torch.reshape(q, (b, t, self.h, self.d_k)).transpose(
|
||||
1, 2
|
||||
) # (batch, head, time1, d_k)
|
||||
k_h = torch.reshape(k, (b, t, self.h, self.d_k)).transpose(
|
||||
1, 2
|
||||
) # (batch, head, time2, d_k)
|
||||
v_h = torch.reshape(v, (b, t, self.h, self.d_k)).transpose(
|
||||
1, 2
|
||||
) # (batch, head, time2, d_k)
|
||||
|
||||
return q_h, k_h, v_h, v
|
||||
|
||||
def forward_attention(self, value, scores, mask, mask_att_chunk_encoder=None):
|
||||
"""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:
|
||||
if mask_att_chunk_encoder is not None:
|
||||
mask = mask * mask_att_chunk_encoder
|
||||
|
||||
mask = mask.unsqueeze(1).eq(0) # (batch, 1, *, time2)
|
||||
|
||||
min_value = float(numpy.finfo(torch.tensor(0, dtype=scores.dtype).numpy().dtype).min)
|
||||
scores = scores.masked_fill(mask, min_value)
|
||||
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, x, mask, mask_att_chunk_encoder=None):
|
||||
"""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_h, k_h, v_h, v = self.forward_qkv(x)
|
||||
q_h = q_h * self.d_k ** (-0.5)
|
||||
scores = torch.matmul(q_h, k_h.transpose(-2, -1))
|
||||
att_outs = self.forward_attention(v_h, scores, mask, mask_att_chunk_encoder)
|
||||
return att_outs
|
||||
@@ -0,0 +1,702 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright FunASR (https://github.com/alibaba-damo-academy/FunASR). All Rights Reserved.
|
||||
# MIT License (https://opensource.org/licenses/MIT)
|
||||
import logging
|
||||
import random
|
||||
from contextlib import contextmanager
|
||||
from distutils.version import LooseVersion
|
||||
from itertools import permutations
|
||||
from typing import Dict
|
||||
from typing import Optional
|
||||
from typing import Tuple, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
|
||||
from funasr.models.transformer.utils.nets_utils import to_device
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.decoder.abs_decoder import AbsDecoder
|
||||
from funasr.models.encoder.abs_encoder import AbsEncoder
|
||||
from funasr.frontends.abs_frontend import AbsFrontend
|
||||
from funasr.models.specaug.abs_specaug import AbsSpecAug
|
||||
from funasr.models.specaug.abs_profileaug import AbsProfileAug
|
||||
from funasr.layers.abs_normalize import AbsNormalize
|
||||
from funasr.train_utils.device_funcs import force_gatherable
|
||||
from funasr.models.base_model import FunASRModel
|
||||
from funasr.losses.label_smoothing_loss import LabelSmoothingLoss, SequenceBinaryCrossEntropy
|
||||
from funasr.utils.misc import int2vec
|
||||
from funasr.utils.hinter import hint_once
|
||||
|
||||
if LooseVersion(torch.__version__) >= LooseVersion("1.6.0"):
|
||||
from torch.cuda.amp import autocast
|
||||
else:
|
||||
# Nothing to do if torch<1.6.0
|
||||
@contextmanager
|
||||
def autocast(enabled=True):
|
||||
"""Autocast.
|
||||
|
||||
Args:
|
||||
enabled: TODO.
|
||||
"""
|
||||
yield
|
||||
|
||||
|
||||
class DiarSondModel(FunASRModel):
|
||||
"""Speaker overlap-aware neural diarization model
|
||||
reference: https://arxiv.org/abs/2211.10243
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
frontend: Optional[AbsFrontend],
|
||||
specaug: Optional[AbsSpecAug],
|
||||
profileaug: Optional[AbsProfileAug],
|
||||
normalize: Optional[AbsNormalize],
|
||||
encoder: torch.nn.Module,
|
||||
speaker_encoder: Optional[torch.nn.Module],
|
||||
ci_scorer: torch.nn.Module,
|
||||
cd_scorer: Optional[torch.nn.Module],
|
||||
decoder: torch.nn.Module,
|
||||
token_list: list,
|
||||
lsm_weight: float = 0.1,
|
||||
length_normalized_loss: bool = False,
|
||||
max_spk_num: int = 16,
|
||||
label_aggregator: Optional[torch.nn.Module] = None,
|
||||
normalize_speech_speaker: bool = False,
|
||||
ignore_id: int = -1,
|
||||
speaker_discrimination_loss_weight: float = 1.0,
|
||||
inter_score_loss_weight: float = 0.0,
|
||||
inputs_type: str = "raw",
|
||||
model_regularizer_weight: float = 0.0,
|
||||
freeze_encoder: bool = False,
|
||||
onfly_shuffle_speaker: bool = True,
|
||||
):
|
||||
|
||||
"""Initialize DiarSondModel.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
frontend: Audio frontend for feature extraction.
|
||||
specaug: TODO.
|
||||
profileaug: TODO.
|
||||
normalize: TODO.
|
||||
encoder: TODO.
|
||||
speaker_encoder: TODO.
|
||||
ci_scorer: TODO.
|
||||
cd_scorer: TODO.
|
||||
decoder: TODO.
|
||||
token_list: TODO.
|
||||
lsm_weight: TODO.
|
||||
length_normalized_loss: TODO.
|
||||
max_spk_num: TODO.
|
||||
label_aggregator: TODO.
|
||||
normalize_speech_speaker: TODO.
|
||||
ignore_id: TODO.
|
||||
speaker_discrimination_loss_weight: TODO.
|
||||
inter_score_loss_weight: TODO.
|
||||
inputs_type: TODO.
|
||||
model_regularizer_weight: TODO.
|
||||
freeze_encoder: TODO.
|
||||
onfly_shuffle_speaker: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.encoder = encoder
|
||||
self.speaker_encoder = speaker_encoder
|
||||
self.ci_scorer = ci_scorer
|
||||
self.cd_scorer = cd_scorer
|
||||
self.normalize = normalize
|
||||
self.frontend = frontend
|
||||
self.specaug = specaug
|
||||
self.profileaug = profileaug
|
||||
self.label_aggregator = label_aggregator
|
||||
self.decoder = decoder
|
||||
self.token_list = token_list
|
||||
self.max_spk_num = max_spk_num
|
||||
self.normalize_speech_speaker = normalize_speech_speaker
|
||||
self.ignore_id = ignore_id
|
||||
self.model_regularizer_weight = model_regularizer_weight
|
||||
self.freeze_encoder = freeze_encoder
|
||||
self.onfly_shuffle_speaker = onfly_shuffle_speaker
|
||||
self.criterion_diar = LabelSmoothingLoss(
|
||||
size=vocab_size,
|
||||
padding_idx=ignore_id,
|
||||
smoothing=lsm_weight,
|
||||
normalize_length=length_normalized_loss,
|
||||
)
|
||||
self.criterion_bce = SequenceBinaryCrossEntropy(normalize_length=length_normalized_loss)
|
||||
self.pse_embedding = self.generate_pse_embedding()
|
||||
self.power_weight = torch.from_numpy(
|
||||
2 ** np.arange(max_spk_num)[np.newaxis, np.newaxis, :]
|
||||
).float()
|
||||
self.int_token_arr = torch.from_numpy(
|
||||
np.array(self.token_list).astype(int)[np.newaxis, np.newaxis, :]
|
||||
).int()
|
||||
self.speaker_discrimination_loss_weight = speaker_discrimination_loss_weight
|
||||
self.inter_score_loss_weight = inter_score_loss_weight
|
||||
self.forward_steps = 0
|
||||
self.inputs_type = inputs_type
|
||||
self.to_regularize_parameters = None
|
||||
|
||||
def get_regularize_parameters(self):
|
||||
"""Get regularize parameters."""
|
||||
to_regularize_parameters, normal_parameters = [], []
|
||||
for name, param in self.named_parameters():
|
||||
if (
|
||||
"encoder" in name
|
||||
and "weight" in name
|
||||
and "bn" not in name
|
||||
and ("conv2" in name or "conv1" in name or "conv_sc" in name or "dense" in name)
|
||||
):
|
||||
to_regularize_parameters.append((name, param))
|
||||
else:
|
||||
normal_parameters.append((name, param))
|
||||
self.to_regularize_parameters = to_regularize_parameters
|
||||
return to_regularize_parameters, normal_parameters
|
||||
|
||||
def generate_pse_embedding(self):
|
||||
"""Generate pse embedding."""
|
||||
embedding = np.zeros((len(self.token_list), self.max_spk_num), dtype=np.float32)
|
||||
for idx, pse_label in enumerate(self.token_list):
|
||||
emb = int2vec(int(pse_label), vec_dim=self.max_spk_num, dtype=np.float32)
|
||||
embedding[idx] = emb
|
||||
return torch.from_numpy(embedding)
|
||||
|
||||
def rand_permute_speaker(self, raw_profile, raw_binary_labels):
|
||||
"""
|
||||
raw_profile: B, N, D
|
||||
raw_binary_labels: B, T, N
|
||||
"""
|
||||
assert (
|
||||
raw_profile.shape[1] == raw_binary_labels.shape[2]
|
||||
), "Num profile: {}, Num label: {}".format(
|
||||
raw_profile.shape[1], raw_binary_labels.shape[-1]
|
||||
)
|
||||
profile = torch.clone(raw_profile)
|
||||
binary_labels = torch.clone(raw_binary_labels)
|
||||
bsz, num_spk = profile.shape[0], profile.shape[1]
|
||||
for i in range(bsz):
|
||||
idx = list(range(num_spk))
|
||||
random.shuffle(idx)
|
||||
profile[i] = profile[i][idx, :]
|
||||
binary_labels[i] = binary_labels[i][:, idx]
|
||||
|
||||
return profile, binary_labels
|
||||
|
||||
def forward(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor = None,
|
||||
profile: torch.Tensor = None,
|
||||
profile_lengths: torch.Tensor = None,
|
||||
binary_labels: torch.Tensor = None,
|
||||
binary_labels_lengths: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor], torch.Tensor]:
|
||||
"""Frontend + Encoder + Speaker Encoder + CI Scorer + CD Scorer + Decoder + Calc loss
|
||||
|
||||
Args:
|
||||
speech: (Batch, samples) or (Batch, frames, input_size)
|
||||
speech_lengths: (Batch,) default None for chunk interator,
|
||||
because the chunk-iterator does not
|
||||
have the speech_lengths returned.
|
||||
see in
|
||||
espnet2/iterators/chunk_iter_factory.py
|
||||
profile: (Batch, N_spk, dim)
|
||||
profile_lengths: (Batch,)
|
||||
binary_labels: (Batch, frames, max_spk_num)
|
||||
binary_labels_lengths: (Batch,)
|
||||
"""
|
||||
assert speech.shape[0] <= binary_labels.shape[0], (speech.shape, binary_labels.shape)
|
||||
batch_size = speech.shape[0]
|
||||
if self.freeze_encoder:
|
||||
hint_once("Freeze encoder", "freeze_encoder", rank=0)
|
||||
self.encoder.eval()
|
||||
self.forward_steps = self.forward_steps + 1
|
||||
if self.pse_embedding.device != speech.device:
|
||||
self.pse_embedding = self.pse_embedding.to(speech.device)
|
||||
self.power_weight = self.power_weight.to(speech.device)
|
||||
self.int_token_arr = self.int_token_arr.to(speech.device)
|
||||
|
||||
if self.onfly_shuffle_speaker:
|
||||
hint_once("On-the-fly shuffle speaker permutation.", "onfly_shuffle_speaker", rank=0)
|
||||
profile, binary_labels = self.rand_permute_speaker(profile, binary_labels)
|
||||
|
||||
# 0a. Aggregate time-domain labels to match forward outputs
|
||||
if self.label_aggregator is not None:
|
||||
binary_labels, binary_labels_lengths = self.label_aggregator(
|
||||
binary_labels, binary_labels_lengths
|
||||
)
|
||||
# 0b. augment profiles
|
||||
if self.profileaug is not None and self.training:
|
||||
speech, profile, binary_labels = self.profileaug(
|
||||
speech,
|
||||
speech_lengths,
|
||||
profile,
|
||||
profile_lengths,
|
||||
binary_labels,
|
||||
binary_labels_lengths,
|
||||
)
|
||||
|
||||
# 1. Calculate power-set encoding (PSE) labels
|
||||
pad_bin_labels = F.pad(
|
||||
binary_labels, (0, self.max_spk_num - binary_labels.shape[2]), "constant", 0.0
|
||||
)
|
||||
raw_pse_labels = torch.sum(pad_bin_labels * self.power_weight, dim=2, keepdim=True)
|
||||
pse_labels = torch.argmax((raw_pse_labels.int() == self.int_token_arr).float(), dim=2)
|
||||
|
||||
# 2. Network forward
|
||||
pred, inter_outputs = self.prediction_forward(
|
||||
speech, speech_lengths, profile, profile_lengths, return_inter_outputs=True
|
||||
)
|
||||
(speech, speech_lengths), (profile, profile_lengths), (ci_score, cd_score) = inter_outputs
|
||||
|
||||
# If encoder uses conv* as input_layer (i.e., subsampling),
|
||||
# the sequence length of 'pred' might be slightly less than the
|
||||
# length of 'spk_labels'. Here we force them to be equal.
|
||||
length_diff_tolerance = 2
|
||||
length_diff = abs(pse_labels.shape[1] - pred.shape[1])
|
||||
if length_diff <= length_diff_tolerance:
|
||||
min_len = min(pred.shape[1], pse_labels.shape[1])
|
||||
pse_labels = pse_labels[:, :min_len]
|
||||
pred = pred[:, :min_len]
|
||||
cd_score = cd_score[:, :min_len]
|
||||
ci_score = ci_score[:, :min_len]
|
||||
|
||||
loss_diar = self.classification_loss(pred, pse_labels, binary_labels_lengths)
|
||||
loss_spk_dis = self.speaker_discrimination_loss(profile, profile_lengths)
|
||||
loss_inter_ci, loss_inter_cd = self.internal_score_loss(
|
||||
cd_score, ci_score, pse_labels, binary_labels_lengths
|
||||
)
|
||||
regularizer_loss = None
|
||||
if self.model_regularizer_weight > 0 and self.to_regularize_parameters is not None:
|
||||
regularizer_loss = self.calculate_regularizer_loss()
|
||||
label_mask = make_pad_mask(binary_labels_lengths, maxlen=pse_labels.shape[1]).to(
|
||||
pse_labels.device
|
||||
)
|
||||
loss = (
|
||||
loss_diar
|
||||
+ self.speaker_discrimination_loss_weight * loss_spk_dis
|
||||
+ self.inter_score_loss_weight * (loss_inter_ci + loss_inter_cd)
|
||||
)
|
||||
# if regularizer_loss is not None:
|
||||
# loss = loss + regularizer_loss * self.model_regularizer_weight
|
||||
|
||||
(
|
||||
correct,
|
||||
num_frames,
|
||||
speech_scored,
|
||||
speech_miss,
|
||||
speech_falarm,
|
||||
speaker_scored,
|
||||
speaker_miss,
|
||||
speaker_falarm,
|
||||
speaker_error,
|
||||
) = self.calc_diarization_error(
|
||||
pred=F.embedding(pred.argmax(dim=2) * (~label_mask), self.pse_embedding),
|
||||
label=F.embedding(pse_labels * (~label_mask), self.pse_embedding),
|
||||
length=binary_labels_lengths,
|
||||
)
|
||||
|
||||
if speech_scored > 0 and num_frames > 0:
|
||||
sad_mr, sad_fr, mi, fa, cf, acc, der = (
|
||||
speech_miss / speech_scored,
|
||||
speech_falarm / speech_scored,
|
||||
speaker_miss / speaker_scored,
|
||||
speaker_falarm / speaker_scored,
|
||||
speaker_error / speaker_scored,
|
||||
correct / num_frames,
|
||||
(speaker_miss + speaker_falarm + speaker_error) / speaker_scored,
|
||||
)
|
||||
else:
|
||||
sad_mr, sad_fr, mi, fa, cf, acc, der = 0, 0, 0, 0, 0, 0, 0
|
||||
|
||||
stats = dict(
|
||||
loss=loss.detach(),
|
||||
loss_diar=loss_diar.detach() if loss_diar is not None else None,
|
||||
loss_spk_dis=loss_spk_dis.detach() if loss_spk_dis is not None else None,
|
||||
loss_inter_ci=loss_inter_ci.detach() if loss_inter_ci is not None else None,
|
||||
loss_inter_cd=loss_inter_cd.detach() if loss_inter_cd is not None else None,
|
||||
regularizer_loss=regularizer_loss.detach() if regularizer_loss is not None else None,
|
||||
sad_mr=sad_mr,
|
||||
sad_fr=sad_fr,
|
||||
mi=mi,
|
||||
fa=fa,
|
||||
cf=cf,
|
||||
acc=acc,
|
||||
der=der,
|
||||
forward_steps=self.forward_steps,
|
||||
)
|
||||
|
||||
loss, stats, weight = force_gatherable((loss, stats, batch_size), loss.device)
|
||||
return loss, stats, weight
|
||||
|
||||
def calculate_regularizer_loss(self):
|
||||
"""Calculate regularizer loss."""
|
||||
regularizer_loss = 0.0
|
||||
for name, param in self.to_regularize_parameters:
|
||||
regularizer_loss = regularizer_loss + torch.norm(param, p=2)
|
||||
return regularizer_loss
|
||||
|
||||
def classification_loss(
|
||||
self, predictions: torch.Tensor, labels: torch.Tensor, prediction_lengths: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Classification loss.
|
||||
|
||||
Args:
|
||||
predictions: TODO.
|
||||
labels: TODO.
|
||||
prediction_lengths: Lengths of prediction.
|
||||
"""
|
||||
mask = make_pad_mask(prediction_lengths, maxlen=labels.shape[1])
|
||||
pad_labels = labels.masked_fill(mask.to(predictions.device), value=self.ignore_id)
|
||||
loss = self.criterion_diar(predictions.contiguous(), pad_labels)
|
||||
|
||||
return loss
|
||||
|
||||
def speaker_discrimination_loss(
|
||||
self, profile: torch.Tensor, profile_lengths: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Speaker discrimination loss.
|
||||
|
||||
Args:
|
||||
profile: TODO.
|
||||
profile_lengths: Lengths of profile.
|
||||
"""
|
||||
profile_mask = (
|
||||
torch.linalg.norm(profile, ord=2, dim=2, keepdim=True) > 0
|
||||
).float() # (B, N, 1)
|
||||
mask = torch.matmul(profile_mask, profile_mask.transpose(1, 2)) # (B, N, N)
|
||||
mask = mask * (1.0 - torch.eye(self.max_spk_num).unsqueeze(0).to(mask))
|
||||
|
||||
eps = 1e-12
|
||||
coding_norm = (
|
||||
torch.linalg.norm(
|
||||
profile * profile_mask + (1 - profile_mask) * eps, dim=2, keepdim=True
|
||||
)
|
||||
* profile_mask
|
||||
)
|
||||
# profile: Batch, N, dim
|
||||
cos_theta = (
|
||||
F.cosine_similarity(profile.unsqueeze(2), profile.unsqueeze(1), dim=-1, eps=eps) * mask
|
||||
)
|
||||
cos_theta = torch.clip(cos_theta, -1 + eps, 1 - eps)
|
||||
loss = (F.relu(mask * coding_norm * (cos_theta - 0.0))).sum() / mask.sum()
|
||||
|
||||
return loss
|
||||
|
||||
def calculate_multi_labels(self, pse_labels, pse_labels_lengths):
|
||||
"""Calculate multi labels.
|
||||
|
||||
Args:
|
||||
pse_labels: TODO.
|
||||
pse_labels_lengths: Lengths of pse_labels.
|
||||
"""
|
||||
mask = make_pad_mask(pse_labels_lengths, maxlen=pse_labels.shape[1])
|
||||
padding_labels = pse_labels.masked_fill(mask.to(pse_labels.device), value=0).to(pse_labels)
|
||||
multi_labels = F.embedding(padding_labels, self.pse_embedding)
|
||||
|
||||
return multi_labels
|
||||
|
||||
def internal_score_loss(
|
||||
self,
|
||||
cd_score: torch.Tensor,
|
||||
ci_score: torch.Tensor,
|
||||
pse_labels: torch.Tensor,
|
||||
pse_labels_lengths: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Internal score loss.
|
||||
|
||||
Args:
|
||||
cd_score: TODO.
|
||||
ci_score: TODO.
|
||||
pse_labels: TODO.
|
||||
pse_labels_lengths: Lengths of pse_labels.
|
||||
"""
|
||||
multi_labels = self.calculate_multi_labels(pse_labels, pse_labels_lengths)
|
||||
ci_loss = self.criterion_bce(ci_score, multi_labels, pse_labels_lengths)
|
||||
cd_loss = self.criterion_bce(cd_score, multi_labels, pse_labels_lengths)
|
||||
return ci_loss, cd_loss
|
||||
|
||||
def collect_feats(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor,
|
||||
profile: torch.Tensor = None,
|
||||
profile_lengths: torch.Tensor = None,
|
||||
binary_labels: torch.Tensor = None,
|
||||
binary_labels_lengths: torch.Tensor = None,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Collect feats.
|
||||
|
||||
Args:
|
||||
speech: Speech audio tensor, shape (batch, time).
|
||||
speech_lengths: Length of each speech sample.
|
||||
profile: TODO.
|
||||
profile_lengths: Lengths of profile.
|
||||
binary_labels: TODO.
|
||||
binary_labels_lengths: Lengths of binary_labels.
|
||||
"""
|
||||
feats, feats_lengths = self._extract_feats(speech, speech_lengths)
|
||||
return {"feats": feats, "feats_lengths": feats_lengths}
|
||||
|
||||
def encode_speaker(
|
||||
self,
|
||||
profile: torch.Tensor,
|
||||
profile_lengths: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Encode speaker.
|
||||
|
||||
Args:
|
||||
profile: TODO.
|
||||
profile_lengths: Lengths of profile.
|
||||
"""
|
||||
with autocast(False):
|
||||
if profile.shape[1] < self.max_spk_num:
|
||||
profile = F.pad(
|
||||
profile, [0, 0, 0, self.max_spk_num - profile.shape[1], 0, 0], "constant", 0.0
|
||||
)
|
||||
profile_mask = (torch.linalg.norm(profile, ord=2, dim=2, keepdim=True) > 0).float()
|
||||
profile = F.normalize(profile, dim=2)
|
||||
if self.speaker_encoder is not None:
|
||||
profile = self.speaker_encoder(profile, profile_lengths)[0]
|
||||
return profile * profile_mask, profile_lengths
|
||||
else:
|
||||
return profile, profile_lengths
|
||||
|
||||
def encode_speech(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Encode speech.
|
||||
|
||||
Args:
|
||||
speech: Speech audio tensor, shape (batch, time).
|
||||
speech_lengths: Length of each speech sample.
|
||||
"""
|
||||
if self.encoder is not None and self.inputs_type == "raw":
|
||||
speech, speech_lengths = self.encode(speech, speech_lengths)
|
||||
speech_mask = ~make_pad_mask(speech_lengths, maxlen=speech.shape[1])
|
||||
speech_mask = speech_mask.to(speech.device).unsqueeze(-1).float()
|
||||
return speech * speech_mask, speech_lengths
|
||||
else:
|
||||
return speech, speech_lengths
|
||||
|
||||
@staticmethod
|
||||
def concate_speech_ivc(speech: torch.Tensor, ivc: torch.Tensor) -> torch.Tensor:
|
||||
"""Concate speech ivc.
|
||||
|
||||
Args:
|
||||
speech: Speech audio tensor, shape (batch, time).
|
||||
ivc: TODO.
|
||||
"""
|
||||
nn, tt = ivc.shape[1], speech.shape[1]
|
||||
speech = speech.unsqueeze(dim=1) # B x 1 x T x D
|
||||
speech = speech.expand(-1, nn, -1, -1) # B x N x T x D
|
||||
ivc = ivc.unsqueeze(dim=2) # B x N x 1 x D
|
||||
ivc = ivc.expand(-1, -1, tt, -1) # B x N x T x D
|
||||
sd_in = torch.cat([speech, ivc], dim=3) # B x N x T x 2D
|
||||
return sd_in
|
||||
|
||||
def calc_similarity(
|
||||
self,
|
||||
speech_encoder_outputs: torch.Tensor,
|
||||
speaker_encoder_outputs: torch.Tensor,
|
||||
seq_len: torch.Tensor = None,
|
||||
spk_len: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Calc similarity.
|
||||
|
||||
Args:
|
||||
speech_encoder_outputs: TODO.
|
||||
speaker_encoder_outputs: TODO.
|
||||
seq_len: TODO.
|
||||
spk_len: TODO.
|
||||
"""
|
||||
bb, tt = speech_encoder_outputs.shape[0], speech_encoder_outputs.shape[1]
|
||||
d_sph, d_spk = speech_encoder_outputs.shape[2], speaker_encoder_outputs.shape[2]
|
||||
if self.normalize_speech_speaker:
|
||||
speech_encoder_outputs = F.normalize(speech_encoder_outputs, dim=2)
|
||||
speaker_encoder_outputs = F.normalize(speaker_encoder_outputs, dim=2)
|
||||
ge_in = self.concate_speech_ivc(speech_encoder_outputs, speaker_encoder_outputs)
|
||||
ge_in = torch.reshape(ge_in, [bb * self.max_spk_num, tt, d_sph + d_spk])
|
||||
ge_len = seq_len.unsqueeze(1).expand(-1, self.max_spk_num)
|
||||
ge_len = torch.reshape(ge_len, [bb * self.max_spk_num])
|
||||
cd_simi = self.cd_scorer(ge_in, ge_len)[0]
|
||||
cd_simi = torch.reshape(cd_simi, [bb, self.max_spk_num, tt, 1])
|
||||
cd_simi = cd_simi.squeeze(dim=3).permute([0, 2, 1])
|
||||
|
||||
if isinstance(self.ci_scorer, AbsEncoder):
|
||||
ci_simi = self.ci_scorer(ge_in, ge_len)[0]
|
||||
ci_simi = torch.reshape(ci_simi, [bb, self.max_spk_num, tt]).permute([0, 2, 1])
|
||||
else:
|
||||
ci_simi = self.ci_scorer(speech_encoder_outputs, speaker_encoder_outputs)
|
||||
|
||||
return ci_simi, cd_simi
|
||||
|
||||
def post_net_forward(self, simi, seq_len):
|
||||
"""Post net forward.
|
||||
|
||||
Args:
|
||||
simi: TODO.
|
||||
seq_len: TODO.
|
||||
"""
|
||||
logits = self.decoder(simi, seq_len)[0]
|
||||
|
||||
return logits
|
||||
|
||||
def prediction_forward(
|
||||
self,
|
||||
speech: torch.Tensor,
|
||||
speech_lengths: torch.Tensor,
|
||||
profile: torch.Tensor,
|
||||
profile_lengths: torch.Tensor,
|
||||
return_inter_outputs: bool = False,
|
||||
) -> [torch.Tensor, Optional[list]]:
|
||||
# speech encoding
|
||||
"""Prediction forward.
|
||||
|
||||
Args:
|
||||
speech: Speech audio tensor, shape (batch, time).
|
||||
speech_lengths: Length of each speech sample.
|
||||
profile: TODO.
|
||||
profile_lengths: Lengths of profile.
|
||||
return_inter_outputs: TODO.
|
||||
"""
|
||||
speech, speech_lengths = self.encode_speech(speech, speech_lengths)
|
||||
# speaker encoding
|
||||
profile, profile_lengths = self.encode_speaker(profile, profile_lengths)
|
||||
# calculating similarity
|
||||
ci_simi, cd_simi = self.calc_similarity(speech, profile, speech_lengths, profile_lengths)
|
||||
similarity = torch.cat([cd_simi, ci_simi], dim=2)
|
||||
# post net forward
|
||||
logits = self.post_net_forward(similarity, speech_lengths)
|
||||
|
||||
if return_inter_outputs:
|
||||
return logits, [
|
||||
(speech, speech_lengths),
|
||||
(profile, profile_lengths),
|
||||
(ci_simi, cd_simi),
|
||||
]
|
||||
return logits
|
||||
|
||||
def encode(
|
||||
self, speech: torch.Tensor, speech_lengths: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Frontend + Encoder
|
||||
|
||||
Args:
|
||||
speech: (Batch, Length, ...)
|
||||
speech_lengths: (Batch,)
|
||||
"""
|
||||
with autocast(False):
|
||||
# 1. Extract feats
|
||||
feats, feats_lengths = self._extract_feats(speech, speech_lengths)
|
||||
|
||||
# 2. Data augmentation
|
||||
if self.specaug is not None and self.training:
|
||||
feats, feats_lengths = self.specaug(feats, feats_lengths)
|
||||
|
||||
# 3. Normalization for feature: e.g. Global-CMVN, Utterance-CMVN
|
||||
if self.normalize is not None:
|
||||
feats, feats_lengths = self.normalize(feats, feats_lengths)
|
||||
|
||||
# 4. Forward encoder
|
||||
# feats: (Batch, Length, Dim)
|
||||
# -> encoder_out: (Batch, Length2, Dim)
|
||||
encoder_outputs = self.encoder(feats, feats_lengths)
|
||||
encoder_out, encoder_out_lens = encoder_outputs[:2]
|
||||
|
||||
assert encoder_out.size(0) == speech.size(0), (
|
||||
encoder_out.size(),
|
||||
speech.size(0),
|
||||
)
|
||||
assert encoder_out.size(1) <= encoder_out_lens.max(), (
|
||||
encoder_out.size(),
|
||||
encoder_out_lens.max(),
|
||||
)
|
||||
|
||||
return encoder_out, encoder_out_lens
|
||||
|
||||
def _extract_feats(
|
||||
self, speech: torch.Tensor, speech_lengths: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Internal: extract feats.
|
||||
|
||||
Args:
|
||||
speech: Speech audio tensor, shape (batch, time).
|
||||
speech_lengths: Length of each speech sample.
|
||||
"""
|
||||
batch_size = speech.shape[0]
|
||||
speech_lengths = (
|
||||
speech_lengths
|
||||
if speech_lengths is not None
|
||||
else torch.ones(batch_size).int() * speech.shape[1]
|
||||
)
|
||||
|
||||
assert speech_lengths.dim() == 1, speech_lengths.shape
|
||||
|
||||
# for data-parallel
|
||||
speech = speech[:, : speech_lengths.max()]
|
||||
|
||||
if self.frontend is not None:
|
||||
# Frontend
|
||||
# e.g. STFT and Feature extract
|
||||
# data_loader may send time-domain signal in this case
|
||||
# speech (Batch, NSamples) -> feats: (Batch, NFrames, Dim)
|
||||
feats, feats_lengths = self.frontend(speech, speech_lengths)
|
||||
else:
|
||||
# No frontend and no feature extract
|
||||
feats, feats_lengths = speech, speech_lengths
|
||||
return feats, feats_lengths
|
||||
|
||||
@staticmethod
|
||||
def calc_diarization_error(pred, label, length):
|
||||
# Note (jiatong): Credit to https://github.com/hitachi-speech/EEND
|
||||
|
||||
"""Calc diarization error.
|
||||
|
||||
Args:
|
||||
pred: TODO.
|
||||
label: TODO.
|
||||
length: TODO.
|
||||
"""
|
||||
(batch_size, max_len, num_output) = label.size()
|
||||
# mask the padding part
|
||||
mask = ~make_pad_mask(length, maxlen=label.shape[1]).unsqueeze(-1).numpy()
|
||||
|
||||
# pred and label have the shape (batch_size, max_len, num_output)
|
||||
label_np = label.data.cpu().numpy().astype(int)
|
||||
pred_np = (pred.data.cpu().numpy() > 0).astype(int)
|
||||
label_np = label_np * mask
|
||||
pred_np = pred_np * mask
|
||||
length = length.data.cpu().numpy()
|
||||
|
||||
# compute speech activity detection error
|
||||
n_ref = np.sum(label_np, axis=2)
|
||||
n_sys = np.sum(pred_np, axis=2)
|
||||
speech_scored = float(np.sum(n_ref > 0))
|
||||
speech_miss = float(np.sum(np.logical_and(n_ref > 0, n_sys == 0)))
|
||||
speech_falarm = float(np.sum(np.logical_and(n_ref == 0, n_sys > 0)))
|
||||
|
||||
# compute speaker diarization error
|
||||
speaker_scored = float(np.sum(n_ref))
|
||||
speaker_miss = float(np.sum(np.maximum(n_ref - n_sys, 0)))
|
||||
speaker_falarm = float(np.sum(np.maximum(n_sys - n_ref, 0)))
|
||||
n_map = np.sum(np.logical_and(label_np == 1, pred_np == 1), axis=2)
|
||||
speaker_error = float(np.sum(np.minimum(n_ref, n_sys) - n_map))
|
||||
correct = float(1.0 * np.sum((label_np == pred_np) * mask) / num_output)
|
||||
num_frames = np.sum(length)
|
||||
return (
|
||||
correct,
|
||||
num_frames,
|
||||
speech_scored,
|
||||
speech_miss,
|
||||
speech_falarm,
|
||||
speaker_scored,
|
||||
speaker_miss,
|
||||
speaker_falarm,
|
||||
speaker_error,
|
||||
)
|
||||
@@ -0,0 +1,46 @@
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
|
||||
|
||||
class DotScorer(torch.nn.Module):
|
||||
def __init__(self):
|
||||
"""Initialize DotScorer."""
|
||||
super().__init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
spk_emb: torch.Tensor,
|
||||
):
|
||||
# xs_pad: B, T, D
|
||||
# spk_emb: B, N, D
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
spk_emb: TODO.
|
||||
"""
|
||||
scores = torch.matmul(xs_pad, spk_emb.transpose(1, 2))
|
||||
return scores
|
||||
|
||||
|
||||
class CosScorer(torch.nn.Module):
|
||||
def __init__(self):
|
||||
"""Initialize CosScorer."""
|
||||
super().__init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
spk_emb: torch.Tensor,
|
||||
):
|
||||
# xs_pad: B, T, D
|
||||
# spk_emb: B, N, D
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
spk_emb: TODO.
|
||||
"""
|
||||
scores = F.cosine_similarity(xs_pad.unsqueeze(2), spk_emb.unsqueeze(1), dim=-1)
|
||||
return scores
|
||||
@@ -0,0 +1,224 @@
|
||||
from typing import List
|
||||
from typing import Optional
|
||||
from typing import Sequence
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
import logging
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
import numpy as np
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
from funasr.models.encoder.abs_encoder import AbsEncoder
|
||||
import math
|
||||
from funasr.models.transformer.utils.repeat import repeat
|
||||
|
||||
|
||||
class EncoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_units,
|
||||
num_units,
|
||||
kernel_size=3,
|
||||
activation="tanh",
|
||||
stride=1,
|
||||
include_batch_norm=False,
|
||||
residual=False,
|
||||
):
|
||||
"""Initialize EncoderLayer.
|
||||
|
||||
Args:
|
||||
input_units: TODO.
|
||||
num_units: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
activation: TODO.
|
||||
stride: TODO.
|
||||
include_batch_norm: TODO.
|
||||
residual: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
left_padding = math.ceil((kernel_size - stride) / 2)
|
||||
right_padding = kernel_size - stride - left_padding
|
||||
self.conv_padding = nn.ConstantPad1d((left_padding, right_padding), 0.0)
|
||||
self.conv1d = nn.Conv1d(
|
||||
input_units,
|
||||
num_units,
|
||||
kernel_size,
|
||||
stride,
|
||||
)
|
||||
self.activation = self.get_activation(activation)
|
||||
if include_batch_norm:
|
||||
self.bn = nn.BatchNorm1d(num_units, momentum=0.99, eps=1e-3)
|
||||
self.residual = residual
|
||||
self.include_batch_norm = include_batch_norm
|
||||
self.input_units = input_units
|
||||
self.num_units = num_units
|
||||
self.stride = stride
|
||||
|
||||
@staticmethod
|
||||
def get_activation(activation):
|
||||
"""Get activation.
|
||||
|
||||
Args:
|
||||
activation: TODO.
|
||||
"""
|
||||
if activation == "tanh":
|
||||
return nn.Tanh()
|
||||
else:
|
||||
return nn.ReLU()
|
||||
|
||||
def forward(self, xs_pad, ilens=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
"""
|
||||
outputs = self.conv1d(self.conv_padding(xs_pad))
|
||||
if self.residual and self.stride == 1 and self.input_units == self.num_units:
|
||||
outputs = outputs + xs_pad
|
||||
|
||||
if self.include_batch_norm:
|
||||
outputs = self.bn(outputs)
|
||||
|
||||
# add parenthesis for repeat module
|
||||
return self.activation(outputs), ilens
|
||||
|
||||
|
||||
class ConvEncoder(AbsEncoder):
|
||||
"""
|
||||
Author: Speech Lab of DAMO Academy, Alibaba Group
|
||||
Convolution encoder in OpenNMT framework
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_layers,
|
||||
input_units,
|
||||
num_units,
|
||||
kernel_size=3,
|
||||
dropout_rate=0.3,
|
||||
position_encoder=None,
|
||||
activation="tanh",
|
||||
auxiliary_states=True,
|
||||
out_units=None,
|
||||
out_norm=False,
|
||||
out_residual=False,
|
||||
include_batchnorm=False,
|
||||
regularization_weight=0.0,
|
||||
stride=1,
|
||||
tf2torch_tensor_name_prefix_torch: str = "speaker_encoder",
|
||||
tf2torch_tensor_name_prefix_tf: str = "EAND/speaker_encoder",
|
||||
):
|
||||
"""Initialize ConvEncoder.
|
||||
|
||||
Args:
|
||||
num_layers: TODO.
|
||||
input_units: TODO.
|
||||
num_units: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dropout_rate: TODO.
|
||||
position_encoder: TODO.
|
||||
activation: TODO.
|
||||
auxiliary_states: TODO.
|
||||
out_units: TODO.
|
||||
out_norm: TODO.
|
||||
out_residual: TODO.
|
||||
include_batchnorm: TODO.
|
||||
regularization_weight: TODO.
|
||||
stride: TODO.
|
||||
tf2torch_tensor_name_prefix_torch: TODO.
|
||||
tf2torch_tensor_name_prefix_tf: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self._output_size = num_units
|
||||
|
||||
self.num_layers = num_layers
|
||||
self.input_units = input_units
|
||||
self.num_units = num_units
|
||||
self.kernel_size = kernel_size
|
||||
self.dropout_rate = dropout_rate
|
||||
self.position_encoder = position_encoder
|
||||
self.out_units = out_units
|
||||
self.auxiliary_states = auxiliary_states
|
||||
self.out_norm = out_norm
|
||||
self.activation = activation
|
||||
self.out_residual = out_residual
|
||||
self.include_batch_norm = include_batchnorm
|
||||
self.regularization_weight = regularization_weight
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
if isinstance(stride, int):
|
||||
self.stride = [stride] * self.num_layers
|
||||
else:
|
||||
self.stride = stride
|
||||
self.downsample_rate = 1
|
||||
for s in self.stride:
|
||||
self.downsample_rate *= s
|
||||
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.cnn_a = repeat(
|
||||
self.num_layers,
|
||||
lambda lnum: EncoderLayer(
|
||||
input_units if lnum == 0 else num_units,
|
||||
num_units,
|
||||
kernel_size,
|
||||
activation,
|
||||
self.stride[lnum],
|
||||
include_batchnorm,
|
||||
residual=True if lnum > 0 else False,
|
||||
),
|
||||
)
|
||||
|
||||
if self.out_units is not None:
|
||||
left_padding = math.ceil((kernel_size - stride) / 2)
|
||||
right_padding = kernel_size - stride - left_padding
|
||||
self.out_padding = nn.ConstantPad1d((left_padding, right_padding), 0.0)
|
||||
self.conv_out = nn.Conv1d(
|
||||
num_units,
|
||||
out_units,
|
||||
kernel_size,
|
||||
)
|
||||
|
||||
if self.out_norm:
|
||||
self.after_norm = LayerNorm(out_units)
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self.num_units
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
ilens: torch.Tensor,
|
||||
prev_states: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
inputs = xs_pad
|
||||
if self.position_encoder is not None:
|
||||
inputs = self.position_encoder(inputs)
|
||||
|
||||
if self.dropout_rate > 0:
|
||||
inputs = self.dropout(inputs)
|
||||
|
||||
outputs, _ = self.cnn_a(inputs.transpose(1, 2), ilens)
|
||||
|
||||
if self.out_units is not None:
|
||||
outputs = self.conv_out(self.out_padding(outputs))
|
||||
|
||||
outputs = outputs.transpose(1, 2)
|
||||
if self.out_norm:
|
||||
outputs = self.after_norm(outputs)
|
||||
|
||||
if self.out_residual:
|
||||
outputs = outputs + inputs
|
||||
|
||||
return outputs, ilens, None
|
||||
@@ -0,0 +1,843 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class _BatchNorm1d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_shape=None,
|
||||
input_size=None,
|
||||
eps=1e-05,
|
||||
momentum=0.1,
|
||||
affine=True,
|
||||
track_running_stats=True,
|
||||
combine_batch_time=False,
|
||||
skip_transpose=False,
|
||||
):
|
||||
"""Initialize _BatchNorm1d.
|
||||
|
||||
Args:
|
||||
input_shape: TODO.
|
||||
input_size: Size/dimension parameter.
|
||||
eps: TODO.
|
||||
momentum: TODO.
|
||||
affine: TODO.
|
||||
track_running_stats: TODO.
|
||||
combine_batch_time: TODO.
|
||||
skip_transpose: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.combine_batch_time = combine_batch_time
|
||||
self.skip_transpose = skip_transpose
|
||||
|
||||
if input_size is None and skip_transpose:
|
||||
input_size = input_shape[1]
|
||||
elif input_size is None:
|
||||
input_size = input_shape[-1]
|
||||
|
||||
self.norm = nn.BatchNorm1d(
|
||||
input_size,
|
||||
eps=eps,
|
||||
momentum=momentum,
|
||||
affine=affine,
|
||||
track_running_stats=track_running_stats,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
shape_or = x.shape
|
||||
if self.combine_batch_time:
|
||||
if x.ndim == 3:
|
||||
x = x.reshape(shape_or[0] * shape_or[1], shape_or[2])
|
||||
else:
|
||||
x = x.reshape(shape_or[0] * shape_or[1], shape_or[3], shape_or[2])
|
||||
|
||||
elif not self.skip_transpose:
|
||||
x = x.transpose(-1, 1)
|
||||
|
||||
x_n = self.norm(x)
|
||||
|
||||
if self.combine_batch_time:
|
||||
x_n = x_n.reshape(shape_or)
|
||||
elif not self.skip_transpose:
|
||||
x_n = x_n.transpose(1, -1)
|
||||
|
||||
return x_n
|
||||
|
||||
|
||||
class _Conv1d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
input_shape=None,
|
||||
in_channels=None,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
padding="same",
|
||||
groups=1,
|
||||
bias=True,
|
||||
padding_mode="reflect",
|
||||
skip_transpose=False,
|
||||
):
|
||||
"""Initialize _Conv1d.
|
||||
|
||||
Args:
|
||||
out_channels: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
input_shape: TODO.
|
||||
in_channels: TODO.
|
||||
stride: TODO.
|
||||
dilation: TODO.
|
||||
padding: TODO.
|
||||
groups: TODO.
|
||||
bias: TODO.
|
||||
padding_mode: TODO.
|
||||
skip_transpose: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.stride = stride
|
||||
self.dilation = dilation
|
||||
self.padding = padding
|
||||
self.padding_mode = padding_mode
|
||||
self.unsqueeze = False
|
||||
self.skip_transpose = skip_transpose
|
||||
|
||||
if input_shape is None and in_channels is None:
|
||||
raise ValueError("Must provide one of input_shape or in_channels")
|
||||
|
||||
if in_channels is None:
|
||||
in_channels = self._check_input_shape(input_shape)
|
||||
|
||||
self.conv = nn.Conv1d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
self.kernel_size,
|
||||
stride=self.stride,
|
||||
dilation=self.dilation,
|
||||
padding=0,
|
||||
groups=groups,
|
||||
bias=bias,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
if not self.skip_transpose:
|
||||
x = x.transpose(1, -1)
|
||||
|
||||
if self.unsqueeze:
|
||||
x = x.unsqueeze(1)
|
||||
|
||||
if self.padding == "same":
|
||||
x = self._manage_padding(x, self.kernel_size, self.dilation, self.stride)
|
||||
|
||||
elif self.padding == "causal":
|
||||
num_pad = (self.kernel_size - 1) * self.dilation
|
||||
x = F.pad(x, (num_pad, 0))
|
||||
|
||||
elif self.padding == "valid":
|
||||
pass
|
||||
|
||||
else:
|
||||
raise ValueError("Padding must be 'same', 'valid' or 'causal'. Got " + self.padding)
|
||||
|
||||
wx = self.conv(x)
|
||||
|
||||
if self.unsqueeze:
|
||||
wx = wx.squeeze(1)
|
||||
|
||||
if not self.skip_transpose:
|
||||
wx = wx.transpose(1, -1)
|
||||
|
||||
return wx
|
||||
|
||||
def _manage_padding(
|
||||
self,
|
||||
x,
|
||||
kernel_size: int,
|
||||
dilation: int,
|
||||
stride: int,
|
||||
):
|
||||
# Detecting input shape
|
||||
"""Internal: manage padding.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dilation: TODO.
|
||||
stride: TODO.
|
||||
"""
|
||||
L_in = x.shape[-1]
|
||||
|
||||
# Time padding
|
||||
padding = get_padding_elem(L_in, stride, kernel_size, dilation)
|
||||
|
||||
# Applying padding
|
||||
x = F.pad(x, padding, mode=self.padding_mode)
|
||||
|
||||
return x
|
||||
|
||||
def _check_input_shape(self, shape):
|
||||
"""Checks the input shape and returns the number of input channels."""
|
||||
|
||||
if len(shape) == 2:
|
||||
self.unsqueeze = True
|
||||
in_channels = 1
|
||||
elif self.skip_transpose:
|
||||
in_channels = shape[1]
|
||||
elif len(shape) == 3:
|
||||
in_channels = shape[2]
|
||||
else:
|
||||
raise ValueError("conv1d expects 2d, 3d inputs. Got " + str(len(shape)))
|
||||
|
||||
# Kernel size must be odd
|
||||
if self.kernel_size % 2 == 0:
|
||||
raise ValueError(
|
||||
"The field kernel size must be an odd number. Got %s." % (self.kernel_size)
|
||||
)
|
||||
return in_channels
|
||||
|
||||
|
||||
def get_padding_elem(L_in: int, stride: int, kernel_size: int, dilation: int):
|
||||
"""Get padding elem.
|
||||
|
||||
Args:
|
||||
L_in: TODO.
|
||||
stride: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dilation: TODO.
|
||||
"""
|
||||
if stride > 1:
|
||||
n_steps = math.ceil(((L_in - kernel_size * dilation) / stride) + 1)
|
||||
L_out = stride * (n_steps - 1) + kernel_size * dilation
|
||||
padding = [kernel_size // 2, kernel_size // 2]
|
||||
|
||||
else:
|
||||
L_out = (L_in - dilation * (kernel_size - 1) - 1) // stride + 1
|
||||
|
||||
padding = [(L_in - L_out) // 2, (L_in - L_out) // 2]
|
||||
return padding
|
||||
|
||||
|
||||
# Skip transpose as much as possible for efficiency
|
||||
class Conv1d(_Conv1d):
|
||||
def __init__(self, *args, **kwargs):
|
||||
"""Initialize Conv1d.
|
||||
|
||||
Args:
|
||||
*args: Variable positional arguments.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__(skip_transpose=True, *args, **kwargs)
|
||||
|
||||
|
||||
class BatchNorm1d(_BatchNorm1d):
|
||||
def __init__(self, *args, **kwargs):
|
||||
"""Initialize BatchNorm1d.
|
||||
|
||||
Args:
|
||||
*args: Variable positional arguments.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super().__init__(skip_transpose=True, *args, **kwargs)
|
||||
|
||||
|
||||
def length_to_mask(length, max_len=None, dtype=None, device=None):
|
||||
"""Length to mask.
|
||||
|
||||
Args:
|
||||
length: TODO.
|
||||
max_len: TODO.
|
||||
dtype: TODO.
|
||||
device: Target device ("cuda:0", "cpu", etc.).
|
||||
"""
|
||||
assert len(length.shape) == 1
|
||||
|
||||
if max_len is None:
|
||||
max_len = length.max().long().item() # using arange to generate mask
|
||||
mask = torch.arange(max_len, device=length.device, dtype=length.dtype).expand(
|
||||
len(length), max_len
|
||||
) < length.unsqueeze(1)
|
||||
|
||||
if dtype is None:
|
||||
dtype = length.dtype
|
||||
|
||||
if device is None:
|
||||
device = length.device
|
||||
|
||||
mask = torch.as_tensor(mask, dtype=dtype, device=device)
|
||||
return mask
|
||||
|
||||
|
||||
class TDNNBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
dilation,
|
||||
activation=nn.ReLU,
|
||||
groups=1,
|
||||
):
|
||||
"""Initialize TDNNBlock.
|
||||
|
||||
Args:
|
||||
in_channels: TODO.
|
||||
out_channels: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dilation: TODO.
|
||||
activation: TODO.
|
||||
groups: TODO.
|
||||
"""
|
||||
super(TDNNBlock, self).__init__()
|
||||
self.conv = Conv1d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
)
|
||||
self.activation = activation()
|
||||
self.norm = BatchNorm1d(input_size=out_channels)
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
return self.norm(self.activation(self.conv(x)))
|
||||
|
||||
|
||||
class Res2NetBlock(torch.nn.Module):
|
||||
"""An implementation of Res2NetBlock w/ dilation.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
in_channels : int
|
||||
The number of channels expected in the input.
|
||||
out_channels : int
|
||||
The number of output channels.
|
||||
scale : int
|
||||
The scale of the Res2Net block.
|
||||
kernel_size: int
|
||||
The kernel size of the Res2Net block.
|
||||
dilation : int
|
||||
The dilation of the Res2Net block.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> inp_tensor = torch.rand([8, 120, 64]).transpose(1, 2)
|
||||
>>> layer = Res2NetBlock(64, 64, scale=4, dilation=3)
|
||||
>>> out_tensor = layer(inp_tensor).transpose(1, 2)
|
||||
>>> out_tensor.shape
|
||||
torch.Size([8, 120, 64])
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, out_channels, scale=8, kernel_size=3, dilation=1):
|
||||
"""Initialize Res2NetBlock.
|
||||
|
||||
Args:
|
||||
in_channels: TODO.
|
||||
out_channels: TODO.
|
||||
scale: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dilation: TODO.
|
||||
"""
|
||||
super(Res2NetBlock, self).__init__()
|
||||
assert in_channels % scale == 0
|
||||
assert out_channels % scale == 0
|
||||
|
||||
in_channel = in_channels // scale
|
||||
hidden_channel = out_channels // scale
|
||||
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
TDNNBlock(
|
||||
in_channel,
|
||||
hidden_channel,
|
||||
kernel_size=kernel_size,
|
||||
dilation=dilation,
|
||||
)
|
||||
for i in range(scale - 1)
|
||||
]
|
||||
)
|
||||
self.scale = scale
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
y = []
|
||||
for i, x_i in enumerate(torch.chunk(x, self.scale, dim=1)):
|
||||
if i == 0:
|
||||
y_i = x_i
|
||||
elif i == 1:
|
||||
y_i = self.blocks[i - 1](x_i)
|
||||
else:
|
||||
y_i = self.blocks[i - 1](x_i + y_i)
|
||||
y.append(y_i)
|
||||
y = torch.cat(y, dim=1)
|
||||
return y
|
||||
|
||||
|
||||
class SEBlock(nn.Module):
|
||||
"""An implementation of squeeze-and-excitation block.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
in_channels : int
|
||||
The number of input channels.
|
||||
se_channels : int
|
||||
The number of output channels after squeeze.
|
||||
out_channels : int
|
||||
The number of output channels.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> inp_tensor = torch.rand([8, 120, 64]).transpose(1, 2)
|
||||
>>> se_layer = SEBlock(64, 16, 64)
|
||||
>>> lengths = torch.rand((8,))
|
||||
>>> out_tensor = se_layer(inp_tensor, lengths).transpose(1, 2)
|
||||
>>> out_tensor.shape
|
||||
torch.Size([8, 120, 64])
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, se_channels, out_channels):
|
||||
"""Initialize SEBlock.
|
||||
|
||||
Args:
|
||||
in_channels: TODO.
|
||||
se_channels: TODO.
|
||||
out_channels: TODO.
|
||||
"""
|
||||
super(SEBlock, self).__init__()
|
||||
|
||||
self.conv1 = Conv1d(in_channels=in_channels, out_channels=se_channels, kernel_size=1)
|
||||
self.relu = torch.nn.ReLU(inplace=True)
|
||||
self.conv2 = Conv1d(in_channels=se_channels, out_channels=out_channels, kernel_size=1)
|
||||
self.sigmoid = torch.nn.Sigmoid()
|
||||
|
||||
def forward(self, x, lengths=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
lengths: TODO.
|
||||
"""
|
||||
L = x.shape[-1]
|
||||
if lengths is not None:
|
||||
mask = length_to_mask(lengths * L, max_len=L, device=x.device)
|
||||
mask = mask.unsqueeze(1)
|
||||
total = mask.sum(dim=2, keepdim=True)
|
||||
s = (x * mask).sum(dim=2, keepdim=True) / total
|
||||
else:
|
||||
s = x.mean(dim=2, keepdim=True)
|
||||
|
||||
s = self.relu(self.conv1(s))
|
||||
s = self.sigmoid(self.conv2(s))
|
||||
|
||||
return s * x
|
||||
|
||||
|
||||
class AttentiveStatisticsPooling(nn.Module):
|
||||
"""This class implements an attentive statistic pooling layer for each channel.
|
||||
It returns the concatenated mean and std of the input tensor.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
channels: int
|
||||
The number of input channels.
|
||||
attention_channels: int
|
||||
The number of attention channels.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> inp_tensor = torch.rand([8, 120, 64]).transpose(1, 2)
|
||||
>>> asp_layer = AttentiveStatisticsPooling(64)
|
||||
>>> lengths = torch.rand((8,))
|
||||
>>> out_tensor = asp_layer(inp_tensor, lengths).transpose(1, 2)
|
||||
>>> out_tensor.shape
|
||||
torch.Size([8, 1, 128])
|
||||
"""
|
||||
|
||||
def __init__(self, channels, attention_channels=128, global_context=True):
|
||||
"""Initialize AttentiveStatisticsPooling.
|
||||
|
||||
Args:
|
||||
channels: TODO.
|
||||
attention_channels: TODO.
|
||||
global_context: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.eps = 1e-12
|
||||
self.global_context = global_context
|
||||
if global_context:
|
||||
self.tdnn = TDNNBlock(channels * 3, attention_channels, 1, 1)
|
||||
else:
|
||||
self.tdnn = TDNNBlock(channels, attention_channels, 1, 1)
|
||||
self.tanh = nn.Tanh()
|
||||
self.conv = Conv1d(in_channels=attention_channels, out_channels=channels, kernel_size=1)
|
||||
|
||||
def forward(self, x, lengths=None):
|
||||
"""Calculates mean and std for a batch (input tensor).
|
||||
|
||||
Arguments
|
||||
---------
|
||||
x : torch.Tensor
|
||||
Tensor of shape [N, C, L].
|
||||
"""
|
||||
L = x.shape[-1]
|
||||
|
||||
def _compute_statistics(x, m, dim=2, eps=self.eps):
|
||||
"""Internal: compute statistics.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
m: TODO.
|
||||
dim: TODO.
|
||||
eps: TODO.
|
||||
"""
|
||||
mean = (m * x).sum(dim)
|
||||
std = torch.sqrt((m * (x - mean.unsqueeze(dim)).pow(2)).sum(dim).clamp(eps))
|
||||
return mean, std
|
||||
|
||||
if lengths is None:
|
||||
lengths = torch.ones(x.shape[0], device=x.device)
|
||||
|
||||
# Make binary mask of shape [N, 1, L]
|
||||
mask = length_to_mask(lengths * L, max_len=L, device=x.device)
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
# Expand the temporal context of the pooling layer by allowing the
|
||||
# self-attention to look at global properties of the utterance.
|
||||
if self.global_context:
|
||||
# torch.std is unstable for backward computation
|
||||
# https://github.com/pytorch/pytorch/issues/4320
|
||||
total = mask.sum(dim=2, keepdim=True).float()
|
||||
mean, std = _compute_statistics(x, mask / total)
|
||||
mean = mean.unsqueeze(2).repeat(1, 1, L)
|
||||
std = std.unsqueeze(2).repeat(1, 1, L)
|
||||
attn = torch.cat([x, mean, std], dim=1)
|
||||
else:
|
||||
attn = x
|
||||
|
||||
# Apply layers
|
||||
attn = self.conv(self.tanh(self.tdnn(attn)))
|
||||
|
||||
# Filter out zero-paddings
|
||||
attn = attn.masked_fill(mask == 0, float("-inf"))
|
||||
|
||||
attn = F.softmax(attn, dim=2)
|
||||
mean, std = _compute_statistics(x, attn)
|
||||
# Append mean and std of the batch
|
||||
pooled_stats = torch.cat((mean, std), dim=1)
|
||||
pooled_stats = pooled_stats.unsqueeze(2)
|
||||
|
||||
return pooled_stats
|
||||
|
||||
|
||||
class SERes2NetBlock(nn.Module):
|
||||
"""An implementation of building block in ECAPA-TDNN, i.e.,
|
||||
TDNN-Res2Net-TDNN-SEBlock.
|
||||
|
||||
Arguments
|
||||
----------
|
||||
out_channels: int
|
||||
The number of output channels.
|
||||
res2net_scale: int
|
||||
The scale of the Res2Net block.
|
||||
kernel_size: int
|
||||
The kernel size of the TDNN blocks.
|
||||
dilation: int
|
||||
The dilation of the Res2Net block.
|
||||
activation : torch class
|
||||
A class for constructing the activation layers.
|
||||
groups: int
|
||||
Number of blocked connections from input channels to output channels.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> x = torch.rand(8, 120, 64).transpose(1, 2)
|
||||
>>> conv = SERes2NetBlock(64, 64, res2net_scale=4)
|
||||
>>> out = conv(x).transpose(1, 2)
|
||||
>>> out.shape
|
||||
torch.Size([8, 120, 64])
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
res2net_scale=8,
|
||||
se_channels=128,
|
||||
kernel_size=1,
|
||||
dilation=1,
|
||||
activation=torch.nn.ReLU,
|
||||
groups=1,
|
||||
):
|
||||
"""Initialize SERes2NetBlock.
|
||||
|
||||
Args:
|
||||
in_channels: TODO.
|
||||
out_channels: TODO.
|
||||
res2net_scale: TODO.
|
||||
se_channels: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
dilation: TODO.
|
||||
activation: TODO.
|
||||
groups: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.out_channels = out_channels
|
||||
self.tdnn1 = TDNNBlock(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
dilation=1,
|
||||
activation=activation,
|
||||
groups=groups,
|
||||
)
|
||||
self.res2net_block = Res2NetBlock(
|
||||
out_channels, out_channels, res2net_scale, kernel_size, dilation
|
||||
)
|
||||
self.tdnn2 = TDNNBlock(
|
||||
out_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
dilation=1,
|
||||
activation=activation,
|
||||
groups=groups,
|
||||
)
|
||||
self.se_block = SEBlock(out_channels, se_channels, out_channels)
|
||||
|
||||
self.shortcut = None
|
||||
if in_channels != out_channels:
|
||||
self.shortcut = Conv1d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
)
|
||||
|
||||
def forward(self, x, lengths=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
lengths: TODO.
|
||||
"""
|
||||
residual = x
|
||||
if self.shortcut:
|
||||
residual = self.shortcut(x)
|
||||
|
||||
x = self.tdnn1(x)
|
||||
x = self.res2net_block(x)
|
||||
x = self.tdnn2(x)
|
||||
x = self.se_block(x, lengths)
|
||||
|
||||
return x + residual
|
||||
|
||||
|
||||
class ECAPA_TDNN(torch.nn.Module):
|
||||
"""An implementation of the speaker embedding model in a paper.
|
||||
"ECAPA-TDNN: Emphasized Channel Attention, Propagation and Aggregation in
|
||||
TDNN Based Speaker Verification" (https://arxiv.org/abs/2005.07143).
|
||||
|
||||
Arguments
|
||||
---------
|
||||
activation : torch class
|
||||
A class for constructing the activation layers.
|
||||
channels : list of ints
|
||||
Output channels for TDNN/SERes2Net layer.
|
||||
kernel_sizes : list of ints
|
||||
List of kernel sizes for each layer.
|
||||
dilations : list of ints
|
||||
List of dilations for kernels in each layer.
|
||||
lin_neurons : int
|
||||
Number of neurons in linear layers.
|
||||
groups : list of ints
|
||||
List of groups for kernels in each layer.
|
||||
|
||||
Example
|
||||
-------
|
||||
>>> input_feats = torch.rand([5, 120, 80])
|
||||
>>> compute_embedding = ECAPA_TDNN(80, lin_neurons=192)
|
||||
>>> outputs = compute_embedding(input_feats)
|
||||
>>> outputs.shape
|
||||
torch.Size([5, 1, 192])
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size,
|
||||
lin_neurons=192,
|
||||
activation=torch.nn.ReLU,
|
||||
channels=[512, 512, 512, 512, 1536],
|
||||
kernel_sizes=[5, 3, 3, 3, 1],
|
||||
dilations=[1, 2, 3, 4, 1],
|
||||
attention_channels=128,
|
||||
res2net_scale=8,
|
||||
se_channels=128,
|
||||
global_context=True,
|
||||
groups=[1, 1, 1, 1, 1],
|
||||
window_size=20,
|
||||
window_shift=1,
|
||||
):
|
||||
|
||||
"""Initialize ECAPA_TDNN.
|
||||
|
||||
Args:
|
||||
input_size: Size/dimension parameter.
|
||||
lin_neurons: TODO.
|
||||
activation: TODO.
|
||||
channels: TODO.
|
||||
kernel_sizes: TODO.
|
||||
dilations: TODO.
|
||||
attention_channels: TODO.
|
||||
res2net_scale: TODO.
|
||||
se_channels: TODO.
|
||||
global_context: TODO.
|
||||
groups: TODO.
|
||||
window_size: Size/dimension parameter.
|
||||
window_shift: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
assert len(channels) == len(kernel_sizes)
|
||||
assert len(channels) == len(dilations)
|
||||
self.channels = channels
|
||||
self.blocks = nn.ModuleList()
|
||||
self.window_size = window_size
|
||||
self.window_shift = window_shift
|
||||
|
||||
# The initial TDNN layer
|
||||
self.blocks.append(
|
||||
TDNNBlock(
|
||||
input_size,
|
||||
channels[0],
|
||||
kernel_sizes[0],
|
||||
dilations[0],
|
||||
activation,
|
||||
groups[0],
|
||||
)
|
||||
)
|
||||
|
||||
# SE-Res2Net layers
|
||||
for i in range(1, len(channels) - 1):
|
||||
self.blocks.append(
|
||||
SERes2NetBlock(
|
||||
channels[i - 1],
|
||||
channels[i],
|
||||
res2net_scale=res2net_scale,
|
||||
se_channels=se_channels,
|
||||
kernel_size=kernel_sizes[i],
|
||||
dilation=dilations[i],
|
||||
activation=activation,
|
||||
groups=groups[i],
|
||||
)
|
||||
)
|
||||
|
||||
# Multi-layer feature aggregation
|
||||
self.mfa = TDNNBlock(
|
||||
channels[-1],
|
||||
channels[-1],
|
||||
kernel_sizes[-1],
|
||||
dilations[-1],
|
||||
activation,
|
||||
groups=groups[-1],
|
||||
)
|
||||
|
||||
# Attentive Statistical Pooling
|
||||
self.asp = AttentiveStatisticsPooling(
|
||||
channels[-1],
|
||||
attention_channels=attention_channels,
|
||||
global_context=global_context,
|
||||
)
|
||||
self.asp_bn = BatchNorm1d(input_size=channels[-1] * 2)
|
||||
|
||||
# Final linear transformation
|
||||
self.fc = Conv1d(
|
||||
in_channels=channels[-1] * 2,
|
||||
out_channels=lin_neurons,
|
||||
kernel_size=1,
|
||||
)
|
||||
|
||||
def windowed_pooling(self, x, lengths=None):
|
||||
# x: Batch, Channel, Time
|
||||
"""Windowed pooling.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
lengths: TODO.
|
||||
"""
|
||||
tt = x.shape[2]
|
||||
num_chunk = int(math.ceil(tt / self.window_shift))
|
||||
pad = self.window_size // 2
|
||||
x = F.pad(x, (pad, pad, 0, 0), "reflect")
|
||||
stat_list = []
|
||||
|
||||
for i in range(num_chunk):
|
||||
# B x C
|
||||
st, ed = i * self.window_shift, i * self.window_shift + self.window_size
|
||||
x = self.asp(
|
||||
x[:, :, st:ed],
|
||||
lengths=(
|
||||
torch.clamp(lengths - i, 0, self.window_size) if lengths is not None else None
|
||||
),
|
||||
)
|
||||
x = self.asp_bn(x)
|
||||
x = self.fc(x)
|
||||
stat_list.append(x)
|
||||
|
||||
return torch.cat(stat_list, dim=2)
|
||||
|
||||
def forward(self, x, lengths=None):
|
||||
"""Returns the embedding vector.
|
||||
|
||||
Arguments
|
||||
---------
|
||||
x : torch.Tensor
|
||||
Tensor of shape (batch, time, channel).
|
||||
lengths: torch.Tensor
|
||||
Tensor of shape (batch, )
|
||||
"""
|
||||
# Minimize transpose for efficiency
|
||||
x = x.transpose(1, 2)
|
||||
|
||||
xl = []
|
||||
for layer in self.blocks:
|
||||
try:
|
||||
x = layer(x, lengths=lengths)
|
||||
except TypeError:
|
||||
x = layer(x)
|
||||
xl.append(x)
|
||||
|
||||
# Multi-layer feature aggregation
|
||||
x = torch.cat(xl[1:], dim=1)
|
||||
x = self.mfa(x)
|
||||
|
||||
if self.window_size is None:
|
||||
# Attentive Statistical Pooling
|
||||
x = self.asp(x, lengths=lengths)
|
||||
x = self.asp_bn(x)
|
||||
# Final linear transformation
|
||||
x = self.fc(x)
|
||||
# x = x.transpose(1, 2)
|
||||
x = x.squeeze(2) # -> B, C
|
||||
else:
|
||||
x = self.windowed_pooling(x, lengths)
|
||||
x = x.transpose(1, 2) # -> B, T, C
|
||||
return x
|
||||
@@ -0,0 +1,218 @@
|
||||
from typing import List
|
||||
from typing import Optional
|
||||
from typing import Sequence
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
import logging
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
import numpy as np
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.transformer.layer_norm import LayerNorm
|
||||
from funasr.models.encoder.abs_encoder import AbsEncoder
|
||||
import math
|
||||
from funasr.models.transformer.utils.repeat import repeat
|
||||
from funasr.models.transformer.utils.multi_layer_conv import FsmnFeedForward
|
||||
|
||||
|
||||
class FsmnBlock(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_feat,
|
||||
dropout_rate,
|
||||
kernel_size,
|
||||
fsmn_shift=0,
|
||||
):
|
||||
"""Initialize FsmnBlock.
|
||||
|
||||
Args:
|
||||
n_feat: TODO.
|
||||
dropout_rate: TODO.
|
||||
kernel_size: Size/dimension parameter.
|
||||
fsmn_shift: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.dropout = nn.Dropout(p=dropout_rate)
|
||||
self.fsmn_block = nn.Conv1d(
|
||||
n_feat, n_feat, kernel_size, stride=1, padding=0, groups=n_feat, bias=False
|
||||
)
|
||||
# padding
|
||||
left_padding = (kernel_size - 1) // 2
|
||||
if fsmn_shift > 0:
|
||||
left_padding = left_padding + fsmn_shift
|
||||
right_padding = kernel_size - 1 - left_padding
|
||||
self.pad_fn = nn.ConstantPad1d((left_padding, right_padding), 0.0)
|
||||
|
||||
def forward(self, inputs, mask, mask_shfit_chunk=None):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
inputs: TODO.
|
||||
mask: TODO.
|
||||
mask_shfit_chunk: TODO.
|
||||
"""
|
||||
b, t, d = inputs.size()
|
||||
if mask is not None:
|
||||
mask = torch.reshape(mask, (b, -1, 1))
|
||||
if mask_shfit_chunk is not None:
|
||||
mask = mask * mask_shfit_chunk
|
||||
|
||||
inputs = inputs * mask
|
||||
x = inputs.transpose(1, 2)
|
||||
x = self.pad_fn(x)
|
||||
x = self.fsmn_block(x)
|
||||
x = x.transpose(1, 2)
|
||||
x = x + inputs
|
||||
x = self.dropout(x)
|
||||
return x * mask
|
||||
|
||||
|
||||
class EncoderLayer(torch.nn.Module):
|
||||
def __init__(self, in_size, size, feed_forward, fsmn_block, dropout_rate=0.0):
|
||||
"""Initialize EncoderLayer.
|
||||
|
||||
Args:
|
||||
in_size: Size/dimension parameter.
|
||||
size: TODO.
|
||||
feed_forward: TODO.
|
||||
fsmn_block: TODO.
|
||||
dropout_rate: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.in_size = in_size
|
||||
self.size = size
|
||||
self.ffn = feed_forward
|
||||
self.memory = fsmn_block
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
|
||||
def forward(
|
||||
self, xs_pad: torch.Tensor, mask: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# xs_pad in Batch, Time, Dim
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
mask: TODO.
|
||||
"""
|
||||
context = self.ffn(xs_pad)[0]
|
||||
memory = self.memory(context, mask)
|
||||
|
||||
memory = self.dropout(memory)
|
||||
if self.in_size == self.size:
|
||||
return memory + xs_pad, mask
|
||||
|
||||
return memory, mask
|
||||
|
||||
|
||||
class FsmnEncoder(AbsEncoder):
|
||||
"""Encoder using Fsmn"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_units,
|
||||
filter_size,
|
||||
fsmn_num_layers,
|
||||
dnn_num_layers,
|
||||
num_memory_units=512,
|
||||
ffn_inner_dim=2048,
|
||||
dropout_rate=0.0,
|
||||
shift=0,
|
||||
position_encoder=None,
|
||||
sample_rate=1,
|
||||
out_units=None,
|
||||
tf2torch_tensor_name_prefix_torch="post_net",
|
||||
tf2torch_tensor_name_prefix_tf="EAND/post_net",
|
||||
):
|
||||
"""Initializes the parameters of the encoder.
|
||||
|
||||
Args:
|
||||
filter_size: the total order of memory block
|
||||
fsmn_num_layers: The number of fsmn layers.
|
||||
dnn_num_layers: The number of dnn layers
|
||||
num_units: The number of memory units.
|
||||
ffn_inner_dim: The number of units of the inner linear transformation
|
||||
in the feed forward layer.
|
||||
dropout_rate: The probability to drop units from the outputs.
|
||||
shift: left padding, to control delay
|
||||
position_encoder: The :class:`opennmt.layers.position.PositionEncoder` to
|
||||
apply on inputs or ``None``.
|
||||
"""
|
||||
super(FsmnEncoder, self).__init__()
|
||||
self.in_units = in_units
|
||||
self.filter_size = filter_size
|
||||
self.fsmn_num_layers = fsmn_num_layers
|
||||
self.dnn_num_layers = dnn_num_layers
|
||||
self.num_memory_units = num_memory_units
|
||||
self.ffn_inner_dim = ffn_inner_dim
|
||||
self.dropout_rate = dropout_rate
|
||||
self.shift = shift
|
||||
if not isinstance(shift, list):
|
||||
self.shift = [shift for _ in range(self.fsmn_num_layers)]
|
||||
self.sample_rate = sample_rate
|
||||
if not isinstance(sample_rate, list):
|
||||
self.sample_rate = [sample_rate for _ in range(self.fsmn_num_layers)]
|
||||
self.position_encoder = position_encoder
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.out_units = out_units
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
|
||||
self.fsmn_layers = repeat(
|
||||
self.fsmn_num_layers,
|
||||
lambda lnum: EncoderLayer(
|
||||
in_units if lnum == 0 else num_memory_units,
|
||||
num_memory_units,
|
||||
FsmnFeedForward(
|
||||
in_units if lnum == 0 else num_memory_units,
|
||||
ffn_inner_dim,
|
||||
num_memory_units,
|
||||
1,
|
||||
dropout_rate,
|
||||
),
|
||||
FsmnBlock(num_memory_units, dropout_rate, filter_size, self.shift[lnum]),
|
||||
),
|
||||
)
|
||||
|
||||
self.dnn_layers = repeat(
|
||||
dnn_num_layers,
|
||||
lambda lnum: FsmnFeedForward(
|
||||
num_memory_units,
|
||||
ffn_inner_dim,
|
||||
num_memory_units,
|
||||
1,
|
||||
dropout_rate,
|
||||
),
|
||||
)
|
||||
if out_units is not None:
|
||||
self.conv1d = nn.Conv1d(num_memory_units, out_units, 1, 1)
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self.num_memory_units
|
||||
|
||||
def forward(
|
||||
self, xs_pad: torch.Tensor, ilens: torch.Tensor, prev_states: torch.Tensor = None
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
inputs = xs_pad
|
||||
if self.position_encoder is not None:
|
||||
inputs = self.position_encoder(inputs)
|
||||
|
||||
inputs = self.dropout(inputs)
|
||||
masks = (~make_pad_mask(ilens)[:, None, :]).to(xs_pad.device)
|
||||
inputs = self.fsmn_layers(inputs, masks)[0]
|
||||
inputs = self.dnn_layers(inputs)[0]
|
||||
|
||||
if self.out_units is not None:
|
||||
inputs = self.conv1d(inputs.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
return inputs, ilens, None
|
||||
@@ -0,0 +1,554 @@
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from funasr.models.encoder.abs_encoder import AbsEncoder
|
||||
from typing import Tuple, Optional
|
||||
from funasr.models.pooling.statistic_pooling import statistic_pooling, windowed_statistic_pooling
|
||||
from collections import OrderedDict
|
||||
import logging
|
||||
import numpy as np
|
||||
|
||||
|
||||
class BasicLayer(torch.nn.Module):
|
||||
|
||||
def __init__(self, in_filters: int, filters: int, stride: int, bn_momentum: float = 0.5):
|
||||
|
||||
"""Initialize BasicLayer.
|
||||
|
||||
Args:
|
||||
in_filters: TODO.
|
||||
filters: TODO.
|
||||
stride: TODO.
|
||||
bn_momentum: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.stride = stride
|
||||
self.in_filters = in_filters
|
||||
self.filters = filters
|
||||
|
||||
self.bn1 = torch.nn.BatchNorm2d(in_filters, eps=1e-3, momentum=bn_momentum, affine=True)
|
||||
self.relu1 = torch.nn.ReLU()
|
||||
self.conv1 = torch.nn.Conv2d(in_filters, filters, 3, stride, bias=False)
|
||||
|
||||
self.bn2 = torch.nn.BatchNorm2d(filters, eps=1e-3, momentum=bn_momentum, affine=True)
|
||||
self.relu2 = torch.nn.ReLU()
|
||||
self.conv2 = torch.nn.Conv2d(filters, filters, 3, 1, bias=False)
|
||||
|
||||
if in_filters != filters or stride > 1:
|
||||
self.conv_sc = torch.nn.Conv2d(in_filters, filters, 1, stride, bias=False)
|
||||
self.bn_sc = torch.nn.BatchNorm2d(filters, eps=1e-3, momentum=bn_momentum, affine=True)
|
||||
|
||||
def proper_padding(self, x, stride):
|
||||
# align padding mode to tf.layers.conv2d with padding_mod="same"
|
||||
"""Proper padding.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
stride: TODO.
|
||||
"""
|
||||
if stride == 1:
|
||||
return F.pad(x, (1, 1, 1, 1), "constant", 0)
|
||||
elif stride == 2:
|
||||
h, w = x.size(2), x.size(3)
|
||||
# (left, right, top, bottom)
|
||||
return F.pad(x, (w % 2, 1, h % 2, 1), "constant", 0)
|
||||
|
||||
def forward(self, xs_pad, ilens):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
"""
|
||||
identity = xs_pad
|
||||
if self.in_filters != self.filters or self.stride > 1:
|
||||
identity = self.conv_sc(identity)
|
||||
identity = self.bn_sc(identity)
|
||||
|
||||
xs_pad = self.relu1(self.bn1(xs_pad))
|
||||
xs_pad = self.proper_padding(xs_pad, self.stride)
|
||||
xs_pad = self.conv1(xs_pad)
|
||||
|
||||
xs_pad = self.relu2(self.bn2(xs_pad))
|
||||
xs_pad = self.proper_padding(xs_pad, 1)
|
||||
xs_pad = self.conv2(xs_pad)
|
||||
|
||||
if self.stride == 2:
|
||||
ilens = (ilens + 1) // self.stride
|
||||
|
||||
return xs_pad + identity, ilens
|
||||
|
||||
|
||||
class BasicBlock(torch.nn.Module):
|
||||
def __init__(self, in_filters, filters, num_layer, stride, bn_momentum=0.5):
|
||||
"""Initialize BasicBlock.
|
||||
|
||||
Args:
|
||||
in_filters: TODO.
|
||||
filters: TODO.
|
||||
num_layer: TODO.
|
||||
stride: TODO.
|
||||
bn_momentum: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
self.num_layer = num_layer
|
||||
|
||||
for i in range(num_layer):
|
||||
layer = BasicLayer(
|
||||
in_filters if i == 0 else filters, filters, stride if i == 0 else 1, bn_momentum
|
||||
)
|
||||
self.add_module("layer_{}".format(i), layer)
|
||||
|
||||
def forward(self, xs_pad, ilens):
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
"""
|
||||
for i in range(self.num_layer):
|
||||
xs_pad, ilens = self._modules["layer_{}".format(i)](xs_pad, ilens)
|
||||
|
||||
return xs_pad, ilens
|
||||
|
||||
|
||||
class ResNet34(AbsEncoder):
|
||||
def __init__(
|
||||
self,
|
||||
input_size,
|
||||
use_head_conv=True,
|
||||
batchnorm_momentum=0.5,
|
||||
use_head_maxpool=False,
|
||||
num_nodes_pooling_layer=256,
|
||||
layers_in_block=(3, 4, 6, 3),
|
||||
filters_in_block=(32, 64, 128, 256),
|
||||
):
|
||||
"""Initialize ResNet34.
|
||||
|
||||
Args:
|
||||
input_size: Size/dimension parameter.
|
||||
use_head_conv: TODO.
|
||||
batchnorm_momentum: TODO.
|
||||
use_head_maxpool: TODO.
|
||||
num_nodes_pooling_layer: TODO.
|
||||
layers_in_block: TODO.
|
||||
filters_in_block: TODO.
|
||||
"""
|
||||
super(ResNet34, self).__init__()
|
||||
|
||||
self.use_head_conv = use_head_conv
|
||||
self.use_head_maxpool = use_head_maxpool
|
||||
self.num_nodes_pooling_layer = num_nodes_pooling_layer
|
||||
self.layers_in_block = layers_in_block
|
||||
self.filters_in_block = filters_in_block
|
||||
self.input_size = input_size
|
||||
|
||||
pre_filters = filters_in_block[0]
|
||||
if use_head_conv:
|
||||
self.pre_conv = torch.nn.Conv2d(
|
||||
1, pre_filters, 3, 1, 1, bias=False, padding_mode="zeros"
|
||||
)
|
||||
self.pre_conv_bn = torch.nn.BatchNorm2d(
|
||||
pre_filters, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
if use_head_maxpool:
|
||||
self.head_maxpool = torch.nn.MaxPool2d(3, 1, padding=1)
|
||||
|
||||
for i in range(len(layers_in_block)):
|
||||
if i == 0:
|
||||
in_filters = pre_filters if self.use_head_conv else 1
|
||||
else:
|
||||
in_filters = filters_in_block[i - 1]
|
||||
|
||||
block = BasicBlock(
|
||||
in_filters,
|
||||
filters=filters_in_block[i],
|
||||
num_layer=layers_in_block[i],
|
||||
stride=1 if i == 0 else 2,
|
||||
bn_momentum=batchnorm_momentum,
|
||||
)
|
||||
self.add_module("block_{}".format(i), block)
|
||||
|
||||
self.resnet0_dense = torch.nn.Conv2d(filters_in_block[-1], num_nodes_pooling_layer, 1)
|
||||
self.resnet0_bn = torch.nn.BatchNorm2d(
|
||||
num_nodes_pooling_layer, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.time_ds_ratio = 8
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self.num_nodes_pooling_layer
|
||||
|
||||
def forward(
|
||||
self, xs_pad: torch.Tensor, ilens: torch.Tensor, prev_states: torch.Tensor = None
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
features = xs_pad
|
||||
assert (
|
||||
features.size(-1) == self.input_size
|
||||
), "Dimension of features {} doesn't match the input_size {}.".format(
|
||||
features.size(-1), self.input_size
|
||||
)
|
||||
features = torch.unsqueeze(features, dim=1)
|
||||
if self.use_head_conv:
|
||||
features = self.pre_conv(features)
|
||||
features = self.pre_conv_bn(features)
|
||||
features = F.relu(features)
|
||||
|
||||
if self.use_head_maxpool:
|
||||
features = self.head_maxpool(features)
|
||||
|
||||
resnet_outs, resnet_out_lens = features, ilens
|
||||
for i in range(len(self.layers_in_block)):
|
||||
block = self._modules["block_{}".format(i)]
|
||||
resnet_outs, resnet_out_lens = block(resnet_outs, resnet_out_lens)
|
||||
|
||||
features = self.resnet0_dense(resnet_outs)
|
||||
features = F.relu(features)
|
||||
features = self.resnet0_bn(features)
|
||||
|
||||
return features, resnet_out_lens
|
||||
|
||||
|
||||
# Note: For training, this implement is not equivalent to tf because of the kernel_regularizer in tf.layers.
|
||||
# TODO: implement kernel_regularizer in torch with munal loss addition or weigth_decay in the optimizer
|
||||
class ResNet34_SP_L2Reg(AbsEncoder):
|
||||
def __init__(
|
||||
self,
|
||||
input_size,
|
||||
use_head_conv=True,
|
||||
batchnorm_momentum=0.5,
|
||||
use_head_maxpool=False,
|
||||
num_nodes_pooling_layer=256,
|
||||
layers_in_block=(3, 4, 6, 3),
|
||||
filters_in_block=(32, 64, 128, 256),
|
||||
tf2torch_tensor_name_prefix_torch="encoder",
|
||||
tf2torch_tensor_name_prefix_tf="EAND/speech_encoder",
|
||||
tf_train_steps=720000,
|
||||
):
|
||||
"""Initialize ResNet34_SP_L2Reg.
|
||||
|
||||
Args:
|
||||
input_size: Size/dimension parameter.
|
||||
use_head_conv: TODO.
|
||||
batchnorm_momentum: TODO.
|
||||
use_head_maxpool: TODO.
|
||||
num_nodes_pooling_layer: TODO.
|
||||
layers_in_block: TODO.
|
||||
filters_in_block: TODO.
|
||||
tf2torch_tensor_name_prefix_torch: TODO.
|
||||
tf2torch_tensor_name_prefix_tf: TODO.
|
||||
tf_train_steps: TODO.
|
||||
"""
|
||||
super(ResNet34_SP_L2Reg, self).__init__()
|
||||
|
||||
self.use_head_conv = use_head_conv
|
||||
self.use_head_maxpool = use_head_maxpool
|
||||
self.num_nodes_pooling_layer = num_nodes_pooling_layer
|
||||
self.layers_in_block = layers_in_block
|
||||
self.filters_in_block = filters_in_block
|
||||
self.input_size = input_size
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
self.tf_train_steps = tf_train_steps
|
||||
|
||||
pre_filters = filters_in_block[0]
|
||||
if use_head_conv:
|
||||
self.pre_conv = torch.nn.Conv2d(
|
||||
1, pre_filters, 3, 1, 1, bias=False, padding_mode="zeros"
|
||||
)
|
||||
self.pre_conv_bn = torch.nn.BatchNorm2d(
|
||||
pre_filters, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
if use_head_maxpool:
|
||||
self.head_maxpool = torch.nn.MaxPool2d(3, 1, padding=1)
|
||||
|
||||
for i in range(len(layers_in_block)):
|
||||
if i == 0:
|
||||
in_filters = pre_filters if self.use_head_conv else 1
|
||||
else:
|
||||
in_filters = filters_in_block[i - 1]
|
||||
|
||||
block = BasicBlock(
|
||||
in_filters,
|
||||
filters=filters_in_block[i],
|
||||
num_layer=layers_in_block[i],
|
||||
stride=1 if i == 0 else 2,
|
||||
bn_momentum=batchnorm_momentum,
|
||||
)
|
||||
self.add_module("block_{}".format(i), block)
|
||||
|
||||
self.resnet0_dense = torch.nn.Conv1d(
|
||||
filters_in_block[-1] * input_size // 8, num_nodes_pooling_layer, 1
|
||||
)
|
||||
self.resnet0_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_pooling_layer, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.time_ds_ratio = 8
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
return self.num_nodes_pooling_layer
|
||||
|
||||
def forward(
|
||||
self, xs_pad: torch.Tensor, ilens: torch.Tensor, prev_states: torch.Tensor = None
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
features = xs_pad
|
||||
assert (
|
||||
features.size(-1) == self.input_size
|
||||
), "Dimension of features {} doesn't match the input_size {}.".format(
|
||||
features.size(-1), self.input_size
|
||||
)
|
||||
features = torch.unsqueeze(features, dim=1)
|
||||
if self.use_head_conv:
|
||||
features = self.pre_conv(features)
|
||||
features = self.pre_conv_bn(features)
|
||||
features = F.relu(features)
|
||||
|
||||
if self.use_head_maxpool:
|
||||
features = self.head_maxpool(features)
|
||||
|
||||
resnet_outs, resnet_out_lens = features, ilens
|
||||
for i in range(len(self.layers_in_block)):
|
||||
block = self._modules["block_{}".format(i)]
|
||||
resnet_outs, resnet_out_lens = block(resnet_outs, resnet_out_lens)
|
||||
|
||||
# B, C, T, F
|
||||
bb, cc, tt, ff = resnet_outs.shape
|
||||
resnet_outs = torch.reshape(resnet_outs.permute(0, 3, 1, 2), [bb, ff * cc, tt])
|
||||
features = self.resnet0_dense(resnet_outs)
|
||||
features = F.relu(features)
|
||||
features = self.resnet0_bn(features)
|
||||
|
||||
return features, resnet_out_lens
|
||||
|
||||
|
||||
class ResNet34Diar(ResNet34):
|
||||
def __init__(
|
||||
self,
|
||||
input_size,
|
||||
embedding_node="resnet1_dense",
|
||||
use_head_conv=True,
|
||||
batchnorm_momentum=0.5,
|
||||
use_head_maxpool=False,
|
||||
num_nodes_pooling_layer=256,
|
||||
layers_in_block=(3, 4, 6, 3),
|
||||
filters_in_block=(32, 64, 128, 256),
|
||||
num_nodes_resnet1=256,
|
||||
num_nodes_last_layer=256,
|
||||
pooling_type="window_shift",
|
||||
pool_size=20,
|
||||
stride=1,
|
||||
tf2torch_tensor_name_prefix_torch="encoder",
|
||||
tf2torch_tensor_name_prefix_tf="seq2seq/speech_encoder",
|
||||
):
|
||||
"""
|
||||
Author: Speech Lab, Alibaba Group, China
|
||||
SOND: Speaker Overlap-aware Neural Diarization for Multi-party Meeting Analysis
|
||||
https://arxiv.org/abs/2211.10243
|
||||
"""
|
||||
|
||||
super(ResNet34Diar, self).__init__(
|
||||
input_size,
|
||||
use_head_conv=use_head_conv,
|
||||
batchnorm_momentum=batchnorm_momentum,
|
||||
use_head_maxpool=use_head_maxpool,
|
||||
num_nodes_pooling_layer=num_nodes_pooling_layer,
|
||||
layers_in_block=layers_in_block,
|
||||
filters_in_block=filters_in_block,
|
||||
)
|
||||
|
||||
self.embedding_node = embedding_node
|
||||
self.num_nodes_resnet1 = num_nodes_resnet1
|
||||
self.num_nodes_last_layer = num_nodes_last_layer
|
||||
self.pooling_type = pooling_type
|
||||
self.pool_size = pool_size
|
||||
self.stride = stride
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
|
||||
self.resnet1_dense = torch.nn.Linear(num_nodes_pooling_layer * 2, num_nodes_resnet1)
|
||||
self.resnet1_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_resnet1, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.resnet2_dense = torch.nn.Linear(num_nodes_resnet1, num_nodes_last_layer)
|
||||
self.resnet2_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_last_layer, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
if self.embedding_node.startswith("resnet1"):
|
||||
return self.num_nodes_resnet1
|
||||
elif self.embedding_node.startswith("resnet2"):
|
||||
return self.num_nodes_last_layer
|
||||
|
||||
return self.num_nodes_pooling_layer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
ilens: torch.Tensor,
|
||||
prev_states: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
endpoints = OrderedDict()
|
||||
res_out, ilens = super().forward(xs_pad, ilens)
|
||||
endpoints["resnet0_bn"] = res_out
|
||||
if self.pooling_type == "frame_gsp":
|
||||
features = statistic_pooling(res_out, ilens, (3,))
|
||||
else:
|
||||
features, ilens = windowed_statistic_pooling(
|
||||
res_out, ilens, (2, 3), self.pool_size, self.stride
|
||||
)
|
||||
features = features.transpose(1, 2)
|
||||
endpoints["pooling"] = features
|
||||
|
||||
features = self.resnet1_dense(features)
|
||||
endpoints["resnet1_dense"] = features
|
||||
features = F.relu(features)
|
||||
endpoints["resnet1_relu"] = features
|
||||
features = self.resnet1_bn(features.transpose(1, 2)).transpose(1, 2)
|
||||
endpoints["resnet1_bn"] = features
|
||||
|
||||
features = self.resnet2_dense(features)
|
||||
endpoints["resnet2_dense"] = features
|
||||
features = F.relu(features)
|
||||
endpoints["resnet2_relu"] = features
|
||||
features = self.resnet2_bn(features.transpose(1, 2)).transpose(1, 2)
|
||||
endpoints["resnet2_bn"] = features
|
||||
|
||||
return endpoints[self.embedding_node], ilens, None
|
||||
|
||||
|
||||
class ResNet34SpL2RegDiar(ResNet34_SP_L2Reg):
|
||||
def __init__(
|
||||
self,
|
||||
input_size,
|
||||
embedding_node="resnet1_dense",
|
||||
use_head_conv=True,
|
||||
batchnorm_momentum=0.5,
|
||||
use_head_maxpool=False,
|
||||
num_nodes_pooling_layer=256,
|
||||
layers_in_block=(3, 4, 6, 3),
|
||||
filters_in_block=(32, 64, 128, 256),
|
||||
num_nodes_resnet1=256,
|
||||
num_nodes_last_layer=256,
|
||||
pooling_type="window_shift",
|
||||
pool_size=20,
|
||||
stride=1,
|
||||
tf2torch_tensor_name_prefix_torch="encoder",
|
||||
tf2torch_tensor_name_prefix_tf="seq2seq/speech_encoder",
|
||||
):
|
||||
"""
|
||||
Author: Speech Lab, Alibaba Group, China
|
||||
TOLD: A Novel Two-Stage Overlap-Aware Framework for Speaker Diarization
|
||||
https://arxiv.org/abs/2303.05397
|
||||
"""
|
||||
|
||||
super(ResNet34SpL2RegDiar, self).__init__(
|
||||
input_size,
|
||||
use_head_conv=use_head_conv,
|
||||
batchnorm_momentum=batchnorm_momentum,
|
||||
use_head_maxpool=use_head_maxpool,
|
||||
num_nodes_pooling_layer=num_nodes_pooling_layer,
|
||||
layers_in_block=layers_in_block,
|
||||
filters_in_block=filters_in_block,
|
||||
)
|
||||
|
||||
self.embedding_node = embedding_node
|
||||
self.num_nodes_resnet1 = num_nodes_resnet1
|
||||
self.num_nodes_last_layer = num_nodes_last_layer
|
||||
self.pooling_type = pooling_type
|
||||
self.pool_size = pool_size
|
||||
self.stride = stride
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
|
||||
self.resnet1_dense = torch.nn.Linear(num_nodes_pooling_layer * 2, num_nodes_resnet1)
|
||||
self.resnet1_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_resnet1, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.resnet2_dense = torch.nn.Linear(num_nodes_resnet1, num_nodes_last_layer)
|
||||
self.resnet2_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_last_layer, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
def output_size(self) -> int:
|
||||
"""Output size."""
|
||||
if self.embedding_node.startswith("resnet1"):
|
||||
return self.num_nodes_resnet1
|
||||
elif self.embedding_node.startswith("resnet2"):
|
||||
return self.num_nodes_last_layer
|
||||
|
||||
return self.num_nodes_pooling_layer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
xs_pad: torch.Tensor,
|
||||
ilens: torch.Tensor,
|
||||
prev_states: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
prev_states: TODO.
|
||||
"""
|
||||
endpoints = OrderedDict()
|
||||
res_out, ilens = super().forward(xs_pad, ilens)
|
||||
endpoints["resnet0_bn"] = res_out
|
||||
if self.pooling_type == "frame_gsp":
|
||||
features = statistic_pooling(res_out, ilens, (2,))
|
||||
else:
|
||||
features, ilens = windowed_statistic_pooling(
|
||||
res_out, ilens, (2,), self.pool_size, self.stride
|
||||
)
|
||||
features = features.transpose(1, 2)
|
||||
endpoints["pooling"] = features
|
||||
|
||||
features = self.resnet1_dense(features)
|
||||
endpoints["resnet1_dense"] = features
|
||||
features = F.relu(features)
|
||||
endpoints["resnet1_relu"] = features
|
||||
features = self.resnet1_bn(features.transpose(1, 2)).transpose(1, 2)
|
||||
endpoints["resnet1_bn"] = features
|
||||
|
||||
features = self.resnet2_dense(features)
|
||||
endpoints["resnet2_dense"] = features
|
||||
features = F.relu(features)
|
||||
endpoints["resnet2_relu"] = features
|
||||
features = self.resnet2_bn(features.transpose(1, 2)).transpose(1, 2)
|
||||
endpoints["resnet2_bn"] = features
|
||||
|
||||
return endpoints[self.embedding_node], ilens, None
|
||||
@@ -0,0 +1,358 @@
|
||||
from typing import List
|
||||
from typing import Optional
|
||||
from typing import Sequence
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
import logging
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from funasr.models.scama.chunk_utilis import overlap_chunk
|
||||
import numpy as np
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
from funasr.models.sond.attention import MultiHeadSelfAttention
|
||||
from funasr.models.transformer.embedding import SinusoidalPositionEncoder
|
||||
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.positionwise_feed_forward import (
|
||||
PositionwiseFeedForward, # noqa: H301
|
||||
)
|
||||
from funasr.models.transformer.utils.repeat import repeat
|
||||
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.models.ctc import CTC
|
||||
from funasr.models.encoder.abs_encoder import AbsEncoder
|
||||
|
||||
|
||||
class EncoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_size,
|
||||
size,
|
||||
self_attn,
|
||||
feed_forward,
|
||||
dropout_rate,
|
||||
normalize_before=True,
|
||||
concat_after=False,
|
||||
stochastic_depth_rate=0.0,
|
||||
):
|
||||
"""Construct an EncoderLayer object."""
|
||||
super(EncoderLayer, self).__init__()
|
||||
self.self_attn = self_attn
|
||||
self.feed_forward = feed_forward
|
||||
self.norm1 = LayerNorm(in_size)
|
||||
self.norm2 = LayerNorm(size)
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.in_size = in_size
|
||||
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
|
||||
self.dropout_rate = dropout_rate
|
||||
|
||||
def forward(self, x, mask, cache=None, mask_att_chunk_encoder=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 self.concat_after:
|
||||
x_concat = torch.cat(
|
||||
(x, self.self_attn(x, mask, mask_att_chunk_encoder=mask_att_chunk_encoder)), dim=-1
|
||||
)
|
||||
if self.in_size == self.size:
|
||||
x = residual + stoch_layer_coeff * self.concat_linear(x_concat)
|
||||
else:
|
||||
x = stoch_layer_coeff * self.concat_linear(x_concat)
|
||||
else:
|
||||
if self.in_size == self.size:
|
||||
x = residual + stoch_layer_coeff * self.dropout(
|
||||
self.self_attn(x, mask, mask_att_chunk_encoder=mask_att_chunk_encoder)
|
||||
)
|
||||
else:
|
||||
x = stoch_layer_coeff * self.dropout(
|
||||
self.self_attn(x, mask, mask_att_chunk_encoder=mask_att_chunk_encoder)
|
||||
)
|
||||
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)
|
||||
|
||||
return x, mask, cache, mask_att_chunk_encoder
|
||||
|
||||
|
||||
class SelfAttentionEncoder(AbsEncoder):
|
||||
"""
|
||||
Author: Speech Lab of DAMO Academy, Alibaba Group
|
||||
Self attention encoder in OpenNMT framework
|
||||
"""
|
||||
|
||||
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=SinusoidalPositionEncoder,
|
||||
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,
|
||||
tf2torch_tensor_name_prefix_torch: str = "encoder",
|
||||
tf2torch_tensor_name_prefix_tf: str = "seq2seq/encoder",
|
||||
out_units=None,
|
||||
):
|
||||
"""Initialize SelfAttentionEncoder.
|
||||
|
||||
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.
|
||||
tf2torch_tensor_name_prefix_torch: TODO.
|
||||
tf2torch_tensor_name_prefix_tf: TODO.
|
||||
out_units: 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),
|
||||
SinusoidalPositionEncoder(),
|
||||
)
|
||||
elif input_layer is None:
|
||||
if input_size == output_size:
|
||||
self.embed = None
|
||||
else:
|
||||
self.embed = torch.nn.Linear(input_size, output_size)
|
||||
elif input_layer == "pe":
|
||||
self.embed = SinusoidalPositionEncoder()
|
||||
elif input_layer == "null":
|
||||
self.embed = None
|
||||
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,
|
||||
output_size,
|
||||
MultiHeadSelfAttention(
|
||||
attention_heads,
|
||||
output_size,
|
||||
output_size,
|
||||
attention_dropout_rate,
|
||||
),
|
||||
positionwise_layer(*positionwise_layer_args),
|
||||
dropout_rate,
|
||||
normalize_before,
|
||||
concat_after,
|
||||
)
|
||||
if lnum > 0
|
||||
else EncoderLayer(
|
||||
input_size,
|
||||
output_size,
|
||||
MultiHeadSelfAttention(
|
||||
attention_heads,
|
||||
input_size if input_layer == "pe" or input_layer == "null" else output_size,
|
||||
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
|
||||
self.dropout = nn.Dropout(dropout_rate)
|
||||
self.tf2torch_tensor_name_prefix_torch = tf2torch_tensor_name_prefix_torch
|
||||
self.tf2torch_tensor_name_prefix_tf = tf2torch_tensor_name_prefix_tf
|
||||
self.out_units = out_units
|
||||
if out_units is not None:
|
||||
self.output_linear = nn.Linear(output_size, out_units)
|
||||
|
||||
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)
|
||||
xs_pad = xs_pad * self.output_size() ** 0.5
|
||||
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)
|
||||
|
||||
xs_pad = self.dropout(xs_pad)
|
||||
# encoder_outs = self.encoders0(xs_pad, masks)
|
||||
# xs_pad, masks = encoder_outs[0], encoder_outs[1]
|
||||
intermediate_outs = []
|
||||
if len(self.interctc_layer_idx) == 0:
|
||||
encoder_outs = self.encoders(xs_pad, masks)
|
||||
xs_pad, masks = encoder_outs[0], encoder_outs[1]
|
||||
else:
|
||||
for layer_idx, encoder_layer in enumerate(self.encoders):
|
||||
encoder_outs = encoder_layer(xs_pad, masks)
|
||||
xs_pad, masks = encoder_outs[0], encoder_outs[1]
|
||||
|
||||
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)
|
||||
|
||||
if self.out_units is not None:
|
||||
xs_pad = self.output_linear(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,127 @@
|
||||
import torch
|
||||
from typing import Optional
|
||||
from typing import Tuple
|
||||
from torch.nn import functional as F
|
||||
from funasr.models.transformer.utils.nets_utils import make_pad_mask
|
||||
|
||||
|
||||
class LabelAggregate(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
win_length: int = 512,
|
||||
hop_length: int = 128,
|
||||
center: bool = True,
|
||||
):
|
||||
"""Initialize LabelAggregate.
|
||||
|
||||
Args:
|
||||
win_length: TODO.
|
||||
hop_length: TODO.
|
||||
center: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.win_length = win_length
|
||||
self.hop_length = hop_length
|
||||
self.center = center
|
||||
|
||||
def extra_repr(self):
|
||||
"""Extra repr."""
|
||||
return (
|
||||
f"win_length={self.win_length}, "
|
||||
f"hop_length={self.hop_length}, "
|
||||
f"center={self.center}, "
|
||||
)
|
||||
|
||||
def forward(
|
||||
self, input: torch.Tensor, ilens: torch.Tensor = None
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""LabelAggregate forward function.
|
||||
|
||||
Args:
|
||||
input: (Batch, Nsamples, Label_dim)
|
||||
ilens: (Batch)
|
||||
Returns:
|
||||
output: (Batch, Frames, Label_dim)
|
||||
|
||||
"""
|
||||
bs = input.size(0)
|
||||
max_length = input.size(1)
|
||||
label_dim = input.size(2)
|
||||
|
||||
# NOTE(jiatong):
|
||||
# The default behaviour of label aggregation is compatible with
|
||||
# torch.stft about framing and padding.
|
||||
|
||||
# Step1: center padding
|
||||
if self.center:
|
||||
pad = self.win_length // 2
|
||||
max_length = max_length + 2 * pad
|
||||
input = torch.nn.functional.pad(input, (0, 0, pad, pad), "constant", 0)
|
||||
input[:, :pad, :] = input[:, pad : (2 * pad), :]
|
||||
input[:, (max_length - pad) : max_length, :] = input[
|
||||
:, (max_length - 2 * pad) : (max_length - pad), :
|
||||
]
|
||||
nframe = (max_length - self.win_length) // self.hop_length + 1
|
||||
|
||||
# Step2: framing
|
||||
output = input.as_strided(
|
||||
(bs, nframe, self.win_length, label_dim),
|
||||
(max_length * label_dim, self.hop_length * label_dim, label_dim, 1),
|
||||
)
|
||||
|
||||
# Step3: aggregate label
|
||||
output = torch.gt(output.sum(dim=2, keepdim=False), self.win_length // 2)
|
||||
output = output.float()
|
||||
|
||||
# Step4: process lengths
|
||||
if ilens is not None:
|
||||
if self.center:
|
||||
pad = self.win_length // 2
|
||||
ilens = ilens + 2 * pad
|
||||
|
||||
olens = (ilens - self.win_length) // self.hop_length + 1
|
||||
output.masked_fill_(make_pad_mask(olens, output, 1), 0.0)
|
||||
else:
|
||||
olens = None
|
||||
|
||||
return output.to(input.dtype), olens
|
||||
|
||||
|
||||
class LabelAggregateMaxPooling(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hop_length: int = 8,
|
||||
):
|
||||
"""Initialize LabelAggregateMaxPooling.
|
||||
|
||||
Args:
|
||||
hop_length: TODO.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
self.hop_length = hop_length
|
||||
|
||||
def extra_repr(self):
|
||||
"""Extra repr."""
|
||||
return f"hop_length={self.hop_length}, "
|
||||
|
||||
def forward(
|
||||
self, input: torch.Tensor, ilens: torch.Tensor = None
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""LabelAggregate forward function.
|
||||
|
||||
Args:
|
||||
input: (Batch, Nsamples, Label_dim)
|
||||
ilens: (Batch)
|
||||
Returns:
|
||||
output: (Batch, Frames, Label_dim)
|
||||
|
||||
"""
|
||||
|
||||
output = F.max_pool1d(input.transpose(1, 2), self.hop_length, self.hop_length).transpose(
|
||||
1, 2
|
||||
)
|
||||
olens = ilens // self.hop_length
|
||||
|
||||
return output.to(input.dtype), olens
|
||||
@@ -0,0 +1,144 @@
|
||||
# Copyright 3D-Speaker (https://github.com/alibaba-damo-academy/3D-Speaker). All Rights Reserved.
|
||||
# Licensed under the Apache License, Version 2.0 (http://www.apache.org/licenses/LICENSE-2.0)
|
||||
|
||||
""" This implementation is adapted from https://github.com/wenet-e2e/wespeaker."""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class TAP(nn.Module):
|
||||
"""
|
||||
Temporal average pooling, only first-order mean is considered
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
"""Initialize TAP.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super(TAP, self).__init__()
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
pooling_mean = x.mean(dim=-1)
|
||||
# To be compatable with 2D input
|
||||
pooling_mean = pooling_mean.flatten(start_dim=1)
|
||||
return pooling_mean
|
||||
|
||||
|
||||
class TSDP(nn.Module):
|
||||
"""
|
||||
Temporal standard deviation pooling, only second-order std is considered
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
"""Initialize TSDP.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super(TSDP, self).__init__()
|
||||
|
||||
def forward(self, x):
|
||||
# The last dimension is the temporal axis
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
pooling_std = torch.sqrt(torch.var(x, dim=-1) + 1e-8)
|
||||
pooling_std = pooling_std.flatten(start_dim=1)
|
||||
return pooling_std
|
||||
|
||||
|
||||
class TSTP(nn.Module):
|
||||
"""
|
||||
Temporal statistics pooling, concatenate mean and std, which is used in
|
||||
x-vector
|
||||
Comment: simple concatenation can not make full use of both statistics
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
"""Initialize TSTP.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
super(TSTP, self).__init__()
|
||||
|
||||
def forward(self, x):
|
||||
# The last dimension is the temporal axis
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
x: TODO.
|
||||
"""
|
||||
pooling_mean = x.mean(dim=-1)
|
||||
pooling_std = torch.sqrt(torch.var(x, dim=-1) + 1e-8)
|
||||
pooling_mean = pooling_mean.flatten(start_dim=1)
|
||||
pooling_std = pooling_std.flatten(start_dim=1)
|
||||
|
||||
stats = torch.cat((pooling_mean, pooling_std), 1)
|
||||
return stats
|
||||
|
||||
|
||||
class ASTP(nn.Module):
|
||||
"""Attentive statistics pooling: Channel- and context-dependent
|
||||
statistics pooling, first used in ECAPA_TDNN.
|
||||
"""
|
||||
|
||||
def __init__(self, in_dim, bottleneck_dim=128, global_context_att=False):
|
||||
"""Initialize ASTP.
|
||||
|
||||
Args:
|
||||
in_dim: Size/dimension parameter.
|
||||
bottleneck_dim: Size/dimension parameter.
|
||||
global_context_att: TODO.
|
||||
"""
|
||||
super(ASTP, self).__init__()
|
||||
self.global_context_att = global_context_att
|
||||
|
||||
# Use Conv1d with stride == 1 rather than Linear, then we don't
|
||||
# need to transpose inputs.
|
||||
if global_context_att:
|
||||
self.linear1 = nn.Conv1d(
|
||||
in_dim * 3, bottleneck_dim, kernel_size=1
|
||||
) # equals W and b in the paper
|
||||
else:
|
||||
self.linear1 = nn.Conv1d(
|
||||
in_dim, bottleneck_dim, kernel_size=1
|
||||
) # equals W and b in the paper
|
||||
self.linear2 = nn.Conv1d(
|
||||
bottleneck_dim, in_dim, kernel_size=1
|
||||
) # equals V and k in the paper
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
x: a 3-dimensional tensor in tdnn-based architecture (B,F,T)
|
||||
or a 4-dimensional tensor in resnet architecture (B,C,F,T)
|
||||
0-dim: batch-dimension, last-dim: time-dimension (frame-dimension)
|
||||
"""
|
||||
if len(x.shape) == 4:
|
||||
x = x.reshape(x.shape[0], x.shape[1] * x.shape[2], x.shape[3])
|
||||
assert len(x.shape) == 3
|
||||
|
||||
if self.global_context_att:
|
||||
context_mean = torch.mean(x, dim=-1, keepdim=True).expand_as(x)
|
||||
context_std = torch.sqrt(torch.var(x, dim=-1, keepdim=True) + 1e-10).expand_as(x)
|
||||
x_in = torch.cat((x, context_mean, context_std), dim=1)
|
||||
else:
|
||||
x_in = x
|
||||
|
||||
# DON'T use ReLU here! ReLU may be hard to converge.
|
||||
alpha = torch.tanh(self.linear1(x_in)) # alpha = F.relu(self.linear1(x_in))
|
||||
alpha = torch.softmax(self.linear2(alpha), dim=2)
|
||||
mean = torch.sum(alpha * x, dim=2)
|
||||
var = torch.sum(alpha * (x**2), dim=2) - mean**2
|
||||
std = torch.sqrt(var.clamp(min=1e-10))
|
||||
return torch.cat([mean, std], dim=1)
|
||||
@@ -0,0 +1,126 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from typing import Union
|
||||
from funasr.models.transformer.utils.nets_utils import make_non_pad_mask
|
||||
from torch.nn import functional as F
|
||||
import math
|
||||
|
||||
VAR2STD_EPSILON = 1e-12
|
||||
|
||||
|
||||
class StatisticPooling(torch.nn.Module):
|
||||
def __init__(self, pooling_dim: Union[int, Tuple] = 2, eps=1e-12):
|
||||
"""Initialize StatisticPooling.
|
||||
|
||||
Args:
|
||||
pooling_dim: Size/dimension parameter.
|
||||
eps: TODO.
|
||||
"""
|
||||
super(StatisticPooling, self).__init__()
|
||||
if isinstance(pooling_dim, int):
|
||||
pooling_dim = (pooling_dim,)
|
||||
self.pooling_dim = pooling_dim
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, xs_pad, ilens=None):
|
||||
# xs_pad in (Batch, Channel, Time, Frequency)
|
||||
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
"""
|
||||
if ilens is None:
|
||||
masks = torch.ones_like(xs_pad).to(xs_pad)
|
||||
else:
|
||||
masks = make_non_pad_mask(ilens, xs_pad, length_dim=2).to(xs_pad)
|
||||
mean = torch.sum(xs_pad, dim=self.pooling_dim, keepdim=True) / torch.sum(
|
||||
masks, dim=self.pooling_dim, keepdim=True
|
||||
)
|
||||
squared_difference = torch.pow(xs_pad - mean, 2.0)
|
||||
variance = torch.sum(squared_difference, dim=self.pooling_dim, keepdim=True) / torch.sum(
|
||||
masks, dim=self.pooling_dim, keepdim=True
|
||||
)
|
||||
for i in reversed(self.pooling_dim):
|
||||
mean, variance = torch.squeeze(mean, dim=i), torch.squeeze(variance, dim=i)
|
||||
|
||||
mask = torch.less_equal(variance, self.eps).float()
|
||||
variance = (1.0 - mask) * variance + mask * self.eps
|
||||
stddev = torch.sqrt(variance)
|
||||
|
||||
stat_pooling = torch.cat([mean, stddev], dim=1)
|
||||
|
||||
return stat_pooling
|
||||
|
||||
|
||||
def statistic_pooling(
|
||||
xs_pad: torch.Tensor, ilens: torch.Tensor = None, pooling_dim: Tuple = (2, 3)
|
||||
) -> torch.Tensor:
|
||||
# xs_pad in (Batch, Channel, Time, Frequency)
|
||||
|
||||
"""Statistic pooling.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
pooling_dim: Size/dimension parameter.
|
||||
"""
|
||||
if ilens is None:
|
||||
seq_mask = torch.ones_like(xs_pad).to(xs_pad)
|
||||
else:
|
||||
seq_mask = make_non_pad_mask(ilens, xs_pad, length_dim=2).to(xs_pad)
|
||||
mean = torch.sum(xs_pad, dim=pooling_dim, keepdim=True) / torch.sum(
|
||||
seq_mask, dim=pooling_dim, keepdim=True
|
||||
)
|
||||
squared_difference = torch.pow(xs_pad - mean, 2.0)
|
||||
variance = torch.sum(squared_difference, dim=pooling_dim, keepdim=True) / torch.sum(
|
||||
seq_mask, dim=pooling_dim, keepdim=True
|
||||
)
|
||||
for i in reversed(pooling_dim):
|
||||
mean, variance = torch.squeeze(mean, dim=i), torch.squeeze(variance, dim=i)
|
||||
|
||||
value_mask = torch.less_equal(variance, VAR2STD_EPSILON).float()
|
||||
variance = (1.0 - value_mask) * variance + value_mask * VAR2STD_EPSILON
|
||||
stddev = torch.sqrt(variance)
|
||||
|
||||
stat_pooling = torch.cat([mean, stddev], dim=1)
|
||||
|
||||
return stat_pooling
|
||||
|
||||
|
||||
def windowed_statistic_pooling(
|
||||
xs_pad: torch.Tensor,
|
||||
ilens: torch.Tensor = None,
|
||||
pooling_dim: Tuple = (2, 3),
|
||||
pooling_size: int = 20,
|
||||
pooling_stride: int = 1,
|
||||
) -> Tuple[torch.Tensor, int]:
|
||||
# xs_pad in (Batch, Channel, Time, Frequency)
|
||||
|
||||
"""Windowed statistic pooling.
|
||||
|
||||
Args:
|
||||
xs_pad: TODO.
|
||||
ilens: TODO.
|
||||
pooling_dim: Size/dimension parameter.
|
||||
pooling_size: Size/dimension parameter.
|
||||
pooling_stride: TODO.
|
||||
"""
|
||||
tt = xs_pad.shape[2]
|
||||
num_chunk = int(math.ceil(tt / pooling_stride))
|
||||
pad = pooling_size // 2
|
||||
if len(xs_pad.shape) == 4:
|
||||
features = F.pad(xs_pad, (0, 0, pad, pad), "replicate")
|
||||
else:
|
||||
features = F.pad(xs_pad, (pad, pad), "replicate")
|
||||
stat_list = []
|
||||
|
||||
for i in range(num_chunk):
|
||||
# B x C
|
||||
st, ed = i * pooling_stride, i * pooling_stride + pooling_size
|
||||
stat = statistic_pooling(features[:, :, st:ed], pooling_dim=pooling_dim)
|
||||
stat_list.append(stat.unsqueeze(2))
|
||||
|
||||
# B x C x T
|
||||
return torch.cat(stat_list, dim=2), ilens / pooling_stride
|
||||
@@ -0,0 +1,55 @@
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
from funasr.models.decoder.abs_decoder import AbsDecoder
|
||||
|
||||
|
||||
class DenseDecoder(AbsDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size,
|
||||
encoder_output_size,
|
||||
num_nodes_resnet1: int = 256,
|
||||
num_nodes_last_layer: int = 256,
|
||||
batchnorm_momentum: float = 0.5,
|
||||
):
|
||||
"""Initialize DenseDecoder.
|
||||
|
||||
Args:
|
||||
vocab_size: Size/dimension parameter.
|
||||
encoder_output_size: Size/dimension parameter.
|
||||
num_nodes_resnet1: TODO.
|
||||
num_nodes_last_layer: TODO.
|
||||
batchnorm_momentum: TODO.
|
||||
"""
|
||||
super(DenseDecoder, self).__init__()
|
||||
self.resnet1_dense = torch.nn.Linear(encoder_output_size, num_nodes_resnet1)
|
||||
self.resnet1_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_resnet1, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.resnet2_dense = torch.nn.Linear(num_nodes_resnet1, num_nodes_last_layer)
|
||||
self.resnet2_bn = torch.nn.BatchNorm1d(
|
||||
num_nodes_last_layer, eps=1e-3, momentum=batchnorm_momentum
|
||||
)
|
||||
|
||||
self.output_dense = torch.nn.Linear(num_nodes_last_layer, vocab_size, bias=False)
|
||||
|
||||
def forward(self, features):
|
||||
"""Forward pass for training.
|
||||
|
||||
Args:
|
||||
features: TODO.
|
||||
"""
|
||||
embeddings = {}
|
||||
features = self.resnet1_dense(features)
|
||||
embeddings["resnet1_dense"] = features
|
||||
features = F.relu(features)
|
||||
features = self.resnet1_bn(features)
|
||||
|
||||
features = self.resnet2_dense(features)
|
||||
embeddings["resnet2_dense"] = features
|
||||
features = F.relu(features)
|
||||
features = self.resnet2_bn(features)
|
||||
|
||||
features = self.output_dense(features)
|
||||
return features, embeddings
|
||||
Reference in New Issue
Block a user