WavLM 全栈语音预训练模型解析与 Transformers 实战指南
【免费下载链接】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
导读:本文围绕 Hugging Face Transformers 中 WavLM 模型的官方文档(见 docs/source/en/model_doc/wavlm.md),系统讲解 WavLM 的模型背景、在 Transformers 中的源码架构、核心配置参数,以及如何基于原始波形完成语音识别(CTC)、音色/说话人验证、说话人日志与音频帧分类等下游任务。读完本文,你将掌握 WavLMConfig 各字段的作用、五大模型类(WavLMModel / WavLMForCTC / WavLMForSequenceClassification / WavLMForAudioFrameClassification / WavLMForXVector)的选择依据,以及从特征提取到微调推理的完整代码路径。
一、WavLM 是什么:面向"全栈语音处理"的自监督预训练模型
WavLM 由微软研究院 Sanyuan Chen、Chengyi Wang 等人提出,论文题为WavLM: Large-Scale Self-Supervised Pre-Training for Full Stack Speech Processing。该模型于 2021-10-26 在 Hugging Face Papers 发布,并于 2021-12-16 合入 Transformers 仓库。其设计目标是:用一个统一的预训练模型去覆盖语音内容建模(spoken content modeling)与说话人身份保持(speaker identity preservation)这两类相互竞争的需求。
论文摘要指出了几个关键技术点,它们也是理解 WavLM 架构的钥匙:
- 基于 HuBERT 框架构建:WavLM 沿用 HuBERT 的自监督学习范式,通过预测掩码语音片段的隐藏单元来学习通用语音表征。
- 引入门控相对位置偏置(gated relative position bias):在 Transformer 结构中增强对语音内容/识别类任务的长程建模能力。
- 话语混合训练策略(utterance mixing training strategy):无监督地构造"叠加说话人"的语音样本参与训练,从而在模型内部同时建模说话人身份,显著提升说话人区分能力。
- 数据集从 60k 小时扩充到 94k 小时:大规模数据(包含真实世界噪声干扰的音频)支撑了更强的泛化能力。
在仓库文档中,官方对 WavLM 的能力定位非常明确:它在说话人验证(speaker verification)、说话人识别(speaker identification)与说话人日志/分割(speaker diarization)任务上表现尤为突出;同时它也天然支持语音识别等内容类任务——这正是"full-stack speech processing"的含义。
二、仓库中的实现概览:从模块化源码看架构
WavLM 在仓库内的实现主要位于 src/transformers/models/wavlm 目录,包含:
configuration_wavlm.py:定义WavLMConfig,model_type = "wavlm";modeling_wavlm.py:定义全部模型类与网络组件(1559 行);modular_wavlm.py:模块化生成源文件(Transformers 新版引入的 modular 规范,modeling_wavlm.py由它自动生成);convert_wavlm_original_pytorch_checkpoint_to_pytorch.py:将微软官方 fairseq 权重转换为 Hugging Face 格式;convert_wavlm_original_s3prl_checkpoint_to_pytorch.py:转换 s3prl 下游任务(分类/日志/xvector)checkpoint;- 测试文件位于 tests/models/wavlm/test_modeling_wavlm.py。
一个非常重要的源码事实:在 modular_wavlm.py 中可以看到,WavLM 的WavLMPositionalConvEmbedding、WavLMFeatureProjection、WavLMFeedForward以及多个任务头(CTC / SequenceClassification / AudioFrameClassification / XVector)直接继承自同目录的 wav2vec2 对应类,而 WavLMAttention 被独立重写为带相对位置偏置的注意力。也就是说:WavLM ≈ Wav2Vec2/HuBERT 骨干 + 门控相对位置偏置注意力 + 面向说话人任务的输出头,这与论文中"基于 HuBERT 框架构建"的描述完全一致。因此在使用层面,WavLM 与 Wav2Vec2 家族共享大量 API 习惯。
从modeling_wavlm.py的组件清单看,其网络结构自上而下为:
WavLMFeatureEncoder:7 层 1D 卷积构成的 CNN 特征编码器(配合 GroupNorm/LayerNorm 三种卷积变体);WavLMFeatureProjection:将卷积特征投影到hidden_size维度;WavLMPositionalConvEmbedding:卷积式位置编码;WavLMEncoder / WavLMEncoderStableLayerNorm:Transformer 编码器,其中注意力使用WavLMAttention(门控相对位置偏置);WavLMAdapter:可选的下采样适配网络(add_adapter=True时启用);- 面向预训练的
WavLMGumbelVectorQuantizer(Gumbel 量化码本)保留自 Wav2Vec2/HuBERT 的自监督预训练所需模块。
数据流(见 modeling_wavlm.py 中WavLMModel.forward)为:input_values(原始波形)→ 特征编码器 → 特征投影 → SpecAugment 掩码 → Transformer 编码器 → 输出last_hidden_state与extract_features。
三、快速上手:用法要点(Usage tips)
官方文档给出了三条使用要点,是避免踩坑的关键:
输入是原始波形的 float 数组。WavLM 不接受频谱图或 token,只接受语音信号的原始波形(1D float 数组)。特征提取请使用 [
Wav2Vec2Processor](WavLM 没有独立 processor,其音频预处理能力由 Wav2Vec2 处理器提供,仓库源码中 WavLM 的文档亦明确指出 "Please useWav2Vec2Processorfor the feature extraction")。在现代 API 中,通常直接使用AutoProcessor.from_pretrained(...)获取对应的 Wav2Vec2 processor/feature extractor。CTC 微调与解码约定:WavLM 可用连接时序分类(connectionist temporal classification, CTC)做语音识别微调,此时模型输出必须用 [
Wav2Vec2CTCTokenizer] 解码。CTC 的 blank 索引、损失规约等由WavLMConfig中的pad_token_id、ctc_loss_reduction、ctc_zero_infinity控制。强项任务:说话人验证、说话人识别、说话人日志(分割)任务建议优先使用
WavLMForXVector与WavLMForAudioFrameClassification;通用表示/特征抽取则使用WavLMModel。
一个最小化的配置-建模示例(摘自WavLMConfig的 docstring 示例):
from transformers import WavLMConfig, WavLMModel # 初始化一个 facebook/wavlm-base-960h 风格配置 configuration = WavLMConfig() # 用随机权重初始化模型 model = WavLMModel(configuration) # 访问模型配置 configuration = model.config四、WavLMConfig 配置参数全解析
WavLMConfig定义于 configuration_wavlm.py,默认值即microsoft/wavlm-base(见该文件@auto_docstring(checkpoint="microsoft/wavlm-base"))。以下按功能分组梳理其核心参数:
4.1 Transformer 主干参数
| 参数 | 默认值 | 含义 |
|---|---|---|
hidden_size | 768 | 编码器隐层维度 |
num_hidden_layers | 12 | Transformer 层数 |
num_attention_heads | 12 | 注意力头数 |
intermediate_size | 3072 | FFN 中间层维度 |
hidden_act | "gelu" | 激活函数 |
hidden_dropout/activation_dropout/attention_dropout | 0.1 | 三类 dropout |
layerdrop | 0.1 | LayerDrop 概率 |
initializer_range | 0.02 | 参数初始化范围 |
layer_norm_eps | 1e-5 | LayerNorm epsilon |
do_stable_layer_norm | False | 为True时在注意力前做 LayerNorm;否则注意力后做 LayerNorm |
此外还有 token 相关字段:vocab_size=32(CTC 头词表)、pad_token_id=0、bos_token_id=1、eos_token_id=2。
4.2 CNN 特征编码器与位置编码参数
特征编码器由多层 1D 卷积构成,其层数由len(conv_dim)决定。默认 7 层:
conv_dim=(512×7):每层输入/输出通道数;conv_stride=(5, 2, 2, 2, 2, 2, 2):每层步长,总下采样率 = 5×2⁶ = 320;conv_kernel=(10, 3, 3, 3, 3, 2, 2):每层卷积核;conv_bias=False:卷积是否带偏置;feat_extract_norm="group":特征编码器归一化方式,"group"表示仅第一层卷积使用 GroupNorm,"layer"表示对所有卷积层使用 LayerNorm;feat_extract_activation="gelu":卷积层激活函数(支持"gelu"、"relu"、"selu"、"gelu_new");feat_proj_dropout=0.0:特征投影输出 dropout。
在__post_init__中会依据conv_dim自动推导num_feat_extract_layers,并在validate_architecture中强制校验conv_stride、conv_kernel、conv_dim三者长度一致(configuration_wavlm.py)。inputs_to_logits_ratio属性返回conv_stride的乘积(默认 320),即波形采样点与模型输出帧之间的下采样比例——计算 CTC 输入长度时正是用它换算。
卷积位置编码相关:
num_conv_pos_embeddings=128:卷积位置编码核大小(即WavLMPositionalConvEmbedding中nn.Conv1d的 kernel,配合 weight-norm 使用);num_conv_pos_embedding_groups=16:该卷积的分组数。
4.3 WavLM 核心差异化:门控相对位置偏置与掩码(SpecAugment)
这是 WavLM 区别于 Wav2Vec2/HuBERT 的关键。在WavLMAttention(modeling_wavlm.py)中:
num_buckets=320:相对位置分桶数,决定rel_attn_embed(nn.Embedding(num_buckets, num_heads))的大小;max_bucket_distance=800:相对位置距离上限;- 每个注意力头还带有一组可学习门控参数:
gru_rel_pos_const(形状(1, heads, 1, 1))与gru_rel_pos_linear(将head_dim投影到 8 维),实现对相对位置偏置的门控调制。compute_bias会把相对位置先"分桶"再映射为嵌入,进而融合成每头的位置偏置。
自监督预训练相关的 SpecAugment 掩码参数(官方文档明确指出参考SpecAugment: A Simple Data Augmentation Method for Automatic Speech Recognition,论文号 1904.08779):
apply_spec_augment=True:是否在特征编码器输出上做 SpecAugment;mask_time_prob=0.05:沿时间轴每个特征向量作为掩码起点的概率,实际掩码数量约mask_time_prob × sequence_length // mask_time_length;mask_time_length=10:时间轴掩码跨度;mask_time_min_masks=2:时间轴最少掩码段数(当按概率算出的掩码数过少时兜底);mask_feature_prob=0.0:沿特征维掩码概率(默认为 0,即默认只做时间轴掩码);mask_feature_length=10:特征维掩码跨度。
量化与对比学习(预训练阶段使用,直接调用WavLMModel做推理时无影响):
num_codevectors_per_group=320、num_codevector_groups=2:乘积量化码本配置;codevector_dim=256:量化向量维度;proj_codevector_dim=256:量化特征与 Transformer 特征统一投影后的维度;contrastive_logits_temperature=0.1:对比损失温度 κ;num_negatives=100:负样本数;diversity_loss_weight=0.1:码本多样性损失权重。
4.4 任务头相关参数
final_dropout=0.1:CTC 头前 dropout;ctc_loss_reduction="mean"、ctc_zero_infinity=False:CTC 损失规约方式,以及是否将无穷损失/梯度置零(输入过短无法对齐标签时易出现无穷损失,仅WavLMForCTC训练相关);use_weighted_layer_sum=False:是否使用可学习权重对各层输出做加权求和(分类类任务头可选);classifier_proj_size=256:序列分类投影维度;- TDNN 模块(XVector 头):
tdnn_dim=(512,512,512,512,1500)、tdnn_kernel=(5,3,3,1,1)、tdnn_dilation=(1,2,3,1,1); xvector_output_dim=512:XVector 嵌入维度;num_ctc_classes=80:音素级 CTC 类别数(文档注明主要用于 UniSpeechForPreTraining 场景,WavLM 中保留);- 适配器:
add_adapter=False、adapter_kernel_size=3、adapter_stride=2、num_adapter_layers=3、output_hidden_size=None(为 None 时默认等于hidden_size)。开启后可叠加小型卷积网络用于 SpeechEncoderDecoder 场景的热启动。
五、五大模型类:选型与实战
modeling_wavlm.py导出的类(见__all__)为WavLMModel、WavLMForCTC、WavLMForSequenceClassification、WavLMForAudioFrameClassification、WavLMForXVector。
5.1 WavLMModel —— 基础骨干模型
仅包含特征编码器、特征投影、位置编码与 Transformer 编码器。前向参数为input_values(原始波形)与可选的attention_mask、mask_time_indices。返回WavLMBaseModelOutput(字段last_hidden_state与extract_features)。它是最底层的表示学习模型,适合自行接自定义任务头或做特征抽取。其前向还会执行 SpecAugment 掩码(仅在训练状态且mask_time_prob>0时生效)。
5.2 WavLMForCTC —— 语音识别(ASR)
结构为WavLMModel + dropout(final_dropout) + Linear(output_hidden_size, vocab_size)。注意:
- 若
config.vocab_size未定义会直接抛错,并提示用WavLMForCTC.from_pretrained(..., vocab_size=vocab_size)实例化; - 支持语言适配器:传入
target_lang="eng"等参数可加载adapter.<lang>权重(要求配置中定义adapter_attn_dim),其加载逻辑被复用在tie_weights中(见 modeling_wavlm.py); - 计算 CTC 损失时,
input_lengths由_get_feat_extract_output_lengths(attention_mask.sum(-1))得出(即用 320 倍下采样换算),labels中-100位置被忽略,blank使用config.pad_token_id,并显式用 fp32 计算以避免 fp16 下 CTC 数值问题。
预训练解码示例:
import torch from transformers import AutoProcessor, WavLMForCTC import soundfile as sf # 加载音频为 16kHz 原始波形 speech, sr = sf.read("/path/to/sample.wav") assert sr == 16000, "WavLM 期望 16kHz 采样率" processor = AutoProcessor.from_pretrained("facebook/wavlm-base-960h") model = WavLMForCTC.from_pretrained("facebook/wavlm-base-960h") inputs = processor(speech, sampling_rate=sr, return_tensors="pt") with torch.no_grad(): logits = model(**inputs).logits # 用 Wav2Vec2CTCTokenizer 解码输出(CTC 贪心/beam 解码) predicted_ids = torch.argmax(logits, dim=-1) transcription = processor.batch_decode(predicted_ids) print(transcription)微调时建议先冻结特征编码器(见下文第六节),并通过model.freeze_base_model()冻结骨干只训练头部。
5.3 WavLMForSequenceClassification —— 序列级分类(如 SUPERB 关键词唤醒 Keyword Spotting)
结构:WavLMModel + projector(Linear(hidden_size, classifier_proj_size)) + classifier(Linear(classifier_proj_size, num_labels))。前向支持两种池化路径(modeling_wavlm.py):
- 当
use_weighted_layer_sum=True:收集所有 Transformer 层(含输入嵌入共num_hidden_layers+1层)的输出,用 softmax 归一化的可学习权重加权求和; - 否则直接使用最后一层隐藏状态。
随后做带掩码的 token 均值池化(把 padding 位置清零后求均值),再送入分类器。labels为标量序列分类标签:当num_labels == 1时算 MSE 回归损失,num_labels > 1时算 CrossEntropy。注意:该任务头不支持 WavLM 适配器(add_adapter=True会抛错)。
5.4 WavLMForAudioFrameClassification —— 帧级音频分类(如说话人日志/事件检测)
结构类似:WavLMModel + (可选加权层求和) + Linear(hidden_size, num_labels),直接对每个时间帧输出标签 logits,不做时间池化。典型用途是逐帧预测"谁在说话"的说话人日志(diarization),labels形状为帧级 one-hot((batch, num_frames, num_labels)),损失内部取 argmax 计算 CrossEntropy。同样不支持适配器。
5.5 WavLMForXVector —— 说话人嵌入(说话人验证/识别)
这是 WavLM 最具特色的任务头(modeling_wavlm.py),复刻自说话人识别领域经典的X-Vector架构:
WavLMModel提取帧级特征;projector线性投影到tdnn_dim[0](默认 512);- 串联 5 层TDNN(时延神经网络),其核以
nn.Linear存储、前向用F.conv1d加速计算,并带膨胀因子tdnn_dilation; - 统计池化(statistic pooling):对 TDNN 输出(按
attention_mask折算后的有效帧)计算均值与标准差并拼接; feature_extractor把拼接后的统计量投影为xvector_output_dim=512维的说话人嵌入(embeddings);classifier得到 logits,训练时用AMSoftmax 损失(scale=30.0、margin=0.4,见AMSoftmaxLoss)驱动,损失函数内对嵌入与类中心权重做 L2 归一化并施加 margin。
说话人验证的典型流程:
import torch from transformers import AutoProcessor, WavLMForXVector import soundfile as sf processor = AutoProcessor.from_pretrained("microsoft/wavlm-base-plus-sv") model = WavLMForXVector.from_pretrained("microsoft/wavlm-base-plus-sv") def extract_xvector(path): speech, sr = sf.read(path) inputs = processor(speech, sampling_rate=sr, return_tensors="pt") with torch.no_grad(): emb = model(**inputs).embeddings[0] # (xvector_output_dim,) return torch.nn.functional.normalize(emb, dim=-1) # 两段音频余弦相似度 > 阈值 → 同一说话人 cos_sim = torch.matmul(extract_xvector("a.wav"), extract_xvector("b.wav")) print(float(cos_sim))六、微调实操要点:冻结特征编码器
无论哪个任务头,文档与源码都强调一个小技巧:微调时先冻结 CNN 特征编码器。WavLMModel及各任务头统一暴露了两个方法:
freeze_feature_encoder():调用底层feature_extractor._freeze_parameters(),关闭特征编码器参数的梯度(CNN 编码器为低层通用声学特征,冻结后更快且不易过拟合);freeze_base_model():冻结整个wavlm骨干,只保留头部可训练——在标注数据有限的 SUPERB 类任务中非常实用。
典型微调骨架:
from transformers import WavLMForSequenceClassification, WavLMConfig model = WavLMForSequenceClassification.from_pretrained( "microsoft/wavlm-base", num_labels=12, # 例如 12 类关键词 ) model.freeze_feature_encoder() # 冻结 CNN 特征编码器 # 训练时,配合 transformers.Trainer 或自写训练循环七、输入预处理与"帧长度"换算
由于 CNN 特征编码器将 16kHz 波形以 320 倍下采样(inputs_to_logits_ratio),模型输出的"帧序列长度"远小于输入采样点长度。WavLMModel内部通过_get_feat_extract_output_lengths(modeling_wavlm.py 中的_conv_out_length递推公式)与_get_feature_vector_attention_mask将原始attention_mask折算到特征帧粒度——这正是 5.2/5.5 节中各种input_lengths计算的依据,也是构造帧级标签(如音频帧分类的逐帧 one-hot)时必须对齐的尺度。
预处理侧务必注意:WavLM 期望16kHz 单声道波形;用soundfile、librosa或torchaudio读取后,统一走AutoProcessor/Wav2Vec2Processor完成重采样对齐、padding 与张量化,不要手动做短时傅里叶变换或加窗,处理器会负责把波形转成input_values和attention_mask。
八、官方资源与进一步阅读
官方文档将 WavLM 指向两条任务指南(链接已转为仓库根目录相对路径):
- Audio classification task guide
- Automatic speech recognition task guide
想深入源码的读者,建议按此顺序阅读:
- configuration_wavlm.py:全部默认超参与校验逻辑;
- modeling_wavlm.py:重点看
WavLMAttention.compute_bias的相对位置分桶与门控参数、WavLMFeatureEncoder的卷积堆叠、以及各任务头forward的池化/损失细节; - modular_wavlm.py:对比它与 Wav2Vec2 组件类的继承关系,理解"WavLM = Wav2Vec2 骨干 + 门控相对位置偏置";
convert_wavlm_original_pytorch_checkpoint_to_pytorch.py与convert_wavlm_original_s3prl_checkpoint_to_pytorch.py:官方权重与 s3prl 下游权重向 Transformers 格式的迁移逻辑;- tests/models/wavlm/test_modeling_wavlm.py:模型正确性、输出维度与集成测试,是理解 API 契约最直接的样例。
九、小结
WavLM 通过"基于 HuBERT 的通用语音框架 + 门控相对位置偏置注意力 + utterance mixing 训练策略 + 94k 小时大规模数据",把内容识别与说话人建模统一进一个模型中。在 Transformers 仓库中,它围绕WavLMConfig与五个模型类提供完整链路:WavLMModel负责通用表示,WavLMForCTC负责语音识别,WavLMForSequenceClassification负责关键词/意图类序列分类,WavLMForAudioFrameClassification负责帧级标签(如说话人日志),WavLMForXVector则面向说话人验证/识别输出统计池化 + AMSoftmax 训练出的说话人嵌入。无论是加载官方预训练权重直接推理,还是冻结特征编码器后在自有数据集上微调,都可以按本文第三节到第六节的路径快速落地。
【免费下载链接】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),仅供参考