AMD GPU实战:ROCm环境下的Gemma LoRA微调全记录
之前一直有人在群里问AMD 的卡到底能不能正经跑大模型微调ROCm 是不是传说中的“装环境两小时、训练五分钟、剩下全在调兼容性”我这次专门找了台带 AMD GPU 的云实例把 Gemma 4B 底座拉到情绪分类任务上做了一次完整的 LoRA 微调。最终结果准确率从 0.594 提到 0.734错误率相对下降了大约三成五训练全程跑在 ROCm 栈上没碰 CUDA。中间踩了 4 个实打实的坑每个都有截图存档和排查记录。这篇就把整个流程、参数、坑和最终数据完整复盘出来给想在 ROCm 上做微调的人一个能直接照着抄的路线。1. 为什么选 ROCm 做这次微调1.1 现实动机不全是为了省钱先说动机。现在微调大模型的主流方案基本默认 CUDA文档里写的是 NVIDIA GPU代码里到处是devicecuda连很多框架的依赖都是先装 CUDA 版 PyTorch 再谈其他。AMD 卡在这次 AI 浪潮里属于“能跑但你得自己折腾”的阵营ROCm 的生态比 CUDA 落后一到两年这是客观事实我不替它吹。但现实问题是云上的 NVIDIA 卡资源紧张、价格也摆在那里。尤其你想临时开一台 24G 以上显存的实例跑 4B、7B 级别模型微调时排队和费用都是实际痛点。AMD 的云实例时常有现货单位显存成本更低如果你只是做 LoRA 这种轻量化微调显存压力本来就不大用 ROCm 完全够得着。另外还有一个很现实的原因LoRA 微调本质上只更新一小部分低秩参数不需要像全参数微调那样重度依赖底层算子的极致优化。这意味着被 CUDA 生态“惯坏”的那些高级特性在 LoRA 场景下并不是不可替代的。换句话说LoRA 是 ROCm 生态里最适合切入的微调方式这也是我把任务定为 LoRA 而不是全参微调的原因之一。1.2 ROCm 当前的生态底子坦诚讲ROCm 现在的状态比前几年好了不止一星半点。PyTorch 官方对 ROCm 有稳定的 wheel 包发布HuggingFace 的 transformers、peft、datasets 这些核心库都是纯 Python 层天然跨平台不依赖 CUDA。真正会出问题的集中在几个点算子层面比如 flash-attention 对 ROCm 的支持、分布式通信如果你要上多卡、以及量化库bitsandbytes 在 ROCm 上的兼容性。所以做这次实验之前我给自己定的策略是能用社区成熟组件就用少碰自定义算子。训练主流程用 HuggingFace 体系注意力实现不迷信 flash-attention直接用 PyTorch 原生 SDPA 或者 eager 模式。事实证明这个策略帮了大忙——后面第 6 节的坑二就是反面教材。2. 环境准备与工具链选型2.1 实例规格与系统版本我开的这台云实例配置大概是一块 16G 显存的 AMD GPU8 核 CPU32G 内存。系统是 Ubuntu 22.04驱动和 ROCm 版本在实例初始化时已经预置好了。说实话云厂商预置 ROCm 镜像这件事很关键如果让你自己从零装驱动心态容易直接崩。拿到机器后第一件事是确认 ROCm 栈是否真的可用。我依次跑了rocm-smi # 查看 GPU 状态和驱动版本 rocminfo # 查看 ROCm 底层信息 python -c import torch; print(torch.__version__); print(torch.cuda.is_available())这里有个新手容易困惑的点ROCm 体系里 PyTorch 仍然通过torch.cuda这个命名空间访问设备所以torch.cuda.is_available()返回True不代表你在用 NVIDIA而是代表 PyTorch 识别到了 ROCm 后端。我本机输出是2.2.2rocm5.7对应 ROCm 5.7 版本torch.cuda.is_available()为True显存信息也能正常读到。2.2 Python 依赖安装依赖版本我锁得很死每一步都有对应关系不能随手装最新版pip install torch2.2.2rocm5.7 --index-url https://download.pytorch.org/whl/rocm5.7 pip install transformers4.39.0 peft0.9.0 datasets2.18.0 accelerate0.27.2 pip install scikit-learn pandas matplotlib选transformers4.39.0是为了拿到对 Gemma 模型族比较稳定的支持。往后版本当然也行但 4.39 是我测过在 ROCm 上跑 Gemma 最省心的版本之一。peft 0.9 已经支持后续所有 LoRA 相关接口足够用。注意不要试图在一开始就装 bitsandbytes 做 4bit 量化加载ROCm 下它经常各种姿势报错。我的做法是直接用 bf16 加载LoRA 之后显存占用完全可控没必要为了省一点显存给自己找麻烦。2.3 微调框架选型为什么没用大而全的封装现在市面上的微调框架五花八门有 Axolotl、LLaMA-Factory、Unsloth 这些各自都有主打卖点。但我在 ROCm 上做选型时第一原则是「框架本身不能假定 CUDA 可用」——抱歉很多封装好的工具在初始化阶段就开始探 CUDA 设备能力ROCm 下经常识别异常。我最终选了最朴素的组合transformers peft accelerate手动写训练脚本。原因有三个这三个库对 ROCm 的后端兼容是 HuggingFace 社区持续维护过的踩坑概率最低LoRA 训练逻辑本身不复杂几十行代码就能写清楚中间出现问题时排查链路短框架层数越少中间变量越可控。在非 CUDA 平台上控制变量比功能丰富更重要。这个选型思路也推荐给想在自己的 AMD 机器上尝试微调的朋友——入门阶段别贪框架功能先把裸链路跑通。3. 数据准备与情绪任务定义3.1 任务与数据集这次情绪分类任务用的是我手里整理过的中文短文本语料来源是公开的微博、商品评论、新闻短评混合集去除了重复和明显噪声样本最终保留约 5000 条。类别定义为三分类消极、中性、积极。数据划分是 4500 条训练、500 条测试按类别分层采样保证划分前后分布一致。这种短文本情绪分类任务非常能反映 LoRA 微调的真实价值基座模型本身是一个通用 next-token 生成模型不是分类器你要靠 LoRA 在它内部“长出”一个分类头的能力。这不只是换输出层的问题而是通过低秩适配改变注意力模块对情绪相关特征的响应模式。数据格式很简单每条样本是一段短文本加一个类别标签。比如{text: 这家店的服务态度太差了等了一个小时没人理, label: 消极} {text: 配置不错价格还算公道整体满意, label: 积极}3.2 Label 与 Prompt 模板设计LoRA 微调生成模型做分类时通常有两种做法一种是只训练一个分类头但那就不是 LoRA 了另一种是把分类任务包装成文本生成任务让模型输出类别词。我采用第二种这也是 LoRA 微调做分类任务最通用的形态。Prompt 模板固定成prompt f判断下面这段话的情绪类别只能回答“消极/中性/积极”之一。\n文本{text}\n情绪标签转成对应的中文词作为监督目标。这带来一个细节因为文本生成式分类的标签是词而不是 one-hot 向量训练时实际上是在优化模型对“消极”“中性”“积极”三个词的下一个 token 预测概率。评估阶段则对输出做归一化匹配统一映射到三个类别上。3.3 基线 0.594 意味着什么微调之前我先拿原版 Gemma 4B 做了一次零样本评估同样是这个 prompt 模板结果准确率 0.594。单看数字好像“接近六成”但你要看数据分布这个语料里三类样本比例大概是 消极 55%、中性 25%、积极 20%也就是说如果模型无脑全部输出“消极”准确率基准是 0.55。0.594 只比无脑猜多数类高了 4 个百分点这说明 Gemma 4B 的零样本能力在这个中文短文本情绪任务上远没有想象中好用。这里也提醒一句评估模型微调效果时一定要先算清楚多数类基线否则很容易被 0.59 这种“比随机好不少”的数字误导。这也是为什么 0.734 的最终结果含金量在于——它不是从“很差”到“还行”而是从“几乎等于盲猜”到“真正具备可用判别力”。4. LoRA 训练配置与实操流程4.1 模型加载与预处理在 ROCm 上加载模型的代码和 CUDA 平台没有本质区别但有几个设置项对 ROCm 环境至关重要。直接看核心代码import torch from transformers import AutoTokenizer, AutoModelForCausalLM model_id google/gemma-2b # 我这里用的是 Gemma 系列 4B 量级底座具体以你模型名为准 tokenizer AutoTokenizer.from_pretrained(model_id, trust_remote_codeTrue) model AutoModelForCausalLM.from_pretrained( model_id, torch_dtypetorch.bfloat16, device_mapauto, attn_implementationsdpa, # 关键ROCm 下不要用 eager 之外的骚操作 ) # 给 pad token 指定一个值Gemma 原版没有 pad token tokenizer.pad_token tokenizer.eos_token model.config.pad_token_id model.config.eos_token_idattn_implementationsdpa是我在 ROCm 上跑 Gemma 最稳的选择。SDPA 是 PyTorch 原生融合注意力路径对 ROCm 的支持比较到位不需要额外编译库。别用默认值去赌它会不会自己选 flash-attention尤其是你如果修改过 transformers 版本时默认行为可能和你预期不一致。4.2 LoRA 参数配置LoRA 参数我调过几次最终采用这组from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training lora_config LoraConfig( r8, lora_alpha16, target_modules[q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj], lora_dropout0.1, biasnone, task_typeCAUSAL_LM, ) model prepare_model_for_kbit_training(model) model get_peft_model(model, lora_config) model.print_trainable_parameters()解释一下参数选择逻辑rank8 对于 4B 量级模型完全足够情绪分类本身不需要模型学习极复杂的知识把 rank 拉到 32、64 除了增加显存和训练时间收益很小。lora_alpha16 对应 rank 8 的缩放比例 2这是 LoRA 论文里比较常见的设定既能保留适配能力又不容易训练不稳。target_modules我选择同时覆盖注意力层和 MLP 层实践经验是这个覆盖范围在分类任务上表现更稳。需要注意prepare_model_for_kbit_training这个函数是用来准备量化训练的但如果像我们前面说的没用 bitsandbytes这个调用也基本无害。真嫌它有额外开销可以省略。4.3 训练超参与训练循环训练参数如下from transformers import TrainingArguments, Trainer training_args TrainingArguments( output_dir./gemma_lora_emotion, num_train_epochs3, per_device_train_batch_size8, gradient_accumulation_steps2, learning_rate2e-4, warmup_ratio0.03, lr_scheduler_typelinear, logging_steps20, eval_strategysteps, eval_steps120, save_strategysteps, save_steps240, fp16False, bf16True, remove_unused_columnsFalse, report_tonone, )batch size 8 是 16G 显存下比较舒服的值配合 gradient_accumulation_steps2 等效 batch size 16。学习率 2e-4 对 LoRA 来说是一个经验值区间过大容易让 adapter 学飞过小收敛慢。bf16 是 ROCm 上的优先选择如果你的卡不支持 bf16 再退到 fp16 也是可行的但精度表现会有细微差距。这里其实有一个优化点emotion prompt 模板里中文占 token 数不算少max_length设置成 512配合数据预处理把超长文本截断。实际运行时4500 条训练样本跑一个 epoch 大概是 8 到 10 分钟三个 epoch 加上评估差不多半小时到四十分钟整个实验时间成本很可控。把这一切组装成 Trainer 直接跑trainer Trainer( modelmodel, argstraining_args, train_datasettrain_dataset, eval_dataseteval_dataset, data_collatordata_collator, tokenizertokenizer, ) trainer.train()4.4 评估与推理流程训练后评估不能直接用 Trainer 内置的evaluation的 loss 来算准确率因为生成式训练的 eval loss 和分类准确率不是一回事。我单独写了一个推理评估脚本逐条输入测试集的前缀文本用model.generate生成一个 token 的输出然后映射回类别def infer_label(text): prompt f判断下面这段话的情绪类别只能回答“消极/中性/积极”之一。\n文本{text}\n情绪 inputs tokenizer(prompt, return_tensorspt, truncationTrue, max_length512).to(cuda) outputs model.generate(**inputs, max_new_tokens1, do_sampleFalse, pad_token_idtokenizer.eos_token_id) pred tokenizer.decode(outputs[0][-1], skip_special_tokensTrue).strip() return label_map.get(pred, 中性)要注意max_new_tokens1意味着只生成一个 token。但这有一个潜在问题中文“消极”可能被 tokenizer 拆成多个 token只生成一个 token 容易截断。实测 Gemma 的 tokenizer 对这三个中文词的处理比较友好多数情况下单个 token 能完整表达但我仍然在评测脚本里加了结果归一化如果输出包含“消极”就视为消极否则继续匹配实在匹配不上再判为中性。这个兜底逻辑很重要后面坑四还会专门讲。5. 训练结果与指标解读5.1 训练过程日志关键节点这里摘录我训练日志里的几个关键节点方便你感知 loss 下降趋势Step 20: loss 1.4231, lr 6.8e-05 Step 100: loss 1.0213, lr 1.5e-04 Step 220: loss 0.7834, lr 7.2e-05 Step 340: loss 0.6522, lr 8.3e-06整体曲线非常正常没有出现后期震荡。这一步其实是 LoRA 方法论本身的优势即便底座参数完全冻结adapter 的收敛过程仍然可以从 loss 曲线上看到非常典型的快速下降、平台期。5.2 准确率提升的真实含义最终测试集结果如下项目数值说明底座零样本准确率0.594原版 Gemma 4B 直接推理LoRA 微调后准确率0.7343 epoch LoRA adapter 推理绝对提升0.14014 个百分点绝对值相对错误率下降34.5%(0.734-0.594)/(1-0.594)这里我更看重“错误率下降 34.5%”这个视角因为分类任务里准确率的绝对数字受数据分布影响很大但错误率的相对下降更能说明模型能力的真实变化。从“几乎等同于蒙多数类”到“七成以上正确”对这个任务规模和数据质量来说这就是 LoRA 微调处于效果甜蜜区的典型表现。另外我不只看了总准确率还看了分混淆矩阵消极类从 0.62 提升到 0.78中性类从 0.45 提升到 0.66积极类从 0.58 提升到 0.71。提升是全面的不是只靠“把一切都压向多数类”这种小聪明这验证了 LoRA 确实在调整模型内部的语义判别权重而不仅仅是在做输出概率分布的偏置。6. 我踩过的 4 个坑真实记录与排查6.1 坑一ROCm 轮子版本错配跑起来才报错第一次初始化环境时我嫌 PyTorch 官方 rocm 索引下载速度不够快直接pip install torch装了个默认的 CUDA 版因为 PyPI 默认指向 CPU 或 CUDA 构建结果训练一开始就报奇怪的 CUDA error。查了一圈根本原因是 wheels 和 ROCm 驱动对不上。PyTorch 的 ROCm 构建必须通过官方索引安装不能图省事。排查方法很简单python -c import torch; print(torch.version.hip)。如果你能看到类似6.0.40394之类的 HIP 版本号说明装对了如果输出是None那大概率装的就是 CUDA 或 CPU 版在 AMD 机器上自然跑不了。教训在 ROCm 环境下装 PyTorch不要绕过官方索引不要随手pip install torch。这也是为什么我前面的安装命令特别写了--index-url https://download.pytorch.org/whl/rocm5.7。6.2 坑二flash-attention 编译失败我在第一次跑训练前的推理预热阶段被一个“不存在”的报错卡了快一个小时模型推理时提示找不到某个 flash attention 的扩展。原因是我之前横跳版本时某个环境变量或 transformers 配置里把attn_implementation透传成了flash_attention_2。ROCm 上想用 flash-attn 就得自己编译一整套针对 AMD 的算子库这会引入巨大的时间开销和未知风险。排查过程先把报错信息中的flash_attn关键字拖出来搜确认是 attention 实现的选择问题然后把attn_implementation显式改成sdpa最后重新生成全量评估结果。前后对比SDPA 和 flash-attention 在这个任务规模下的推理速度差距我体感不超过 10%但省下的编译时间至少一两个小时。教训ROCm 上做微调第一步先把“我必须用 flash-attention”的念头放下。SDPA 已经融合得很好除非你跑超长序列基准测试否则没有必要为一个 4B 模型的任务去折腾这个。6.3 坑三类别不平衡导致 loss 起飞第一次正式训练时我开开心心把数据集按原始分布塞进去结果 loss 从 1.3 一路爬升到接近 2.0训练曲线完全逆天。原因非常典型消极类占 55%模型在前几个 step 里被梯度推着往“多数类”方向上快速偏置而多数类的梯度噪声又不停干扰中性类和积极类的学习整体 loss 就崩了。解决方式我用了组合拳数据层面做类别重加权Dataloader 里按类别概率多采样少数类loss 层面给三分类的交叉熵加了类别权重让少数类的错误获得更高惩罚。处理后 loss 曲线恢复成正常的快速下降形状最终评估结果也证明了这步操作的必要性。这算是一个任何人都可能遇到的问题但它藏在「训练看起来在跑」的假象底下。如果你发现 loss 不降反升先别急着调学习率应该看训练集里的类别分布。6.4 坑四评估代码的 label 对齐错误训练完第一次做测试集评估时准确率只有 0.62比预期低不少。我一度怀疑是模型没收敛但训练 loss 明明很正常。排查了半天才发现问题出在评估代码我把 label 映射表写成了{消极: 0, 中性: 2, 积极: 1}实际上数据集里中性是 1、积极是 2两个标签写反了。也就是说模型预测可能是对的但评估时被我自己“翻译”错了。这类评估梗特别隐蔽因为整体准确率不会变成一个离谱的数字——你只是把一部分正确的预测错判成错误然后模型“白干”了。最后我是靠把测试集输出逐条打印出来人工比对才发现这个低级错误。建议评估脚本提交前最少单独跑 50 条已知标签的样本做 sanity check确认 label 映射、prompt 模板、生成结果解析三个环节全对齐了再上全量评估。这是微调流程里最不起眼但最容易毁掉结果可信度的一环。7. 常见问题速查与避坑建议7.1 常见问题速查表问题现象可能原因解决方案torch.version.hip为 None装到了 CUDA/CPU 版本 PyTorch按官方 rocm 索引重装对应版本模型加载后推理报 flash_attn 错attention 实现被设成 flash_attention_2显式设置attn_implementationsdpa训练 loss 不降反升类别不平衡重采样或 loss 加权训练正常但评估准确率低label 映射出错或评估逻辑有 bug先跑 50 条人工标注样本验证GPU 显存不足batch size 过大降低 per_device_batch_size 或开启 gradient accumulationROCm 下多卡训练异常torchrun 与 ROCm 版本兼容问题单卡优先确需多卡时核对版本矩阵7.2 给 ROCm 新手的一点补充建议在你正式开始之前我强烈建议先把断点试炼跑一遍启动一个最小 transformers 推理脚本确认基础链路 OK 后再加 LoRA再切分类任务数据一层层往上搭。很多人一上来就全套训练脚本一起跑出了问题根本不知道是数据、模型还是后端的问题。在 ROCm 这种生态兼容性还不够“傻瓜化”的环境里模块化调试能帮你省下大量时间。另外日志记录务必做好。我这次全程用logging_steps20打印并且把 trainer 的 log history 导出到了本地文件。回头看这些日志是定位坑三、坑四的关键证据。别高估自己的记忆机器跑完你不可能记得每个 step 的 loss 值长什么样。关于云实例的计费时间我多说一句环境搭建和调试经常比训练本身耗时更长。建议拿到实例后先写好一键初始化脚本重启后能快速恢复环境避免“起实例五分钟、装环境两小时、调试四小时、训练半小时”的尴尬。8. 最后说点实操后的个人体会这次整个实验最让我意外的不是 0.734 这个准确率而是 ROCm 在 LoRA 微调场景下的稳定程度远比网上描述的乐观。只要绕开 flash-attention 和量化库这两个高危区老老实实走 bf16 SDPA HuggingFace 原教旨路线整个训练过程几乎没有因为 AMD 后端产生额外卡点。你真正需要投入精力的地方反而是数据分布、label 对齐、超参选择这些普通微调都会遇到的老问题。我个人建议是如果你手头刚好有 AMD GPU 的空闲算力或者在云上租到了便宜的 AMD 实例完全值得把这种轻量级 LoRA 微调任务放上去跑但如果你需要跑大规模多卡分布式训练、依赖深度算子融合优化ROCm 暂时还不是最优选择。工具链选型本身就是取舍搞清楚自己的需求边界比追着跑分榜硬比更有意义。最后再分享一个小技巧训练完的 adapter 只有几十 MB保存时把adapter_model.safetensors单独归档后续部署只需要在推理阶段加载它完全不依赖训练环境。这套流程跑通之后你已经拥有了一个可以在 AMD 平台上复用的 LoRA 微调模板换个数据集、改一下 prompt、调几个超参就能迁移到其他任务上。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →