☰
AMD ROCm 云上 Gemma4 情绪 LoRA 微调实战:准确率提升14%与避坑指南
2026/10/1 3:15:12 网站建设 项目流程

1. 为什么我选择在 AMD ROCm 云上折腾 Gemma4 情绪 LoRA

先说结论:这次实验的起点其实很朴素——我手里有一个情绪分类任务,数据量不大,大概几千条带标注的短文本,标签是六类情绪。用 GPT 类 API 跑推理当然可以,但成本随调用量线性上涨,而且延迟不可控。于是我想试试用开源模型做 LoRA 微调,把一个小模型调到“够用”的水平,然后自己部署。

选 Gemma4 的原因很直接:它的底座能力在同参数量级里比较均衡,指令跟随和语义理解都不差,而且社区里已经有比较成熟的 transformers 加载方案。选 LoRA 而不是全量微调,是因为我的数据量撑不起全参数训练,全量微调不仅显存吃紧,还容易过拟合。LoRA 只训练低秩适配矩阵,参数量能压到原模型的百分之几甚至千分之几,训练快、显存省、产物小,非常适合我这种“单卡、小数据、快速迭代”的场景。

那为什么是 AMD ROCm 云而不是更常见的 CUDA 环境?坦白说,一开始是出于成本考虑。我对比了几家云厂商的 GPU 实例价格,同等级别显存下,AMD 的 MI 系列实例单价确实更有吸引力。但真正让我下定决心的是想验证一件事:ROCm 生态到底能不能撑起一次完整的 LoRA 微调流程。网上关于 ROCm 的讨论,很多停留在“能跑推理”或者“装环境很痛苦”的层面,真正把微调全流程跑通并给出准确率对比的案例并不多。我想自己踩一遍,把坑记下来。

这里先交代一下我的实验配置,方便你对照复现:

项目配置
云平台AMD ROCm 云实例
GPUAMD Instinct MI 系列(显存 48GB 级别)
ROCm 版本6.x
Python3.10
核心框架PyTorch (ROCm 版) + transformers + peft
底座模型Gemma4 指令版
微调方法LoRA (r=8, alpha=16)
任务六分类情绪识别
训练数据约 4000 条短文本
评估指标准确率 (accuracy)

最终结果:微调前基线准确率 0.594,微调后 0.734,提升了 14 个百分点。这个提升幅度不算惊艳,但对于一个几千条数据的小任务来说,已经足够说明 LoRA 在这个底座上是有效的。下面我把整个流程拆开讲,包括我踩的四个坑。

2. 环境搭建:ROCm 云上的第一道坎

2.1 ROCm 环境确认与 PyTorch 安装

拿到云实例后,第一件事不是急着装 transformers,而是确认 ROCm 本身是否正常。很多人一上来就 pip install,结果后面报错根本分不清是 ROCm 没配好还是 Python 包冲突。

先跑这两条命令:

rocm-smi rocminfo | grep -i "gfx"

rocm-smi会列出 GPU 的显存占用、温度、功耗等信息。如果这条命令都跑不出来,后面不用继续了,先找云厂商确认驱动。rocminfo里的 gfx 架构代号很关键,比如 gfx90a、gfx942 之类,它决定了你后面装 PyTorch 时要用哪个版本的 wheel。

确认 ROCm 正常后,装 PyTorch 的 ROCm 版本。注意,不要用默认的 pip 源装 torch,那样装出来的是 CUDA 版或者 CPU 版。正确做法是去 PyTorch 官网找对应 ROCm 版本的安装命令,类似:

pip install torch torchvision --index-url https://download.pytorch.org/whl/rocm6.x

装完验证:

import torch print(torch.__version__) print(torch.cuda.is_available()) # ROCm 环境下这个返回 True print(torch.cuda.get_device_name(0))

这里有个容易困惑的点:ROCm 版的 PyTorch 依然使用torch.cuda这个命名空间,这是历史遗留,不代表它在用 CUDA。只要is_available()返回 True 且设备名是你的 AMD 卡,就说明环境通了。

注意:ROCm 版本、PyTorch 版本、gfx 架构三者必须匹配。我见过有人用 gfx90a 的卡装了只支持 gfx942 的 wheel,结果is_available()一直是 False,排查了半天。

2.2 依赖安装顺序与版本锁定

环境通了之后,装 transformers、peft、datasets、accelerate 这几个核心包。我的建议是先把版本锁死,不要用最新版。原因是 ROCm 生态的兼容性窗口比 CUDA 窄,最新版 transformers 可能引入了某些算子,在 ROCm 上还没适配。

我这次用的组合大致是:

pip install transformers==4.4x.x pip install peft==0.1x.x pip install datasets accelerate

具体小版本号我建议你根据自己底座模型的要求去查,但原则是:transformers 和 peft 的版本要互相兼容,peft 的版本要支持你用的模型架构。装完之后跑一个最小加载测试:

