AMD ROCm 上基于 LoRA 微调 Gemma 4 情绪分类实操手记

上一次我们已经完成 GPU 核验 → 权重下载 → vLLM 起服 → API 冒烟。接下来我们开始学习怎么利用特定的数据集进行模型的微调。这里我们也在上一篇的基础上,介绍「调模型」多出来的操作——额外包装什么、训练数据怎么造、加载方式与推理有何不同、LoRA 怎么配,并提供一个完整的微调流程,简历一个整体的实操框架。

段末注释:LoRA(Low-Rank Adaptation,低秩适配)在冻结基座权重的前提下,为部分线性层插入低秩可训练矩阵;SFT(Supervised Fine-Tuning,监督微调)用标注样本直接优化模型输出。

这次微调,也是基于一个成熟的参考项目:把 google/gemma-4-E4B-it 用 LoRA 拉到 AI-ModelScope/emotion 六类情绪分类上,并做微调前后可量化对比。
我们先对微调建立一个整体的执行概念,一个LLM微调大概过程可以参考下图:
AMD ROCm 单卡 LoRA 微调 Gemma 4 情绪分类流程概览


一、部署 / 推理 vs 微调:增量在哪里?

我们之前进行了模型的部署推理,关注的是「模型能不能快而稳地生成」),而现在我们要做的微调,还要更进一步,要解决「用什么数据、更新哪部分参数、怎么证明变好了」。所以先来看看为了实现更进一步的目标,和之前的单纯部署会有哪些不同?

维度 部署篇(vLLM 推理) 本篇(LoRA 微调)
入口 vllm serve trl.SFTTrainer + peft
模型形态 只读权重,KV Cache 为主 前向 + 反向,要存激活与优化器状态
额外依赖 基本已有 transformers / vllm trlpeftdatasetsscikit-learn
数据 无(在线 prompt) 标注数据集 + prompt-completion 构造
权重产物 原样 15 GB 基座 ~百 MB 级 adapteradapter_model.safetensors
评估 人工对话冒烟 生成式分类指标(accuracy / macro F1 / 混淆矩阵)
ROCm 注意点 起服 OOM、上下文长度 不用 QLoRAadamw_torch 优化器gradient_checkpointing

二、额外依赖:相对部署多装什么、解决什么

部署环境通常已有 torchtransformersvllmmodelscope。为了满足微调的业务需求,我们需要再补:

推理有没有 微调为什么需要
trl 提供 SFTTrainer / SFTConfig:tokenize、训练循环、completion_only_loss(只对 assistant 段算 loss)
peft LoraConfig 注入低秩适配器,避免全量更新 8B 参数
datasets 一般无 管理 train/val/test split、map 转格式、ClassLabel 映射标签名
scikit-learn 微调前后 classification report、混淆矩阵、macro F1
pandas 可选 预测明细、前后对比表落 CSV

accelerate 多由 trl 间接依赖,单卡 SFTTrainer 也会用到其设备抽象。不必为微调重装 vLLM;若 Step 1 已用部署篇环境,只补上面几行即可。


三、微调专用配置(推理里没有的项)

1
2
3
4
5
6
7
8
9
10
11
12
13
MODELSCOPE_DATASET_ID = "AI-ModelScope/emotion"
OUTPUT_DIR = "./gemma4-it-emotion-lora-ms-single-gpu"

TRAIN_LIMIT = 4000 # 先小子集跑通;None = 全量 16000
VALIDATION_LIMIT = 400
TEST_LIMIT = 400
EVAL_LIMIT = 400

MODEL_DTYPE = torch.bfloat16
BF16, FP16 = True, False # 训练精度;与 SFTConfig.bf16 一致

SYSTEM_PROMPT = """You are an emotion classification assistant.
...只输出 sadness/joy/love/anger/fear/surprise 之一..."""
配置 目的
TRAIN_LIMIT 控制迭代成本;情绪集全量不大,但生成式评估慢,故 EVAL_LIMIT 单独限
SYSTEM_PROMPT 训练与评估必须同一份,否则微调学的分布和打分用的分布不一致
BF16 全精度 ROCm 下 bitsandbytes 4bit QLoRA 不稳定,本篇用 BF16 基座 + LoRA,用显存换兼容性
OUTPUT_DIR 存 adapter、checkpoint、评估 CSV,不是完整 15 GB 权重

基座路径 LOCAL_MODEL_DIR 可直接复用部署篇 ./models/google/gemma-4-E4B-it;若目录已在,跳过重复 snapshot_download


四、任务数据:微调独有的数据链路

推理只需用户输入;微调必须准备 (text, label) 监督信号。

4.1 下载与加载

1
2
3
4
5
6
7
8
9
10
dataset_dir = dataset_snapshot_download("AI-ModelScope/emotion", cache_dir="./datasets")

raw_dataset = load_dataset("parquet", data_files={
"train": glob(".../data/train-*.parquet"),
"validation": glob(".../data/validation-*.parquet"),
"test": glob(".../data/test-*.parquet"),
})
raw_dataset[split] = raw_dataset[split].cast_column(
"label", ClassLabel(names=["sadness","joy","love","anger","fear","surprise"])
)

4.2 什么是 SFT 格式

SFT(Supervised Fine-Tuning,监督微调)格式,指把标注样本组织成「模型该看到什么输入 → 该生成什么标准答案」的结构,供 SFTTrainer 按因果语言模型(Causal LM)目标训练。

段末注释:Causal LM 按从左到右预测下一个 token;SFT 即在基座预训练之上,用标注的「输入—回答」对继续优化这一预测目标。

对本 Notebook 而言,采用的是 TRL 支持的 chat 版 prompt-completion 格式——每条样本两个字段:

字段 内容 训练时的角色
prompt system + user 消息列表 条件上下文:任务规则 + 待分类文本
completion assistant 消息列表 监督目标:应生成的标准标签

tokenizer.apply_chat_template 会把上述消息列表编成模型实际读入的 token 序列;配合 completion_only_loss=Trueloss 主要落在 completion(assistant 回复)上,避免对 system/user 前缀做无意义的梯度更新。

与几种常见数据形态对比:

形态 示例 能否直接用于本篇
原始分类 text + 整数 label 否,需先映射为标签文本
Alpaca 三字段 instruction / input / output 需再套 chat template
纯文本拼接 "User: ... Assistant: joy" 可行但易与官方模板不一致
SFT(chat) prompt + completion 消息列表 ,与 Gemma-it 推理格式一致

选用 chat 版 SFT 的核心原因:gemma-4-E4B-it 是指令对齐的对话模型,训练和推理都应走同一套 chat template;若仍用 (text, label) 整数标签,相当于任务形式与模型接口不匹配。

4.3 本任务中的样本构造

1
2
3
4
5
6
7
{
"prompt": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": "Classify the emotion of this text:\n\n{原文}"},
],
"completion": [{"role": "assistant", "content": "fear"}],
}

原始 emotion 集一条记录是 {"text": "...", "label": 4}to_prompt_completion()label=4 查表为 "fear",写入 completion。模型因此学习的是:在给定 system 约束与用户文本后,生成且仅生成一个合法情绪词——与第六节评估用的 generate_label() 完全同分布。

1
2
3
4
5
6
7
8
9
10
11
def to_prompt_completion(example):
label = label_names[example["label"]]
return {
"prompt": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": f"Classify the emotion of this text:\n\n{example['text']}"},
],
"completion": [{"role": "assistant", "content": label}],
}

sft_dataset = dataset.map(to_prompt_completion, remove_columns=dataset["train"].column_names)

注意SYSTEM_PROMPT 在训练、微调前评估、微调后评估中必须保持一致;改 prompt 等于改任务定义,前后指标将不可比。


五、训练态加载:和 vLLM 推理的关键差异

部署篇用 vLLM 托管生成;微调必须用 transformers.AutoModelForCausalLM 拿可求导的计算图:

