Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled

Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
freedakgmail
2026-07-09 22:38:58 +08:00
commit 6116b1f3c6
3683 changed files with 990984 additions and 0 deletions
View File
+322
View File
@@ -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
+702
View File
@@ -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,
)
+46
View File
@@ -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
+224
View File
@@ -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
+218
View File
@@ -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
+127
View File
@@ -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
+55
View File
@@ -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