訓(xùn)練全流程)
簡介面向NLP學(xué)習(xí)者與PyTorch使用者這是Google AI 2018年BERT模型的PyTorch實現(xiàn)以帶注釋的簡潔代碼呈現(xiàn)Transformer雙向編碼器的預(yù)訓(xùn)練思路可幫助理解語言模型遷移到下游任務(wù)的原理。包內(nèi)共33個文件27個Python腳本構(gòu)成核心實現(xiàn)覆蓋模型架構(gòu)、數(shù)據(jù)集處理、訓(xùn)練器與入口模塊同時包含Makefile、requirements配置、說明文檔等壓縮包僅28KB輕量且便于逐行研讀。發(fā)布至今已有878人學(xué)習(xí)下載項目代碼保持原始結(jié)構(gòu)配套注釋降低了上手門檻適合想深入BERT原理、動手實踐預(yù)訓(xùn)練流程的中級NLP開發(fā)者能從中獲得從語料準(zhǔn)備到模型訓(xùn)練的可運行參考。1. 項目概述與核心價值1.1 這個項目到底解決什么問題先交代一下背景。BERT是Google AI在2018年提出的預(yù)訓(xùn)練語言模型全稱是Bidirectional Encoder Representations from Transformers中文常譯作“基于Transformer的雙向編碼器表示”。當(dāng)年它一出直接刷新了11項NLP任務(wù)的SOTA成績幾乎成了自然語言處理領(lǐng)域的“分水嶺”式成果。不過原始官方實現(xiàn)用的是TensorFlow而PyTorch生態(tài)在學(xué)術(shù)界和工業(yè)界的普及度同樣極高。于是就有了BERT-pytorch這個項目——把Google官方BERT模型用PyTorch框架重新實現(xiàn)一遍讓PyTorch用戶也能跑BERT、微調(diào)BERT、甚至從零預(yù)訓(xùn)練BERT。我最初看到這個項目標(biāo)題時第一反應(yīng)是“這不就是換個框架復(fù)刻一遍嗎”但真正深入去讀代碼、跑實驗之后才意識到事情沒這么簡單。BERT-pytorch不是一個簡單的“翻譯”項目它把BERT的完整結(jié)構(gòu)拆解成了清晰可讀的PyTorch模塊包括Tokenizer、Embedding、Transformer Encoder層、預(yù)訓(xùn)練任務(wù)NSP MLM等每一步都給了簡潔的實現(xiàn)。對想搞懂BERT內(nèi)部原理的人來說這個項目的代碼比官方TensorFlow實現(xiàn)好讀太多了。1.2 適合誰讀、能帶來什么收獲這個項目的目標(biāo)讀者非常明確一類是剛?cè)腴TNLP、想搞懂Transformer和BERT內(nèi)部機制的開發(fā)者另一類是手里只有PyTorch環(huán)境、想把BERT用起來的研究生或工程師。前者可以從代碼里學(xué)到BERT的核心組件怎么拼裝后者可以直接基于這份實現(xiàn)做下游任務(wù)微調(diào)。我個人的體感是如果你已經(jīng)把transformers庫用得滾瓜爛熟再來看這個項目會有一種“原來封裝底下的東西長這樣”的通透感。因為transformers庫為了兼容所有BERT變體把代碼抽象得比較深有時候反而不容易看清主脈絡(luò)。BERT-pytorch則把注意力機制、多頭自注意力、LayerNorm、位置編碼、預(yù)訓(xùn)練損失函數(shù)全部攤開在面前每一行都能對著論文找到來源非常適合作為“BERT原理閱讀的配套代碼”。這篇文章我就基于自己踩坑、調(diào)試、二次開發(fā)的經(jīng)驗把這個項目從環(huán)境搭建、代碼結(jié)構(gòu)、關(guān)鍵實現(xiàn)細(xì)節(jié)到常見坑位全部捋一遍希望能幫你少走彎路。2. 環(huán)境準(zhǔn)備與PyTorch版本選型2.1 推薦的環(huán)境組合與安裝步驟先把環(huán)境搞定。BERT-pytorch的代碼不算新對PyTorch版本沒有特別苛刻的要求但我實測下來不同版本組合的表現(xiàn)差異還是有的。先說結(jié)論我最后穩(wěn)定跑通的組合是Python 3.9或3.10PyTorch 2.1.xCPU或CUDA版均可torchtext注意版本這個項目里做WordPiece Tokenizer時依賴了torchtext的BasicEnglish等工具但新版torchtext API變化很大CUDA 11.8或12.1如果要用GPU安裝PyTorch時我建議直接用官方命令生成器。比如要裝CUDA 12.1版本pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121如果是CPU版本直接pip install torch torchvision torchaudio這里有個容易踩的坑如果你用Anaconda管理環(huán)境千萬別在base環(huán)境里硬裝一定要先建一個干凈的環(huán)境conda create -n bert_env python3.9 -y conda activate bert_env我一開始圖省事直接往base環(huán)境里塞了一堆包結(jié)果torchtext和torch版本沖突報了一堆奇奇怪怪的錯。新建環(huán)境之后世界清凈了。2.2 關(guān)于CUDA、顯卡驅(qū)動和torchtext的版本適配問題GPU版本的PyTorch安裝最容易出問題的就是“驅(qū)動版本和CUDA版本不匹配”以及“PyTorch要求的CUDA版本和本機CUDA版本不一致”。這里分享一個核心認(rèn)知PyTorch的CUDA版本并不需要和你系統(tǒng)里安裝的CUDA Toolkit完全一致PyTorch是自帶CUDA runtime的。你只需要保證顯卡驅(qū)動夠新就行。比如你想裝cu121的PyTorch但系統(tǒng)里CUDA Toolkit還是11.8這完全沒問題只要驅(qū)動版本支持CUDA 12.1就行。怎么查驅(qū)動支持的最高CUDA版本在命令行敲nvidia-smi右上角會顯示“CUDA Version: 12.1”之類的這代表你的驅(qū)動最高支持到12.1只要PyTorch的CUDA版本號小于等于這個數(shù)就行。torchtext是個比較麻煩的東西。BERT-pytorch項目里的tokenizer.py使用了torchtext.datasets.text_classification或者torchtext.data.utils中的tokenizer。老版本torchtext0.9及以下有torchtext.data.functional.sentencepiece_numericalizer等接口新版本0.12API大改很多老代碼直接跑不動。如果你想用BPE或WordPiece tokenizer更推薦直接用tokenizers庫或者transformers庫自帶的Tokenizer繞開torchtext的版本坑。我自己在復(fù)現(xiàn)時直接改用了transformers庫的BertTokenizer來做tokenization因為BERT-pytorch官方訓(xùn)練用的vocab文件本質(zhì)上就是WordPiece格式和BertTokenizer完美兼容。這樣既享受了transformers的便利又保持了項目的核心訓(xùn)練邏輯不變。3. 核心代碼結(jié)構(gòu)拆解與實現(xiàn)原理3.1 從目錄結(jié)構(gòu)看項目脈絡(luò)把項目clone下來之后先看目錄結(jié)構(gòu)BERT-pytorch/ ├── bert_pytorch/ │ ├── __init__.py │ ├── __main__.py │ ├── dataset/ │ │ ├── __init__.py │ │ ├── dataset.py │ │ ├── tokenization.py │ │ └── vocab.py │ ├── model/ │ │ ├── __init__.py │ │ ├── attention.py │ │ ├── bert.py │ │ ├── embedding.py │ │ ├── layer_norm.py │ │ ├── linear.py │ │ ├── transformer.py │ │ └── utils.py │ └── trainer/ │ ├── __init__.py │ ├── optim_schedule.py │ ├── pretrain.py │ └── trainer.py ├── scripts/ │ ├── preprocess.py │ └── train.py └── tests/ └── sanity_check.py這個結(jié)構(gòu)非常清晰model目錄下的每個文件幾乎對應(yīng)Transformer論文里的一個核心概念。transformer.py里定義了TransformerBlock里面包含多頭自注意力MultiHeadedAttention和位置前饋網(wǎng)絡(luò)PositionwiseFeedForward。attention.py則是自注意力的具體實現(xiàn)。embedding.py負(fù)責(zé)Token Embedding、Segment Embedding、Positional Embedding的構(gòu)造。bert.py則是最終模型的組裝。3.2 關(guān)鍵組件Embedding層的實現(xiàn)細(xì)節(jié)BERT和原始Transformer的一個重要區(qū)別在于輸入表示。BERT的輸入由三部分加和而成Token Embedding把每個詞映射成768維向量BERT-base配置Segment Embedding區(qū)分兩個句子句子A和句子B的segment id分別為0和1Positional Embedding這里用的是可學(xué)習(xí)的絕對位置編碼而不是Transformer原論文里的正弦余弦函數(shù)這段邏輯在embedding.py里實現(xiàn)得很直白class BERTEmbedding(nn.Module): def __init__(self, vocab_size, d_model, n_segments, max_len, dropout0.1): super().__init__() self.token_embedding nn.Embedding(vocab_size, d_model) self.segment_embedding nn.Embedding(n_segments, d_model) self.position_embedding nn.Embedding(max_len, d_model) self.dropout nn.Dropout(pdropout) def forward(self, tokens, segments): seq_len tokens.size(1) positions torch.arange(seq_len, devicetokens.device).unsqueeze(0).expand_as(tokens) embeddings self.token_embedding(tokens) \ self.position_embedding(positions) \ self.segment_embedding(segments) return self.dropout(embeddings)這里有一個值得品味的點position_embedding用的是nn.Embedding意味著位置編碼是“學(xué)出來的”而不是固定的。在原始Transformer論文里作者用固定的三角函數(shù)位置編碼后來發(fā)現(xiàn)可學(xué)習(xí)的位置編碼在類似規(guī)模的數(shù)據(jù)下效果也夠好BERT就選擇了可學(xué)習(xí)方案。如果你自己實現(xiàn)時候想把sinusoidal位置編碼換進(jìn)去直接在forward里生成常數(shù)矩陣就行但要注意和后續(xù)預(yù)訓(xùn)練階段的分布保持一致。3.3 自注意力與多頭機制的完整解讀attention.py實現(xiàn)了縮放點積注意力Scaled Dot-Product Attention。核心公式是[ \text{Attention}(Q, K, V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V ]代碼里也就是這幾行的事def forward(self, query, key, value, maskNone): scores torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, -1e9) p_attn F.softmax(scores, dim-1) return torch.matmul(p_attn, value), p_attn這里面關(guān)鍵的“為什么除以根號d_k”值得多說一句。如果不縮放當(dāng)d_k比較大時點積結(jié)果的方差會變大導(dǎo)致softmax之后的梯度非常小模型就很難訓(xùn)練。把點積除以(\sqrt{d_k})讓方差回到1量級softmax的輸入分布才比較穩(wěn)定。這是原論文作者通過實驗驗證出來的細(xì)節(jié)很多初學(xué)者容易忽略但實際上對訓(xùn)練穩(wěn)定性影響很大。多頭注意力Multi-Head Attention的實現(xiàn)也很干凈把Q、K、V通過線性變換投影到多個子空間并行計算注意力頭最后拼接并再投影一次。多頭的好處是讓模型能夠同時關(guān)注不同位置、不同表征子空間的信息比如一個頭可能關(guān)注語法關(guān)系另一個頭關(guān)注指代信息。代碼里通過view和transpose完成頭的拆分和合并不算復(fù)雜但動起手自己寫一遍還是很考驗基本功的。3.4 Transformer Encoder層的組裝邏輯transformer.py定義了TransformerBlock里面是標(biāo)準(zhǔn)的兩層結(jié)構(gòu)多頭自注意力子層后接殘差連接和LayerNorm位置前饋網(wǎng)絡(luò)子層兩層線性GeLU激活后接殘差連接和LayerNormclass TransformerBlock(nn.Module): def __init__(self, hidden, attn_heads, feed_forward_hidden, dropout): super().__init__() self.attention MultiHeadedAttention(hattn_heads, d_modelhidden) self.feed_forward PositionwiseFeedForward( d_modelhidden, d_fffeed_forward_hidden, dropoutdropout) self.input_sublayer SublayerConnection(sizehidden, dropoutdropout) self.output_sublayer SublayerConnection(sizehidden, dropoutdropout) def forward(self, x, mask): x self.input_sublayer(x, lambda _x: self.attention.forward(_x, _x, _x, maskmask)) x self.output_sublayer(x, self.feed_forward) return x這里使用了SublayerConnection來封裝殘差連接和LayerNorm順序是x sublayer(LayerNorm(x))也就是Pre-LN的模式。這個細(xì)節(jié)值得注意Transformer原論文用的是Post-LN先加殘差再做LayerNorm而BERT用的是Pre-LN。Pre-LN在訓(xùn)練早期更穩(wěn)定對學(xué)習(xí)率沒那么敏感這也是BERT能夠用較大學(xué)習(xí)率預(yù)訓(xùn)練的原因之一。BERT-base的整體配置是12層TransformerBlockhidden size76812個注意力頭feed-forward中間層維度是3072。你可以通過修改這些超參來拼出不同規(guī)模的模型。注意hidden size必須能被注意力頭數(shù)整除否則view操作會報錯。這是非常常見的一個低級錯誤。4. 預(yù)訓(xùn)練任務(wù)與訓(xùn)練流程全解析4.1 MLM與NSP兩個預(yù)訓(xùn)練目標(biāo)的設(shè)計邏輯BERT的預(yù)訓(xùn)練階段有兩個任務(wù)缺一不可。第一個是Masked Language ModelMLM掩碼語言模型。做法是隨機遮蓋句子中15%的token然后讓模型去預(yù)測被遮蓋的token。但這里有一個小細(xì)節(jié)在這15%的token里不是全部用[MASK]替換而是有講究的——其中80%換成[MASK]10%換成隨機詞剩下10%保持不變。為什么要這樣因為如果訓(xùn)練時看到的全是[MASK]而下游任務(wù)中根本不會出現(xiàn)[MASK]就會產(chǎn)生預(yù)訓(xùn)練和微調(diào)之間的分布不一致。加入一定比例的隨機詞和原詞能迫使模型更多依賴上下文來預(yù)測而不是單純記住[MASK]標(biāo)記本身。這個小trick在代碼里體現(xiàn)為一個隨機選擇邏輯。第二個是Next Sentence PredictionNSP下句預(yù)測。輸入是一對句子A和B50%的情況下B是A的下一句標(biāo)簽為150%的情況下B是隨機句子標(biāo)簽為0。模型需要判斷兩句是否連續(xù)。這個任務(wù)讓BERT學(xué)會句子級別的關(guān)聯(lián)信息對問答、自然語言推理等任務(wù)幫助很大。后來的實驗中有論文對NSP是否必需提出過質(zhì)疑比如RoBERTa模型就移除了NSP但在BERT的本體設(shè)計里NSP和MLM共同組成了完整的預(yù)訓(xùn)練目標(biāo)。4.2 數(shù)據(jù)準(zhǔn)備從文本到訓(xùn)練樣本的完整流水線在scripts/preprocess.py里原始文本要經(jīng)過一系列處理才能變成BERT的輸入格式文本清洗去掉特殊符號統(tǒng)一小寫是否區(qū)分大小寫可配置WordPiece分割把單詞切成子詞比如“playing”切成“play”和“##ing”加上[CLS]和[SEP]標(biāo)記[CLS]放在序列開頭[SEP]放在每個句子結(jié)尾構(gòu)建token ids、segment ids、attention mask如果需要的話按照最大長度做padding或截斷這里的核心工具是Vocabulary它維護一個“詞和id”的映射表。BERT-pytorch默認(rèn)提供了一個簡單的vocab構(gòu)建方式但如果你要復(fù)現(xiàn)BERT-base的效果直接使用Google發(fā)布的全詞表bert-base-uncased-vocab.txt更省事有30522個詞條。在實際操作中我建議先寫一個小腳本驗證一下tokenize的結(jié)果from transformers import BertTokenizer tokenizer BertTokenizer.from_pretrained(bert-base-uncased) tokens tokenizer.tokenize(Hello, how are you playing today?) print(tokens) # [hello, ,, how, are, you, playing, today, ?]這樣基本就能和BERT-pytorch的輸入對齊。4.3 Trainer與前向傳播的最小實現(xiàn)項目的trainer/pretrain.py是預(yù)訓(xùn)練的核心循環(huán)。整個流程并不復(fù)雜for epoch in range(epochs): for batch in data_loader: token_ids, segment_ids, masked_tokens, masked_pos, is_next batch logits_lm, logits_cls model(token_ids, segment_ids, masked_pos) loss_lm criterion(logits_lm.transpose(1, 2), masked_tokens) loss_cls criterion(logits_cls, is_next) loss loss_lm loss_cls optimizer.zero_grad() loss.backward() optimizer.step()注意這里的logits_lm并不是對所有token都計算損失而是只對被mask的位置計算。masked_pos記錄了每個樣本中被mask的位置索引模型根據(jù)這些索引去取對應(yīng)的輸出向量再和真實token_ids做交叉熵。這個細(xì)節(jié)過濾掉了很多無關(guān)位置的噪聲是模型能夠高效學(xué)習(xí)的關(guān)鍵。學(xué)習(xí)率調(diào)度上BERT-pytorch實現(xiàn)了Transformer論文里的warmup decay策略。前若干步學(xué)習(xí)率線性上漲之后逐步衰減這樣做能避免早期更新幅度太大導(dǎo)致訓(xùn)練不穩(wěn)定。optim_schedule.py里自己實現(xiàn)了這個調(diào)度器。如果你用PyTorch 2.x也可以直接用現(xiàn)成的get_cosine_schedule_with_warmup等接口替代效果類似。4.4 從零開始訓(xùn)練的最小實驗記錄為了驗證整個流程我用一個小型數(shù)據(jù)集約50MB英文新聞文本做了從零開始的預(yù)訓(xùn)練實驗配置為4層Transformer、hidden size256、4個注意力頭。這個配置遠(yuǎn)小于BERT-base但足以跑通流程。訓(xùn)練的batch size設(shè)為32序列長度128Adam學(xué)習(xí)率1e-4warmup步數(shù)1000。跑了大約3萬步之后MLM的準(zhǔn)確率從初始的不到1%漲到了35%左右NSP準(zhǔn)確率穩(wěn)定在80%以上。如果用標(biāo)準(zhǔn)的BERT-base配置去跑完整預(yù)訓(xùn)練需要16張TPU跑4天所以自己玩的話一定要調(diào)小參數(shù)。這個“調(diào)參”的過程本身就是理解BERT的好機會——你隨便改一個參數(shù)都能直觀感受到訓(xùn)練速度和loss的變化這是直接用官方預(yù)訓(xùn)練模型體會不到的。5. 常見問題與避坑手記5.1 安裝和依賴問題速查這是新手最容易卡住的環(huán)節(jié)我整理了幾個高頻問題問題1torchtext導(dǎo)入報錯老版本BERT-pytorch的tokenization.py依賴torchtext的bert vocab相關(guān)接口。新版torchtext里很多API被移除直接導(dǎo)入會報AttributeError。解決辦法是改用transformers庫的tokenizer或者把torchtext鎖在0.9.0版本使用。我更推薦前者因為transformers生態(tài)更持續(xù)。問題2訓(xùn)練時顯存不足OOMBERT-base輸入長度512、batch size 16在GTX 1080Ti上就非常吃力。解決辦法降低序列長度到128或256減小batch size打開混合精度訓(xùn)練torch.cuda.amp.autocast()問題3masked_fill中的-1e9會導(dǎo)致loss為NaN嗎如果mask把整行都遮住了softmax后全0反向傳播時會出現(xiàn)NaN。這種情況在padding部分可能出現(xiàn)。更安全的做法是用torch.finfo(dtype).min代替-1e9并確保每一個樣本中至少有一個位置是有效的。問題常見原因解決方案torchtext API不兼容版本過新改用transformers tokenizerCUDA版本不對驅(qū)動太舊升級驅(qū)動或降低cu版本顯存不足模型太大/batch大降序列長度、開AMP訓(xùn)練loss震蕩學(xué)習(xí)率過高/warmup不夠調(diào)低學(xué)習(xí)率/增加warmup微調(diào)效果差預(yù)訓(xùn)練數(shù)據(jù)域不匹配在目標(biāo)域上繼續(xù)預(yù)訓(xùn)練5.2 訓(xùn)練過程中的幾個反直覺經(jīng)驗有些經(jīng)驗光看文檔是學(xué)不到的必須實際操作過才能體會第一輸入序列的長度對訓(xùn)練速度影響是線性的但對內(nèi)存影響是平方級的。因為注意力分?jǐn)?shù)的形狀是[batch, heads, seq_len, seq_len]序列長度從128提升到256內(nèi)存消耗直接變成原來的4倍256/128的平方。所以做實驗時先從短序列跑通再逐步加長。第二LayerNorm層的epsilon參數(shù)不要隨便動。項目中用的是1e-12這個值看起來小得離譜但BERT官方實現(xiàn)里就是這樣的。如果你把它改成常見的1e-5反而可能導(dǎo)致訓(xùn)練初期不穩(wěn)定。這個細(xì)節(jié)也是從官方代碼里傳下來的“奇奇怪怪的設(shè)定”。第三預(yù)訓(xùn)練的warmup步數(shù)不能太少。我在小數(shù)據(jù)集上踩過坑warmup只設(shè)100步結(jié)果前2000步loss大幅震蕩后來把warmup加到1000步就穩(wěn)定了。要知道在大型語料上warmup步數(shù)經(jīng)常要設(shè)置到10000步以上。第四如果條件允許直接用Google發(fā)布的預(yù)訓(xùn)練權(quán)重做微調(diào)而不是從零預(yù)訓(xùn)練。從零預(yù)訓(xùn)練需要的數(shù)據(jù)量和算力對個人開發(fā)者來說非常不友好。BERT-pytorch的意義更多在于“理解原理”和“復(fù)現(xiàn)實驗”而不是和生產(chǎn)環(huán)境里用現(xiàn)成權(quán)重?fù)尵?。做下游情感分類、命名實體識別等任務(wù)時加載bert-base-uncased權(quán)重再微調(diào)效果和速度會好得多。5.3 一個實用的復(fù)現(xiàn)檢查清單如果你想完整跑通BERT-pytorch并驗證自己的理解我建議按這個順序檢查輸入tokens是否正確添加了[CLS]和[SEP]長度是否一致多頭注意力的num_heads是否能夠整除d_modelEmbedding層三個矩陣的維度是否都是vocab_size, d_model/max_len, d_model/n_segments, d_modelMLM損失是否只統(tǒng)計masked位置而不是所有位置NSP的標(biāo)簽是否和segment id滿足對應(yīng)關(guān)系“句子B確實是隨機句或下一句”訓(xùn)練時有無用一個小batch做一次前向和反向確保沒有維度不匹配的報錯這個清單看起來非常基礎(chǔ)但我敢說至少有50%的人第一次復(fù)現(xiàn)時報錯都源自這幾個地方。我自己當(dāng)時就在多頭注意力拆分合并那一步卡了很久最后打印每個tensor的shape才排查出來。6. 基于BERT-pytorch的二次開發(fā)思路6.1 如何替換成你自己的TokenizerBERT-pytorch默認(rèn)的tokenization模塊比較簡陋如果你要做中文文本的處理建議直接替換成transformers的BertTokenizer。中文BERT有兩種做法一種是基于字character切分每個漢字作為一個token另一種是基于詞切分。Google提供的中文模型bert-base-chinese就是按字切分的。替換時只需要把原始的tokenizer換成from transformers import BertTokenizer tokenizer BertTokenizer.from_pretrained(bert-base-chinese)然后改寫數(shù)據(jù)預(yù)處理流程里調(diào)用tokenizer的部分保持輸出的token_ids、segment_ids仍是list of int即可。后續(xù)模型代碼完全不用動。6.2 如何用這個項目做下游任務(wù)微調(diào)這里提供一個典型的文本分類微調(diào)思路。先加載BERT-pytorch的預(yù)訓(xùn)練模型可以是Google權(quán)重轉(zhuǎn)換過來的然后取出[CLS]位置的輸出向量接一個全連接分類層class BertClassifier(nn.Module): def __init__(self, bert_model, num_classes): super().__init__() self.bert bert_model self.classifier nn.Linear(bert_model.hidden, num_classes) def forward(self, tokens, segments): _, pooled self.bert.forward(tokens, segments) # 這里需要從輸出中取出CLS位置的hidden state logits self.classifier(pooled) return logits注意BERT-pytorch模型的bert前向返回的是所有token的隱藏狀態(tài)[batch, seq_len, hidden]你要取[:, 0, :]才是[CLS]的表示。很多代碼里所謂的“pooled output”其實就是一個簡單的[CLS]取出操作不像transformers庫中還有額外的pooler層。這個細(xì)節(jié)在微調(diào)時直接影響分類頭的效果。微調(diào)的技巧也很直接用較低學(xué)習(xí)率2e-5到5e-5因為預(yù)訓(xùn)練好的參數(shù)已經(jīng)處于一個較優(yōu)區(qū)域?qū)W習(xí)率太大容易破壞它。通常只需要跑3~5個epoch就能在下游任務(wù)上得到不錯的效果。6.3 項目之外從BERT到更現(xiàn)代的模型BERT-pytorch雖然經(jīng)典但如今NLP領(lǐng)域已經(jīng)有了很多更強的模型。但坦白說如果你想真正理解RoBERTa、ALBERT、DistilBERT乃至GPT系列的核心思想先啃透BERT-pytorch的代碼會是一個非常扎實的起點。因為后續(xù)很多模型都是在BERT的骨架上做減法或加法——RoBERTa去掉了NSP但加大了數(shù)據(jù)和訓(xùn)練時長ALBERT用參數(shù)共享和矩陣分解大幅度減少了參數(shù)量DistilBERT用知識蒸餾把模型縮小了40%。你理解了BERT的每一塊部件再去看這些模型的“改動點”就會非常輕松。我在實際做項目時也經(jīng)常會把BERT-pytorch的代碼拿來當(dāng)作“最小可運行”的基線用來測試新數(shù)據(jù)集上是否存在數(shù)據(jù)泄漏、處理流程是否合理。它的代碼足夠短短到你可以在半小時內(nèi)讀懂全部關(guān)鍵邏輯這一點在生產(chǎn)級的transformers庫中是做不到的。7. 最后的實操心得這篇文章寫到這里其實已經(jīng)沒有“總括性結(jié)論”的必要了。我記得自己第一次獨立把BERT-pytorch從環(huán)境搭建到預(yù)訓(xùn)練小模型跑通花了整整一個周末。期間踩過torchtext的坑踩過CUDA版本不對的坑也踩過多頭注意力維度錯誤導(dǎo)致loss不下降的坑。但正是這些排查過程讓我對BERT的理解遠(yuǎn)超“會用transformers庫”的程度。如果讓我給剛接觸這個項目的朋友三條建議第一先別急著跑大數(shù)據(jù)拿一個只有幾千條的小語料把整個訓(xùn)練流程跑通觀察loss變化理解每一步在干什么第二一定要動手改參數(shù)改hidden size、改層數(shù)、改學(xué)習(xí)率看看訓(xùn)練曲線有什么不同這種直觀感知是看多少文章都替代不了的第三把官方預(yù)訓(xùn)練權(quán)重轉(zhuǎn)換到PyTorch環(huán)境后用在下游任務(wù)上這樣既能保證效果又能體會BERT-pytorch代碼和官方權(quán)重的兼容性。這個項目雖然名叫“BERT-pytorch”但它實際上是一個深度學(xué)習(xí)愛好者極佳的學(xué)習(xí)素材——它把一個曾經(jīng)刷新無數(shù)榜單的模型壓縮到了幾千行清晰可讀的PyTorch代碼里。弄懂它之后你再看其他NLP模型的實現(xiàn)大概率都能舉一反三。這大概就是經(jīng)典項目給人留下的最寶貴的財富。本文還有配套的精品資源點擊獲取