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())

更多推荐