情感分析實(shí)戰(zhàn):PyTorch實(shí)現(xiàn)文本語(yǔ)音視頻三模態(tài)融合)
簡(jiǎn)介基于PyTorch框架的多模態(tài)情感分析系統(tǒng)源碼包面向希望掌握情感計(jì)算、多模態(tài)特征融合與PyTorch工程實(shí)踐的計(jì)算機(jī)方向開(kāi)發(fā)者。系統(tǒng)針對(duì)文本-圖像配對(duì)數(shù)據(jù)設(shè)計(jì)完成積極、中性、消極三分類(lèi)任務(wù)文本模塊基于預(yù)訓(xùn)練BERT提取語(yǔ)義向量圖像模塊使用輕量神經(jīng)網(wǎng)絡(luò)抽取視覺(jué)特征并將兩類(lèi)特征融合后送入分類(lèi)器工程內(nèi)還包含數(shù)據(jù)預(yù)處理、劃分訓(xùn)練/驗(yàn)證/測(cè)試集、訓(xùn)練循環(huán)、驗(yàn)證評(píng)估與預(yù)測(cè)結(jié)果保存等完整環(huán)節(jié)。壓縮包共22個(gè)文件以7個(gè)Python源碼文件為核心配套JSON格式的標(biāo)注數(shù)據(jù)、JPG圖像樣本、TXT數(shù)據(jù)列表、依賴清單、模型目錄及說(shuō)明文檔整體僅325KB便于快速下載部署。項(xiàng)目結(jié)構(gòu)清晰集中管理參數(shù)并支持命令行靈活調(diào)整訓(xùn)練與測(cè)試設(shè)置。目前已有86人學(xué)習(xí)下載適合入門(mén)多模態(tài)情感分析并搭建可運(yùn)行基線后續(xù)可在源碼基礎(chǔ)上替換骨干模型、增加數(shù)據(jù)增強(qiáng)或擴(kuò)展多任務(wù)。1. 多模態(tài)情感分析系統(tǒng)是什么為什么 PyTorch 成了這套工程的地基多模態(tài)情感分析系統(tǒng)要把同一段話里的文本、語(yǔ)音、視頻幀一起消費(fèi)輸出積極、消極、中性這類(lèi)標(biāo)簽。PyTorch 是這個(gè)方向最常見(jiàn)的落地框架動(dòng)態(tài)計(jì)算圖讓三個(gè)分支各自 forward 后再合并的處理接近直覺(jué)torchaudio、transformers、torchvision 三個(gè)生態(tài)庫(kù)正好覆蓋音頻、語(yǔ)言和圖像三路數(shù)據(jù)。系統(tǒng)要解決的是純文本模型在諷刺、矛盾表達(dá)上翻車(chē)的問(wèn)題——嘴上說(shuō)“挺好的”但語(yǔ)氣低落、畫(huà)面灰暗時(shí)單模態(tài)模型很容易判錯(cuò)多模態(tài)系統(tǒng)靠融合信息拿準(zhǔn)。它適合做輿情監(jiān)測(cè)、客服質(zhì)檢、視頻人物情感分析這類(lèi)任務(wù)。如果你手里已有多模態(tài)標(biāo)注數(shù)據(jù)這套源碼整理出來(lái)的工程方案可以直接作為第二版基線上線。2. 多模態(tài)特征提取與 PyTorch 數(shù)據(jù)管線的三個(gè)分支怎么選2.1 文本分支BERT 還是詞向量中文場(chǎng)景怎么選文本是所有模態(tài)里最容易出信號(hào)的一路工程上建議直接用預(yù)訓(xùn)練編碼器而不是從零訓(xùn)練詞向量。常見(jiàn)做法是加載bert-base-chinese輸入用AutoTokenizer做截?cái)嗪?padding輸出取[CLS]向量作為整句話的語(yǔ)義表示。詞向量加 BiLSTM 的方案只有在推理機(jī)器非常老舊、不能帶 transformer 時(shí)才值得考慮否則它在遇到一個(gè)詞被同音字寫(xiě)錯(cuò)時(shí)會(huì)直接帶偏整句話的情感極性。參數(shù)選擇上有一套默認(rèn)值max_text_len64對(duì)短評(píng)、客服對(duì)話足夠新聞長(zhǎng)文本提到 128paddingmax_length盡量做右側(cè) padding因?yàn)橹形那楦信袛喑R蕾嚲渥游膊緽ERT 是雙向編碼左側(cè) padding 會(huì)干擾序列的語(yǔ)義落點(diǎn)凍結(jié)策略在上線初期很關(guān)鍵剛跑通時(shí)只放開(kāi)第 11 層和 pooler顯存能省接近一半等數(shù)據(jù)量上來(lái)了再打開(kāi)全部層微調(diào)。這里提醒一句AutoTokenizer的詞典對(duì)繁體字支持不好中文語(yǔ)料先做繁轉(zhuǎn)簡(jiǎn)再進(jìn) tokenizer否則會(huì)在錯(cuò)別字和異體字上浪費(fèi)訓(xùn)練時(shí)間。多模態(tài)數(shù)據(jù)集的文本字段也往往帶特殊表情符號(hào)清洗階段先統(tǒng)一去掉或替換成占位符否則 tokenizer 會(huì)產(chǎn)生一堆無(wú)用片段。2.2 音頻分支log-mel 頻譜比原始波形更穩(wěn)音頻分支常見(jiàn)有兩種輸入。一是直接把波形點(diǎn)送進(jìn)一維卷積理論上是信息無(wú)損的但實(shí)際對(duì)采樣率和噪聲很敏感情感短句往往只有一兩秒卷積下采樣后特征不夠穩(wěn)定。二是先算 log-mel 頻譜把波形壓成 64 路梅爾濾波器再送一維卷積這個(gè)路徑在語(yǔ)音情感數(shù)據(jù)集上表現(xiàn)更穩(wěn)也是多數(shù)視頻人物情感分析工程的主流選擇。mel_spec torchaudio.transforms.MelSpectrogram( sample_rate16000, n_fft400, hop_length160, n_mels64) log_mel torch.log(mel_spec(wav) 1e-6)這里的參數(shù)不要隨手改大n_fft400對(duì)應(yīng) 25ms 窗長(zhǎng)hop_length160對(duì)應(yīng) 10ms 幀移這套配置在語(yǔ)音識(shí)別和情感識(shí)別里都是通用配合n_mels64是性價(jià)比比較高的值提到 128 會(huì)讓 mel 行數(shù)翻倍準(zhǔn)確率通常只漲不到 0.5 個(gè)點(diǎn)。1e-6是給 log 運(yùn)算墊底的小量不能省否則靜音段會(huì)算出負(fù)無(wú)窮。音頻的文件格式盡量統(tǒng)一成 wavmp3 在 torchaudio 加載時(shí)容易因編碼器缺失報(bào)錯(cuò)統(tǒng)一轉(zhuǎn)碼是最穩(wěn)的做法。音頻長(zhǎng)度問(wèn)題在第 4 章會(huì)重點(diǎn)講這里先記住一個(gè)原則音頻尾部往往是情緒落點(diǎn)pad 時(shí)寧可截頭部也不要截尾部反過(guò)來(lái)對(duì)視頻抽幀情緒高峰多在中后段所以要均勻抽幀而不是只取開(kāi)頭。2.3 視頻分支直接用預(yù)提取特征別在訓(xùn)練時(shí)跑整份視頻視頻分支最常見(jiàn)的翻車(chē)點(diǎn)是直接把 ResNet 甚至 3D CNN 放進(jìn)模型去抽幀16G 顯存在 batch_size 8 下就會(huì)告急。工程上更常見(jiàn)的是離線階段先抽幀用預(yù)訓(xùn)練模型把每段視頻轉(zhuǎn)成固定維度的向量緩存成.npy或.pt文件訓(xùn)練階段只加載這些向量特征。我一般用 ResNet18 或 CLIP 的 image encoder 做特征抽取每個(gè)視頻均勻抽 8 幀取平均。CLIP 的好處是它的圖像編碼空間和文本語(yǔ)義空間天然對(duì)齊后續(xù)和 BERT 特征融合時(shí)線性投影的負(fù)擔(dān)更小。具體操作注意三件事視頻按均勻時(shí)間間隔抽幀不要只抽第一幀人物表情變化集中在中間段。每段視頻輸出一個(gè)[D]向量D 常見(jiàn)為 512來(lái)自 ResNet18 倒數(shù)第二層或 CLIP ViT-B/32 的 CLS 輸出。特征文件命名必須和音頻、文本樣本一一對(duì)應(yīng)同一個(gè)樣本 id 貫穿三個(gè)模態(tài)否則對(duì)齊會(huì)亂成一團(tuán)。復(fù)雜場(chǎng)景下多模態(tài)情感預(yù)測(cè)的數(shù)學(xué)建模和算法設(shè)計(jì)最后能不能落地很大程度取決于這三路特征在時(shí)間軸上是否對(duì)齊。視頻和音頻必須先按原視頻時(shí)間戳對(duì)齊后再切樣本不能音頻一段是 0 到 3 秒、視頻抽幀是 5 到 8 秒那種數(shù)據(jù)喂給模型只會(huì)得到一堆隨機(jī)噪聲。2.4 融合層級(jí)怎么選直接拼接、加權(quán)求和還是跨模態(tài)注意力融合層的設(shè)計(jì)直接決定這套系統(tǒng)值不值得做。最樸素的是 early fusion把文本[CLS]、音頻特征、視頻特征拼在一起送全連接。實(shí)現(xiàn)最短但三個(gè)模態(tài)特征分布差異大拼接后讓全連接自己學(xué)內(nèi)部關(guān)系小數(shù)據(jù)集上容易過(guò)擬合。更穩(wěn)的是 late fusion每個(gè)模態(tài)先各自過(guò)一個(gè)分類(lèi)頭得到 logits最后對(duì) logits 加權(quán)求和。梯度更新穩(wěn)定但模態(tài)之間沒(méi)有交互遇到“文本說(shuō)開(kāi)心、音頻卻很低沉”這類(lèi)矛盾表達(dá)學(xué)不到聯(lián)合判斷。多模態(tài)融合論文里出現(xiàn)頻率最高、落地也最實(shí)用的還是跨模態(tài)注意力。常見(jiàn)做法以文本[CLS]作為 query視頻特征作為 key/value 進(jìn)入nn.MultiheadAttention讓文本去挑選和它語(yǔ)義相關(guān)的視頻信息再把注意力輸出和音頻特征拼接起來(lái)。這個(gè)結(jié)構(gòu)多不了多少顯存但能額外產(chǎn)出一個(gè)注意力權(quán)重后面做樣本級(jí)排查時(shí)非常好用。融合方式參數(shù)量梯度穩(wěn)定性可解釋性適合場(chǎng)景直接拼接early fusion低一般差數(shù)據(jù)量大、特征維度一致logits 加權(quán)l(xiāng)ate fusion最低最好中數(shù)據(jù)少、快速出基線跨模態(tài)注意力中較好好數(shù)據(jù)量中等以上需要分析樣本三種方式在驗(yàn)證集上差距通常不超過(guò)兩三個(gè)點(diǎn)但有注意力權(quán)重的那一版會(huì)讓你在線上定位“為什么把這條差評(píng)判錯(cuò)”時(shí)多一條路。我建議先出 late fusion 基線再上跨模態(tài)注意力作為正式方案。3. 用 PyTorch 把多模態(tài)情感分析模型跑通Dataset、模型與損失權(quán)重3.1 數(shù)據(jù)集封裝文本、音頻、視頻在同一個(gè) Dataset 里對(duì)齊源碼包里最核心的是 Dataset。它要完成三件事讀出元數(shù)據(jù)里的文本和標(biāo)簽、把音頻轉(zhuǎn)成固定長(zhǎng)度的 log-mel 頻譜、把離線抽好的視頻特征加載進(jìn)來(lái)。以下是我常用的一套結(jié)構(gòu)import json import numpy as np import torch import torchaudio from torch.utils.data import Dataset from transformers import AutoTokenizer class MultimodalSentimentDataset(Dataset): def __init__(self, meta_json, max_text_len64, max_mel_frames96): with open(meta_json, encodingutf-8) as f: self.meta json.load(f) # 每一項(xiàng)含 text / wav_path / video_npy / label self.tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) self.max_text_len max_text_len self.max_mel_frames max_mel_frames self.mel_spec torchaudio.transforms.MelSpectrogram( sample_rate16000, n_fft400, hop_length160, n_mels64) def __len__(self): return len(self.meta) def __getitem__(self, idx): item self.meta[idx] # 1) 文本BERT tokenizer 截?cái)? padding text_ids self.tokenizer( item[text], truncationTrue, max_lengthself.max_text_len, paddingmax_length, return_tensorspt) # 2) 音頻resample - log-mel - 固定幀數(shù) wav, sr torchaudio.load(item[wav_path]) if sr ! 16000: wav torchaudio.functional.resample(wav, sr, 16000) mel torch.log(self.mel_spec(wav) 1e-6).squeeze(0) # [64, T] if mel.shape[1] self.max_mel_frames: mel mel[:, :self.max_mel_frames] # 截尾部會(huì)丟情緒落點(diǎn)慎用 else: mel torch.nn.functional.pad(mel, (0, self.max_mel_frames - mel.shape[1])) # 3) 視頻讀取離線預(yù)提取特征向量已經(jīng)是固定維度 video_feat torch.tensor(np.load(item[video_npy]), dtypetorch.float32) return { text_ids: text_ids[input_ids].squeeze(0), text_mask: text_ids[attention_mask].squeeze(0), mel: mel, video: video_feat, label: torch.tensor(item[label], dtypetorch.long), }這段代碼的關(guān)鍵是三個(gè)返回張量的形狀。text_ids和text_mask是[max_text_len]mel是[64, max_mel_frames]video是[D]。三者并不強(qiáng)制同長(zhǎng)到模型里才在特征維度上相遇所以 Dataset 階段不要強(qiáng)行把它們 pad 成同一種長(zhǎng)度那樣只會(huì)浪費(fèi)讀寫(xiě)時(shí)間。max_mel_frames怎么定按hop_length160算1 秒音頻產(chǎn)生 100 幀頻譜3 秒就是 300 幀。先統(tǒng)計(jì)訓(xùn)練集音頻長(zhǎng)度分布取 95 分位作為這個(gè)參數(shù)比拍腦袋定 96 更穩(wěn)。另一個(gè)常見(jiàn)選擇是離線把所有音頻統(tǒng)一轉(zhuǎn)成固定幀數(shù)的 npy訓(xùn)練時(shí)直接讀矩陣而不現(xiàn)算 mel。這樣能縮短每個(gè) epoch 的時(shí)間但每次改n_mels或hop_length就要重新生成一遍緩存。我一般在小數(shù)據(jù)集上現(xiàn)算數(shù)據(jù)量過(guò)萬(wàn)段再做離線緩存。3.2 模型結(jié)構(gòu)三個(gè)編碼分支加一個(gè)融合分類(lèi)頭模型側(cè)的關(guān)鍵不是把 BERT、CNN、線性投影簡(jiǎn)單堆起來(lái)而是給每個(gè)分支都留一個(gè)分類(lèi)頭讓單模態(tài)也能輸出 logits。這樣訓(xùn)練時(shí)可以同時(shí)監(jiān)督模態(tài)分支避免某個(gè)分支變成看不見(jiàn)的黑匣子。import torch.nn as nn from transformers import BertModel class MultimodalSentimentNet(nn.Module): def __init__(self, num_classes3, video_dim512, text_dim768, audio_dim128): super().__init__() # 文本分支BERT 中文預(yù)訓(xùn)練 self.bert BertModel.from_pretrained(bert-base-chinese) for name, p in self.bert.named_parameters(): # 初期只放開(kāi)最后兩層和 pooler省顯存 if encoder.layer.11 not in name and pooler not in name: p.requires_grad False # 音頻分支一維卷積 全局池化輸入 [B,64,T] self.audio_conv nn.Sequential( nn.Conv1d(64, 128, kernel_size3, padding1), nn.ReLU(), nn.Conv1d(128, 128, kernel_size3, padding1), nn.ReLU(), nn.AdaptiveAvgPool1d(1), ) self.audio_proj nn.Linear(128, audio_dim) # 視頻分支離線特征已是向量線性投影到文本維度 self.video_proj nn.Linear(video_dim, text_dim) # 跨模態(tài)注意力文本為 query視頻為 key/value self.cross_attn nn.MultiheadAttention(text_dim, num_heads4, batch_firstTrue) # 四個(gè)分類(lèi)頭 self.text_head nn.Linear(text_dim, num_classes) self.audio_head nn.Linear(audio_dim, num_classes) self.video_head nn.Linear(text_dim, num_classes) self.fusion_head nn.Sequential( nn.Linear(text_dim audio_dim text_dim, 256), nn.ReLU(), nn.Dropout(0.3), nn.Linear(256, num_classes), ) def forward(self, text_ids, text_mask, mel, video): bert_out self.bert(text_ids, attention_masktext_mask).last_hidden_state text_cls bert_out[:, 0] text_logits self.text_head(text_cls) audio_feat self.audio_conv(mel).squeeze(-1) # [B,128] audio_feat self.audio_proj(audio_feat) audio_logits self.audio_head(audio_feat) video_feat self.video_proj(video) # [B,768] video_logits self.video_head(video_feat) attn_out, attn_weight self.cross_attn( text_cls.unsqueeze(1), video_feat.unsqueeze(1), video_feat.unsqueeze(1)) fusion_in torch.cat([text_cls, audio_feat, attn_out.squeeze(1)], dim-1) fusion_logits self.fusion_head(fusion_in) return { fusion_logits: fusion_logits, text_logits: text_logits, audio_logits: audio_logits, video_logits: video_logits, attn_weight: attn_weight, }音頻分支的AdaptiveAvgPool1d(1)會(huì)把任意長(zhǎng)度 mel 壓成[B,128,1]訓(xùn)練和推理時(shí)允許音頻長(zhǎng)短不一進(jìn)入模型這是對(duì)付音頻長(zhǎng)度抖動(dòng)的一個(gè)很實(shí)用的手段比大段 pad 邏輯省事得多。視頻分支只有一個(gè)Linear因?yàn)轭A(yù)處理階段已經(jīng)把整段視頻濃縮成了單一向量如果你選擇每段抽 8 幀、輸出[8,512]的特征就需要在進(jìn)入融合前先做幀間聚合。直接 mean pooling 已經(jīng)能打好基線我沒(méi)有在初版里加幀級(jí)時(shí)序建模。凍結(jié)參數(shù)的邏輯要看清named_parameters()里的encoder.layer.11會(huì)同時(shí)命中第 11 層的子參數(shù)pooler是 BERT 的池化層習(xí)慣上和第 11 層一起放開(kāi)。初期這樣為了省顯存跑通后建議逐步解凍全部層重新訓(xùn)練。提示正式訓(xùn)練前先跑一個(gè) batch 的前向確認(rèn)返回 dict 里每個(gè) key 的 shape 都符合預(yù)期再進(jìn)循環(huán)。多模態(tài)模型的報(bào)錯(cuò)信息經(jīng)常在多層嵌套里被吞掉前向驗(yàn)證能省半天時(shí)間。3.3 訓(xùn)練循環(huán)融合主損失加三個(gè)分支輔助損失訓(xùn)練時(shí)不能只算融合 logits 的交叉熵三個(gè)分支的 loss 也要帶進(jìn)來(lái)這就是多模態(tài)項(xiàng)目里常說(shuō)的 loss weighting。完整代碼from torch.utils.data import DataLoader from torch.optim import AdamW device torch.device(cuda if torch.cuda.is_available() else cpu) model MultimodalSentimentNet(num_classes3).to(device) train_loader DataLoader(dataset, batch_size16, shuffleTrue, num_workers2) optimizer AdamW([p for p in model.parameters() if p.requires_grad], lr2e-5) criterion nn.CrossEntropyLoss() aux_weight 0.3 # 分支輔助損失權(quán)重正式實(shí)驗(yàn)前先跑一組 0 做對(duì)比 for epoch in range(15): model.train() epoch_loss 0.0 for batch in train_loader: text_ids batch[text_ids].to(device) text_mask batch[text_mask].to(device) mel batch[mel].to(device) video batch[video].to(device) label batch[label].to(device) out model(text_ids, text_mask, mel, video) loss criterion(out[fusion_logits], label) loss loss aux_weight * criterion(out[text_logits], label) loss loss aux_weight * criterion(out[audio_logits], label) loss loss aux_weight * criterion(out[video_logits], label) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() epoch_loss loss.item() print(fepoch{epoch:02d} train_loss{epoch_loss / len(train_loader):.4f})aux_weight0.3的含義是讓融合頭為主導(dǎo)三個(gè)分支保持“能單獨(dú)判斷但不帶偏主路”的狀態(tài)。如果某個(gè)模態(tài)數(shù)據(jù)噪聲特別大例如音頻里大量靜音或環(huán)境音把對(duì)應(yīng)分支的權(quán)重降到 0.1如果某個(gè)分支 loss 明顯不降說(shuō)明該分支特征質(zhì)量太差先修數(shù)據(jù)不要硬調(diào) lr。clip_grad_norm_(max_norm1.0)在融合模型里幾乎必加。BERT 和隨機(jī)初始化的 CNN 分支梯度尺度能差一到兩個(gè)數(shù)量級(jí)不加裁剪前幾個(gè) step 一次大梯度就可能毀掉整個(gè)融合頭。AdamW 的lr2e-5是 BERT 微調(diào)常用值但要看到新初始化的層學(xué)得慢常見(jiàn)做法是給音頻 CNN 和視頻投影單獨(dú)開(kāi)lr5e-4BERT 保持2e-5。想快速出第一個(gè)有效模型先把 BERT 全部?jī)鼋Y(jié)只訓(xùn)分支和融合頭穩(wěn)定后再解凍微調(diào)。3.4 保存、加載與驗(yàn)證拆分保存時(shí)統(tǒng)一存state_dict不要存整個(gè)模型對(duì)象。加載前先創(chuàng)建同結(jié)構(gòu)的模型實(shí)例再load_state_dict。驗(yàn)證拆分最好提前固定一個(gè)val_loader每個(gè) epoch 結(jié)束用驗(yàn)證集算一次準(zhǔn)確率然后做最樸素的早停連續(xù) 3 個(gè) epoch 驗(yàn)證指標(biāo)不漲就回退到之前最好的 checkpoint。best_val_acc 0.0 patience 0 for epoch in range(15): # 訓(xùn)練循環(huán)略... val_acc evaluate(model, val_loader, device) if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_model.pt) patience 0 else: patience 1 if patience 3: breakmap_location在加載時(shí)寫(xiě)清楚torch.load(best_model.pt, map_locationdevice)否則在 CPU 機(jī)器上訓(xùn)練、GPU 上推理時(shí)會(huì)因?yàn)槿鄙?CUDA 報(bào)錯(cuò)。checkpoint 文件名帶 epoch 號(hào)每次保留前一個(gè)版本晚點(diǎn)你就知道這有多重要。4. 多模態(tài)情感分析避坑記錄五個(gè)常見(jiàn)報(bào)錯(cuò)與排查順序4.1 現(xiàn)象DataLoader 在 batch 階段報(bào) shape 不一致的 RuntimeError常見(jiàn)于換了一批數(shù)據(jù)后某條樣本音頻較長(zhǎng)mel 計(jì)算出來(lái)超過(guò)max_mel_frames但__getitem__里沒(méi)有走截?cái)噙壿嫽蛘咭曨l特征文件輸出的是[8,512]而不是預(yù)期的[512]batch 內(nèi)維度對(duì)不上。原因Dataset 返回張量形狀必須一致默認(rèn)collate_fn用torch.stack堆疊任何一維不一致都直接崩。而且 DataLoader 的 worker 進(jìn)程會(huì)吞掉原始 traceback只在終端打一行 RuntimeError。解決在__getitem__里對(duì)音頻和視頻各做一次 shape 防御先打印當(dāng)前樣本 id 和實(shí)際 shape再在__init__里預(yù)掃描所有wav_path和video_npy把形狀異常的樣本直接過(guò)濾掉。這比在訓(xùn)練循環(huán)里 try/except 高效得多。4.2 現(xiàn)象訓(xùn)練正常每次跑驗(yàn)證的 epoch 結(jié)束就 CUDA OOM訓(xùn)練 loss 在下降一到驗(yàn)證就CUDA out of memory。原因驗(yàn)證代碼沒(méi)有包在torch.no_grad()里。model.eval()只改了 dropout 和 BN 行為沒(méi)有關(guān)梯度驗(yàn)證一樣保留計(jì)算圖。另外驗(yàn)證集 DataLoader 仍shuffleTruebatch_size 又和訓(xùn)練一致峰值顯存自然更高。解決驗(yàn)證循環(huán)統(tǒng)一寫(xiě)成model.eval() with torch.no_grad(): # 推理循環(huán)驗(yàn)證 DataLoader 設(shè)shuffleFalsebatch_size 降為訓(xùn)練的一半。還 OOM 就把訓(xùn)練 batch_size 從 16 降到 8再看max_mel_frames。判斷順序先降 batch_size再降音頻幀數(shù)因?yàn)?BERT 激活值占的顯存往往比音頻更大。4.3 現(xiàn)象多模態(tài)融合后的驗(yàn)證指標(biāo)不如單獨(dú)的文本分支這是最打擊人的一種加了音頻和視頻F1 反而比只跑 BERT 還低。原因三個(gè)分支收斂速度不匹配。BERT 是大容量預(yù)訓(xùn)練模型第一輪就能學(xué)到文本信息音頻 CNN 和視頻線性層隨機(jī)初始化早期輸出接近噪聲而融合頭被這些噪聲帶著跑偏。另一個(gè)常見(jiàn)原因是輔助 loss 權(quán)重配比失調(diào)某支噪聲大的分支把融合頭拖住了。解決前兩輪先用aux_weight0只看融合 loss等文本分支穩(wěn)定后再加回輔助權(quán)重。損失權(quán)重建議從w_text0.3, w_audio0.2, w_video0.1起步融合主 loss 權(quán)重保持 1。多模態(tài)模型最忌諱一上來(lái)三個(gè)分支完全平等數(shù)據(jù)質(zhì)量不均衡時(shí)平等訓(xùn)練等于讓短板帶節(jié)奏。4.4 現(xiàn)象import torch后cuda.is_available()為 False加載時(shí)報(bào) CUDA 相關(guān)錯(cuò)誤這個(gè)坑多數(shù)出現(xiàn)在 Anaconda 環(huán)境里。常見(jiàn)流程是先pip install torch裝了 CPU 版后面又用 conda 混裝同一環(huán)境出現(xiàn)多個(gè) torch。原因PyTorch 的 CUDA 支持是通過(guò)預(yù)編譯 wheel 綁定的pip 默認(rèn)裝 CPU 版本conda 如果 torch 和 cudatoolkit 版本不匹配也會(huì)識(shí)別不到卡。解決不要在原環(huán)境里修直接建一個(gè)干凈環(huán)境conda create -n mmsa python3.9 conda activate mmsa conda install pytorch torchvision torchaudio cudatoolkit11.8 -c pytorch -c conda-forge python -c import torch; print(torch.__version__, torch.cuda.is_available())裝完立刻打印驗(yàn)證。Apple 芯片機(jī)器不要裝帶 cu 的 wheel直接用 pip 安裝 macOS 版 PyTorch。CUDA 驅(qū)動(dòng)太舊時(shí)裝新 wheel 也會(huì)在第一步驗(yàn)證處失敗先用nvidia-smi看驅(qū)動(dòng)支持的 CUDA 版本再選對(duì)應(yīng)包。4.5 現(xiàn)象load_state_dict報(bào) missing key 或 unexpected key報(bào)錯(cuò)內(nèi)容常見(jiàn)是Missing key(s): bert.embeddings.word_embeddings.weight或者Unexpected key(s): module.classifier.0.weight。原因訓(xùn)練時(shí)用了nn.DataParallel或 DDP 包裝過(guò)模型保存的權(quán)重帶module.前綴加載的模型是裸模型鍵名對(duì)不上。另一類(lèi)是保存時(shí)不小心存了整個(gè) model 對(duì)象而不是state_dict()。解決統(tǒng)一按state_dict保存加載時(shí)做一次鍵名清理ckpt torch.load(best_model.pt, map_locationcpu) if module. in list(ckpt.keys())[0]: ckpt {k.replace(module., ): v for k, v in ckpt.items()} model.load_state_dict(ckpt)同時(shí)把best_epoch、損失權(quán)重、凍結(jié)層配置單獨(dú)寫(xiě)一份 json 存起來(lái)。多個(gè)實(shí)驗(yàn)之間靠文件名猜配置是后期最耗時(shí)間的事。5. 模型驗(yàn)證與上線前的最后一課用注意力權(quán)重給預(yù)測(cè)上“后悔藥”5.1 先看分類(lèi)報(bào)告別只看準(zhǔn)確率多模態(tài)情感數(shù)據(jù)大多數(shù)不平衡中文評(píng)論尤其明顯中性樣本可能占五成正負(fù)樣本各兩成。驗(yàn)證階段直接打印分類(lèi)報(bào)告不要只盯準(zhǔn)確率from sklearn.metrics import classification_report preds, all_labels [], [] model.eval() with torch.no_grad(): for batch in test_loader: out model(batch[text_ids].to(device), batch[text_mask].to(device), batch[mel].to(device), batch[video].to(device)) preds.extend(out[fusion_logits].argmax(-1).cpu().tolist()) all_labels.extend(batch[label].tolist()) print(classification_report(all_labels, preds, digits4))如果某一個(gè)類(lèi)別 recall 明顯低不要急著調(diào)閾值先看這一類(lèi)樣本是否集中在長(zhǎng)句、音頻帶噪或視頻抽幀不均這幾類(lèi)原因上。5.2 用注意力權(quán)重做樣本級(jí)降級(jí)判斷模型結(jié)構(gòu)里除了 logits 還返回attn_weight它是[B,1,1]的權(quán)重表示文本從視頻里挑出的信息量大小??梢宰鲆粚咏导?jí)保險(xiǎn)融合預(yù)測(cè)置信度低、文本分支置信度高時(shí)信任文本分支fusion_prob torch.softmax(fusion_logits, dim-1).max(dim-1).values text_prob torch.softmax(text_logits, dim-1).max(dim-1).values final_pred fusion_logits.argmax(dim-1) fallback (fusion_prob 0.6) (text_prob 0.8) final_pred[fallback] text_logits[fallback].argmax(dim-1)在驗(yàn)證集上跑一版統(tǒng)計(jì)有多少樣本走了 fallback。如果命中率只有 2% 到 3%說(shuō)明視頻或音頻特征確實(shí)在個(gè)別樣本上拖后腿值得繼續(xù)優(yōu)化如果超過(guò) 10%說(shuō)明融合頭本身沒(méi)訓(xùn)好要回去重調(diào)損失權(quán)重而不是依賴這條兜底邏輯。5.3 上線前的部署檢查固定輸入長(zhǎng)度再導(dǎo) ONNX本地驗(yàn)證通過(guò)后如果目標(biāo)是 CPU 服務(wù)推薦把模型導(dǎo)出成 ONNX 再用 ONNX Runtime 推理。導(dǎo)出前有三個(gè)準(zhǔn)備工作模型置 eval、輸入張量全部固定形狀、文本分支換掉。BERT 是推理耗時(shí)的大頭把bert-base-chinese換成 DistilBERT 中文版本整個(gè)模型推理時(shí)間能壓到原來(lái)的三分之一。ONNX 導(dǎo)出要求輸入動(dòng)態(tài)軸盡量少否則引擎會(huì)頻繁使用動(dòng)態(tài) shape速度反而更差。最后說(shuō)一個(gè)我踩過(guò)好幾輪的教訓(xùn)。多模態(tài)模型的變量實(shí)在太多三個(gè)分支要不要分開(kāi)設(shè)學(xué)習(xí)率、損失權(quán)重怎么配、checkpoint 存的是哪個(gè) epoch、凍結(jié)了哪幾層。不寫(xiě)配置文件的實(shí)驗(yàn)就是黑匣子。我現(xiàn)在訓(xùn)練腳本第一行就讀一個(gè) config.json把所有參數(shù)連同隨機(jī)種子一起記檔每次實(shí)驗(yàn)復(fù)制一份帶時(shí)間戳的配置備份。遇到指標(biāo)退步時(shí)能快速回滾到上一個(gè)能用的權(quán)重而不是靠記憶猜。希望幫到你。本文還有配套的精品資源點(diǎn)擊獲取