- 人工智能
- 语音
- 音频
- 深度学习
- NLP
【免费下载链接】espnet
End-to-End Speech Processing Toolkit
本篇技术指南围绕 ESPnet 仓库中egs2/meld/cls1情感分类(CLS)食谱展开,系统讲解其在 MELD 多模态对话情感数据集上的完整落地路径:从数据准备、特征前处理、模型架构(S3PRL WavLM Base+ 冻结前端 + Transformer 编码器 + 线性分类头)到训练、推理、评分与模型打包上传的十个流水线阶段。读者读完本文后,将掌握如何在 ESPnet2 框架下复现该食谱、读懂其全部配置参数、理解评分指标(mean_acc / mAP / mean_auc)的含义,并知道如何基于同一套 CLS 任务模板迁移到其他分类语料。
一、食谱定位与目录结构
egs2/meld/cls1是 ESPnet2 标准食谱(recipe)体系中面向**语音情感分类(Speech Emotion Classification)**任务的实现。该食谱的官方说明明确将其定位为"首个可工作的实现,而非精心调优的配置"(原文:This configuration is a first working implementation, not a tuned one),这一诚实定位也解释了其指标仍有较大提升空间。
食谱目录结构(见 egs2/meld/cls1):
| 目录 / 文件 | 职责 |
|---|---|
| run.sh | 一键入口脚本,封装数据准备到模型上传的完整流程 |
| cls.sh | 分类任务通用流水线脚本(Stage 1~10),定义全部可调参数 |
| conf/train_cls_wavlm_transformer.yaml | 模型与训练核心配置(前端、编码器、优化器、任务类型) |
| local/data.sh | 数据下载与 Kaldi 风格数据目录构建 |
| local/data_prep.py | 将 MELD 的 CSV 标注转换为text/wav.scp/utt2spk |
| pyscripts/utils/cls_score.py | 评分脚本:计算 mean_acc / mAP / mean_auc |
| scripts/utils/show_cls_result.sh | 汇总环境信息与各 split 的评分结果,生成 Markdown 报告 |
二、数据集:MELD 及其挑战
MELD(Multimodal EmotionLines Dataset)是一个多模态多轮对话数据集,来源于美剧Friends,标注了 7 类情感:neutral(中性)、joy(喜悦)、surprise(惊讶)、anger(愤怒)、sadness(悲伤)、disgust(厌恶)、fear(恐惧)。本食谱仅使用音频轨道做单标签(single-label)多分类。
该数据集有两个被社区公认的难点,直接决定了评估方式:
- 高度类别不平衡:训练集中
neutral占比高达 47%,因此在测试集上"多数类基线"(全部预测为 neutral)的准确率即为 48.2%。这意味着只看准确率(plain accuracy)不足以反映模型真实能力,必须结合加权 F1、mAP、AUC 等多类指标评估。 - utterance 级对齐不完美:MELD 的句子级音频切分与标注对齐存在已知误差,这从数据侧限制了可达到的准确率上限。
数据准备逻辑见 local/data.sh:Stage 1 从原始站点(回退到备份源)下载并解压;Stage 2 由 local/data_prep.py 读取{train,valid,test}_sent_emo.csv,为每个 utterance 生成:
utt2spk:以{speaker}-dia{语轮}-utt{句号}-sea{季}-epi{集}-{split}作为 utterance ID;text:utterance ID + 情感标签(作为分类标签,非转录文本);wav.scp:通过ffmpeg -i ...mp4 -ac 1 -ar 16000 -f wav -vn -的管道方式从视频文件中实时抽取 16 kHz 单声道音频(该管道输入会在 Stage 2 被改写为实体音频文件)。
值得注意的细节:data_prep.py硬编码过滤了 4 条"极长序列"(如Ross-dia125-utt3-sea4-epi18-train),用于剔除标注异常的样本。
三、模型架构:冻结 WavLM 前端 + Transformer + 线性分类头
本食谱的模型由三部分组成,完整配置见 conf/train_cls_wavlm_transformer.yaml:
- frontend:
s3prl,上游模型wavlm_base_plus(WavLM Base+,由 S3PRL 框架加载),并通过freeze_param: frontend.upstream冻结全部参数; - encoder:Transformer,4 个 block、输出维度 128、
input_layer: linear; - decoder:线性分类头,配合mean pooling(对编码器输出的时间维做掩码平均池化后映射到类别数)。
3.1 S3PRL 前端与冻结机制
S3prlFrontend的实现位于 espnet2/asr/frontend/s3prl.py:它通过S3PRLUpstream加载预训练上游模型,并用Featurizer将其输出转化为下游可用的特征序列。其关键参数包括:
fs:输入采样率,默认 16000,S3PRL 全部上游模型目前仅支持 16 kHz 音频;upstream:上游模型名称,本食谱为wavlm_base_plus;download_dir:预训练权重下载目录,本食谱设为./hub;multilayer_feature:是否拼接多层特征(本食谱开启true),若显式指定layer则会关闭多层拼接;normalize:是否对上游输入做归一化,默认关闭。
前端初始化后会立即eval()并保存pretrained_params快照。在 ESPnet 的训练框架中,freeze_param: frontend.upstream会确保该模块在反向传播时梯度不更新,从而把 WavLM 当作固定的特征提取器使用。这也是本食谱 GPU 显存占用与训练成本的主要可控因素之一。
3.2 Transformer 编码器
编码器选用 ESPnet 标准的TransformerEncoder,配置为:
| 参数 | 值 | 说明 |
|---|---|---|
output_size | 128 | 编码器输出维度 |
attention_heads | 4 | 多头注意力头数 |
linear_units | 1024 | FFN 中间层维度 |
num_blocks | 4 | Transformer block 数量 |
dropout_rate | 0.4 | dropout 比率 |
input_layer | linear | 输入投影层类型 |
从 CLS 任务定义(espnet2/tasks/cls.py)可以看到,CLSTask的可选编码器包括transformer、conformer与beats(默认transformer),可选前端包括default、sliding_window、s3prl、fused,本食谱即采用了默认的 Transformer + 显式指定的 S3PRL 前端组合。
3.3 线性解码器与 mean pooling
解码器LinearDecoder(espnet2/cls/decoder/linear_decoder.py)接收编码器输出(B, T, D),支持三种池化方式:
mean:对时间维做掩码均值池化(默认,本食谱采用);max:掩码后取时间维最大值;CLS:直接取序列第一个位置的表示。
池化后的向量经一个nn.Linear(encoder_output_size, n_classes)得到 logits。推理时的score()接口假定 batch size 为 1 且输入为未 padding 的单个序列。
3.4 分类任务类型与损失
模型主体ESPnetClassificationModel(espnet2/cls/espnet_model.py)支持两种分类类型:
multi-class(本食谱):softmax +CrossEntropyLoss(可配合lsm_weight做标签平滑);multi-label:sigmoid +BCEWithLogitsLoss,训练时支持 mixup 增强,但该模式仅支持 PyTorch Lightning 训练器(cls.sh中会强制要求--use_lightning true)。
训练过程中模型会实时统计acc与macro_precision(基于 torcheval),并在log_epoch_metrics: true时缓存每个 epoch 的预测用于 mAP 日志。类别数在CLSTask.build_model中被计算为len(token_list) - 1,即从词表(7 个情感标签)中扣除为兼容性而添加的<unk>占位符。
四、训练配置逐项解析
核心 YAML 配置全文如下(conf/train_cls_wavlm_transformer.yaml):
# ======== Training ======== batch_size: 32 max_epoch: 30 # ======== Optimizer ======== optim: adam optim_conf: lr: 1.0e-3 # ======== Learning rate scheduler ======== scheduler: warmuplr scheduler_conf: warmup_steps: 3180 # 10 epochs (batch size 32) # ======== Checkpointing and logging ======== patience: 5 best_model_criterion: - - valid - acc - max keep_nbest_models: 1 num_att_plot: 0 # ======== Model architecture ======== frontend: s3prl frontend_conf: frontend_conf: upstream: wavlm_base_plus download_dir: ./hub multilayer_feature: true freeze_param: - frontend.upstream encoder: transformer encoder_conf: output_size: 128 attention_heads: 4 linear_units: 1024 num_blocks: 4 dropout_rate: 0.4 input_layer: linear # ======== Classification task settings ======== model_conf: classification_type: multi-class log_epoch_metrics: true参数要点:
- 优化与调度:Adam(lr=1e-3)+ warmup LR 调度器,
warmup_steps: 3180对应约 10 个 epoch(batch size 32 下的估算值),之后学习率逐步衰减。 - 早停与模型选择:
patience: 5,以验证集acc最大化作为最佳模型准则,仅保留 1 个最优 checkpoint(keep_nbest_models: 1),推理默认使用valid.acc.best.pth。 - 特征归一化:
run.sh中通过--feats_normalize uttmvn指定 UtteranceMVN;若改用global_mvn,cls.sh会自动追加--normalize=global_mvn --normalize_conf stats_file=${cls_stats_dir}/train/feats_stats.npz,其中统计量来自 Stage 5 的 collect-stats。 - 时长约束:
run.sh设置--min_wav_duration 0.1、--max_wav_duration 20,超出范围的样本会在 Stage 3 被过滤(该过滤只作用于训练/验证集,测试集保持原始数据)。
五、端到端流水线:run.sh 与十个 Stage
run.sh 是复现入口,其核心调用如下:
./cls.sh \ --cls_tag "${mynametag}" \ --datadir "${storage_dir}/data" \ --dumpdir "${storage_dir}/dump" \ --expdir "${storage_dir}/exp" \ --gpu_inference true \ --feats_normalize uttmvn \ --stage 1 \ --stop_stage 10 \ --nj 10 \ --inference_nj 4 \ --label_fold_length 2 \ --min_wav_duration 0.1 \ --max_wav_duration 20 \ --cls_config "${cls_config}" \ --train_set "${train_set}" \ --valid_set "${valid_set}" \ --test_sets "${test_sets}" "$@"其中train_set="train"、valid_set="valid"、test_sets="test",cls_config=conf/train_cls_wavlm_transformer.yaml,cls_tag默认取当前时间戳。cls.sh(egs2/meld/cls1/cls.sh)将整个流程划分为以下阶段:
| Stage | 内容 | 说明 |
|---|---|---|
| 1 | 数据下载与准备 | 调用local/data.sh,下载 MELD 并生成 Kaldi 数据目录 |
| 2 | 格式化 wav.scp | 将管道式wav.scp落盘为真实音频文件(--audio-format flac --fs 16k),并写feats_type=raw |
| 3 | 长/短数据过滤 | 按min_wav_duration/max_wav_duration过滤训练与验证集 |
| 4 | 生成 token_list | 用espnet2.bin.tokenize_text --token_type word从text_classes构建类别词表,<unk>仅作占位 |
| 5 | 收集统计信息 | espnet2.bin.cls_train --collect_stats true,并行产出 shape 文件并聚合 |
| 6 | 训练 | espnet2.bin.cls_train(或 Lightning 模式),输出到exp/cls_${cls_tag} |
| 7 | 推理 | espnet2.bin.cls_inference,产出score与text(每个 split 一个目录) |
| 8 | 评分 | pyscripts/utils/cls_score.py计算指标,show_cls_result.sh生成RESULTS.md |
| 9 | 打包 | espnet2.bin.pack cls将配置、模型、RESULTS 等打成 zip |
| 10 | 上传 | 上传到 Hugging Face 仓库(需先配置hf_repo与 git-lfs) |
几个值得注意的实现细节:
- collect-stats 与训练可续跑:每个阶段都会在对应目录生成
run.sh,便于从上一个阶段断点续跑(如--stage 4从训练阶段开始)。 - 数据读取类型:
sound类型直接支持 wav/flac;若audio_format带ark后缀则使用kaldi_ark类型。 - 推理并行:
inference_nj 4将 key 文件切分为多份并行推理,随后按 utterance ID 排序拼接输出;--output_all_probabilities true保证评分脚本能拿到完整的 7 维概率向量。 - 多标签限制:
classification_type=multi-label时cls.sh强制要求 Lightning 训练器并校验通过,否则直接报错退出。
六、评估指标与实验结果
6.1 官方评分指标
Stage 8 的评分由 pyscripts/utils/cls_score.py 完成,它读取 ground-truth 标签、预测文本与预测概率,基于 sklearn 计算三个指标:
- mean_acc:平均准确率(对各类别逐一计算 argmax 准确率后的均值);
- mAP:mean Average Precision,逐类计算 AP 后取平均;
- mean_auc:逐类 ROC AUC 的均值。
该脚本输出格式与show_cls_result.sh(scripts/utils/show_cls_result.sh)配合,自动汇总环境信息(python / espnet2 / pytorch 版本、Git hash 与提交时间)并生成 Markdown 报告。
6.2 本食谱结果
README 中记录的实验环境为:python 3.10.14、espnet2 202604、pytorch 2.11.0+cu130,实验标记为cls_20260822.155629。官方评分结果如下:
| Split | mean_acc | mAP | mean_auc | n_labels | n_instances |
|---|---|---|---|---|---|
| cls_test | 50.81 | 28.24 | 70.31 | 7.00 | 2608.00 |
| cls_valid | 48.19 | 30.05 | 69.63 | 7.00 | 1104.00 |
6.3 与既有工作的对比
README 同时给出了与 MELD 原始论文(bcLSTM、DialogueRNN)及 EmoBox 榜单的对比(说明:为对齐既有工作口径,对比表中的指标由模型预测手工复算,而非 Stage 8 自动输出的 mean_acc / mAP / mean_auc)。
与 MELD 原始论文对比(加权 F1):
| 模型 | Weighted F1 |
|---|---|
| 本食谱(WavLM Base+) | 48.15 |
| bcLSTM(audio) | 39.08 |
| DialogueRNN(audio) | 41.79 |
与 EmoBox 榜单对比(WA / UA / Macro F1):
| 模型 | WA | UA | Macro F1 |
|---|---|---|---|
| 本食谱(WavLM Base+) | 50.81 | 28.00 | 27.75 |
| EmoBox WavLM base | 44.71 | 23.44 | 24.25 |
| EmoBox WavLM large | 49.31 | 28.18 | 29.11 |
| EmoBox Whisper large v3 | 51.89 | 31.54 | 32.95 |
可以看到:本食谱在加权准确率(WA 50.81)上已超过多数类基线(48.2%)和 WavLM base 对照,并逼近 Whisper large v3 的水平;但由于 MELD 类别高度不平衡与对齐噪声,UA 与 Macro F1 仍偏低,这正是 README 强调"准确率单独不足以评价"的原因。需要重申:这是首个可工作实现而非调优结果,通过标签平滑、类别加权、多模态融合或更精细的对齐后处理,指标仍有明显提升空间。
七、复现步骤与预训练模型
7.1 本地复现
- 确保已按 ESPnet 安装指南准备好环境(包含
s3prl、torcheval等依赖,S3PRL 可通过tools下的安装脚本启用); - 进入食谱目录并确认
db.sh中MELD=downloads(默认会自动下载); - 执行
./run.sh即可从 Stage 1 跑到 Stage 10。若仅想复现训练与评测,可改用./run.sh --stage 1 --stop_stage 8;跳过上传则加--skip_upload true(默认即跳过)。
7.2 直接使用预训练模型
作者已将模型发布为espnet/meld_cls1_wavlm_base_plus(Hugging Face 仓库)。cls.sh支持--download_model参数:下载后自动将模型文件迁移到exp/目录,并用其执行 Stage 7 推理与 Stage 8 评分,无需本地训练。
八、迁移到其他分类任务的要点
该食谱的 CLS 流水线具备良好通用性,迁移到新语料时重点关注:
- 类别词表:
--text_classes默认取dump/${train_set}/text(第一列是 utterance ID、第二列起是标签),Stage 4 用tokenize_text生成 token_list,类别数即token_list行数减 1; - 分类类型:单标签用
multi-class;多标签/二分类需multi-label且强制 Lightning 训练器; - 前端替换:可在
frontend_conf中更换任意 S3PRL 上游模型(如wavlm_large、hubert_base等),也可改用default前端 +--feats_normalize的经典 FBank 路线; - 编码器扩展:
CLSTask支持将transformer换成conformer或beats编码器,便于对照实验。
九、总结
egs2/meld/cls1是一个结构完整、可复现、可迁移的 ESPnet2 语音情感分类食谱范例:冻结的 WavLM Base+ 提供了强表征,轻量 Transformer 编码器与均值池化线性头将表征映射为 7 类情感概率,十阶段流水线覆盖了从原始视频语料到 Hugging Face 模型发布的全部环节。其代码实现(espnet2/tasks/cls.py、espnet2/cls/espnet_model.py、espnet2/cls/decoder/linear_decoder.py)同时为后续研究提供了清晰的扩展点:无论是调优训练策略、更换预训练前端,还是引入多标签与 mixup 增强,都可以在既有框架内直接进行。
- 人工智能
- 语音
- 音频
- 深度学习
- NLP
【免费下载链接】espnet
End-to-End Speech Processing Toolkit
相关推荐
文本分类与情感分析:Flair分类器深度解析
文本分类与情感分析:Flair分类器深度解析 本文深入解析了Flair框架在文本分类与情感分析任务中的核心架构设计与实现方案。文章系统介绍了Flair的文档分类
NLP深度学习机器学习基于 LangChain4j 的文本分类实战:LLM 情感分析与 Embedding 语义分类
基于 LangChain4j 的文本分类实战:LLM 情感分析与 Embedding 语义分类 LangChain4j 为 Java 开发者提供了一套统一的 L
人工智能AI 应用RAGAI Agent工具调用深度解析deberta-v3-base-zeroshot-v2.0:从模型架构到商用优势
深度解析deberta v3 base zeroshot v2.0:从模型架构到商用优势 deberta v3 base zeroshot v2.0是一款基于D
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考