news 2026/9/25 2:59:56

ESPnet 情感分类食谱深度解析:基于 MELD 数据集与 WavLM Base+ 冻结前端的 Transformer 分类实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
ESPnet 情感分类食谱深度解析:基于 MELD 数据集与 WavLM Base+ 冻结前端的 Transformer 分类实践
  • 人工智能
  • 语音
  • 音频
  • 深度学习
  • NLP

【免费下载链接】espnet

End-to-End Speech Processing Toolkit

项目地址:https://gitcode.com/gh_mirrors/es/espnet
点击查看免费下载

本篇技术指南围绕 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)多分类。

该数据集有两个被社区公认的难点,直接决定了评估方式:

  1. 高度类别不平衡:训练集中neutral占比高达 47%,因此在测试集上"多数类基线"(全部预测为 neutral)的准确率即为 48.2%。这意味着只看准确率(plain accuracy)不足以反映模型真实能力,必须结合加权 F1、mAP、AUC 等多类指标评估。
  2. 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_size128编码器输出维度
attention_heads4多头注意力头数
linear_units1024FFN 中间层维度
num_blocks4Transformer block 数量
dropout_rate0.4dropout 比率
input_layerlinear输入投影层类型

从 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。官方评分结果如下:

Splitmean_accmAPmean_aucn_labelsn_instances
cls_test50.8128.2470.317.002608.00
cls_valid48.1930.0569.637.001104.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):

模型WAUAMacro F1
本食谱(WavLM Base+)50.8128.0027.75
EmoBox WavLM base44.7123.4424.25
EmoBox WavLM large49.3128.1829.11
EmoBox Whisper large v351.8931.5432.95

可以看到:本食谱在加权准确率(WA 50.81)上已超过多数类基线(48.2%)和 WavLM base 对照,并逼近 Whisper large v3 的水平;但由于 MELD 类别高度不平衡与对齐噪声,UA 与 Macro F1 仍偏低,这正是 README 强调"准确率单独不足以评价"的原因。需要重申:这是首个可工作实现而非调优结果,通过标签平滑、类别加权、多模态融合或更精细的对齐后处理,指标仍有明显提升空间。

七、复现步骤与预训练模型

7.1 本地复现

  1. 确保已按 ESPnet 安装指南准备好环境(包含s3prl、torcheval等依赖,S3PRL 可通过tools下的安装脚本启用);
  2. 进入食谱目录并确认db.sh中MELD=downloads(默认会自动下载);
  3. 执行./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

项目地址:https://gitcode.com/gh_mirrors/es/espnet
点击查看免费下载
上一篇:不用先上传再下载:FilePizza 用一条链接让两个浏览器直连传文件
下一篇:RIOT 的 avr-rss2 板卡移植指南:Atmega256RFR2(Radio Sensors)构建、烧录与板级配置解析

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 2:58:28

Multisim 14.0安装教程:环境准备、授权激活与常见报错排查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华