Transformers 中的 Audio Spectrogram Transformer(AST):原理、配置与音视频分类实战指南 📅 发布时间:2026/9/11 22:30:59 👁 浏览次数: Transformers 中的 Audio Spectrogram TransformerAST原理、配置与音视频分类实战指南【免费下载链接】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 TransformerAST 模型文档展开系统讲解这一首个「无卷积、纯注意力」音频分类模型的设计动机、在 Transformers 中的实现结构ASTConfig、ASTFeatureExtractor、ASTModel、ASTForAudioClassification四个核心组件、输入归一化与低学习率等关键使用要点并结合仓库源码与测试用例给出可直接运行的推理与微调方案。读完本文你将掌握 AST 的 patch 化流程、特征提取细节、配置项语义以及如何基于audio-classificationpipeline 和示例脚本在自有音频数据上完成分类任务。一、模型概览把声音当作图像来理解Audio Spectrogram TransformerAST由 Yuan Gong、Yu-An Chung 与 James Glass 在论文 AST: Audio Spectrogram Transformer 中提出。它的核心思想非常直观先将原始音频波形转换为频谱图spectrogram再将其当作一张图像直接应用 Vision TransformerViT架构——这正是仓库文档中「音声を画像スペクトログラムに変換することで、音声に 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定义ASTConfigfeature_extraction_audio_spectrogram_transformer.py定义ASTFeatureExtractormodeling_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、std4.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 需要更低的 lrAST 对学习率非常敏感需要较低的初始学习率。论文作者在与 PSLA 论文提出的 CNN 模型对比时使用了小 10 倍的学习率。同时 AST 收敛速度较快官方文档建议针对自己的任务仔细搜索合适的学习率与学习率调度器scheduler。这一提示在微调章节会再次体现——示例脚本--learning_rate参数的取值应明显低于常规 CNN 语音模型。三、ASTConfig参数语义与默认值ASTConfig继承自PreTrainedConfig模型类型为audio-spectrogram-transformer。仓库中该配置类的完整默认值如下见 configuration_audio_spectrogram_transformer.py参数默认值含义hidden_size768Transformer 隐藏层维度num_hidden_layers12编码器层数num_attention_heads12注意力头数intermediate_size3072MLP 中间层维度hidden_actgelu激活函数hidden_dropout_prob0.0隐藏层 dropout 概率attention_probs_dropout_prob0.0注意力 dropout 概率initializer_range0.02权重初始化范围layer_norm_eps1e-12LayerNorm 的 epsilonpatch_size16patch 尺寸可为标量或(height, width)二元组qkv_biasTrueQ/K/V 投影是否带偏置frequency_stride10频谱图 patch 化时的频率方向步长time_stride10频谱图 patch 化时的时间方向步长max_length1024频谱图的时间维度num_mel_bins128Mel 频带数量配置类文档中给出的标准用法与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_size1提取特征的维度sampling_rate16000音频数字化采样率Hznum_mel_bins128Mel 频带数max_length1024特征 padding/截断的目标长度do_normalizeTrue是否用mean/std归一化 log-Mel 特征mean-4.2677393归一化均值默认取 AudioSet 统计值std4.5689974归一化标准差默认取 AudioSet 统计值return_attention_maskFalse是否在调用时返回attention_mask4.2 特征提取的双后端实现ASTFeatureExtractor的特征提取存在两条路径仓库特意为两条路径都编写了测试TorchAudio 后端当环境安装了torchaudio时调用torchaudio.compliance.kaldi.fbank提取 fbank 特征window_typehanning、num_mel_binsself.num_mel_bins。注意源码注释提醒该后端要求 16-bit 有符号整数输入因此波形在特征提取前不应被归一化。NumPy 后端当torchaudio不可用时退化为transformers.audio_utils中的spectrogram函数使用 400 长度的 Hann 窗periodicFalse、160 的 hop 长度、512 点 FFT、0.97 预加重preemphasis与 Kaldi 风格 mel 滤波器组mel_scalekaldi、triangularize_in_mel_spaceTrue、mel_floor1.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缺失时仅打印警告支持单个样本与 batchnumpy 2D 数组、list 等输入return_tensors可取值pt或np。测试 test_integration 给出了一个可直接对照的数值验证对一段 LibriSpeech 样本ASTFeatureExtractor()输出的input_values形状为(1, 1024, 128)且input_values[0, 0, :30]与期望张量一致容差rtol1e-4, atol1e-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.Conv2dkernel_size(patch_size, patch_size)、stride(frequency_stride, time_stride)见 modeling_audio_spectrogram_transformer.py。frequency_stride与time_stride因此直接决定 patch 的稠密程度与序列长度。5.2 ASTEmbeddingsCLS 令牌、蒸馏令牌与位置编码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在池化输出之上叠加ASTMLPHeadLayerNorm Linear 分类头用于 AudioSet、Speech Commands V2 等数据集见 modeling_audio_spectrogram_transformer.pyconfig.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 pipelineaudio-classification/标注。使用pipeline推理的最小示例from transformers import pipeline # 自动加载 AST 特征提取器与 ASTForAudioClassification classifier pipeline(audio-classification, modelMIT/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_rate16000, return_tensorspt) 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_namedatasets库中的数据集名与配置名--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_maskTrue才能在分类头获得正确的 masked mean-pooling但不一定总能带来更高准确率--ignore_mismatched_sizes当预训练模型分类头维度与数据集标签数不匹配时自动调整分类头--token、--trust_remote_codeHugging 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_length1024 帧 ≈ 10.24 秒 16kHz输入要求。如果只做推断验证仓库测试中MIT/ast-finetuned-audioset-10-10-0.4593的加载与推理流程test_model_from_pretrained可作为最小可复现模板。八、注意事项与限制采样率约束ASTFeatureExtractor默认sampling_rate16000输入采样率不一致会报错音频需为单声道 float32 波形。固定输入长度特征会被截断/补零到max_length1024 帧超长音频务必在训练侧裁剪--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 架构图为论文作者绘制的示意图托管于外部站点本文不再重复引用如需查看架构细节可阅读论文原文或原文档第 2730 行对应的插图说明。【免费下载链接】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),仅供参考