Audio Spectrogram Transformer 入门实战:从零构建音频分类模型

1次阅读
没有评论

共计 2093 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景痛点

传统 CNN 在图像处理上表现出色,但直接套用到音频频谱图时会出现几个明显问题:

Audio Spectrogram Transformer 入门实战:从零构建音频分类模型

  1. 平移不变性假设失效:图像中猫耳朵在左上角还是右下角都是猫,但音频频谱图中频率分量偏移会完全改变语义(如不同音高的音符)
  2. 长距离依赖捕获困难:CNN 的局部感受野难以建模跨时间的音频事件关联(如鸟叫的开始和结束部分)
  3. 频域 - 时域差异处理:标准卷积核同时处理时间和频率轴,忽略了音频信号在这两个维度上的不同特性

技术对比

模型类型 参数量(M) ESC-50 准确率 FLOPS(G) 特点
CRNN 3.2 81.3% 2.1 CNN+RNN 经典组合
CNN-Transformer 24.7 85.6% 8.3 混合架构显存消耗大
AST 86.7 88.9% 16.2 纯 Transformer 端到端建模

核心实现

Mel 频谱图提取

import librosa

def extract_melspectrogram(audio_path, sr=16000):
    # 加载音频并统一长度
    waveform, _ = librosa.load(audio_path, sr=sr, duration=5)  

    # 提取 Mel 特征(建议 128 维)mel = librosa.feature.melspectrogram(
        y=waveform,
        sr=sr,
        n_fft=1024,
        hop_length=512,
        n_mels=128,
        fmax=8000
    )
    # 转换为对数刻度
    log_mel = librosa.power_to_db(mel, ref=np.max)
    return log_mel  # 输出形状:(128, 157)

AST 的 Patch Embedding

import torch
import torch.nn as nn

class ASTPatchEmbed(nn.Module):
    def __init__(self, img_size=(128, 157), patch_size=16, in_chans=1, embed_dim=768):
        super().__init__()
        # 计算分块数量
        num_patches = (img_size[0] // patch_size) * (img_size[1] // patch_size)
        self.proj = nn.Conv2d(
            in_chans, embed_dim,
            kernel_size=patch_size,
            stride=patch_size
        )

    def forward(self, x):
        # x: [B, 1, 128, 157]
        x = self.proj(x)  # [B, 768, 8, 9]
        x = x.flatten(2).transpose(1, 2)  # [B, 72, 768]
        return x

Class Token 处理技巧

音频分类需要在序列前添加特殊分类 token:

  1. 初始化可学习参数:self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
  2. 前向传播时拼接:x = torch.cat((self.cls_token.expand(B, -1, -1), x), dim=1)
  3. 最终取第一个位置作为分类特征:logits = self.head(x[:, 0])

避坑指南

采样率问题解决方案

  • 统一重采样:使用 torchaudio.transforms.Resample 强制所有输入音频到相同采样率
  • 频谱对齐:当原始采样率未知时,可通过峰值检测对齐基频

可变长度音频处理

  1. 动态填充:根据 batch 内最长样本进行零填充(需配合 attention mask)
  2. 分段处理:长音频切分为 5 秒片段分别处理,最后平均预测结果

GPU 显存优化

  • 启用梯度检查点:torch.utils.checkpoint.checkpoint
  • 混合精度训练:scaler = torch.cuda.amp.GradScaler()
  • 减小 batch size 但增大累计步数

实验验证

在 ESC-50 环境声音数据集上的训练曲线显示:

  1. AST 在 100epoch 时验证准确率达到 87.5%
  2. 相比 CNN 模型,AST 的收敛速度更快(约快 2 倍)
  3. 过拟合现象更轻微(验证集与训练集差距 <3%)

生产建议

模型量化部署

  1. ONNX 导出

    torch.onnx.export(
        model,
        dummy_input,
        "ast.onnx",
        opset_version=13,
        input_names=["melspectrogram"],
        output_names=["logits"]
    )

  2. TensorRT 优化

  3. 使用 trtexec 转换 ONNX 到 TensorRT 引擎
  4. 设置 FP16 模式:--fp16
  5. 优化 profile:--minShapes=input:1x1x128x157 --optShapes=input:8x1x128x157 --maxShapes=input:32x1x128x157

开放问题思考

当处理 10 秒以上的长音频时(如 600 帧频谱图),AST 面临:
1. 注意力计算复杂度呈平方增长(O(n²))
2. 内存占用可能超过 GPU 容量
3. 长序列导致注意力权重过于分散

可能的改进方向:
– 局部注意力窗口(如只处理相邻 100 帧)
– 引入跨步注意力(stride attention)
– 层次化处理(先粗粒度后细粒度)

正文完
 0
评论(没有评论)