
SqueezeBERT 源碼解析Transformers 中用分組卷積替代全連接層的高效雙向 Transformer【免費(fèi)下載鏈接】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.項(xiàng)目地址: https://gitcode.com/GitHub_Trending/tra/transformersSqueezeBERT 是 Hugging Face Transformers 中一個(gè)面向高效推理的雙向 Transformer 模型它借鑒計(jì)算機(jī)視覺(jué)中的分組卷積grouped convolutions技術(shù)將 BERT 架構(gòu)中 Q、K、V 投影和 FFN 的全連接層替換為帶分組的 1×1 卷積層。本文基于倉(cāng)庫(kù)中的模型文檔與 PyTorch 實(shí)現(xiàn)源碼完整梳理 SqueezeBERT 的設(shè)計(jì)動(dòng)機(jī)、配置參數(shù)、Tokenizer 與任務(wù)頭結(jié)構(gòu)以及NCW數(shù)據(jù)布局下卷積替代線性層的底層實(shí)現(xiàn)方式幫助讀者理解如何在 Transformers 中構(gòu)建和微調(diào)這類瘦身版 BERT。設(shè)計(jì)動(dòng)機(jī)卷積思想移植到 NLPSqueezeBERT 模型源自論文《SqueezeBERT: What can computer vision teach NLP about efficient neural networks?》由 Forrest N. Iandola 等人提出于 2020 年 6 月 19 日發(fā)布、2020 年 11 月 16 日貢獻(xiàn)到 Hugging Face Transformers。它本質(zhì)上是一個(gè)類似 BERT 的雙向 Transformer與 BERT 架構(gòu)的關(guān)鍵區(qū)別在于SqueezeBERT 在 Q、K、V 層和 FFN 層使用分組卷積grouped convolutions替代全連接fully-connected層。論文摘要中給出的背景值得注意BERT-base 在 Pixel 3 智能手機(jī)上分類一條文本片段需要 1.7 秒而 SqueezeBERT 在保持 GLUE 測(cè)試集競(jìng)爭(zhēng)性精度的前提下在 Pixel 3 上運(yùn)行速度比 BERT-base 快 4.3 倍。作者的觀察是分組卷積等技術(shù)在視覺(jué)網(wǎng)絡(luò)中已帶來(lái)顯著加速但當(dāng)時(shí)許多這類技術(shù)尚未被 NLP 網(wǎng)絡(luò)設(shè)計(jì)者采納。該模型由論文作者之一forresti貢獻(xiàn)到 Transformers相關(guān)實(shí)現(xiàn)位于 src/transformers/models/squeezebert/ 目錄包含配置、建模和分詞三個(gè)模塊文件。從源碼結(jié)構(gòu)看分組卷積替代線性層并非口號(hào)式的改動(dòng)而是貫穿了模型的四個(gè)核心位置自注意力中的 Q、K、V 投影SqueezeBertSelfAttention注意力輸出后的過(guò)渡層post_attentionFFN 的升維層intermediateFFN 的降維層output。每一處的分組數(shù)量都可以獨(dú)立配置見(jiàn)下文q_groups等參數(shù)這正是配置類中存在六個(gè)*_groups字段的原因。使用建議官方 Usage tips模型文檔 docs/source/en/model_doc/squeezebert.md 給出了三條重要的使用建議這里完整繼承并展開(kāi)說(shuō)明右側(cè)填充right paddingSqueezeBERT 使用絕對(duì)位置嵌入absolute position embeddings因此通常建議對(duì)輸入做右側(cè)填充而非左側(cè)填充。從 SqueezeBertEmbeddings.forward 可以看到當(dāng)position_ids未顯式傳入時(shí)模型直接使用預(yù)置的self.position_ids[:, :seq_length]從位置 0 開(kāi)始的連續(xù)位置左填充會(huì)讓真實(shí) token 的位置編碼整體偏移。適合理解、不適合生成SqueezeBERT 類似 BERT依賴掩碼語(yǔ)言建模MLM目標(biāo)訓(xùn)練因此擅長(zhǎng)預(yù)測(cè)被掩碼的 token 和一般 NLU 任務(wù)但不適合文本生成——因果語(yǔ)言建模CLM目標(biāo)訓(xùn)練出的模型在生成任務(wù)上表現(xiàn)更好。倉(cāng)庫(kù)中也沒(méi)有SqueezeBertForCausalLM這一類與其 MLM 定位一致。微調(diào)起點(diǎn)推薦在序列分類任務(wù)上微調(diào)時(shí)文檔建議使用squeezebert/squeezebert-mnli-headlesscheckpoint 作為起點(diǎn)。測(cè)試代碼中也用到了squeezebert/squeezebert-uncased和squeezebert/squeezebert-mnli兩個(gè)官方 checkpoint見(jiàn) tests/models/squeezebert/test_modeling_squeezebert.py。官方文檔同時(shí)提供了以下任務(wù)指南作為下游使用入口已轉(zhuǎn)換為倉(cāng)庫(kù)根目錄相對(duì)路徑文本分類任務(wù)指南Token 分類任務(wù)指南問(wèn)答任務(wù)指南掩碼語(yǔ)言建模任務(wù)指南多選任務(wù)指南SqueezeBertConfig六個(gè)分組參數(shù)是核心配置類定義于 src/transformers/models/squeezebert/configuration_squeezebert.pymodel_type squeezebert并使用了strict裝飾器來(lái)自huggingface_hub.dataclasses意味著配置對(duì)未知字段采取嚴(yán)格處理策略。與 BERT 對(duì)齊的基礎(chǔ)參數(shù)默認(rèn)值與 BERT-base 一致便于從 BERT 生態(tài)直接遷移參數(shù)默認(rèn)值說(shuō)明vocab_size30522詞表大小與 BERT 相同hidden_size768隱層寬度同時(shí)是 Q/K/V 卷積的輸入/輸出通道數(shù)num_hidden_layers12Transformer 層數(shù)num_attention_heads12注意力頭數(shù)要求能整除hidden_size否則SqueezeBertSelfAttention構(gòu)造時(shí)拋出ValueErrorintermediate_size3072FFN 升維后的通道數(shù)hidden_actgelu激活函數(shù)hidden_dropout_prob/attention_probs_dropout_prob0.1 / 0.1兩類 dropout 概率max_position_embeddings512最大位置數(shù)絕對(duì)位置嵌入type_vocab_size2segment type 數(shù)initializer_range0.02權(quán)重初始化標(biāo)準(zhǔn)差layer_norm_eps1e-12LayerNorm 的 epsembedding_size768詞嵌入寬度必須等于hidden_sizetie_word_embeddingsTrueMLM 頭輸出權(quán)重與詞嵌入共享SqueezeBERT 特有的分組參數(shù)這是該配置區(qū)別于 BERT 配置的關(guān)鍵部分參數(shù)默認(rèn)值作用位置q_groups4Q 投影 1×1 卷積的分組數(shù)k_groups4K 投影 1×1 卷積的分組數(shù)v_groups4V 投影 1×1 卷積的分組數(shù)post_attention_groups1注意力后過(guò)渡層第一個(gè) FFN 相關(guān)卷積層的分組數(shù)intermediate_groups4FFN 第二個(gè)卷積層的分組數(shù)output_groups4FFN 第三個(gè)卷積層的分組數(shù)分組數(shù)越大卷積的稀疏化程度越高、參數(shù)量越少nn.Conv1d(cin, cout, kernel_size1, groupsg)的參數(shù)量約為全連接層的1/g。需要注意的是groups必須能整除輸入/輸出通道數(shù)默認(rèn)配置下hidden_size768可被 4 整除、intermediate_size3072亦可被 4 整除因此全部默認(rèn)分組數(shù)都合法。配置類文檔中給出的最小示例from transformers import SqueezeBertConfig, SqueezeBertModel # 初始化 SqueezeBERT 配置 configuration SqueezeBertConfig() # 基于該配置初始化模型隨機(jī)權(quán)重 model SqueezeBertModel(configuration) # 訪問(wèn)模型配置 configuration model.config一個(gè)值得注意的實(shí)現(xiàn)細(xì)節(jié)SqueezeBertEncoder.__init__中有一條硬性斷言modeling_squeezebert.pyassert config.embedding_size config.hidden_size, ( If you want embedding_size ! intermediate hidden_size, please insert a Conv1d layer to adjust the number of channels before the first SqueezeBertModule. )從源碼結(jié)構(gòu)看編碼器在NCW布局下把通道維當(dāng)作特征維處理若嵌入寬度與隱層寬度不一致需要在第一個(gè)SqueezeBertModule前自行插入一個(gè)Conv1d做通道數(shù)調(diào)整——這是自定義配置時(shí)需要遵守的前提。NCW 布局下的卷積前向流程SqueezeBERT 實(shí)現(xiàn)最特別的一點(diǎn)是采用了NCWNbatch、Cchannels、W序列長(zhǎng)度數(shù)據(jù)布局這正是 1×1Conv1d能直接扮演逐位置全連接角色的原因。整體數(shù)據(jù)流如下嵌入層SqueezeBertEmbeddingsL45-L82標(biāo)準(zhǔn) BERT 式三件套——詞嵌入padding_idxpad_token_id、絕對(duì)位置嵌入、token_type 嵌入三者相加后過(guò) LayerNorm 與 Dropout。token_type_ids缺省時(shí)填 0。布局轉(zhuǎn)換SqueezeBertEncoder.forward開(kāi)頭執(zhí)行hidden_states.permute(0, 2, 1)從[B, L, H]轉(zhuǎn)成[B, H, L]讓序列維成為Conv1d的空間維L301-L310。每層SqueezeBertModuleL247-L286由四部分串聯(lián)SqueezeBertSelfAttentionQ/K/V 是三個(gè)nn.Conv1d(cin, cin, kernel_size1, groups...)分別使用q_groups/k_groups/v_groupsConvDropoutLayerNormpost_attentionConv1d → Dropout → 殘差相加 → LayerNorm對(duì)應(yīng) BERT 中 attention 后的輸出投影ConvActivationintermediateConv1d → GELU對(duì)應(yīng) FFN 的升維層ConvDropoutLayerNormoutput對(duì)應(yīng) FFN 的降維層殘差連接到post_attention的輸出而非 attention 輸出。收尾編碼結(jié)束后再permute(0, 2, 1)轉(zhuǎn)回[B, L, H]返回BaseModelOutputSqueezeBertModel.forward再經(jīng)SqueezeBertPooler得到BaseModelOutputWithPooling。兩個(gè)實(shí)現(xiàn)細(xì)節(jié)值得單獨(dú)指出SqueezeBertLayerNormL105-L118是nn.LayerNorm的子類由于NCW布局下歸一化維度是 C 而非最后一維forward中先permute(0, 2, 1)歸一化再換回行為上等價(jià)于標(biāo)準(zhǔn) LayerNorm 作用在hidden_size維。MatMulWrapperL85-L102只是一個(gè)包裝torch.matmul的空模塊源碼注釋明確其目的是讓 FLOPs 計(jì)數(shù)工具能夠統(tǒng)計(jì)到 matmul 的浮點(diǎn)運(yùn)算量——直接調(diào)用torch.matmul通常會(huì)被 FLOPs 計(jì)數(shù)器忽略。這與 SqueezeBERT以計(jì)算量為綱做模型瘦身的定位相呼應(yīng)。注意力前向本身與 BERT 一致attention_score matmul(Q, K^T) / sqrt(head_size)疊加attention_mask由create_bidirectional_mask生成雙向掩碼后 softmax、dropout再與 V 相乘多頭重排通過(guò)transpose_for_scores/transpose_output在[N, C, W]布局上完成。TokenizerBertTokenizer 的別名與模型本體相比SqueezeBERT 的分詞實(shí)現(xiàn)極為輕量。src/transformers/models/squeezebert/tokenization_squeezebert.py 全部實(shí)現(xiàn)只有三行有效代碼from ..bert.tokenization_bert import BertTokenizer # SqueezeBertTokenizer is an alias for BertTokenizer SqueezeBertTokenizer BertTokenizer # SqueezeBertTokenizerFast is an alias for SqueezeBertTokenizer (since BertTokenizer is already a fast tokenizer) SqueezeBertTokenizerFast SqueezeBertTokenizer即SqueezeBertTokenizer直接就是BertTokenizerSqueezeBertTokenizerFast再作為其別名。文檔中列出的get_special_tokens_mask、save_vocabulary等方法都繼承自BertTokenizer。在 Auto 分詞映射中src/transformers/models/auto/tokenization_auto.pysqueezebert在tokenizers庫(kù)可用時(shí)同樣映射到BertTokenizer。這意味著 SqueezeBERT 與 BERT 的 vocab 文件vocab.txt 詞表完全通用squeezebert/squeezebert-uncased使用的就是標(biāo)準(zhǔn) BERT 詞表默認(rèn)vocab_size30522。模型類與任務(wù)頭建模文件共提供六個(gè)模型類__all__見(jiàn) modeling_squeezebert.py 末尾全部繼承自SqueezeBertPreTrainedModelbase_model_prefix transformer即任務(wù)頭內(nèi)部以self.transformer持有SqueezeBertModel模型類任務(wù)頭結(jié)構(gòu)輸出SqueezeBertModel嵌入 12 層編碼器 Pooler取第一個(gè) token 隱狀態(tài) → Linear → TanhBaseModelOutputWithPoolingSqueezeBertForMaskedLMSqueezeBertOnlyMLMHeadDense → GELU → LayerNorm → Linear 到 vocabdecoder 權(quán)重與詞嵌入 tiedMaskedLMOutputSqueezeBertForSequenceClassificationPooler 輸出 → Dropout →Linear(hidden_size, num_labels)SequenceClassifierOutputSqueezeBertForTokenClassification序列輸出 → Dropout →Linear(hidden_size, num_labels)TokenClassifierOutputSqueezeBertForMultipleChoice將[B, choices, L]reshape 為[B*choices, L]后走 Pooler → Linear(1)MultipleChoiceModelOutputSqueezeBertForQuestionAnswering序列輸出 →Linear(hidden_size, num_labels)logits 拆分為 start/end 兩路QuestionAnsweringModelOutput各頭的損失約定與 BERT 一致序列分類支持自動(dòng)推斷problem_typenum_labels 1走回歸 MSE整型標(biāo)簽走交叉熵否則多標(biāo)簽 BCEQA 頭對(duì) start/end loss 取平均并對(duì)越界位置clamp處理。權(quán)重初始化上有兩處 BERT 風(fēng)格的特殊處理_init_weightsL407-L414MLM 預(yù)測(cè)頭的bias置零詞嵌入中對(duì)應(yīng) pad 位置的行也會(huì)走padding_idx語(yǔ)義SqueezeBertForMaskedLM通過(guò)_tied_weights_keys聲明cls.predictions.decoder.weight與transformer.embeddings.word_embeddings.weight共享decoder.bias與獨(dú)立的cls.predictions.bias對(duì)應(yīng)保證resize_token_embeddings時(shí)偏置能隨詞表正確縮放。這些類均已注冊(cè)到 Auto 映射例如 modeling_auto.py 中squeezebert出現(xiàn)在AutoModel、AutoModelForMaskedLM、AutoModelForSequenceClassification、AutoModelForQuestionAnswering、AutoModelForTokenClassification、AutoModelForMultipleChoice等映射表如 L506、L665、L1489、L1579、L1702、L1749因此可以直接用AutoModelForXxx.from_pretrained(squeezebert/...)方式加載。典型使用與測(cè)試驗(yàn)證下面給出與源碼和測(cè)試用例一致的最小可運(yùn)行示例推理/加載權(quán)重來(lái)自官方 checkpointimport torch from transformers import ( AutoTokenizer, SqueezeBertForMaskedLM, SqueezeBertForSequenceClassification, SqueezeBertConfig, ) # 1. MLM 推理與 BERT 詞表通用注意右側(cè)填充 tokenizer AutoTokenizer.from_pretrained(squeezebert/squeezebert-uncased) model SqueezeBertForMaskedLM.from_pretrained(squeezebert/squeezebert-uncased) inputs tokenizer(The capital of France is [MASK]., return_tensorspt) with torch.no_grad(): logits model(**inputs).logits # [1, seq_len, 30522] # 2. 從零定義配置構(gòu)建自定義模型調(diào)整分組數(shù)以壓縮算力 config SqueezeBertConfig( hidden_size768, num_attention_heads12, q_groups4, k_groups4, v_groups4, post_attention_groups1, intermediate_groups4, output_groups4, ) # 隨后可用 SqueezeBertForSequenceClassification(config) 等任務(wù)頭構(gòu)建單元測(cè)試位于 tests/models/squeezebert/test_modeling_squeezebert.py驗(yàn)證要點(diǎn)包括SqueezeBertModelTester用小規(guī)模配置hidden_size32、q_groups2、post_attention_groups2、output_groups1等逐一構(gòu)造六個(gè)模型類并校驗(yàn)各任務(wù) logits 形狀證明分組參數(shù)可任意合法取值L39-L70pipeline_model_mapping聲明了 SqueezeBERT 對(duì)feature-extraction、fill-mask、text-classification、token-classification、zero-shot五條 pipeline 的模型映射L229-L239慢速集成測(cè)試加載squeezebert/squeezebert-mnli分類頭對(duì)固定輸入斷言 logits 為[[0.6401, -0.0349, -0.6041]]容差 1e-4三分類對(duì)應(yīng) MNLI 的矛盾/中立/蘊(yùn)含可作為數(shù)值回歸的參考錨點(diǎn)L284-L295。小結(jié)SqueezeBERT 的本質(zhì)是卷積化的 BERTQ/K/V 投影與 FFN 的四個(gè)全連接位置全部換成kernel_size1的分組Conv1d配合NCW數(shù)據(jù)布局使序列維直接成為卷積的空間維SqueezeBertConfig中六個(gè)*_groups參數(shù)是效率調(diào)節(jié)旋鈕默認(rèn)值4/4/4/1/4/4與 BERT-base 同尺寸12 層、768 隱層、12 頭保持兼容分詞層完全復(fù)用BertTokenizer詞表、填充策略右側(cè)填充與 BERT 生態(tài)通用六個(gè)模型類覆蓋 MLM、序列/Token 分類、多選和 QA 五大理解任務(wù)均已在 Auto 映射中注冊(cè)可用from_pretrained直接加載squeezebert/系列 checkpoint由于采用 MLM 目標(biāo)訓(xùn)練SqueezeBERT 適合文本理解類任務(wù)不建議用于文本生成序列分類微調(diào)推薦以squeezebert/squeezebert-mnli-headless為起點(diǎn)。相關(guān)源碼與文檔索引模型實(shí)現(xiàn) src/transformers/models/squeezebert/modeling_squeezebert.py、配置 src/transformers/models/squeezebert/configuration_squeezebert.py、分詞 src/transformers/models/squeezebert/tokenization_squeezebert.py、原始文檔 docs/source/en/model_doc/squeezebert.md、測(cè)試 tests/models/squeezebert/test_modeling_squeezebert.py?!久赓M(fèi)下載鏈接】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.項(xiàng)目地址: https://gitcode.com/GitHub_Trending/tra/transformers創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考