共计 3576 个字符,预计需要花费 9 分钟才能阅读完成。
背景与痛点
在音频信号处理领域,传统卷积神经网络(CNN)一直是主流方法。然而,CNN 在处理音频频谱图时存在几个明显的局限性:

-
感受野限制 :CNN 的卷积核大小固定,难以捕获长距离的时频依赖关系。对于音频这种具有长时序特性的信号,局部感受野可能导致全局信息丢失。
-
平移不变性假设 :CNN 的平移不变性假设在图像领域很有效,但在音频处理中,频率轴和时间轴的平移具有完全不同的语义,简单假设平移不变性可能导致模型性能下降。
相比之下,RNN/LSTM 虽然能够处理长序列,但也有其自身的问题:
- 计算效率低 :RNN/LSTM 的时序依赖性导致难以并行计算,训练速度慢。
- 梯度消失 / 爆炸 :长序列训练过程中容易出现梯度问题。
AST 核心技术解析
Audio Spectrogram Transformer 通过引入注意力机制,有效解决了上述问题。其核心设计包括:
1. 二维频谱图 patch 嵌入
AST 将频谱图视为 ” 图像 ”,采用类似 ViT 的 patch 划分方式,但针对音频特性做了特殊处理:
# PatchEmbedding 实现示例
class PatchEmbed(nn.Module):
def __init__(self, img_size=(128, 128), patch_size=(16, 16), in_chans=1, embed_dim=768):
super().__init__()
self.img_size = img_size
self.patch_size = patch_size
self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1])
self.num_patches = self.grid_size[0] * self.grid_size[1]
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size,
stride=patch_size)
def forward(self, x):
B, C, H, W = x.shape
# 输出 shape: (B, embed_dim, grid_h, grid_w)
x = self.proj(x)
# 展平: (B, embed_dim, grid_h*grid_w)
x = x.flatten(2).transpose(1, 2)
return x
2. 时频位置编码
AST 采用可学习的位置编码,分别处理时间和频率轴:
$$
PE_{(pos,2i)} = \sin(pos/10000^{2i/d_{model}})
$$
$$
PE_{(pos,2i+1)} = \cos(pos/10000^{2i/d_{model}})
$$
3. 多头注意力机制
标准的多头注意力计算过程:
$$
Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V
$$
[CLS] token 在音频分类中的作用可以通过以下公式表示:
$$
p(y|x) = \text{softmax}(W\cdot h_{[CLS]} + b)
$$
其中 $h_{[CLS]}$ 是 [CLS] token 的最终隐藏状态。
完整 PyTorch 实现
以下是 AST 的关键实现部分:
频谱图生成
import librosa
def generate_spectrogram(audio_path, sr=16000, n_fft=1024, hop_length=160, n_mels=128):
"""
生成梅尔频谱图
参数说明:
sr: 采样率 (必须与训练时一致)
n_fft: FFT 窗口大小
hop_length: 帧移
n_mels: 梅尔带数量
"""
y, _ = librosa.load(audio_path, sr=sr)
S = librosa.feature.melspectrogram(y=y, sr=sr,
n_fft=n_fft,
hop_length=hop_length,
n_mels=n_mels)
S_dB = librosa.power_to_db(S, ref=np.max)
return S_dB
AST 模型主体
import torch
import torch.nn as nn
from einops import rearrange
class ASTModel(nn.Module):
def __init__(self,
input_size=(128, 128),
patch_size=(16, 16),
embed_dim=768,
num_heads=12,
num_layers=12,
num_classes=10):
super().__init__()
# Patch 嵌入
self.patch_embed = PatchEmbed(
img_size=input_size,
patch_size=patch_size,
in_chans=1,
embed_dim=embed_dim
)
# [CLS] token
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
# 位置编码
num_patches = self.patch_embed.num_patches
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
# Transformer 编码器
encoder_layer = nn.TransformerEncoderLayer(
d_model=embed_dim,
nhead=num_heads,
dim_feedforward=embed_dim * 4,
dropout=0.1,
activation="gelu"
)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
# 分类头
self.head = nn.Linear(embed_dim, num_classes)
def forward(self, x):
"""
输入 x: (B, 1, H, W)
输出: (B, num_classes)
"""
B = x.shape[0]
# 生成 patch 嵌入 (B, num_patches, embed_dim)
x = self.patch_embed(x)
# 添加 [CLS] token (B, 1, embed_dim)
cls_tokens = self.cls_token.expand(B, -1, -1)
x = torch.cat((cls_tokens, x), dim=1)
# 添加位置编码
x = x + self.pos_embed
# 通过 Transformer 编码器
x = self.encoder(x)
# 取 [CLS] token 进行分类
cls_token = x[:, 0, :]
out = self.head(cls_token)
return out
生产实践与优化
性能优化技巧
- 混合精度训练 :
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for inputs, labels in train_loader:
inputs = inputs.to(device)
labels = labels.to(device)
with autocast():
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
- ONNX 导出注意事项 :
torch.onnx.export(
model,
dummy_input,
"ast_model.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={"input": {0: "batch_size"},
"output": {0: "batch_size"}
},
opset_version=13
)
常见问题与解决方案
- 梅尔尺度 vs 线性尺度 :
- 梅尔尺度更接近人类听觉感知,适合语音 / 音乐
-
线性尺度保留完整频域信息,适合特殊声音检测
-
采样率不一致 :
- 训练和推理必须使用相同采样率
- 可通过重采样解决不一致问题
延伸思考与改进方向
- 低资源设备部署 :
- 知识蒸馏压缩模型
-
量化到 8 位 / 4 位整数
-
架构改进 :
- 结合 ConvStem 的混合架构
-
分层注意力机制
-
其他改进方向 :
- 自监督预训练
- 多模态融合
总结
AST 通过引入 Transformer 架构到音频领域,有效解决了传统 CNN 和 RNN 的局限性。本文详细介绍了 AST 的核心原理、完整实现以及生产部署中的优化技巧,希望能为读者在实际项目中应用 AST 提供参考。根据具体任务需求,可以进一步探索模型压缩、架构改进等方向,以获得更好的性能表现。
