實(shí)戰(zhàn):準(zhǔn)確率提升14%與避坑指南)
1. 為什么我選擇在 AMD ROCm 云上折騰 Gemma4 情緒 LoRA先說結(jié)論這次實(shí)驗(yàn)的起點(diǎn)其實(shí)很樸素——我手里有一個(gè)情緒分類任務(wù)數(shù)據(jù)量不大大概幾千條帶標(biāo)注的短文本標(biāo)簽是六類情緒。用 GPT 類 API 跑推理當(dāng)然可以但成本隨調(diào)用量線性上漲而且延遲不可控。于是我想試試用開源模型做 LoRA 微調(diào)把一個(gè)小模型調(diào)到“夠用”的水平然后自己部署。選 Gemma4 的原因很直接它的底座能力在同參數(shù)量級里比較均衡指令跟隨和語義理解都不差而且社區(qū)里已經(jīng)有比較成熟的 transformers 加載方案。選 LoRA 而不是全量微調(diào)是因?yàn)槲业臄?shù)據(jù)量撐不起全參數(shù)訓(xùn)練全量微調(diào)不僅顯存吃緊還容易過擬合。LoRA 只訓(xùn)練低秩適配矩陣參數(shù)量能壓到原模型的百分之幾甚至千分之幾訓(xùn)練快、顯存省、產(chǎn)物小非常適合我這種“單卡、小數(shù)據(jù)、快速迭代”的場景。那為什么是 AMD ROCm 云而不是更常見的 CUDA 環(huán)境坦白說一開始是出于成本考慮。我對比了幾家云廠商的 GPU 實(shí)例價(jià)格同等級別顯存下AMD 的 MI 系列實(shí)例單價(jià)確實(shí)更有吸引力。但真正讓我下定決心的是想驗(yàn)證一件事ROCm 生態(tài)到底能不能撐起一次完整的 LoRA 微調(diào)流程。網(wǎng)上關(guān)于 ROCm 的討論很多停留在“能跑推理”或者“裝環(huán)境很痛苦”的層面真正把微調(diào)全流程跑通并給出準(zhǔn)確率對比的案例并不多。我想自己踩一遍把坑記下來。這里先交代一下我的實(shí)驗(yàn)配置方便你對照復(fù)現(xiàn)項(xiàng)目配置云平臺AMD ROCm 云實(shí)例GPUAMD Instinct MI 系列顯存 48GB 級別ROCm 版本6.xPython3.10核心框架PyTorch (ROCm 版) transformers peft底座模型Gemma4 指令版微調(diào)方法LoRA (r8, alpha16)任務(wù)六分類情緒識別訓(xùn)練數(shù)據(jù)約 4000 條短文本評估指標(biāo)準(zhǔn)確率 (accuracy)最終結(jié)果微調(diào)前基線準(zhǔn)確率 0.594微調(diào)后 0.734提升了 14 個(gè)百分點(diǎn)。這個(gè)提升幅度不算驚艷但對于一個(gè)幾千條數(shù)據(jù)的小任務(wù)來說已經(jīng)足夠說明 LoRA 在這個(gè)底座上是有效的。下面我把整個(gè)流程拆開講包括我踩的四個(gè)坑。2. 環(huán)境搭建ROCm 云上的第一道坎2.1 ROCm 環(huán)境確認(rèn)與 PyTorch 安裝拿到云實(shí)例后第一件事不是急著裝 transformers而是確認(rèn) ROCm 本身是否正常。很多人一上來就 pip install結(jié)果后面報(bào)錯(cuò)根本分不清是 ROCm 沒配好還是 Python 包沖突。先跑這兩條命令rocm-smi rocminfo | grep -i gfxrocm-smi會(huì)列出 GPU 的顯存占用、溫度、功耗等信息。如果這條命令都跑不出來后面不用繼續(xù)了先找云廠商確認(rèn)驅(qū)動(dòng)。rocminfo里的 gfx 架構(gòu)代號很關(guān)鍵比如 gfx90a、gfx942 之類它決定了你后面裝 PyTorch 時(shí)要用哪個(gè)版本的 wheel。確認(rèn) ROCm 正常后裝 PyTorch 的 ROCm 版本。注意不要用默認(rèn)的 pip 源裝 torch那樣裝出來的是 CUDA 版或者 CPU 版。正確做法是去 PyTorch 官網(wǎng)找對應(yīng) ROCm 版本的安裝命令類似pip install torch torchvision --index-url https://download.pytorch.org/whl/rocm6.x裝完驗(yàn)證import torch print(torch.__version__) print(torch.cuda.is_available()) # ROCm 環(huán)境下這個(gè)返回 True print(torch.cuda.get_device_name(0))這里有個(gè)容易困惑的點(diǎn)ROCm 版的 PyTorch 依然使用torch.cuda這個(gè)命名空間這是歷史遺留不代表它在用 CUDA。只要is_available()返回 True 且設(shè)備名是你的 AMD 卡就說明環(huán)境通了。注意ROCm 版本、PyTorch 版本、gfx 架構(gòu)三者必須匹配。我見過有人用 gfx90a 的卡裝了只支持 gfx942 的 wheel結(jié)果is_available()一直是 False排查了半天。2.2 依賴安裝順序與版本鎖定環(huán)境通了之后裝 transformers、peft、datasets、accelerate 這幾個(gè)核心包。我的建議是先把版本鎖死不要用最新版。原因是 ROCm 生態(tài)的兼容性窗口比 CUDA 窄最新版 transformers 可能引入了某些算子在 ROCm 上還沒適配。我這次用的組合大致是pip install transformers4.4x.x pip install peft0.1x.x pip install datasets accelerate具體小版本號我建議你根據(jù)自己底座模型的要求去查但原則是transformers 和 peft 的版本要互相兼容peft 的版本要支持你用的模型架構(gòu)。裝完之后跑一個(gè)最小加載測試from transformers import AutoModelForCausalLM, AutoTokenizer model AutoModelForCausalLM.from_pretrained(google/gemma-4-xxx, device_mapauto) tokenizer AutoTokenizer.from_pretrained(google/gemma-4-xxx) print(model.device)如果這一步能順利把模型加載到 GPU 上說明環(huán)境基本沒問題。如果報(bào)顯存不足先檢查是不是用了device_mapauto但顯存被其他進(jìn)程占了。3. 數(shù)據(jù)準(zhǔn)備與 LoRA 配置的核心細(xì)節(jié)3.1 情緒數(shù)據(jù)的格式化處理我的原始數(shù)據(jù)是 CSV兩列text 和 label。label 是六個(gè)情緒類別。做指令微調(diào)時(shí)不能直接把 text 丟進(jìn)去要構(gòu)造成指令格式。我用的模板大致是def format_sample(text, label): instruction 請判斷以下文本的情緒類別只輸出類別名稱。 return { prompt: f{instruction}\n文本{text}\n情緒, response: label }這里有個(gè)細(xì)節(jié)response 只放類別名稱不要放解釋。因?yàn)槲业脑u估是精確匹配如果模型輸出“這段文本的情緒是開心”那和“開心”就不匹配了。訓(xùn)練時(shí)讓模型學(xué)會(huì)只輸出標(biāo)簽推理時(shí)再做后處理提取。數(shù)據(jù)劃分上我按 8:1:1 分訓(xùn)練、驗(yàn)證、測試。驗(yàn)證集用來監(jiān)控訓(xùn)練過程中的過擬合測試集只在最后評估一次。很多人會(huì)把驗(yàn)證集和測試集混用導(dǎo)致最終指標(biāo)虛高。3.2 LoRA 參數(shù)怎么選r、alpha、target_modulesLoRA 的核心參數(shù)有三個(gè)秩 r、縮放系數(shù) alpha、以及作用在哪些模塊上。r 決定低秩矩陣的維度。r 越大可訓(xùn)練參數(shù)越多擬合能力越強(qiáng)但過擬合風(fēng)險(xiǎn)也越高。我的數(shù)據(jù)量只有幾千條所以選了 r8。如果你數(shù)據(jù)量上萬可以試 r16 或 32。alpha 一般設(shè)為 r 的兩倍我設(shè) alpha16。alpha/r 的比值影響適配矩陣的縮放這個(gè)比值比絕對值更重要。target_modules 是最容易被忽略的參數(shù)。Gemma 這類模型里注意力層的 q_proj、k_proj、v_proj、o_proj 是常見選擇。我一開始只加了 q_proj 和 v_proj結(jié)果準(zhǔn)確率只到 0.65 左右。后來把 o_proj 也加進(jìn)去才到了 0.73。原因是輸出投影層也承載了語義信息只調(diào) qv 不夠。from peft import LoraConfig, get_peft_model lora_config LoraConfig( r8, lora_alpha16, target_modules[q_proj, k_proj, v_proj, o_proj], lora_dropout0.05, biasnone, task_typeCAUSAL_LM ) model get_peft_model(model, lora_config) model.print_trainable_parameters()print_trainable_parameters()會(huì)告訴你可訓(xùn)練參數(shù)占比。我這次大概是 0.3% 左右非常輕量。實(shí)操心得target_modules 不要憑感覺加。先用默認(rèn)的 qv 跑一版看驗(yàn)證集準(zhǔn)確率再逐步加模塊。每加一個(gè)模塊訓(xùn)練時(shí)間會(huì)增加但收益不一定線性。4. 訓(xùn)練過程與四個(gè)坑的完整記錄4.1 坑一ROCm 上的 flash attention 不可用第一個(gè)坑出現(xiàn)在訓(xùn)練啟動(dòng)階段。我原本想用 flash attention 加速因?yàn)?transformers 里可以通過attn_implementationflash_attention_2開啟。結(jié)果在 ROCm 上報(bào)錯(cuò)說找不到對應(yīng)的算子。原因很簡單flash attention 的 ROCm 適配版本和 CUDA 版本不是一回事很多預(yù)編譯 wheel 只覆蓋 CUDA。解決辦法是退回默認(rèn)的 eager attention或者用 ROCm 社區(qū)維護(hù)的 flash attention 分支。我為了省事直接用默認(rèn)實(shí)現(xiàn)訓(xùn)練速度慢一些但穩(wěn)定。model AutoModelForCausalLM.from_pretrained( model_name, device_mapauto, attn_implementationeager # ROCm 上先別開 flash )這個(gè)坑的教訓(xùn)是ROCm 生態(tài)里很多 CUDA 上的“默認(rèn)優(yōu)化”并不默認(rèn)可用。遇到算子缺失先退回基礎(chǔ)實(shí)現(xiàn)跑通再考慮優(yōu)化。4.2 坑二混合精度訓(xùn)練在 ROCm 上的表現(xiàn)差異第二個(gè)坑是混合精度。CUDA 上大家習(xí)慣用 fp16 或 bf16 做混合精度訓(xùn)練省顯存又提速。我在 ROCm 上直接開 fp16結(jié)果 loss 出現(xiàn) NaN。排查后發(fā)現(xiàn)ROCm 對 fp16 的支持在某些算子上有差異尤其是 softmax 和 layernorm 相關(guān)。換成 bf16 后問題消失。bf16 的動(dòng)態(tài)范圍比 fp16 大不容易溢出在 AMD 卡上兼容性更好。from transformers import TrainingArguments training_args TrainingArguments( output_dir./output, per_device_train_batch_size4, gradient_accumulation_steps4, learning_rate2e-4, num_train_epochs3, bf16True, # ROCm 上優(yōu)先用 bf16 fp16False, logging_steps20, eval_strategysteps, eval_steps100, save_strategyepoch, report_tonone )注意如果你的卡不支持 bf16那就只能用 fp16但要加 loss scaling。ROCm 下 loss scaling 的配置和 CUDA 略有不同建議先小規(guī)模試跑。4.3 坑三batch size 與梯度累積的顯存平衡第三個(gè)坑是顯存。我一開始把 per_device_train_batch_size 設(shè)成 8結(jié)果 OOM。降到 4 還是緊張最后用 batch_size4 加 gradient_accumulation_steps4等效 batch size 16。這里要理解一個(gè)概念LoRA 雖然只訓(xùn)練少量參數(shù)但前向傳播和激活值依然要占顯存。激活值的大小和 batch size、序列長度成正比。我的文本平均長度 128 token最長 256所以序列長度設(shè) 256。如果文本更長顯存壓力會(huì)明顯上升。顯存估算的粗略公式是模型權(quán)重 激活值 優(yōu)化器狀態(tài)。LoRA 的優(yōu)化器狀態(tài)只針對適配矩陣很小所以大頭是權(quán)重和激活值。48GB 顯存跑 Gemma4 這個(gè)量級batch size 4 到 8 是比較穩(wěn)的區(qū)間。4.4 坑四評估指標(biāo)的計(jì)算方式導(dǎo)致虛高第四個(gè)坑最隱蔽。我一開始用訓(xùn)練框架自帶的 evaluation它計(jì)算的是 token 級別的 loss不是準(zhǔn)確率。loss 下降不代表分類準(zhǔn)確率上升。后來我自己寫了評估函數(shù)對驗(yàn)證集逐條推理提取輸出標(biāo)簽和真實(shí)標(biāo)簽比對。def evaluate(model, tokenizer, dataset): correct 0 total 0 for sample in dataset: inputs tokenizer(sample[prompt], return_tensorspt).to(model.device) with torch.no_grad(): outputs model.generate(**inputs, max_new_tokens8) pred tokenizer.decode(outputs[0], skip_special_tokensTrue) pred_label extract_label(pred) if pred_label sample[response]: correct 1 total 1 return correct / total這個(gè)坑的教訓(xùn)是微調(diào)任務(wù)的評估指標(biāo)必須和業(yè)務(wù)目標(biāo)對齊。分類任務(wù)就看分類準(zhǔn)確率不要被 loss 曲線迷惑。5. 結(jié)果對比與效果分析5.1 基線 vs 微調(diào)后的準(zhǔn)確率微調(diào)前我用 Gemma4 底座直接做 zero-shot 推理準(zhǔn)確率 0.594。這個(gè)數(shù)字說明底座模型對情緒分類有一定理解但不夠精準(zhǔn)尤其是中性、驚訝、厭惡這幾類容易混。微調(diào)后測試集準(zhǔn)確率 0.734。分類別看情緒類別微調(diào)前微調(diào)后開心0.780.86悲傷0.710.82憤怒0.690.80中性0.450.62驚訝0.520.68厭惡0.410.61提升最明顯的是中性、驚訝、厭惡這三類正好是基線表現(xiàn)最差的。說明 LoRA 確實(shí)學(xué)到了數(shù)據(jù)里的判別邊界而不是只強(qiáng)化了原本就會(huì)的類別。5.2 訓(xùn)練曲線與過擬合判斷訓(xùn)練 loss 從 1.2 降到 0.4 左右驗(yàn)證 loss 在前兩個(gè) epoch 下降第三個(gè) epoch 開始持平甚至微升。這是典型的過擬合信號。我最終選了第二個(gè) epoch 的 checkpoint而不是最后一個(gè)。判斷過擬合不能只看 loss還要看驗(yàn)證集準(zhǔn)確率。我的驗(yàn)證準(zhǔn)確率在第二個(gè) epoch 達(dá)到峰值 0.72第三個(gè) epoch 掉到 0.70。所以早停是必要的。實(shí)操心得LoRA 雖然參數(shù)少但小數(shù)據(jù)下依然會(huì)過擬合。建議每個(gè) epoch 都存 checkpoint最后用驗(yàn)證集挑最好的不要默認(rèn)用最后一個(gè)。6. 常見問題速查與避坑清單6.1 ROCm 環(huán)境類問題問題可能原因解決方向torch.cuda.is_available() 為 FalsePyTorch 裝成 CUDA/CPU 版重裝 ROCm 版 wheel算子找不到flash attention 未適配退回 eager attentionloss 出現(xiàn) NaNfp16 溢出換 bf16 或加 loss scaling顯存 OOMbatch size 過大降 batch加梯度累積訓(xùn)練極慢未啟用優(yōu)化算子檢查 ROCm 版本與 PyTorch 匹配6.2 LoRA 配置類問題target_modules 選少了模型學(xué)不動(dòng)選多了訓(xùn)練變慢且容易過擬合。我的建議是從 qv 開始逐步加 k、o。r 和 alpha 不要同時(shí)調(diào)先固定 alpha2r只調(diào) r。數(shù)據(jù)格式上prompt 和 response 的分隔要清晰避免模型把指令也當(dāng)成要生成的內(nèi)容。評估時(shí)一定要做輸出解析不能直接拿生成文本比對。6.3 我個(gè)人的避坑清單第一環(huán)境沒驗(yàn)證通過之前不要碰數(shù)據(jù)。第二先跑一個(gè) 100 條的小子集確認(rèn)整個(gè)流程能走通再上全量。第三每個(gè) epoch 存 checkpoint別省這點(diǎn)磁盤。第四評估函數(shù)自己寫不要完全依賴框架默認(rèn)。第五ROCm 上遇到問題先查 gfx 架構(gòu)和版本匹配再查代碼。這套流程跑下來我對 ROCm 做 LoRA 微調(diào)的信心是有的。它不像 CUDA 那么“開箱即用”但把版本和環(huán)境理順之后穩(wěn)定性是可以接受的。后面我打算試試更大的 r 和更多 target_modules看看準(zhǔn)確率還有沒有上升空間。