制的問(wèn)答摘要生成:從數(shù)據(jù)清洗到推理驗(yàn)證)
簡(jiǎn)介這是一份汽車大師問(wèn)答摘要與推理比賽的參賽源碼與項(xiàng)目說(shuō)明面向自然語(yǔ)言處理初學(xué)者、算法競(jìng)賽愛好者以及需要完成相關(guān)課程設(shè)計(jì)、期末大作業(yè)或畢業(yè)設(shè)計(jì)的計(jì)算機(jī)、數(shù)學(xué)、電子信息類專業(yè)學(xué)生。壓縮包共37個(gè)文件以28個(gè)Python腳本和7個(gè)Jupyter Notebook為主另含1個(gè)Markdown說(shuō)明文檔整體體積僅128KB便于快速下載與本地調(diào)試。源碼涵蓋seq2seq與seq2seq_attention兩套模型實(shí)現(xiàn)并擴(kuò)展了Transformer、PGN等思路的Notebook演示以及數(shù)據(jù)預(yù)處理、Beam Search解碼、訓(xùn)練與測(cè)試腳本等完整流程項(xiàng)目說(shuō)明文檔梳理了從數(shù)據(jù)清洗到模型推理的關(guān)鍵步驟代碼模塊按utils、models等目錄組織結(jié)構(gòu)清晰適合逐模塊研讀。Notebook中逐步演示了訓(xùn)練與推理過(guò)程可幫助讀者理解問(wèn)答摘要生成的細(xì)節(jié)也方便在此基礎(chǔ)上二次開發(fā)或作為畢業(yè)設(shè)計(jì)的基線系統(tǒng)。目前已有99人學(xué)習(xí)瀏覽作為算法類參考資料具備較強(qiáng)的借鑒價(jià)值。1. 汽車大師問(wèn)答摘要與推理本質(zhì)是把長(zhǎng)問(wèn)答壓成一句可落地的結(jié)論做問(wèn)答摘要和推理這類比賽最怕的不是模型跑不起來(lái)而是拿到數(shù)據(jù)后不知道輸入到底該喂什么。汽車大師這個(gè)賽題把場(chǎng)景限制得很具體車主在平臺(tái)上提問(wèn)維修技師給出一長(zhǎng)段回答參賽者要訓(xùn)練模型把這組問(wèn)答壓縮成一條有結(jié)論、可推理的摘要。換句話說(shuō)輸入是“問(wèn)題 回答”兩段長(zhǎng)文本輸出是一句短摘要而摘要里的關(guān)鍵信息還要能被后續(xù)的推理環(huán)節(jié)直接用起來(lái)。這條技術(shù)路線最適合兩類人剛接觸生成式 NLP、想用 seq2seq 練手的工程師以及要在真實(shí)問(wèn)答場(chǎng)景里做信息壓縮的從業(yè)者。下面我會(huì)按數(shù)據(jù)準(zhǔn)備、seq2seq 基線、注意力機(jī)制、訓(xùn)練調(diào)參和推理驗(yàn)證的順序把這一整套方案拆開講清楚。2. 數(shù)據(jù)準(zhǔn)備先做好問(wèn)答對(duì)如何拼裝、清洗和變短2.1 這個(gè)任務(wù)和普通摘要不一樣輸入是“問(wèn)題回答”的雙段文本如果做過(guò)新聞?wù)銜?huì)發(fā)現(xiàn)新聞標(biāo)題和正文的語(yǔ)義是高度對(duì)齊的直接抽第一段往往就能拿個(gè)不錯(cuò)的 ROUGE。汽車大師這類問(wèn)答摘要完全不是這個(gè)邏輯問(wèn)題里帶著車型、年款、故障現(xiàn)象回答里帶著分析過(guò)程、診斷結(jié)論和維修建議而參考摘要通常是把“結(jié)論”和“依據(jù)”融合成一句通順的話。比如問(wèn)題可能是“13 款朗逸1.6L 自動(dòng)擋冷車啟動(dòng)時(shí)發(fā)動(dòng)機(jī)艙吱吱響熱車后消失”回答會(huì)分析皮帶老化、張緊輪磨損、水泵軸承等好幾個(gè)可能原因摘要卻往往只寫“冷車啟動(dòng)異響多為發(fā)電機(jī)皮帶或張緊輪老化建議更換皮帶及張緊輪”。這個(gè)合成過(guò)程既需要從回答里定位結(jié)論又需要回扣問(wèn)題里的車型和癥狀約束所以普通的抽取式摘要很難覆蓋生成式模型才是主流做法。常見的數(shù)據(jù)格式是 JSON Lines每行一條樣本。核心字段就三個(gè)question、answer、summary分別對(duì)應(yīng)用戶提問(wèn)、技師回答和參考摘要。也有些版本會(huì)把問(wèn)題拆成 title 和 content或者額外給一段“推理依據(jù)”。拿到手先不要急著訓(xùn)模型把每一個(gè)字段的真實(shí)長(zhǎng)度分布打印出來(lái)看看回答動(dòng)輒幾百字摘要往往只有幾十個(gè)字輸入輸出長(zhǎng)度差異極大這直接決定了后續(xù)的詞表構(gòu)建、填充策略和截?cái)嗖呗詰?yīng)該怎么做。2.2 JSON 解析與清洗把原始問(wèn)答轉(zhuǎn)成訓(xùn)練三元組數(shù)據(jù)清洗的核心目標(biāo)是把原始 JSON 轉(zhuǎn)成干凈的三元組問(wèn)題、回答、摘要。我一般先把問(wèn)題截?cái)嗟揭粋€(gè)合理長(zhǎng)度再把回答單獨(dú)截?cái)嘧詈蟛抛銎唇印_@樣能避免“問(wèn)題太長(zhǎng)把回答擠掉”這種低級(jí)問(wèn)題。下面這段解析腳本在多個(gè)類似賽題里都能直接用import json import re def tokenize(text): # 中文按漢字切連續(xù)的英文/數(shù)字串作為一個(gè)整體標(biāo)點(diǎn)單獨(dú)成詞 return re.findall(r[\u4e00-\u9fffA-Za-z0-9]|[^\s\u4e00-\u9fffA-Za-z0-9], text) def clean_text(text): text text.replace(\u3000, ).replace(\n, ).strip() text re.sub(r\s, , text) return text def load_samples(data_path, max_q_len80, max_a_len200): samples [] with open(data_path, r, encodingutf-8) as f: for line in f: line line.strip() if not line: continue obj json.loads(line) question clean_text(obj.get(question, ))[:max_q_len] answer clean_text(obj.get(answer, ))[:max_a_len] summary clean_text(obj.get(summary, )) if not question or not answer or not summary: continue src question [SEP] answer samples.append({src: src, tgt: summary}) return samples這里的 max_q_len 和 max_a_len 不是隨便拍的。我一般先把訓(xùn)練集里 question 和 answer 的長(zhǎng)度分布跑一遍question 超過(guò) 80 個(gè)字符的樣本占比很低而 answer 超過(guò) 200 個(gè)字符的樣本非常多但真正參與結(jié)論表達(dá)的往往是前半段。如果你把 answer 全量保留詞表里會(huì)出現(xiàn)大量低頻專業(yè)詞模型訓(xùn)起來(lái)很慢且容易過(guò)擬合。代碼里 clean_text 做了兩個(gè)重要處理把全角空格和換行統(tǒng)一成半角空格再把連續(xù)空白壓縮成單個(gè)空格。這一步不做后續(xù) tokenize 會(huì)把換行符當(dāng)成獨(dú)立詞詞表里多出一堆沒意義的詞條。清洗完建議順手統(tǒng)計(jì)一下過(guò)濾后還剩多少條樣本、摘要的平均長(zhǎng)度、src 的平均長(zhǎng)度。摘要太長(zhǎng)的樣本比如超過(guò) 60 個(gè)字可以在訓(xùn)練時(shí)直接截?cái)嗷蛘咛蕹驗(yàn)檫@種“長(zhǎng)摘要”很多是從多個(gè)回答片段拼出來(lái)的參考質(zhì)量本身就不穩(wěn)定。2.3 詞表構(gòu)建與 mini-batchOOV 和截?cái)嘣趺刺幚韘eq2seq 模型離不開固定詞表。常見的做法是按詞頻過(guò)濾低頻詞把出現(xiàn)次數(shù)少于某個(gè)閾值的詞替換成 UNK。min_freq2 是個(gè)比較穩(wěn)的起點(diǎn)詞頻為 1 的詞通常是車型名、人名或者拼寫變體這些詞對(duì)摘要生成貢獻(xiàn)不大卻把詞表?yè)蔚煤艽?。詞表里四個(gè)特殊符的順序建議固定后續(xù)代碼里到處要用from collections import Counter PAD, SOS, EOS, UNK 0, 1, 2, 3 SPECIAL_TOKENS [pad, sos, eos, unk] def build_vocab(samples, min_freq2): counter Counter() for s in samples: for tok in tokenize(s[src]): counter[tok] 1 for tok in tokenize(s[tgt]): counter[tok] 1 vocab {tok: idx for idx, tok in enumerate(SPECIAL_TOKENS)} for tok, freq in counter.most_common(): if freq min_freq and tok not in vocab: vocab[tok] len(vocab) return vocab這里把 src 和 tgt 的詞合并統(tǒng)計(jì)詞頻好處是摘要里的高頻詞也能進(jìn)入詞表。壞處是 src 里的大量無(wú)關(guān)背景詞會(huì)擠占詞表容量。實(shí)際操作中可以用雙詞表encoder 用 src 統(tǒng)計(jì)出的詞表decoder 用 tgt 統(tǒng)計(jì)出的詞表。對(duì)于初學(xué)者單詞表更省事但我建議至少把 min_freq 從 1 提到 2否則詞表大小會(huì)從一兩萬(wàn)跳到四五萬(wàn)Embedding 層和輸出層的參數(shù)數(shù)量都跟著翻倍。mini-batch 的組裝是另一個(gè)容易翻車的地方。PyTorch 的pack_padded_sequence要求序列按長(zhǎng)度降序排列所以 collate_fn 里除了 padding 和記錄真實(shí)長(zhǎng)度還必須做排序import torch from torch.nn.utils.rnn import pad_sequence def collate_fn(batch, vocab, max_src_len200, max_tgt_len64): src_ids_list, tgt_ids_list [], [] for item in batch: src_tokens tokenize(item[src])[:max_src_len] tgt_tokens tokenize(item[tgt])[:max_tgt_len] src_ids [vocab.get(t, UNK) for t in src_tokens] tgt_ids [SOS] [vocab.get(t, UNK) for t in tgt_tokens] [EOS] src_ids_list.append(torch.tensor(src_ids, dtypetorch.long)) tgt_ids_list.append(torch.tensor(tgt_ids, dtypetorch.long)) src_ids pad_sequence(src_ids_list, batch_firstTrue, padding_valuePAD) tgt_ids pad_sequence(tgt_ids_list, batch_firstTrue, padding_valuePAD) src_lens torch.tensor([len(x) for x in src_ids_list], dtypetorch.long) # pack_padded_sequence 要求降序 src_lens, order torch.sort(src_lens, descendingTrue) src_ids src_ids[order] tgt_ids tgt_ids[order] return src_ids, src_lens, tgt_idsmax_tgt_len64 是因?yàn)閰⒖颊苌俪^(guò) 60 個(gè)詞。這里有幾個(gè)參數(shù)初學(xué)者容易調(diào)錯(cuò)SOS 和 EOS 都要拼進(jìn)去但 padding 值必須用 PAD 而不是 0 以外的東西padding 后 tgt 里會(huì)出現(xiàn)大量 PAD訓(xùn)練時(shí)計(jì)算損失要把 ignore_index 設(shè)為 PAD排序后 tgt_ids 必須跟著 src_ids 一起重排否則 batch 內(nèi)部的配對(duì)就亂了。很多人第一次跑通后 loss 亂跳回頭查基本都是這三處的問(wèn)題。3. 跑通 seq2seq 基線雙向 GRU 編碼器 解碼器的最小訓(xùn)練回路3.1 為什么基線先選 GRU參數(shù)少、收斂快適合比賽試錯(cuò)比賽場(chǎng)景下基線模型的選擇標(biāo)準(zhǔn)不是“效果最好”而是“快速跑通、快速找到 bug”。LSTM 和 GRU 在這個(gè)任務(wù)上的效果差距很小但 GRU 的參數(shù)只有 LSTM 的四分之三訓(xùn)練速度和顯存占用都更友好。汽車大師問(wèn)答摘要這種中等規(guī)模數(shù)據(jù)集GRU 在十幾個(gè)小時(shí)內(nèi)就能收斂到能看到效果的 checkpointLSTM 往往要多跑 30% 的時(shí)間。所以我的建議很直接基線用 GRU如果后續(xù)要刷分再換 LSTM 或者 Transformer。編碼器用雙向 GRU解碼器用單向 GRU這是 seq2seq 摘要任務(wù)最常見的配置。雙向編碼器能讓每個(gè)時(shí)間步的隱狀態(tài)同時(shí)看到前文和后文對(duì)“皮帶老化”這種跨詞組的語(yǔ)義非常有用。雙向帶來(lái)的問(wèn)題是如何把兩個(gè)方向的最終狀態(tài)合并成解碼器的初始狀態(tài)常見做法是把 forward 方向和 backward 方向的最后一個(gè)隱狀態(tài)拼接過(guò)一個(gè)線性層壓縮到單向隱狀態(tài)維度再用 tanh 激活。下面這段 Encoder 實(shí)現(xiàn)可以直接抄import torch import torch.nn as nn import torch.nn.functional as F class Encoder(nn.Module): def __init__(self, vocab_size, embed_size, hidden_size, dropout0.1): super().__init__() self.embedding nn.Embedding(vocab_size, embed_size, padding_idxPAD) self.gru nn.GRU(embed_size, hidden_size, bidirectionalTrue, batch_firstTrue) self.fc nn.Linear(hidden_size * 2, hidden_size) self.dropout nn.Dropout(dropout) def forward(self, src_ids, src_lens): emb self.dropout(self.embedding(src_ids)) packed nn.utils.rnn.pack_padded_sequence( emb, src_lens.cpu(), batch_firstTrue, enforce_sortedTrue ) packed_out, hidden self.gru(packed) encoder_outputs, _ nn.utils.rnn.pad_packed_sequence( packed_out, batch_firstTrue ) # hidden: (2, B, H)分別對(duì)應(yīng) forward 和 backward 最后一個(gè)時(shí)刻 h_fwd, h_bwd hidden[0], hidden[1] hidden_cat torch.cat([h_fwd, h_bwd], dim-1) # B, 2H decoder_init torch.tanh(self.fc(hidden_cat)) return encoder_outputs, decoder_init這段代碼里最容易出錯(cuò)的是 hidden 的取值。PyTorch 的 GRU 返回的 hidden 維度是 (num_layers * num_directions, batch, hidden_size)。單層雙向時(shí)hidden[0] 是 forward 方向的最終狀態(tài)hidden[1] 是 backward 方向的最終狀態(tài)。很多同學(xué)直接取 hidden[-1] 當(dāng)成一個(gè)方向的輸出但在雙向模型里 hidden[-1] 是 backward 的最終狀態(tài)取錯(cuò)了方向整個(gè)初始狀態(tài)就廢了。參數(shù)上embed_size256、hidden_size256 是一個(gè)性價(jià)比很高的起點(diǎn)如果顯存緊張把 hidden_size 降到 128 也能跑但摘要質(zhì)量會(huì)明顯下降。3.2 Decoder 與訓(xùn)練循環(huán)teacher forcing 的時(shí)機(jī)和梯度裁剪沒有 attention 的 baseline Decoder 很簡(jiǎn)單每個(gè)時(shí)間步把上一時(shí)刻的 token 輸入到 GRU用當(dāng)前隱狀態(tài)預(yù)測(cè)下一個(gè) token。上一節(jié)已經(jīng)是完整代碼了這里給出 Decoder 和訓(xùn)練循環(huán)中最關(guān)鍵的train_step。訓(xùn)練時(shí)我不會(huì)一次性把整個(gè)序列解出來(lái)而是逐時(shí)間步循環(huán)這樣能用 teacher forcing 控制模型看到真實(shí)歷史的比例import random def train_step(model, src_ids, src_lens, tgt_ids, teacher_forcing_ratio0.5): # src_ids: B,T_src tgt_ids: B,T_tgt encoder_outputs, decoder_init model.encoder(src_ids, src_lens) batch_size src_ids.size(0) decoder_input tgt_ids[:, 0].unsqueeze(1) # 強(qiáng)制從 sos 開始 hidden decoder_init.unsqueeze(0) # 1,B,H loss 0.0 for t in range(1, tgt_ids.size(1)): logits, hidden model.decoder(decoder_input, hidden) # B,1,V loss F.cross_entropy( logits.squeeze(1), tgt_ids[:, t], ignore_indexPAD ) use_teacher random.random() teacher_forcing_ratio if use_teacher: decoder_input tgt_ids[:, t].unsqueeze(1) else: decoder_input logits.argmax(dim-1) # 模型自己猜 return loss / (tgt_ids.size(1) - 1)teacher_forcing_ratio0.5 的意思是訓(xùn)練時(shí)每個(gè)時(shí)間步有 50% 概率把真實(shí) token 喂回解碼器另外 50% 概率用模型上一時(shí)刻的預(yù)測(cè)作為輸入。這個(gè)比例不能太高太高會(huì)讓模型在推理階段一遇到自己的錯(cuò)誤預(yù)測(cè)就連環(huán)出錯(cuò)也不能太低太低收斂會(huì)很慢。我一般從 0.5 開始訓(xùn)練到后半程降到 0.3。訓(xùn)練主循環(huán)里還缺一個(gè)關(guān)鍵操作梯度裁剪。GRU 在長(zhǎng)序列上很容易梯度爆炸loss 突然變成 NaN 的案例幾乎都是沒做裁剪optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(30): for batch in dataloader: optimizer.zero_grad() src_ids, src_lens, tgt_ids [x.to(device) for x in batch] loss train_step(model, src_ids, src_lens, tgt_ids) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step()Adam 的初始學(xué)習(xí)率 1e-3 在這個(gè)任務(wù)上是安全的clip 值 5.0 是我的個(gè)人偏好太小會(huì)讓訓(xùn)練變慢太大會(huì)失去保護(hù)作用。這里要提醒一句loss 在下降不代表生成質(zhì)量在變好seq2seq 的 loss 是逐 token 的交叉熵模型可能學(xué)會(huì)了復(fù)制高頻詞但沒學(xué)會(huì)組織句子。我一般在每 500 步做一次 greedy 解碼把預(yù)測(cè)的摘要打印出來(lái)和參考摘要對(duì)比靠肉眼判斷模型是不是真的在“變聰明”。4. 給解碼器裝上 attention上下文向量的對(duì)齊邏輯與 PyTorch 實(shí)現(xiàn)4.1 摘要任務(wù)的信息瓶頸固定向量扛不住長(zhǎng)回答純 seq2seq 的問(wèn)題在問(wèn)答摘要任務(wù)里暴露得很明顯解碼器每一步都只能依賴編碼器最后時(shí)刻壓縮出的一個(gè)固定向量。如果車主描述有 300 個(gè)字技師回答有 200 個(gè)字這個(gè)固定向量要承載全部信息早期輸入的內(nèi)容早就被后續(xù) token 沖刷掉了。而汽車大師摘要偏偏需要跨段對(duì)齊“冷車啟動(dòng)異響”在問(wèn)題的開頭“皮帶老化”在回答的中后段模型要建立這兩個(gè)位置之間的關(guān)聯(lián)。沒有 attention解碼器只能“記住個(gè)大概”生成出來(lái)的摘要經(jīng)常張冠李戴。attention 的核心思想是解碼器在第 i 步生成詞之前先計(jì)算當(dāng)前隱狀態(tài)與編碼器每個(gè)位置隱狀態(tài)的相似度把相似度歸一化成權(quán)重再對(duì)編碼器所有隱狀態(tài)做加權(quán)求和得到這一步專屬的上下文向量。這個(gè)機(jī)制最早由 Bahdanau 提出也是標(biāo)題里 seq2seq_attention 對(duì)應(yīng)的一類實(shí)現(xiàn)相當(dāng)于給解碼器裝了一個(gè)“通用注意力模塊”PyTorch 里完全可以自己寫成一個(gè)獨(dú)立的 nn.Module 復(fù)用。4.2 Bahdanau attention 的 PyTorch 實(shí)現(xiàn)mask 是關(guān)鍵Bahdanau 注意力又叫加性注意力它對(duì)編碼器輸出和當(dāng)前解碼器隱狀態(tài)各做一次線性變換再相加經(jīng)過(guò) tanh 和另一個(gè)線性層得到標(biāo)量分?jǐn)?shù)。公式是e_ij v^T tanh(W_h h_j W_s s_{i-1})a_ij softmax(e_ij)c_i Σ_j a_ij h_j代碼實(shí)現(xiàn)時(shí)要注意一個(gè)很容易被忽略的細(xì)節(jié)padding 位置不能參與注意力打分。編碼器輸出經(jīng)過(guò) pad_packed_sequence 后所有短句子的尾部都是 PAD 對(duì)應(yīng)的隱狀態(tài)如果不把這些位置的分?jǐn)?shù)壓成負(fù)無(wú)窮模型會(huì)學(xué)到“注意力放到 PAD 上也無(wú)所謂”訓(xùn)練 loss 照樣降但注意力熱圖完全失去可解釋性。下面是完整實(shí)現(xiàn)class BahdanauAttention(nn.Module): def __init__(self, hidden_size): super().__init__() self.W_h nn.Linear(hidden_size * 2, hidden_size) self.W_s nn.Linear(hidden_size, hidden_size) self.v nn.Linear(hidden_size, 1) def forward(self, decoder_hidden, encoder_outputs, mask): # decoder_hidden: B, H - B, 1, H q self.W_s(decoder_hidden).unsqueeze(1) # encoder_outputs: B, T, 2H - B, T, H k self.W_h(encoder_outputs) scores self.v(torch.tanh(q k)).squeeze(2) # B, T scores scores.masked_fill(mask 0, -1e9) # 屏蔽 padding attn_weights F.softmax(scores, dim-1) context torch.bmm(attn_weights.unsqueeze(1), encoder_outputs).squeeze(1) return context, attn_weightsmask 的來(lái)源是src_ids ! PADshape 為 B, T。把 mask 傳進(jìn) attention 之前要確保它和 encoder_outputs 的 T 維度一致因?yàn)?pad_packed_sequence 返回的序列長(zhǎng)度是整個(gè) batch 里最長(zhǎng)的那條所以 mask 直接用src_ids ! PAD生成即可。這里的 hidden_size 是解碼器的隱狀態(tài)維度encoder_outputs 的最后一維是 hidden_size * 2所以 W_h 輸入維度是 2HW_s 輸入維度是 H。我見過(guò)有人把這兩個(gè)線性層搞反或者都寫成 H結(jié)果 score 計(jì)算的維度對(duì)不上報(bào)錯(cuò)后只能靠猜。先把這個(gè)對(duì)應(yīng)關(guān)系寫清楚抄代碼的時(shí)候就不會(huì)亂。4.3 Decoder 接入 context 后的結(jié)構(gòu)變化加了 attention 之后Decoder 的每個(gè)時(shí)間步不再只吃 token embedding而是把上一時(shí)刻的上下文向量和當(dāng)前 token 的 embedding 拼接在一起再送進(jìn) GRU。這一步的目的是讓 GRU 在生成每個(gè)詞時(shí)都能“看著”原文本的相關(guān)位置。下面是完整 Decoderclass Decoder(nn.Module): def __init__(self, vocab_size, embed_size, hidden_size, dropout0.1): super().__init__() self.embedding nn.Embedding(vocab_size, embed_size, padding_idxPAD) # 輸入維度 embedding encoder_outputs 維度2H self.gru nn.GRU(embed_size hidden_size * 2, hidden_size, batch_firstTrue) self.attention BahdanauAttention(hidden_size) self.fc_out nn.Linear(hidden_size, vocab_size) self.dropout nn.Dropout(dropout) def forward(self, decoder_input, hidden, encoder_outputs, mask): emb self.dropout(self.embedding(decoder_input)) # B,1,E context, attn_weights self.attention( hidden.squeeze(0), encoder_outputs, mask ) # context: B,2H context context.unsqueeze(1) # B,1,2H rnn_input torch.cat([emb, context], dim-1) # B,1,E2H out, hidden self.gru(rnn_input, hidden) # out: B,1,H logits self.fc_out(out) # B,1,V return logits, hidden, attn_weights這里有一個(gè)維度細(xì)節(jié)attention 返回的 context 是 B,2H和 embedding 拼接后 GRU 輸入維度變成 embed_size 2H。GRU 輸出維度是 H最后接一個(gè) Linear(H, V) 預(yù)測(cè)詞表分布。hidden 在 GRU 之間傳遞時(shí)要保持 (1,B,H) 的形狀但 attention 需要的是 (B,H)所以傳入 attention 前要 squeeze拿到新 hidden 后繼續(xù)做循環(huán)。如果你用的是多層解碼器這里要額外處理每一層的 hidden初學(xué)者建議先用單層跑通再加深。訓(xùn)練循環(huán)大部分邏輯和上一章一致只有一個(gè)小改動(dòng)train_step里每步都要把 encoder_outputs 和 mask 傳進(jìn) decoder同時(shí)接收 attention 權(quán)重用于可視化。我習(xí)慣把 attention 權(quán)重存下來(lái)每跑完一個(gè) epoch 挑幾個(gè)樣本畫熱圖看模型是否關(guān)注到了“車型型號(hào)”和“故障結(jié)論”這些位置這一步對(duì)判斷是模型問(wèn)題還是數(shù)據(jù)問(wèn)題非常有效。5. 訓(xùn)練與調(diào)參常見坑seq2seq 摘要最容易翻車的四處細(xì)節(jié)5.1 損失降但輸出全是重復(fù)詞先查 SOS/EOS 和 teacher forcing現(xiàn)象訓(xùn)練集上 loss 一路下降但驗(yàn)證時(shí)生成的摘要里全是“皮帶 皮帶 皮帶”“異響 異響 異響”這種重復(fù)詞。原因分兩類。第一目標(biāo)序列構(gòu)造時(shí)沒有把 EOS 加在末尾或者加了但計(jì)算損失時(shí)沒有對(duì)最后一個(gè) EOS 做預(yù)測(cè)導(dǎo)致模型永遠(yuǎn)學(xué)不到“這句話該結(jié)束了”于是只能一直重復(fù)上一個(gè)高頻詞。第二teacher forcing 比例設(shè)得太高比如 1.0 或 0.9模型從沒見過(guò)自己的錯(cuò)誤輸出推理時(shí)一旦第一步選了個(gè)低頻詞后面就全亂了。解決檢查 tgt_ids 是否[SOS] tokens [EOS]t 從 1 循環(huán)到 tgt_ids.size(1) - 1確保每個(gè) token 都會(huì)被預(yù)測(cè)一遍訓(xùn)練初期 teacher_forcing_ratio 維持 0.5如果還重復(fù)降到 0.3 并增大 dropout。5.2 解碼停不下來(lái)EOS 被 beam search 當(dāng)成了普通候選現(xiàn)象訓(xùn)練正常、驗(yàn)證 loss 正常但 greedy 解碼時(shí)序列都生成了 80 個(gè) token 還在繼續(xù)完全沒有 EOS 的跡象。原因有兩個(gè)層面一是 EOS 在詞表里的初始分?jǐn)?shù)比較低訓(xùn)練樣本里 EOS 只會(huì)出現(xiàn)在句尾模型天然有“多寫幾個(gè)詞”的偏向二是推理時(shí)如果用了 beam searchbeam 會(huì)把 EOS 當(dāng)成一個(gè)普通候選詞它的累計(jì)分?jǐn)?shù)不如普通詞高于是被其他候選擠掉。解決訓(xùn)練時(shí)在 loss 上對(duì) EOS 的 token 加權(quán)比如把 EOS 的權(quán)重設(shè)成 2.0強(qiáng)迫模型重視結(jié)束標(biāo)志推理時(shí)對(duì) beam search 加長(zhǎng)度懲罰對(duì)包含 EOS 的候選做獎(jiǎng)勵(lì)或者當(dāng)某條 beam 生成 EOS 后直接凍結(jié)它不再擴(kuò)展。更省事的辦法是解碼長(zhǎng)度超過(guò) max_len 時(shí)強(qiáng)制截?cái)嗟@只能止血不能解決模型本身不學(xué) EOS 的問(wèn)題。5.3 ROUGE 死活不漲多半是輸入被截?cái)喟殃P(guān)鍵信息切掉了現(xiàn)象換 attention、調(diào) lr、加 dropoutROUGE-L 始終停留在 0.3 左右上不去。我排查這類問(wèn)題會(huì)先看一條具體樣本的輸入輸出。如果發(fā)現(xiàn)參考摘要里的“冷車啟動(dòng)”“發(fā)電機(jī)皮帶”在 src 里根本找不到那問(wèn)題不在模型在數(shù)據(jù)預(yù)處理。前面 2.2 節(jié)我特意把 answer 單獨(dú)截?cái)嗟?200 字符就是因?yàn)橛行┩瑢W(xué)直接把 question answer 拼起來(lái)再截?cái)嘤龅?question 本身很長(zhǎng)的情況answer 會(huì)被砍掉一大半結(jié)論信息全丟了。解決把截?cái)嗖呗愿某伞皅uestion 截?cái)嗟?80、answer 截?cái)嗟?200再拼接”并在預(yù)處理后寫一個(gè)簡(jiǎn)單腳本統(tǒng)計(jì)參考摘要里的 bigram 有多少比例出現(xiàn)在截?cái)嗪蟮?src 里。這個(gè)覆蓋率如果低于 80%說(shuō)明截?cái)嗵菪枰糯?max_a_len。很多調(diào)參調(diào)不動(dòng)的“玄學(xué)”最后查出來(lái)都是數(shù)據(jù)預(yù)處理的問(wèn)題。5.4 注意力熱圖亂成一團(tuán)padding 位置也在打分現(xiàn)象attention 可視化畫出來(lái)每一行的權(quán)重都散布在整個(gè)寬度上甚至句尾的 PAD 位置權(quán)重很高完全看不出對(duì)齊關(guān)系。原因很明確初始化 mask 時(shí)用的是src_ids ! PAD但 encoder_outputs 是從pad_packed_sequence恢復(fù)的它的長(zhǎng)度等于 batch 內(nèi)最長(zhǎng)序列。如果你的 mask 在預(yù)處理時(shí)和 src_ids 一起被排序重排過(guò)順序應(yīng)該沒問(wèn)題但如果 mask 是用原始 batch 生成的而 encoder_outputs 是重排后的順序兩者就錯(cuò)位了。解決mask 必須在 collate_fn 里隨著 src_ids 一起重排或者在 Encoder forward 里根據(jù)傳入的 src_lens 重新生成 mask。我習(xí)慣直接在 forward 里生成mask (src_ids ! PAD).unsqueeze(1) # B,1,T 用于廣播另外attention 的 dropout 也很重要。沒做 attention dropout 時(shí)模型容易極度依賴單個(gè)位置的權(quán)重?zé)釄D看起來(lái)就是一根“細(xì)線”或者一片“亂麻”。我一般單獨(dú)給 attention 的 score 加一個(gè) dropout0.1能讓熱圖更平滑也會(huì)小幅提升 ROUGE。5.5 訓(xùn)練太慢pack 沒有按長(zhǎng)度降序?qū)е麓罅繜o(wú)效計(jì)算現(xiàn)象GPU 利用率很低一個(gè) epoch 要跑很久。原因如果沒用 pack_padded_sequence每個(gè) batch 都會(huì)按最長(zhǎng)序列做 paddingbatch 里短樣本占比高時(shí)浪費(fèi)嚴(yán)重或者用了 pack 但 enforce_sortedFalsePyTorch 內(nèi)部要重新排序額外開銷也不小。解決在 collate_fn 里完成降序排序讓pack_padded_sequence走 enforce_sortedTrue 的快路徑。另一個(gè)常見慢點(diǎn)是詞表太大輸出層 Linear(H, V) 的計(jì)算量隨 V 線性增長(zhǎng)可以把詞表的大小壓縮到 2 萬(wàn)以內(nèi)或者用 weight tying 讓 embedding 和輸出層共享權(quán)重。前者簡(jiǎn)單有效后者省顯存但實(shí)現(xiàn)稍復(fù)雜建議優(yōu)先控制詞表。6. 推理驗(yàn)證與進(jìn)階用 beam search 和 ROUGE 替代肉眼抽查6.1 beam search 的 k 怎么選訓(xùn)練完成后greedy 解碼只能作為錯(cuò)誤排查工具真正用來(lái)評(píng)估和交付應(yīng)該用 beam search。k3 是小規(guī)模摘要任務(wù)的穩(wěn)妥起點(diǎn)比 greedy 少很多重復(fù)詞又比 k5 快不少。k5 在這種中文短摘要上收益遞減而且容易生成“過(guò)于順滑”但偏離原意的句子。beam search 實(shí)現(xiàn)時(shí)記得在每個(gè)候選維護(hù)自己的 decoder hidden state不能所有候選共用一份。長(zhǎng)度歸一化建議用“累計(jì) log 概率除以已生成步數(shù)的 0.7 次方”這個(gè)超參數(shù)比 k 本身更影響質(zhì)量k3 配 length_penalty0.7 是我常用的組合。6.2 ROUGE 評(píng)估腳本和注意力可視化評(píng)估摘要任務(wù)ROUGE 比 BLEU 更貼合語(yǔ)義重疊度因?yàn)?ROUGE 衡量的是參考摘要里的 n-gram 有多少被預(yù)測(cè)出來(lái)了BLEU 更偏向機(jī)器翻譯的流暢度。用rouge-score庫(kù)三行就能跑from rouge_score import rouge_scorer scorer rouge_scorer.RougeScorer([rouge1, rouge2, rougeL], use_stemmerTrue) scores scorer.score(reference_summary, predicted_summary) print(scores[rougeL].fmeasure)注意里的四個(gè)指標(biāo)分開看ROUGE-1 反映關(guān)鍵詞覆蓋ROUGE-2 反映短語(yǔ)層面的連貫性ROUGE-L 反映最長(zhǎng)公共子序列和語(yǔ)序。我評(píng)估時(shí)會(huì)額外加一條“是否包含車型關(guān)鍵數(shù)字”的硬規(guī)則比如摘要里“1.6L”“13 款”這類帶數(shù)字的詞組比單純看 ROUGE 更能體現(xiàn)推理能力。注意力可視化則用 matplotlib 畫 heatmap橫軸是輸入 token縱軸是輸出 token顏色越深表示權(quán)重越高。汽車大師摘要里如果模型生成“皮帶”時(shí)對(duì)輸入中“皮帶老化”位置的權(quán)重很深說(shuō)明 attention 確實(shí)在發(fā)揮作用。最后說(shuō)一點(diǎn)教訓(xùn)我早期跑這類 seq2seq 摘要賽題時(shí)總把評(píng)估腳本放到最后一步結(jié)果是每次訓(xùn)練完才發(fā)現(xiàn) beam search 忘寫了長(zhǎng)度懲罰、ROUGE 統(tǒng)計(jì)的參考文件沒對(duì)齊白白浪費(fèi)好幾個(gè)小時(shí)。后來(lái)我把評(píng)估腳本和訓(xùn)練腳本寫在同一份代碼里每?jī)蓚€(gè) epoch 自動(dòng)跑一次驗(yàn)證集 ROUGE-L把“肉眼抽查”變成“指標(biāo)監(jiān)控”模型好壞一眼就能判斷。這個(gè)習(xí)慣幫我在后續(xù)多個(gè)生成任務(wù)里少踩了很多坑也希望幫到你。本文還有配套的精品資源點(diǎn)擊獲取