1
2
3
4
5
6
tokenizer = AutoTokenizer.from_pretrained(LOCAL_MODEL_DIR, use_fast=True)
base_model = AutoModelForCausalLM.from_pretrained(
LOCAL_MODEL_DIR, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True,
)
base_model.to("cuda")
base_model.config.use_cache = False # 训练必须关 KV cache
推理(vLLM) 微调(本篇)
加载方式 vllm serve from_pretrained + .to(cuda)
use_cache 开启,加速 decode 关闭,否则 backward 与 checkpoint 行为异常
量化 可按需 AWQ/GPTQ 等 不用 bnb 4bit;ROCm 上 QLoRA 坑多
chat template vLLM 内置 必须 apply_chat_template 与训练数据一致;缺失时从魔搭补 chat_template.jinja
显存 权重 + KV 权重 + 激活gradient_checkpointing=True 用算力换显存)+ LoRA 优化器状态

六、评估管线:微调前为什么要先跑基线

这一步要解决什么?

训练 loss 下降,只能说明模型在训练集标注格式上拟合得更好,不能直接回答「情绪分类有没有变好」。若跳过微调前评估,后面只能看到一组绝对分数,无法判断提升来自 LoRA,还是基座 + SYSTEM_PROMPT 本身已经够用。

因此要在动任何 LoRA 权重之前,用与微调后完全相同的推理脚本,在 held-out 的 test 集上打一次分——这就是基线(baseline)

评估怎么做?

本篇不走「取最后一层 logits 做 softmax 六分类」,而是与 SFT 训练目标对齐的 生成式分类

1
2
3
4
5
待测 text
→ apply_chat_template(同一 SYSTEM_PROMPT)
→ model.generate(贪心,max_new_tokens=4)
→ 正则抽取 sadness / joy / … / surprise
→ 与 gold label 比对

对应 Notebook 里的 generate_label()evaluate_model()。核心指标:

指标 含义 为何要看
accuracy 预测标签完全一致的比例 直观总览
macro_f1 六类 F1 的宏平均 情绪集类别不平衡,macro 比 accuracy 更能反映少数类
invalid_predictions 生成结果不在词表内的次数 衡量模型是否「守格式」——分类任务里这类错误通常应趋近 0

每条样本都要走一遍 autoregressive decode,比传统分类头推理慢一个数量级,故默认 EVAL_LIMIT=400,在速度与可信度之间折中。

段末注释:生成式分类把标签当作模型应生成的短文本,评估与 SFT 的「预测下一个 token」目标同构;传统分类则在固定表示上接 softmax 头,与 chat 模型的训练方式不一致。

微调前基线(实测,400 条 test 子集)

指标 数值 简要解读
accuracy 0.625 基座 + prompt 已有一定零样本分类能力,并非「白板模型」
macro_f1 0.482 明显低于 accuracy,说明多数类(如 joy)拉高总分,少数类 recall 偏弱
invalid_predictions 2 格式基本可控,但仍有偶发越界输出

从混淆矩阵看,微调前 sadnessjoy 互混最多;loveangerfear 样本少且 recall 低——这正是 LoRA 需要重点改善的方向,也是后面「微调后对比」的对照锚点。

有了基线,后面怎么比?

微调结束后再跑同一 evaluate_model()(同一 test 子集、同一 prompt、同一 decode 策略),只比较 delta:

对比项 期望
accuracy / macro_f1 上升;macro_f1 涨幅往往比 accuracy 更有说服力
invalid_predictions 下降或归零
少数类 recall 上升(love / anger / fear 等)
混淆矩阵非对角线 减少

若微调后 loss 很低但 macro_f1 几乎不动,常见原因是过拟合训练格式、或评估 prompt 与训练不一致——有基线才能把这种「假提升」筛出来。


七、LoRA 与训练超参:ROCm 单卡怎么配

7.1 LoRA

1
2
3
4
LoraConfig(
r=16, lora_alpha=32, lora_dropout=0.05,
task_type="CAUSAL_LM", target_modules="all-linear",
)

原理:对线性层 (W) 增加 (\Delta W = BA)(秩 (r \ll d)),只训练 (A,B)。实测可训练参数 50,499,584 / 7,991,600,416 ≈ 0.63%

