Transformers 中的 Audio Spectrogram Transformer(AST):原理、配置与音视频分类实战指南
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
导读
本文围绕 🤗 Transformers 仓库中 Audio Spectrogram Transformer(AST) 模型文档展开,系统讲解这一首个「无卷积、纯注意力」音频分类模型的设计动机、在 Transformers 中的实现结构(ASTConfig、ASTFeatureExtractor、ASTModel、ASTForAudioClassification四个核心组件)、输入归一化与低学习率等关键使用要点,并结合仓库源码与测试用例给出可直接运行的推理与微调方案。读完本文,你将掌握 AST 的 patch 化流程、特征提取细节、配置项语义,以及如何基于audio-classificationpipeline 和示例脚本在自有音频数据上完成分类任务。
一、模型概览:把声音当作图像来理解
Audio Spectrogram Transformer(AST)由 Yuan Gong、Yu-An Chung 与 James Glass 在论文 AST: Audio Spectrogram Transformer 中提出。它的核心思想非常直观:先将原始音频波形转换为频谱图(spectrogram),再将其当作一张图像,直接应用 Vision Transformer(ViT)架构——这正是仓库文档中「音声を画像(スペクトログラム)に変換することで、音声に Vision Transformer を適用します」一句的完整含义。该模型在多个音频分类基准上取得了当时的最先进结果(论文报告 AudioSet 上 0.485 mAP、ESC-50 上 95.6% 准确率、Speech Commands V2 上 98.1% 准确率)。
1.1 论文要旨:告别 CNN 的纯注意力路线
过去十年,卷积神经网络(CNN)一直是端到端音频分类模型的主要构件,其目标是从音频频谱图直接学习到对应标签的映射。为了捕捉更长距离的全局上下文,业界普遍的做法是在 CNN 之上叠加自注意力机制,形成「CNN + Attention」的混合模型。论文要回答的问题是:CNN 依赖是否必要?纯注意力神经网络能否在音频分类上取得好成绩?AST 正是为回答这些问题而提出的首个音频分类专用、无卷积纯注意力模型。仓库文档忠实记录了论文的结论:AST 在 AudioSet、ESC-50、Speech Commands V2 等多项基准上刷新了纪录。
1.2 在仓库中的落地位置
AST 在 Transformers 中的实现集中于 src/transformers/models/audio_spectrogram_transformer/ 目录,包含:
- configuration_audio_spectrogram_transformer.py:定义
ASTConfig; - feature_extraction_audio_spectrogram_transformer.py:定义
ASTFeatureExtractor; - modeling_audio_spectrogram_transformer.py:定义
ASTModel、ASTForAudioClassification等核心模型类(由 modular_audio_spectrogram_transformer.py 自动生成); - convert_audio_spectrogram_transformer_original_to_pytorch.py:原论文作者 YuanGongND/ast 官方代码到 Transformers 的权重转换脚本。
对应的测试套件位于 tests/models/audio_spectrogram_transformer/,其中 test_modeling_audio_spectrogram_transformer.py 与 test_feature_extraction_audio_spectrogram_transformer.py 提供了形状、集成与数值一致性的验证。
二、使用要点:归一化与学习率(来自官方文档的核心提示)
原文档给出了两条对实际使用至关重要的建议,必须严格遵守:
2.1 输入归一化:均值 0、标准差 0.5
在自有数据集上微调 AST 时,建议对输入做归一化处理,使输入均值接近 0、标准差接近 0.5。这一工作由ASTFeatureExtractor自动完成。需要特别注意:特征提取器默认使用的是 AudioSet 数据集的均值与标准差(默认mean=-4.2677393、std=4.5689974)。如果你在其它下游数据集上微调,作者在原始代码库的ast/src/get_norm_stats.py中给出了如何计算该数据集自身统计量的方法,可据此覆盖默认值。
从源码看,归一化逻辑在 feature_extraction_audio_spectrogram_transformer.py 中实现为:
def normalize(self, input_values: np.ndarray) -> np.ndarray: return (input_values - (self.mean)) / (self.std * 2)注意这里分母是std * 2(而非std),这正是为了让归一化后的分布近似「均值 0、标准差 0.5」这一官方推荐目标。默认的 AudioSet 均值/标准差在ASTFeatureExtractor.__init__中设定,可通过构造参数mean、std显式覆盖。
2.2 学习率:AST 需要更低的 lr
AST 对学习率非常敏感,需要较低的初始学习率。论文作者在与 PSLA 论文提出的 CNN 模型对比时,使用了小 10 倍的学习率。同时 AST 收敛速度较快,官方文档建议针对自己的任务仔细搜索合适的学习率与学习率调度器(scheduler)。这一提示在微调章节会再次体现——示例脚本--learning_rate参数的取值应明显低于常规 CNN 语音模型。
三、ASTConfig:参数语义与默认值
ASTConfig继承自PreTrainedConfig,模型类型为audio-spectrogram-transformer。仓库中该配置类的完整默认值如下(见 configuration_audio_spectrogram_transformer.py):
| 参数 | 默认值 | 含义 |
|---|---|---|
hidden_size | 768 | Transformer 隐藏层维度 |
num_hidden_layers | 12 | 编码器层数 |
num_attention_heads | 12 | 注意力头数 |
intermediate_size | 3072 | MLP 中间层维度 |
hidden_act | "gelu" | 激活函数 |
hidden_dropout_prob | 0.0 | 隐藏层 dropout 概率 |
attention_probs_dropout_prob | 0.0 | 注意力 dropout 概率 |
initializer_range | 0.02 | 权重初始化范围 |
layer_norm_eps | 1e-12 | LayerNorm 的 epsilon |
patch_size | 16 | patch 尺寸(可为标量或(height, width)二元组) |
qkv_bias | True | Q/K/V 投影是否带偏置 |
frequency_stride | 10 | 频谱图 patch 化时的频率方向步长 |
time_stride | 10 | 频谱图 patch 化时的时间方向步长 |
max_length | 1024 | 频谱图的时间维度 |
num_mel_bins | 128 | Mel 频带数量 |
配置类文档中给出的标准用法(与ASTModel配合):
>>> from transformers import ASTConfig, ASTModel >>> # 初始化一个 AST MIT/ast-finetuned-audioset-10-10-0.4593 风格的配置 >>> configuration = ASTConfig() >>> # 基于该配置初始化一个随机权重的模型 >>> model = ASTModel(configuration) >>> # 访问模型配置 >>> configuration = model.config测试套件中 ASTModelTester.get_config 展示了配置如何被传入模型进行微缩版验证,同时印证了frequency_stride、time_stride、attn_implementation等参数的实际传递路径。
四、ASTFeatureExtractor:从波形到标准化 log-Mel 特征
ASTFeatureExtractor继承自SequenceFeatureExtractor,负责三件事:提取 mel 滤波器组(fbank)特征、padding/截断到固定长度、按均值标准差归一化。
4.1 关键构造参数
| 参数 | 默认值 | 说明 |
|---|---|---|
feature_size | 1 | 提取特征的维度 |
sampling_rate | 16000 | 音频数字化采样率(Hz) |
num_mel_bins | 128 | Mel 频带数 |
max_length | 1024 | 特征 padding/截断的目标长度 |
do_normalize | True | 是否用mean/std归一化 log-Mel 特征 |
mean | -4.2677393 | 归一化均值(默认取 AudioSet 统计值) |
std | 4.5689974 | 归一化标准差(默认取 AudioSet 统计值) |
return_attention_mask | False | 是否在调用时返回attention_mask |
4.2 特征提取的双后端实现
ASTFeatureExtractor的特征提取存在两条路径,仓库特意为两条路径都编写了测试:
- TorchAudio 后端:当环境安装了
torchaudio时,调用torchaudio.compliance.kaldi.fbank提取 fbank 特征(window_type="hanning"、num_mel_bins=self.num_mel_bins)。注意源码注释提醒:该后端要求 16-bit 有符号整数输入,因此波形在特征提取前不应被归一化。 - NumPy 后端:当
torchaudio不可用时,退化为transformers.audio_utils中的spectrogram函数,使用 400 长度的 Hann 窗(periodic=False)、160 的 hop 长度、512 点 FFT、0.97 预加重(preemphasis)与 Kaldi 风格 mel 滤波器组(mel_scale="kaldi"、triangularize_in_mel_space=True、mel_floor=1.192092955078125e-07)计算 log-Mel 频谱。
提取后,特征会被 padding 或截断到max_length(默认 1024 帧)。测试 test_feature_extraction_audio_spectrogram_transformer.py 中专门 mock 掉is_speech_available来验证 NumPy 后端路径,确保两条实现行为一致。
4.3call的输入约束
调用ASTFeatureExtractor时:
- 仅支持单声道音频(
len(raw_speech.shape) > 2时直接抛出ValueError); - 强烈建议传入
sampling_rate,若与构造时的采样率不一致会抛出ValueError,缺失时仅打印警告; - 支持单个样本与 batch(numpy 2D 数组、list 等)输入,
return_tensors可取值"pt"或"np"。
测试 test_integration 给出了一个可直接对照的数值验证:对一段 LibriSpeech 样本,ASTFeatureExtractor()输出的input_values形状为(1, 1024, 128),且input_values[0, 0, :30]与期望张量一致(容差rtol=1e-4, atol=1e-4)——这从测试层面固化了「1024 帧 × 128 mel 频带」这一标准输入形态。
五、模型实现:从频谱图到分类输出的完整链路
5.1 ASTPatchEmbeddings:卷积实现的 patch 化
ASTPatchEmbeddings接收形状为(batch_size, max_length, num_mel_bins)的 mel 频谱图,输出(batch_size, seq_length, hidden_size)的 patch 嵌入。实现上它使用一个单通道nn.Conv2d:kernel_size=(patch_size, patch_size)、stride=(frequency_stride, time_stride)(见 modeling_audio_spectrogram_transformer.py)。frequency_stride与time_stride因此直接决定 patch 的稠密程度与序列长度。
5.2 ASTEmbeddings:CLS 令牌、蒸馏令牌与位置编码
ASTEmbeddings在 patch 嵌入前拼接两个特殊令牌:
cls_token:用于汇聚全局分类信息;distillation_token:蒸馏令牌,是 AST 从 ViT 蒸馏变体继承的设计。
patch 数量由get_shape按卷积输出尺寸公式计算:
frequency_out_dimension = (config.num_mel_bins - config.patch_size) // config.frequency_stride + 1 time_out_dimension = (config.max_length - config.patch_size) // config.time_stride + 1 num_patches = frequency_out_dimension * time_out_dimension以默认配置计算:(128 - 16) // 10 + 1 = 12(频率方向)、(1024 - 16) // 10 + 1 = 101(时间方向),共 1212 个 patch,加上 2 个特殊令牌,位置编码维度为num_patches + 2。测试类注释同样印证了「序列长度 = patch 数 + 2」的约定(test_modeling_audio_spectrogram_transformer.py)。
5.3 Transformer 编码器与池化
ASTLayer采用Pre-LayerNorm 结构:先 LayerNorm → 自注意力 → 残差,再 LayerNorm → MLP → 残差(见 modeling_audio_spectrogram_transformer.py)。注意力为双向(非因果),缩放因子为head_dim ** -0.5,并支持通过ALL_ATTENTION_FUNCTIONS接口切换 eager / SDPA / Flash Attention / Flex Attention 等后端(ASTPreTrainedModel声明了_supports_sdpa、_supports_flash_attn、_supports_flex_attn)。
ASTModel.forward最终返回BaseModelOutputWithPooling,其中池化输出取 CLS 令牌与蒸馏令牌的均值:
pooled_output = (sequence_output[:, 0] + sequence_output[:, 1]) / 2模型的主要输入为input_values((batch_size, max_length, num_mel_bins)的torch.FloatTensor),可通过AutoFeatureExtractor从.flac/.wav波形提取得到。
5.4 ASTForAudioClassification:分类头与损失
ASTForAudioClassification在池化输出之上叠加ASTMLPHead(LayerNorm + Linear 分类头),用于 AudioSet、Speech Commands V2 等数据集(见 modeling_audio_spectrogram_transformer.py):
config.num_labels > 1:计算交叉熵分类损失;config.num_labels == 1:计算均方误差(回归损失)。
集成测试 ASTModelIntegrationTest.test_inference_audio_classification 展示了标准推理流程:使用MIT/ast-finetuned-audioset-10-10-0.4593预训练权重 + 对应ASTFeatureExtractor,输入一段 AudioSet 样本音频,输出logits形状为(1, 527)(对应 AudioSet 的 527 个类别),且前三个 logits 与期望值[-0.8760, -7.0042, -8.6602]严格对齐。
六、快速上手:AST 音视频分类推理
仓库为 AST 提供了开箱即用的 pipeline 支持(文档中以<PipelineTag pipeline="audio-classification"/>标注)。使用pipeline推理的最小示例:
from transformers import pipeline # 自动加载 AST 特征提取器与 ASTForAudioClassification classifier = pipeline("audio-classification", model="MIT/ast-finetuned-audioset-10-10-0.4593") result = classifier("path/to/your/audio.wav") print(result)该 pipeline 的模型映射在测试中定义为{"audio-classification": ASTForAudioClassification, "feature-extraction": ASTModel}(test_modeling_audio_spectrogram_transformer.py)。使用ASTModel做特征提取时也可直接基于ASTFeatureExtractor手动构造输入:
import torch from transformers import ASTFeatureExtractor, ASTModel feature_extractor = ASTFeatureExtractor.from_pretrained("MIT/ast-finetuned-audioset-10-10-0.4593") model = ASTModel.from_pretrained("MIT/ast-finetuned-audioset-10-10-0.4593") # audio 为 16000Hz 采样的单声道波形数组 inputs = feature_extractor(audio, sampling_rate=16000, return_tensors="pt") with torch.no_grad(): outputs = model(**inputs) # outputs.pooler_output 即特征向量七、微调实践:基于官方音频分类示例脚本
ASTForAudioClassification由仓库中的 run_audio_classification.py 示例脚本正式支持(原文档明确注明「[ASTForAudioClassification] は、この[例示スクリプト]と[ノートブック]によってサポートされています」)。脚本基于HfArgumentParser解析三类参数:ModelArguments(模型)、DataTrainingArguments(数据)、TrainingArguments(训练),也可直接传入一个 JSON 配置文件(python run_audio_classification.py path/to/config.json)。
7.1 关键命令行参数
数据侧(DataTrainingArguments):
--dataset_name/--dataset_config_name:datasets库中的数据集名与配置名;--train_split_name(默认train)/--eval_split_name(默认validation):训练/评估划分;--audio_column_name(默认audio)/--label_column_name(默认label):音频列与标签列;--max_length_seconds(默认20):训练时随机将音频裁剪到该时长(秒)。
模型侧(ModelArguments):
--model_name_or_path:预训练模型名或路径(微调 AST 时替换为 AST 检查点);--feature_extractor_name:预处理配置名(默认复用模型名);--freeze_feature_encoder(默认True):是否冻结特征编码器层;--attention_mask(默认True):特征提取器是否生成 attention mask——源码注释提示,return_attention_mask=True才能在分类头获得正确的 masked mean-pooling,但不一定总能带来更高准确率;--ignore_mismatched_sizes:当预训练模型分类头维度与数据集标签数不匹配时,自动调整分类头;--token、--trust_remote_code:Hugging Face Hub 鉴权与远程代码信任选项。
训练侧(TrainingArguments):标准 Trainer 参数,如--learning_rate、--num_train_epochs、--per_device_train_batch_size、--gradient_accumulation_steps、--fp16、--eval_strategy、--save_strategy、--metric_for_best_model、--push_to_hub等。
7.2 微调命令示例
以下命令展示了在单卡 GPU 上以较低学习率微调(结合第二节的建议,AST 学习率应显著低于 CNN 模型,如3e-5量级甚至更低,并配合 warmup 调度器):
python run_audio_classification.py \ --model_name_or_path MIT/ast-finetuned-audioset-10-10-0.4593 \ --dataset_name superb \ --dataset_config_name ks \ --output_dir ast-ft-keyword-spotting \ --remove_unused_columns False \ --do_train \ --do_eval \ --fp16 \ --learning_rate 3e-5 \ --max_length_seconds 10 \ --attention_mask False \ --warmup_ratio 0.1 \ --num_train_epochs 5 \ --per_device_train_batch_size 8 \ --gradient_accumulation_steps 4 \ --per_device_eval_batch_size 8 \ --eval_strategy epoch \ --save_strategy epoch \ --load_best_model_at_end True \ --metric_for_best_model accuracy \ --save_total_limit 3 \ --seed 0脚本内部会通过AutoFeatureExtractor.from_pretrained加载 AST 特征提取器,对每个音频样本提取 fbank 特征,并调用Trainer完成训练与评估;--max_length_seconds对应的random_subsample函数会在训练时对超长音频随机裁剪,以匹配 AST 固定的max_length(1024 帧 ≈ 10.24 秒 @16kHz)输入要求。如果只做推断验证,仓库测试中MIT/ast-finetuned-audioset-10-10-0.4593的加载与推理流程(test_model_from_pretrained)可作为最小可复现模板。
八、注意事项与限制
- 采样率约束:
ASTFeatureExtractor默认sampling_rate=16000,输入采样率不一致会报错;音频需为单声道 float32 波形。 - 固定输入长度:特征会被截断/补零到
max_length(1024 帧),超长音频务必在训练侧裁剪(--max_length_seconds)。 - 归一化统计量:默认使用 AudioSet 的
mean/std,换数据集微调时建议按原论文get_norm_stats.py的逻辑重算并覆盖。 - 学习率敏感性:低学习率 + 合适的 scheduler 是 AST 微调成功的关键。
- 依赖:
torchaudio为可选依赖;未安装时特征提取自动回退到 NumPy 实现,行为由 test_feature_extraction_audio_spectrogram_transformer.py 中的 mock 测试保证。
参考资料
- 模型文档:docs/source/ja/model_doc/audio-spectrogram-transformer.md(本文依据)
- 配置实现:configuration_audio_spectrogram_transformer.py
- 特征提取实现:feature_extraction_audio_spectrogram_transformer.py
- 模型实现:modeling_audio_spectrogram_transformer.py
- 测试套件:tests/models/audio_spectrogram_transformer/
- 微调示例:run_audio_classification.py 及其 README
- 任务文档:音视频分类任务(
docs/source/en/tasks/audio_classification.md)
说明:原文档中展示的 AST 架构图为论文作者绘制的示意图,托管于外部站点,本文不再重复引用;如需查看架构细节,可阅读论文原文或原文档第 27~30 行对应的插图说明。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考