使用 Conformer 模型实现自动语音识别,ASR
·
import torch
import torch.nn.functional as F
from torch import Tensor
from torch import nn
from torch.nn import CTCLoss
from torchaudio.models import Conformer
class SpeechConformer(nn.Module):
def __init__(self,
num_classes: int,
feature_dim: int = 80,
conformer_params: dict = None) -> None:
"""
特征提取层 -> Conformer -> 输出层
Args:
num_classes: 标签类数
feature_dim: 特征维度,默认为 80
conformer_params: torchaudio.models.Conformer 参数,默认为 None
"""
super(SpeechConformer, self).__init__()
# 默认的 Conformer 参数
default_conformer_params = {
"input_dim": 512, # 输入维度
"num_heads": 8, # 注意力头数
"ffn_dim": 2048, # 前馈层的隐藏层维度
"num_layers": 6, # 层数
"depthwise_conv_kernel_size": 31, # 卷积核大小
"dropout": 0.1 # 丢弃率
}
if conformer_params is not None:
default_conformer_params.update(conformer_params)
# 特征提取层
input_dim = default_conformer_params["input_dim"]
self.feature_layer = nn.Sequential(
nn.Linear(feature_dim, input_dim),
nn.ReLU(),
nn.LayerNorm(input_dim)
)
# Conformer
self.conformer = Conformer(**default_conformer_params)
# 输出层
self.output_layer = nn.Linear(input_dim, num_classes)
# 损失函数
self.ctc_loss = CTCLoss(zero_infinity=True)
def forward(self,
x: Tensor,
x_len: Tensor,
tgt: Tensor = None,
tgt_len: Tensor = None) -> Tensor | tuple[Tensor, Tensor]:
"""
当 tgt 和 tgt_len 不为 None 时,返回输出和损失
Args:
x: 输入 shape=(N, T, C)
x_len: 输入序列的真实长度 shape=(N,)
tgt: 标签 shape=(N, T)
tgt_len: 标签序列的真实长度 shape=(N,)
Returns:
y: 输出 shape=(N, T, num_classes)
loss: 损失
"""
# (N, T, feature_dim) -> (N, T, input_dim)
y = self.feature_layer(x)
# (N, T, input_dim)
y, y_len = self.conformer(y, x_len)
# (N, T, input_dim) -> (N, T, num_classes)
y = self.output_layer(y)
# 计算损失
if tgt is not None and tgt_len is not None:
# CTCLoss 接收经过 log_softmax 处理后的输出
log_probs = F.log_softmax(y.transpose(0, 1), -1) # (T, N, C)
loss = self.ctc_loss(log_probs, tgt, y_len, tgt_len)
return y, loss
return y
def predict(self, x: Tensor, x_len: Tensor) -> Tensor:
"""
预测
Args:
x: 输入 shape=(N, T, C)
x_len: 输入序列的真实长度 shape=(N,)
Returns:
y: 预测结果 shape=(N, T)
"""
with torch.no_grad(): # 推理时使用
y = self.forward(x, x_len) # (N, T, num_classes)
y = torch.argmax(y, dim=-1) # (N, T)
return y
if __name__ == '__main__':
model = SpeechConformer(num_classes=21128)
out, loss = model(
torch.randn((5, 20, 80)),
torch.tensor([18, 20, 15, 19, 10]),
torch.randn((5, 15)),
torch.tensor([12, 15, 14, 15, 7])
)
print(out.shape, loss.item())
更多推荐


所有评论(0)