训练前必查Trainable LoRA parameters 若为 0,说明 target_modules 未命中,不要开训。Gemma 4 含 vision tower,all-linear 会挂到视觉侧;纯文本任务可后续收窄到语言模型 attention/MLP 以省显存。

7.2 SFTConfig 里值得盯住的项

参数 微调场景下的含义
per_device_train_batch_size × gradient_accumulation_steps 4 × 4 = 16 等效 batch;OOM 先降 batch
learning_rate 1e-4 LoRA 常用量级
completion_only_loss True 只对 assistant 标签算 loss,不对 system/user 浪费梯度
gradient_checkpointing True 反向时重算激活,训练显存的主要手段之一
max_length 256 情绪句子短,截断 mainly 防异常长文本吃显存
optim adamw_torch 避开 ROCm 下 bitsandbytes 优化器兼容问题
eval_strategy / save_steps steps / 25 看 val loss 是否过拟合
1
2
3
4
5
6
7
8
9
trainer = SFTTrainer(
model=base_model,
train_dataset=sft_dataset["train"],
eval_dataset=sft_dataset["validation"],
peft_config=lora_config,
args=training_args,
processing_class=tokenizer,
)
train_result = trainer.train()

训练日志(节选):250 step 内 train loss 约 0.58 → 0.12,val loss 同步下降,mean token accuracy 升至 ~0.94——对「只生成一个标签词」的任务,该指标与学没学会强相关。


八、产物、微调后验证与接回部署

8.1 保存什么

1
2
3
4
5
gemma4-it-emotion-lora-ms-single-gpu/
├── adapter_model.safetensors
├── adapter_config.json
├── checkpoint-*/
└── *.csv(评估明细)

微调落盘的是 adapter,不是 15 GB 全量。推理时 PeftModel.from_pretrained(base_model, OUTPUT_DIR) 挂载;或合并权重后再 vllm serve

8.2 微调后评估

同一 evaluate_model() 在 test 集再跑一遍,与第六节基线对比 accuracy / macro_f1 / invalid,并导出 gemma4_emotion_before_after_metrics.csv 等。训练后优先在内存里评估,避免立刻重载导致显存碎片 OOM。

8.3 接回 vLLM(部署篇延续)

方式 说明
合并全量权重 体积回到 ~15 GB,vLLM 路径与部署篇相同
LoRA 热加载 取决于 vLLM 版本是否支持;轻量但配置更复杂

九、微调专属踩坑(部署篇未覆盖)

现象 原因 对策
MsDataset / verification_mode 报错 modelscope ↔ datasets 版本桥接 用 parquet 直读(本篇默认路径)
训练 OOM 激活 + 优化器 > 推理仅权重 batch_sizemax_lengthTRAIN_LIMIT;开 gradient_checkpointing
Trainable params = 0 target_modules 未匹配 target_modules 或查模型层名
微调后指标没涨 SYSTEM_PROMPT 训练/评估不一致 两边共用同一常量
生成式评估太慢 每条 autoregressive decode 降低 EVAL_LIMIT;正式实验再放大
想多卡训练 Notebook 多进程难调试 .py + accelerate launch,别在 Jupyter 里硬上

显存调参优先顺序:per_device_train_batch_sizemax_lengthTRAIN_LIMIT → 收窄 target_modules


十、小结

相对部署篇「把模型跑成 API」,微调增量可以收成五块:

  1. 依赖trl + peft + datasets + sklearn——训练循环、低秩更新、数据管道、效果量化。
  2. 数据:标注集 → prompt-completion;completion_only_loss 只对标签算梯度。
  3. 加载transformers 训练态(关 use_cache、BF16、不用 bnb 4bit);权重目录可复用部署篇下载结果。
  4. 训练:LoRA ~0.63% 参数 + adamw_torch + gradient_checkpointing;先跑微调前基线再对比。
  5. 产物:adapter + CSV 评估;合并或热加载后接回 vLLM。

建议先 TRAIN_LIMIT=4000、1 epoch 跑通全链路,再放开数据量与 epoch 做正式实验。

-------------本文结束感谢您的阅读-------------