from transformers import AutoModelForCausalLM, AutoTokenizer model = AutoModelForCausalLM.from_pretrained("google/gemma-4-xxx", device_map="auto") tokenizer = AutoTokenizer.from_pretrained("google/gemma-4-xxx") print(model.device)

如果这一步能顺利把模型加载到 GPU 上,说明环境基本没问题。如果报显存不足,先检查是不是用了device_map="auto"但显存被其他进程占了。

3. 数据准备与 LoRA 配置的核心细节

3.1 情绪数据的格式化处理

我的原始数据是 CSV,两列:text 和 label。label 是六个情绪类别。做指令微调时,不能直接把 text 丢进去,要构造成指令格式。我用的模板大致是:

def format_sample(text, label): instruction = "请判断以下文本的情绪类别,只输出类别名称。" return { "prompt": f"{instruction}\n文本:{text}\n情绪:", "response": label }

这里有个细节:response 只放类别名称,不要放解释。因为我的评估是精确匹配,如果模型输出“这段文本的情绪是开心”,那和“开心”就不匹配了。训练时让模型学会只输出标签,推理时再做后处理提取。

数据划分上,我按 8:1:1 分训练、验证、测试。验证集用来监控训练过程中的过拟合,测试集只在最后评估一次。很多人会把验证集和测试集混用,导致最终指标虚高。

3.2 LoRA 参数怎么选:r、alpha、target_modules

LoRA 的核心参数有三个:秩 r、缩放系数 alpha、以及作用在哪些模块上。

r 决定低秩矩阵的维度。r 越大,可训练参数越多,拟合能力越强,但过拟合风险也越高。我的数据量只有几千条,所以选了 r=8。如果你数据量上万,可以试 r=16 或 32。alpha 一般设为 r 的两倍,我设 alpha=16。alpha/r 的比值影响适配矩阵的缩放,这个比值比绝对值更重要。

target_modules 是最容易被忽略的参数。Gemma 这类模型里,注意力层的 q_proj、k_proj、v_proj、o_proj 是常见选择。我一开始只加了 q_proj 和 v_proj,结果准确率只到 0.65 左右。后来把 o_proj 也加进去,才到了 0.73。原因是输出投影层也承载了语义信息,只调 qv 不够。

from peft import LoraConfig, get_peft_model lora_config = LoraConfig( r=8, lora_alpha=16, target_modules=["q_proj", "k_proj", "v_proj", "o_proj"], lora_dropout=0.05, bias="none", task_type="CAUSAL_LM" ) model = get_peft_model(model, lora_config) model.print_trainable_parameters()

print_trainable_parameters()会告诉你可训练参数占比。我这次大概是 0.3% 左右,非常轻量。

实操心得:target_modules 不要凭感觉加。先用默认的 qv 跑一版,看验证集准确率,再逐步加模块。每加一个模块,训练时间会增加,但收益不一定线性。

4. 训练过程与四个坑的完整记录

4.1 坑一:ROCm 上的 flash attention 不可用

第一个坑出现在训练启动阶段。我原本想用 flash attention 加速,因为 transformers 里可以通过attn_implementation="flash_attention_2"开启。结果在 ROCm 上报错,说找不到对应的算子。

原因很简单:flash attention 的 ROCm 适配版本和 CUDA 版本不是一回事,很多预编译 wheel 只覆盖 CUDA。解决办法是退回默认的 eager attention,或者用 ROCm 社区维护的 flash attention 分支。我为了省事,直接用默认实现,训练速度慢一些但稳定。

