ESPnet 实战:基于 BEATs 编码器在 ESC-50 上训练音频分类任务的完整 Recipe 解析
人工智能语音音频深度学习NLP【免费下载链接】espnetEnd-to-End Speech Processing Toolkit项目地址https://gitcode.com/gh_mirrors/es/espnet点击查看免费下载导读本文以 egs2/esc50/asr1/README.md 为核心骨架系统讲解如何在 ESPnet 中以 BEATsAudio Pre-Training with Acoustic Tokenizers论文 arXiv:2212.09058作为冻结/微调编码器、配合线性分类解码器在 ESC-50 环境声音数据集上完成 5 折交叉验证的音频分类训练与评测。读完本文你将掌握ESC-50 数据集下载与 db.sh 配置、BEATs_iter3 预训练权重获取、beats_classification.yaml 全量参数语义、run.sh 五折并行训练与平均精度计算流程以及底层 BeatsEncoder 与 LinearDecoder 的源码级工作机制。一、Recipe 概况用 ASR 管线跑音频分类该 Recipe 属于egs2下的esc50/asr1任务但实现的是**音频分类audio classification**任务而非传统语音识别。它复用了asr.sh这一 ESPnet 标准训练脚本通过--token_type word、--use_lm false、--feats_type raw等参数将整条管线改造成序列输入 → 分类输出的形态并将模型结构替换为编码器encoder: beats对应 BeatsEncoder解码器decoder: linear_decoder对应 LinearDecoder。该实现复用了微软 BEATs 仓库论文 Table 1 倒数第二行BEATS-iter3的微调配置与结果在本 Recipe 中均有对应实现。当前仓库中 espnet2/tasks/cls.py 同样注册了BeatsEncoder与LinearDecoder可作为分类任务的官方任务入口参考而本 Recipe 通过 asr 管线复用在工程上实现了同一套模型能力。二、数据集与实验设计ESC-50 的 5 折交叉验证2.1 数据集规模ESC-50 是环境声音分类基准数据集本 Recipe 按官方划分执行5 折交叉验证5-fold cross-validation共2000 个样本每折400 个样本。每个 fold 中该 fold 的 400 条作为验证集其余 1600 条作为训练集对应 yaml 注释 12.5 steps per epoch with 1600 samples。2.2 训练资源与超参说明官方文档给出明确提示单折微调约需 1 块 33 GB 显存的 GPU在 L40S 上运行约 4.5 小时。需要特别注意的是本 Recipe 的超参数与 BEATs 论文附录 A.1 不完全一致——the ones used here gave us best results文档原文这些超参先在fold 5 上调优再复用到其他 fold。这意味着不要机械照抄论文附录的超参若要复现论文数字应使用本 yaml 中给出的调优后取值若更换数据集建议沿用先固定一折调参、再推广其余折的流程。三、三步启动数据、权重、训练3.1 下载 ESC-50 并配置 db.sh从 ESC-50 官方仓库下载数据集后将数据集根目录路径填入 db.sh 中的ESC50变量ESC50/your/path/to/ESC-50注意官方说明推荐手动下载并指向本地路径若留空local/data.sh 会提示 Fill the value of ESC50 of db.sh 并退出stage 1 的下载分支仅在ESC50/LICENSE不存在时触发。3.2 下载 BEATs 预训练权重从 BEATs 仓库下载BEATs_iter3检查点.pt文件并把路径写入 beats_classification.yaml 的encoder_conf.beats_ckpt_pathencoder_conf: beats_ckpt_path: /path/to/models/BEATs/BEATs_iter3.pt从源码看BeatsEncoder 加载该权重时使用 safe_torch_load 进行安全加载并在 generate_beats_checkpoint.py 中提供了权重转换/适配脚本说明该路径既支持官方原始权重也可经过转换后使用。3.3 一键启动cd egs2/esc50/asr1 ./run.shrun.sh会并行启动 5 个 fold 的训练任务n_folds5脚本注释明确提醒 This runs all 5 folds in parallel, take care因此建议在资源充足的集群/多卡机器上运行。四、核心配置文件 beats_classification.yaml 全参数解析yaml 是本 Recipe 的灵魂以下是逐段拆解4.1 训练基础配置token_type: word optim: adamw optim_conf: lr: 1.0e-4 weight_decay: 1.0e-2 betas: [0.9, 0.98] accum_grad: 1 batch_size: 128 # 12.5 steps per epoch with 1600 samples max_epoch: 1000 scheduler: CosineAnnealingWarmupRestarts scheduler_conf: first_cycle_steps: 6000 warmup_steps: 300 max_lr: 1.0e-4 min_lr: 5.0e-6optimadamw优化器学习率1e-4权重衰减1e-2beta 为(0.9, 0.98)batch_size128每个 epoch 约 12.5 步1600 训练样本 / 128max_epoch1000配合余弦退火重启调度器在长时间尺度上逐步收敛schedulerCosineAnnealingWarmupRestarts首轮周期 6000 步、预热 300 步、最大学习率1e-4、最小学习率5e-6——从源码角度该调度器在 espnet2/schedulers/ 中实现预热后按余弦曲线衰减并支持周期重启。4.2 前端与归一化交给 BEATs 内部处理# BEATs implementation takes care of generating mel spectrogram, normalization and specaug frontend: none input_size: 1 # important to set input_size to 1 if frontend is none normalize: none # BEATs code does global mean and variance normalization这是理解本 Recipe 的关键设计BEATs 编码器内部自行完成 mel 频谱生成、全局均值和方差归一化以及 SpecAug 数据增强因此外部 frontend 和 normalize 都置为none。input_size必须设为1注释强调 important to set input_size to 1 if frontend is none否则输入维度推断会出错。这一设计可从源码印证BeatsEncoder内部实现了 fbank 提取基于 torchaudio Kaldi 接口见 encoder.py并使用 DEFAULT_FBANK_MEAN / DEFAULT_FBANK_STD 做全局归一化同时集成了 SpecAug。4.3 模型与解码器# Initialization for the decoder init: xavier_normal model_conf: ctc_weight: 0.0 # No CTC, no attention. lsm_weight: 0.1 # label smoothing weight length_normalized_loss: true decoder: linear_decoder decoder_conf: pooling: mean dropout: 0.1init: xavier_normal解码器线性层采用 Xavier 正态初始化ctc_weight: 0.0既不使用 CTC 也不使用注意力解码纯分类输出lsm_weight: 0.1标签平滑权重 0.1decoder: linear_decoder线性分类头pooling: mean对编码器输出的时间维做均值池化dropout: 0.1。源码 LinearDecoder 表明它先将vocab_size - 3作为输出类别数ESPnet 文本处理中unk/blank/sos/eos等特殊 token 占 3 个位置支持mean/max/CLS三种池化方式前向时对hs_pad按hlens构建 mask做均值池化后过nn.Linear输出(B, n_classes)。4.4 BEATs 编码器微调细节encoder: beats encoder_conf: # Please download the BEATs model from the BEATs repo (iter3) and update the path below beats_ckpt_path: /compute/babel-13-33/sbharad2/models/BEATs/BEATs_iter3.pt # Most values from Appendix A.1 of the BEATs paper or tuned on fold 5. # Please also check the README.md fbank_mean: 11.72215 fbank_std: 10.60431 beats_config: layer_wise_gradient_decay_ratio: 0.2 encoder_layerdrop: 0.1 dropout: 0.0 specaug_config: apply_time_warp: true apply_freq_mask: false apply_time_mask: true time_mask_width_ratio_range: - 0 - 0.06 num_time_mask: 1 roll_augment: true roll_interval: 16000 # 1 second, only 5 possible augmentations per sample use_weighted_representation: false逐项说明参数取值含义beats_ckpt_pathBEATs_iter3.pt 路径BEATs-iter3 预训练权重位置必须替换为本地路径fbank_mean/fbank_std11.72215 / 10.60431全局 fbank 统计量用于输入归一化layer_wise_gradient_decay_ratio0.2层间梯度衰减比例底层学习率按 0.2 比例递减稳定微调encoder_layerdrop0.1训练时随机丢弃 transformer 层的概率一种正则化手段dropout0.0编码器主体 dropout注意解码器另有 0.1apply_time_warp/apply_time_masktrue / trueSpecAug 时间维增强时间弯曲 时间掩码apply_freq_maskfalse不应用频率掩码time_mask_width_ratio_range[0, 0.06]时间掩码宽度占序列长度的比例范围num_time_mask1每条样本时间掩码数量roll_augmenttrue时域滚动roll增强roll_interval16000滚动步长 16000 采样点注释 1 second, only 5 possible augmentations per sample即 5 秒音频共有 5 种滚动位置use_weighted_representationfalse不使用加权表示从源码印证BeatsConfigencoder.py中定义了layer_wise_gradient_decay_ratio、encoder_layerdrop、dropout等默认值用户 yaml 中的beats_config会通过cfg.update覆盖默认值roll_augment在实现上复用了 roll_tensor 完成时域滚动。4.5 训练调度与杂项batch_type: folded unused_parameters: true grad_clip: 1 patience: none best_model_criterion: - - valid - acc - max keep_nbest_models: 1 use_amp: false # whether to use automatic mixed precision num_att_plot: 0 num_workers: 2 # dataloader workersbest_model_criterion以验证集acc准确率最大化为选模标准keep_nbest_models: 1仅保留最优模型与inference_modelvalid.acc.best.pth对应use_amp: false默认关闭自动混合精度若显存紧张可尝试开启num_workers: 2DataLoader 工作进程数。五、run.sh五折并行训练与结果聚合run.sh 是整个流程的调度中枢5.1 关键前置变量asr_speech_fold_length1000 # 6.25 sec, because audio is 5 sec each. inference_modelvalid.acc.best.pth n_folds5 # This runs all 5 folds in parallel, take care. asr_configconf/beats_classification.yaml mynametagfast.foldasr_speech_fold_length1000按帧frame长度对齐注释说明音频每条 5 秒取 6.25 秒上限以保证安全inference_model指向按valid.acc.best准则保存的最优权重与 yaml 中keep_nbest_models: 1呼应。5.2 五折并行训练循环for fold in $(seq 1 $n_folds); do train_settrain${fold} valid_setval${fold} test_setval${fold} ./asr.sh \ --local_data_opts ${fold} \ --asr_tag ${mynametag}${fold} \ --lang ${fold} \ --ngpu 1 \ --stage 15 \ --inference_args --ctc_weight 0.0 --maxlenratio -1 \ --token_type word \ --asr_speech_fold_length ${asr_speech_fold_length} \ --use_lm false \ --feats_type raw \ --max_wav_duration 6 \ --feats_normalize utterance_mvn\ --inference_nj 8 \ --inference_asr_model ${inference_model} \ --asr_config ${asr_config} \ --train_set ${train_set} \ --valid_set ${valid_set} \ --test_sets ${test_set} $ done wait要点解析--local_data_opts ${fold}把当前 fold 号传给local/data.sh后者以FOLD${1:-1}接收进而调用 data_prep_multi_folds.py 生成该折的train{fold}/val{fold}数据目录--lang ${fold}脚本注释明确指出 Abusing variable lang to store fold number——lang变量被借用来存 fold 编号这是工程上的取巧写法--stage 15从 stage 15训练与推理阶段开始跳过特征提取等前期 stage——因为 BEATs 内部自管 fbank且数据准备已由 local 脚本完成--ngpu 1每个 fold 用 1 块 GPU5 折并行则共需 5 块 GPU--inference_args --ctc_weight 0.0 --maxlenratio -1推理时 CTC 权重为 0、无长度约束分类任务不需要解码搜索长度控制--feats_type raw/--max_wav_duration 6直接以原始波形训练单条音频最长 6 秒--feats_normalize utterance_mvn训练侧使用语句级 MVN尽管 yaml 中 normalize 为 none此处针对的是 asr.sh 管线的外层特征处理--use_lm false不训练/不使用语言模型。5.3 五折平均精度计算训练结束后脚本从每个 fold 的exp/asr_fast.fold${i}/RESULTS.md中解析验证集准确率并求平均total_sum0 total_count0 for i in $(seq 1 $n_folds); do values$(grep val${i} exp/asr_${mynametag}${i}/RESULTS.md | head -n 1 | awk -F| {print $(NF-1)}) for value in $values; do total_sum$(echo $total_sum $value | bc) total_count$((total_count 1)) break done done if [ $total_count -gt 0 ]; then average$(echo scale2; $total_sum / $total_count | bc) echo Avg. acc: $(echo 100 - $average | bc) echo Over $total_count folds. fi运行前请确保各 fold 的RESULTS.md已生成且为空模板脚本注释提醒 Please ensure that the RESULTS.md file is empty before running this script否则grep可能命中历史残留行。show_asr_result.sh见 scripts/utils/负责把各 fold 结果汇总写入 RESULTS 表。六、数据准备链路local 脚本与 5 折划分实现6.1 data.sh 的流水线local/data.sh 接收 fold 号默认 1执行stage 1检查${ESC50}/LICENSE是否存在缺失时提示下载stage 2调用 data_prep_multi_folds.py 生成train${FOLD}/val${FOLD}的wav.scp、text、utt2spk随后用utils/utt2spk_to_spk2utt.pl生成spk2utt并调用utils/validate_data_dir.sh --no-feats做数据目录校验。fold 1 的特殊分支当FOLD为 1 时还会额外调用data_prep.py生成data/{train,valid,test}三个常规目录保留 SLU 任务的默认路径兼容性。6.2 data_prep_multi_folds.py 的划分逻辑data_prep_multi_folds.py 的核心逻辑非常简洁读取meta/esc50.csvpandasval{fold}meta_data[fold] fold_num的行400 条train{fold}其余所有行1600 条每条样本生成三行记录wav.scputt_id /path/to/audio/xxx.wavtextutt_id audio_class:target类别标签以audio_class:前缀写入文本域供分类任务使用utt2spkutt_id utt_id每条音频视作独立 speaker这种把分类标签编码进 text 文件的写法正是该 Recipe 复用 asr 管线实现分类任务的关键技巧。七、实验环境与基准结果RESULTS文档中的 RESULTS 段由 show_asr_result.sh 自动生成记录了训练环境与各折准确率。7.1 复现环境项值日期Sat Dec 14 19:04:56 EST 2024Python3.9.20 (GCC 11.2.0)ESPnetespnet 202412PyTorchpytorch 2.4.0Git hashcb80e61a15d6a13dc342ae5a413d2b870dd869c6commit dateFri Dec 13 11:57:16 2024 -05007.2 各折验证准确率datasetSntWrdAccorg/val140040094.3org/val240040097.0org/val340040094.8org/val440040096.3org/val540040091.8Average94.8文档同时给出各 fold 已发布权重Fold-1 至 Fold-5对应精度 94.3 / 97.0 / 94.8 / 96.3 / 91.8平均精度94.8%。需要说明这些 checkpoint 托管在外部模型托管平台仓库内不包含权重文件本身如需使用请按其对应说明下载后放入本地路径。八、误差分析fold-5 的主要混淆模式文档指出在 fold-5 上观察到的主要混淆来自helicopter直升机类它主要与washing machine洗衣机和airplane飞机混淆。这一观察符合该类别的声学直觉——机械噪声、螺旋桨声与家电运转声在频谱上存在重叠。对后续改进的启示若目标类间声学相似度高可尝试在 SpecAug 中启用apply_freq_mask增加频率维正则可对比pooling: max/CLS与mean池化在混淆类上的表现可结合混淆矩阵做细粒度调参例如针对性地提高易混类别样本的增强强度。九、从源码看实现边界与可扩展点9.1 BEATs 编码器内部结构BeatsConfig 展示了 BEATs 编码器默认架构patch embeddingpatch size 16、embed dim 512、12 层 transformerembed dim 768、FFN 3072、12 头注意力、GELU 激活、conv 位置编码128 滤波器、以及面向预训练任务的 codebook1024 词与 mask ratio 0.75。微调时用户通过beats_config覆盖layer_wise_gradient_decay_ratio、encoder_layerdrop、dropout即可无需改动其余默认结构。9.2 线性分类头的类别数推导LinearDecoder 中output_dim vocab_size - 3ESPnet 的 word 级 token 系统默认包含unk、blank、sos/eos三个特殊 token减去后即为真实类别数ESC-50 为 50 类。这也解释了为什么token_type: word与分类任务能天然兼容。9.3 与 cls 任务的对照当前仓库 espnet2/tasks/cls.py 已把BeatsEncoder、LinearDecoder 注册进分类任务CLS体系说明 BEATs 线性头的组合在 ESPnet 内是一条正式支持的模型路径本 esc50 Recipe 则提供了以 asr.sh 快速上手的工程范本两者模型成分一致、入口不同可按需选择。十、快速排障清单现象排查方向脚本提示Fill the value of ESC50 of db.sh确认 db.sh 中ESC50已指向数据集根目录报错找不到 BEATs 权重检查beats_ckpt_path是否指向已下载的BEATs_iter3.pt必要时参考 generate_beats_checkpoint.py 转换权重格式显存不足五折并行共需 5×33 GB 显存可减小n_folds串行执行、开启use_amp: true或降低batch_size平均精度脚本读不到值确保各 fold 的RESULTS.md已生成且为干净模板exp/asr_fast.fold${i}目录路径与mynametag一致训练侧与推理侧输入不一致确认input_size: 1、frontend: none、normalize: none未被覆盖BEATs 内部 fbank 与归一化参数fbank_mean/fbank_std与预训练时一致结语本文围绕 egs2/esc50/asr1/README.md 完整还原了BEATs 编码器 线性解码器在 ESC-50 上做 5 折交叉验证分类的实战链路从数据与权重准备、yaml 全量参数语义到 run.sh 五折并行调度与平均精度聚合再到BeatsEncoder/LinearDecoder的源码级印证。README 报告的平均准确率 94.8% 可作为本地复现的对照基线若需深挖模型内部实现可继续阅读 espnet2/beats/encoder.py、espnet2/beats/utils.py 与 espnet2/asr/decoder/linear_decoder.py。赞分享人工智能语音音频深度学习NLP【免费下载链接】espnetEnd-to-End Speech Processing Toolkit项目地址https://gitcode.com/gh_mirrors/es/espnet点击查看免费下载相关推荐Acton蓝绿部署无停机部署策略Acton蓝绿部署无停机部署策略 Acton作为TON智能合约开发的完整工具链提供了强大的蓝绿部署能力帮助开发者实现零停机时间的智能合约更新。本文将详细介人工智能语音音频Wazuh Engine conf 模块深度解析三级优先级配置系统的架构、文件格式与源码实现Wazuh Engine conf 模块深度解析三级优先级配置系统的架构、文件格式与源码实现 本文为 Wazuh 引擎Wazuh Engine的 conf人工智能语音音频深度学习NLP在 Temporal 工作流中编排 SkyPilot 训练任务基于 skypilot 的完整实战指南在 Temporal 工作流中编排 SkyPilot 训练任务基于 skypilot 的完整实战指南 SkyPilot 提供了跨云的一站式集群生命周期管理后端任务调度MLOps集群管理上一篇3步搞定微信公众号RSS订阅告别信息碎片化打造专属阅读体验下一篇gh_mirrors/pkg/pkg 函数参考lib/index.ts 导出 API 文档创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考