Transformers SEW-D 深度解析:squeeze 降采样与解耦相对位置注意力的语音预训练模型

📅 发布时间:2026/9/8 22:01:22
Transformers SEW-D 深度解析:squeeze 降采样与解耦相对位置注意力的语音预训练模型
Transformers SEW-D 深度解析squeeze 降采样与解耦相对位置注意力的语音预训练模型【免费下载链接】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 仓库中的 SEW-DSqueezed and Efficient Wav2Vec with Disentangled attention模型文档展开覆盖其论文背景、squeeze_factor 降采样编码器的设计动机、解耦自注意力c2p/p2c 相对位置偏置的实现原理以及SEWDConfig的完整参数说明、SEWDModel/SEWDForCTC/SEWDForSequenceClassification三个模型类的调用方式。读完本文你可以直接基于asapp/sew-d-tiny-100k等预训练检查点完成 ASR 推理、CTC 微调并理解其底层卷积-Transformer 管线。SEW-D 是什么SEW-D 由 Felix Wu、Kwangyoun Kim、Jing Pan、Kyu Han、Kilian Q. Weinberger、Yoav Artzi 在论文《Performance-Efficiency Trade-offs in Unsupervised Pre-training for Speech Recognition》中提出论文于 2021-09-14 发布2021-10-15 由 anton-l 贡献到 Hugging Face Transformers。论文摘要的核心结论是本文研究自动语音识别ASR预训练模型中性能与效率的权衡。作者以 wav2vec 2.0 为基线系统化了若干同时影响模型性能与效率的架构设计。综合全部观察提出 SEWSqueezed and Efficient Wav2vec——在多种训练配置下性能与效率均有显著提升的预训练架构。例如在 LibriSpeech 的 100h-960h 半监督设置下SEW 相比 wav2vec 2.0 获得 1.9 倍推理加速词错率相对下降 13.5%在推理时间相近的情况下SEW 在不同模型规模下词错率降低 25-50%。“Disentangled”解耦指的是 SEW-D 在 SEW 的 squeeze 高效架构之上引入了源自 DeBERTa 风格的解耦相对位置注意力用c2pcontent-to-position和p2cposition-to-content两个独立的注意力偏置项来建模相对位置信息。仓库实现位于 modeling_sew_d.py配置文件位于 configuration_sew_d.py模型类型为sew-d。架构总览从 16kHz 波形到 25Hz 表征从源码结构看SEWDModel的前向管线是对应 modeling_sew_d.py#L1257-L1294input_values (batch, seq_len16kHz 采样点) → SEWDFeatureEncoder13 层 1D 卷积320 倍下采样 → LayerNorm 特征投影可选 dropout → SpecAugment 掩码仅训练 → SEWDEncodersqueeze 降采样 2 倍 位置卷积嵌入 DeBERTa 风格解耦注意力 Transformer → 上采样upsample恢复原始帧率 → last_hidden_state (batch, 卷积输出帧数, hidden_size)特征编码器13 层卷积实现 320 倍下采样SEWDFeatureEncodermodeling_sew_d.py#L397-L432完全沿用 wav2vec 2.0 的卷积堆叠设计默认 13 层 1D 卷积conv_dim(64, 128, 128, 128, 128, 256, 256, 256, 256, 512, 512, 512, 512)conv_stride(5, 2, 1, 2, 1, 2, 1, 2, 1, 2, 1, 2, 1)conv_kernel(10, 3, 1, 3, 1, 3, 1, 3, 1, 2, 1, 2, 1)步长乘积为 320即 16kHz 语音被压缩为每秒 50 帧的表征——这正是SEWDConfig.inputs_to_logits_ratio属性所计算的值configuration_sew_d.py#L191-L193。配置类在validate_architecture中强制校验len(conv_dim) len(conv_stride) len(conv_kernel)三者共同决定卷积层数configuration_sew_d.py#L177-L189。归一化方式由feat_extract_norm控制group默认时第一层用GroupNorm、其余层用无归一化卷积layer时全部使用LayerNorm卷积层分别对应SEWDGroupNormConvLayer和SEWDLayerNormConvLayermodeling_sew_d.py#L396-L432。SEW 的关键创新squeeze_factor 再降 2 倍SEWDEncodermodeling_sew_d.py#L1070-L1125是 SEW/SEW-D 效率的来源前向过程squeeze用nn.AvgPool1d(squeeze_factor, squeeze_factor)把 50Hz 的特征再平均池化 2 倍默认squeeze_factor2Transformer 只需处理 25Hz 的序列注意力计算量近似降至 1/4位置嵌入SEWDPositionalConvEmbeddingmodeling_sew_d.py#L318-L358是一个深度可分 1D 卷积kernelnum_conv_pos_embeddings128groupsnum_conv_pos_embedding_groups16且stridesqueeze_factor——它同时承担位置卷积与再降采样两种角色输出与池化结果相加解耦注意力 TransformerSEWDTransformerEncoder内是num_hidden_layers个SEWDLayer解耦注意力 BERT 风格 FFN恢复帧率SEWDUpsamplingmodeling_sew_d.py#L374-L393先做hidden_size → hidden_size * squeeze_factor的线性投影再通过 reshape 把通道维拆回时间维将 25Hz 表征还原为 50Hz 帧率必要时右侧补零对齐输入帧数。attention_mask也会用max_pool1d同步池化到 squeeze 后的长度modeling_sew_d.py#L1096-L1107源码注释特别说明选择max_pool1d而非arange广播是为了避免 ONNX 导出时把序列长度固化为常量。解耦自注意力c2p 与 p2c 相对位置偏置DisentangledSelfAttentionmodeling_sew_d.py#L630-L837实现了 DeBERTa 风格的相对位置注意力核心逻辑build_relative_position构建 query 到 key 的相对位置矩阵R_{q→k} P_q - P_kmodeling_sew_d.py#L178-L205当position_buckets 0默认 256时相对位置经make_log_bucket_position做对数分桶把连续相对位置映射到有界桶区间远端位置的嵌入表大小固定为position_buckets注意力分数在标准 QK 点积之外叠加两个可选偏置c2pcontent-to-positionQ · PosKey再按相对位置relative_pos att_span收集gather即内容查询位置嵌入p2cposition-to-contentKey · PosQuery按-r_pos att_span收集即位置查询内容缩放因子sqrt(head_dim * scale_factor)会随启用的偏置项数量1、2 或 3动态调整保证方差稳定share_att_keyTrue默认时c2p/p2c 与内容注意力共享 query/key 投影省参数False时可为 c2p、p2c 单独配置pos_key_proj/pos_query_projpos_att_type默认为(p2c, c2p)即两个偏置全开可通过配置关闭其中一项softmax 使用自定义XSoftmax、dropout 使用StableDropout/XDropout均为借鉴 DeBERTa 的内存优化实现用 mask 操作代替逐元素乘法并提供 ONNXsymbolic导出支持。相对位置嵌入表在SEWDTransformerEncoder.__init__中创建大小为position_buckets * 2当norm_rel_ebd含layer_norm默认时先过一层 LayerNorm 再参与计算modeling_sew_d.py#L967-L1001。SpecAugment 数据增强训练时self.training且apply_spec_augmentTrueSEWDModel._mask_hidden_statesmodeling_sew_d.py#L1210-L1255沿时间轴与特征轴对卷积特征做掩码时间轴按mask_time_prob默认 0.05、mask_time_length默认 10生成连续掩码片段被掩位置替换为可学习参数masked_spec_embed特征轴mask_feature_prob默认 0.0即默认不做特征轴掩码_compute_mask_indicesmodeling_sew_d.py#L44-L160按mask_prob * len(axis) / mask_length计算掩码片段数并受mask_time_min_masks默认 2下限约束测试套件 test_modeling_sew_d.py#L410-L433 专门验证了掩码数量与重叠行为的正确性。SEWDConfig完整参数说明SEWDConfig是严格校验strict的配置类model_type sew-d。以下参数说明继承自 configuration_sew_d.py 的 docstring默认值来自类属性定义SEW-D 特有参数参数默认值说明squeeze_factor2编码器之后的序列长度下采样因子、Transformer 之后的上采样因子position_buckets256相对位置嵌入的最大桶数share_att_keyTruec2p 与 p2c 是否共享注意力 key 投影relative_attentionTrue是否使用相对位置编码pos_att_type(p2c, c2p)相对位置注意力类型可组合取(p2c)、(c2p)或两者norm_rel_ebdlayer_norm相对位置嵌入是否先做 LayerNormfeat_proj_dropout0.0特征编码器输出的 dropout 概率final_dropout0.1SEWDForCTC最终投影层的 dropout 概率feature_layer_norm_eps1e-5特征编码器之后 LayerNorm 的 epsilon卷积特征提取参数参数默认值说明feat_extract_normgroup1D 卷积归一化方式group仅第一层 GroupNorm或layer全部 LayerNormfeat_extract_activationgelu卷积激活函数支持gelu、relu、selu、gelu_newconv_dim(64, 128, 128, 128, 128, 256, 256, 256, 256, 512, 512, 512, 512)每层卷积输入/输出通道数长度即卷积层数conv_stride(5, 2, 1, 2, 1, 2, 1, 2, 1, 2, 1, 2, 1)每层步长长度须与conv_dim一致conv_kernel(10, 3, 1, 3, 1, 3, 1, 3, 1, 2, 1, 2, 1)每层核大小长度须与conv_dim一致conv_biasFalse卷积是否带偏置num_conv_pos_embeddings128位置卷积嵌入的核大小num_conv_pos_embedding_groups16位置卷积嵌入的分组数SpecAugment 参数仅当apply_spec_augmentTrue时生效参数默认值说明apply_spec_augmentTrue是否对特征编码器输出做 SpecAugmentmask_time_prob0.05时间轴掩码比例掩码片段数为mask_time_prob*len(time_axis)/mask_time_lengthmask_time_length10时间轴掩码片段长度mask_time_min_masks2时间轴最少掩码片段数当概率计算不足该值时生效mask_feature_prob0.0特征轴掩码比例mask_feature_length10特征轴掩码片段长度mask_feature_min_masks0特征轴最少掩码片段数CTC 与分类头参数参数默认值说明ctc_loss_reductionmeanCTC 损失的 reduction 方式ctc_zero_infinityFalse是否将torch.nn.CTCLoss的无穷损失置零输入过短无法对齐目标时会产生 infuse_weighted_layer_sumFalse是否用可学习权重对各层输出加权平均仅对序列分类头有意义classifier_proj_size256分类头 token 池化前的投影维度vocab_size32CTC 头词表大小默认值很小实际加载预训练权重时会覆盖骨干 Transformer 参数hidden_size768、num_hidden_layers12、num_attention_heads12、intermediate_size3072、max_position_embeddings512、hidden_actgelu_python、layer_norm_eps1e-7、各类 dropout 默认 0.1、initializer_range0.02以及pad_token_id0/bos_token_id1/eos_token_id2。配置示例与配置类 docstring 保持一致from transformers import SEWDConfig, SEWDModel # 初始化一个 asapp/sew-d-tiny-100k 风格的配置默认参数 configuration SEWDConfig() # 用该配置初始化随机权重模型 model SEWDModel(configuration) # 访问模型配置 configuration model.config三个模型类及 forward模块导出见init.py共四个符号SEWDPreTrainedModel、SEWDModel、SEWDForCTC、SEWDForSequenceClassification。SEWDModel基础编码器接收原始波形浮点数组返回BaseModelOutput。核心输入参数modeling_sew_d.py#L1260-L1294input_values(batch_size, sequence_length)的原始语音波形16kHz 采样attention_mask可选右填充掩码会经_get_feature_vector_attention_mask折算为卷积输出帧级的掩码modeling_sew_d.py#L1179-L1185mask_time_indices训练时可选的、预先生成的时间轴掩码位置。若最后一层卷积通道conv_dim[-1]默认 512与hidden_size默认 768不一致会插入feature_projection线性层project_featuresTrue。输出帧长为卷积输出长度(L - 10) // 5 // 2 // ...测试套件中用output_seq_length ceil逐步折叠(L-(k-1))/s验证该形状test_modeling_sew_d.py#L116-L120。SEWDForCTCASR 任务在SEWDModel之上叠加nn.Dropout(final_dropout)与lm_head Linear(hidden_size, vocab_size)用 CTC 损失训练modeling_sew_d.py#L1303-L1432。forward 要点labels(batch_size, target_length)取值-100表示忽略labels.max() vocab_size时抛ValueError训练时从attention_mask推导input_lengths _get_feat_extract_output_lengths(...)目标长度取labels 0的计数再调用nn.functional.ctc_lossblankconfig.pad_token_id、reductionconfig.ctc_loss_reduction、zero_infinityconfig.ctc_zero_infinity并显式禁用 cuDNN 后端CTC 不支持 fp16logits 以 float32 计算提供freeze_feature_encoder()只冻结 13 层卷积modeling_sew_d.py#L1357-L1362与freeze_base_model()冻结整个sew_d主干两个微调常用方法测试check_ctc_training验证了冻结特征编码器 ctc_zero_infinityTrue时反向传播不产生 inftest_modeling_sew_d.py#L197-L224另外该类支持target_lang参数与load_adapter从源码结构看这是为asapp/sew-d-100h-ft这类带语言适配器权重的检查点准备的加载机制tie_weights被重新用于触发适配器加载见 modeling_sew_d.py#L1333-L1355。SEWDForSequenceClassification音频分类任务在SEWDModel之上接projector Linear(hidden_size, classifier_proj_size)与classifier Linear(classifier_proj_size, num_labels)适用于 SUPERB 关键词唤醒等任务modeling_sew_d.py#L1442-L1535。forward 要点池化前先用_get_feature_vector_attention_mask屏蔽填充帧再做带掩码的均值池化无attention_mask时退化为普通均值池化use_weighted_layer_sumTrue时对num_hidden_layers 1个 hidden states 做 softmax 加权求和需output_hidden_statesTrue与 CTC 类相同提供freeze_feature_encoder()/freeze_base_model()config.num_labels 1时计算交叉熵损失num_labels 1时为回归MSE。使用要点官方文档给出的 Usage tips继承自 sew-d.mdSEW-D 是语音模型接受对应原始语音波形的 float 数组作为输入SEWDForCTC使用连接时序分类CTC微调因此模型输出必须用 CTC 解码器解码——即Wav2Vec2CTCTokenizerSEW-D 没有独立的 tokenizer 类复用 wav2vec2 的 CTC 解码器做贪心解码、去 blank、合并重复 token。与仓库集成测试一致的端到端推理流程test_modeling_sew_d.py#L505-L525import torch from transformers import SEWDForCTC, Wav2Vec2Processor model SEWDForCTC.from_pretrained(asapp/sew-d-tiny-100k-ft-ls100h) # do_lower_case 与集成测试保持一致 processor Wav2Vec2Processor.from_pretrained( asapp/sew-d-tiny-100k-ft-ls100h, do_lower_caseTrue ) # input_speech: list[np.ndarray]16kHz 原始波形 inputs processor(input_speech, return_tensorspt, paddingTrue) input_values inputs.input_values.to(model.device) with torch.no_grad(): logits model(input_values).logits predicted_ids torch.argmax(logits, dim-1) predicted_trans processor.batch_decode(predicted_ids)processor内部即Wav2Vec2FeatureExtractor负责 16kHz 重采样、padding、张量化加Wav2Vec2CTCTokenizer负责 CTC 贪心解码的组合仅取特征时也可单独用Wav2Vec2FeatureExtractor见 test_modeling_sew_d.py#L451-L503 的基础编码测试该测试还硬编码了与原始 SEW-D 实现数值对齐的期望输出。微调时的实践建议均有源码/测试依据数据量有限时先model.freeze_feature_encoder()只训 Transformer 与 CTC 头设置config.ctc_zero_infinity True避免极短音频样本产生 inf 损失通过config.ctc_loss_reduction在mean默认与sum之间切换两者损失行为由check_ctc_loss测试验证需要 ONNX 导出时源码中max_pool1d、广播比较构建特征掩码等写法专门规避了序列长度被固化为常量的问题可放心使用动态形状。检查点转换与测试官方 fairseq 检查点转换脚本convert_sew_d_original_pytorch_checkpoint_to_pytorch.py通过MAPPING表把 ASAPP 原版权重如post_extract_proj→feature_projection、encoder.pos_conv.0→encoder.pos_conv_embed.conv、w2v_encoder.proj→lm_head逐一映射到 HF 结构并生成对应的Wav2Vec2FeatureExtractor/Wav2Vec2CTCTokenizer/Wav2Vec2Processor文件单元测试与集成测试tests/models/sew_d/test_modeling_sew_d.py涵盖形状检查、CTC 损失/训练、分类损失、越界 label 报错、asapp/sew-d-tiny-100k数值对齐以及 LibriSpeech 转写集成测试相关任务文档音频分类任务指南、自动语音识别任务指南。小结SEW-D 用两处结构性改动换取了显著的效率提升卷积编码后再以squeeze_factor配合AvgPool1d与 stride 位置卷积把序列再压缩一半让 Transformer 处理更短的 25Hz 序列注意力端引入可开关pos_att_type、可对数分桶position_buckets、可共享投影share_att_key的解耦相对位置偏置兼顾位置感知与参数效率。三者均为配置开关可在 configuration_sew_d.py 中按需调整配合SEWDForCTC的 CTC 微调流程即可落地 ASR 场景。【免费下载链接】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),仅供参考