model = AutoModelForCausalLM.from_pretrained( model_name, device_map="auto", attn_implementation="eager" # ROCm 上先别开 flash )

这个坑的教训是:ROCm 生态里,很多 CUDA 上的“默认优化”并不默认可用。遇到算子缺失,先退回基础实现,跑通再考虑优化。

4.2 坑二:混合精度训练在 ROCm 上的表现差异

第二个坑是混合精度。CUDA 上大家习惯用 fp16 或 bf16 做混合精度训练,省显存又提速。我在 ROCm 上直接开 fp16,结果 loss 出现 NaN。

排查后发现,ROCm 对 fp16 的支持在某些算子上有差异,尤其是 softmax 和 layernorm 相关。换成 bf16 后问题消失。bf16 的动态范围比 fp16 大,不容易溢出,在 AMD 卡上兼容性更好。

from transformers import TrainingArguments training_args = TrainingArguments( output_dir="./output", per_device_train_batch_size=4, gradient_accumulation_steps=4, learning_rate=2e-4, num_train_epochs=3, bf16=True, # ROCm 上优先用 bf16 fp16=False, logging_steps=20, eval_strategy="steps", eval_steps=100, save_strategy="epoch", report_to="none" )

注意:如果你的卡不支持 bf16,那就只能用 fp16,但要加 loss scaling。ROCm 下 loss scaling 的配置和 CUDA 略有不同,建议先小规模试跑。

4.3 坑三:batch size 与梯度累积的显存平衡

第三个坑是显存。我一开始把 per_device_train_batch_size 设成 8,结果 OOM。降到 4 还是紧张,最后用 batch_size=4 加 gradient_accumulation_steps=4,等效 batch size 16。

这里要理解一个概念:LoRA 虽然只训练少量参数,但前向传播和激活值依然要占显存。激活值的大小和 batch size、序列长度成正比。我的文本平均长度 128 token,最长 256,所以序列长度设 256。如果文本更长,显存压力会明显上升。

显存估算的粗略公式是:模型权重 + 激活值 + 优化器状态。LoRA 的优化器状态只针对适配矩阵,很小,所以大头是权重和激活值。48GB 显存跑 Gemma4 这个量级,batch size 4 到 8 是比较稳的区间。

4.4 坑四:评估指标的计算方式导致虚高

第四个坑最隐蔽。我一开始用训练框架自带的 evaluation,它计算的是 token 级别的 loss,不是准确率。loss 下降不代表分类准确率上升。后来我自己写了评估函数,对验证集逐条推理,提取输出标签,和真实标签比对。

def evaluate(model, tokenizer, dataset): correct = 0 total = 0 for sample in dataset: inputs = tokenizer(sample["prompt"], return_tensors="pt").to(model.device) with torch.no_grad(): outputs = model.generate(**inputs, max_new_tokens=8) pred = tokenizer.decode(outputs[0], skip_special_tokens=True) pred_label = extract_label(pred) if pred_label == sample["response"]: correct += 1 total += 1 return correct / total

这个坑的教训是:微调任务的评估指标必须和业务目标对齐。分类任务就看分类准确率,不要被 loss 曲线迷惑。

5. 结果对比与效果分析

5.1 基线 vs 微调后的准确率

微调前,我用 Gemma4 底座直接做 zero-shot 推理,准确率 0.594。这个数字说明底座模型对情绪分类有一定理解,但不够精准,尤其是中性、惊讶、厌恶这几类容易混。

微调后,测试集准确率 0.734。分类别看:

情绪类别微调前微调后
开心0.780.86
悲伤0.710.82
愤怒0.690.80
中性0.450.62
惊讶0.520.68
厌恶0.410.61

提升最明显的是中性、惊讶、厌恶这三类,正好是基线表现最差的。说明 LoRA 确实学到了数据里的判别边界,而不是只强化了原本就会的类别。

5.2 训练曲线与过拟合判断

训练 loss 从 1.2 降到 0.4 左右,验证 loss 在前两个 epoch 下降,第三个 epoch 开始持平甚至微升。这是典型的过拟合信号。我最终选了第二个 epoch 的 checkpoint,而不是最后一个。

判断过拟合不能只看 loss,还要看验证集准确率。我的验证准确率在第二个 epoch 达到峰值 0.72,第三个 epoch 掉到 0.70。所以早停是必要的。

实操心得:LoRA 虽然参数少,但小数据下依然会过拟合。建议每个 epoch 都存 checkpoint,最后用验证集挑最好的,不要默认用最后一个。

6. 常见问题速查与避坑清单

6.1 ROCm 环境类问题

问题可能原因解决方向
torch.cuda.is_available() 为 FalsePyTorch 装成 CUDA/CPU 版重装 ROCm 版 wheel
算子找不到flash attention 未适配退回 eager attention
loss 出现 NaNfp16 溢出换 bf16 或加 loss scaling
显存 OOMbatch size 过大降 batch,加梯度累积
训练极慢未启用优化算子检查 ROCm 版本与 PyTorch 匹配

6.2 LoRA 配置类问题

target_modules 选少了,模型学不动;选多了,训练变慢且容易过拟合。我的建议是从 qv 开始,逐步加 k、o。r 和 alpha 不要同时调,先固定 alpha=2r,只调 r。

数据格式上,prompt 和 response 的分隔要清晰,避免模型把指令也当成要生成的内容。评估时一定要做输出解析,不能直接拿生成文本比对。

6.3 我个人的避坑清单

第一,环境没验证通过之前,不要碰数据。第二,先跑一个 100 条的小子集,确认整个流程能走通,再上全量。第三,每个 epoch 存 checkpoint,别省这点磁盘。第四,评估函数自己写,不要完全依赖框架默认。第五,ROCm 上遇到问题,先查 gfx 架构和版本匹配,再查代码。

这套流程跑下来,我对 ROCm 做 LoRA 微调的信心是有的。它不像 CUDA 那么“开箱即用”,但把版本和环境理顺之后,稳定性是可以接受的。后面我打算试试更大的 r 和更多 target_modules,看看准确率还有没有上升空间。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询