在 AMD 云环境上微调 Gemma 4:情绪分类 LoRA 实战学习笔记

任务:1 小时极速体验 AMD 云环境模型微调 模型:google/gemma-4-E4B-it(魔搭 ModelScope) 数据集:AI-ModelScope/emotion(6 类情绪分类) 微调方式:LoRA(单卡,BF16) 最终结果:accuracy 0.625 → 0.915 ,invalid predictions 2 → 0


一、核心概念速览

在动手之前,先把几个贯穿全程的关键概念过一遍,跑代码时才不会"知其然不知其所以然"。

1. 微调 ≠ 预训练,但本质是同一件事

模型的能力都藏在内部参数里,"训练"就是反复"猜答案 → 对照正确答案 → 调参数"的过程。预训练和微调的区别只在于:

用什么数据 从哪开始 练出什么
预训练 海量通用数据 从零开始 什么都懂一点的"通才"
微调 少量专门数据 在通才基础上接着练 某个领域的"专才"

一句话理解:微调 = 用你的数据,把"通才"模型调成你需要的"专才",改的是模型参数本身。

2. 为什么不直接用提示词(Prompt)解决?

临时任务用提示词足够,但遇到 大量、重复、要求格式高度稳定 的任务,微调比每次写提示词更可靠、更省事,效果也更好。例如:

  • 让输出格式固定(本教程:只回一个情绪词,不啰嗦)

  • 让风格/语气统一(客服口吻、公司模板)

  • 让某个垂直领域更准(医疗、法律、小语种等)

3. LoRA:省显存的微调方法

全量微调要把模型几十亿参数全部重调,单卡扛不住。LoRA 的做法是:

把原模型参数全部"冻住"不动,只额外加一小撮新参数来训练。

打个比方:不是把整本书重写一遍,而是在关键页面贴"便利贴(批注)"。训练产出不是完整新模型,而是这一小撮参数,叫 adapter(适配器)

  • adapter 可以和原模型一起加载,也可以直接 融合(Merge) 成一个完整新模型。

本次实测数据(来自训练日志):

复制代码
Trainable LoRA parameters: 50,499,584
Total parameters:          7,991,600,416
Trainable ratio:            0.6319%

也就是说:只训练了全部参数的 0.63% ,单卡跑 1 个 epoch 用了约 17 分钟(train_runtime ≈ 1027s)。

4. epoch(训练轮数)------ 不是越多越好

把整套"教材"完整学一遍叫一个 epoch。本教程只训练 1 轮,目标是先把流程跑通,不是冲最高分。

⚠️ 危险警告 :epoch 太高会导致 过拟合(Overfitting) ------模型把教材"死记硬背",遇到没见过的新句子反而答不出来。本质是失去了举一反三的能力

5. 评估指标:怎么判断微调有没有效果

指标 含义 方向
accuracy(准确率) 答对的比例 越高越好
macro F1 综合各类情绪表现的得分 越高越好
invalid predictions(无效预测) 模型输出不在 6 个标签内的次数(答非所问) 越低越好

关键区分(容易混淆)

  • 真实标签 fear,模型输出 surprise → 这是 预测错误(按规矩答了,但答错了),算入 accuracy 的"不对"部分。

  • 真实标签 love,模型输出"这句话表达了一种积极的情感" → 这是 无效预测(没在 6 个标签内,答非所问),单独统计。

  • 真实标签 anger,模型输出 anger预测正确

两者分开看的意义:无效预测高 → 模型连"格式要求"都没学会;准确率低 → 格式对了但判断不准。


二、整体流程八步法

跑通这个 Notebook,相当于完整走完了一次工业级微调标准流程,共 8 步:

  1. 安装依赖modelscope(下模型/数据集)、transformers(加载模型/tokenizer)、datasets(读数据)、trl(SFTTrainer 指令微调)、peft(配置 LoRA)、scikit-learn(算指标)。

  2. 检查 GPUtorch.cuda.is_available()。在 AMD ROCm 环境下依然写 cuda,这是 PyTorch 的统一接口命名,不代表用的是 NVIDIA 显卡

  3. 下载模型和数据集:模型 = "待培训的学生",数据集 = "教材"。本版本全部从魔搭 ModelScope 下载,无需 Hugging Face 登录。

  4. 改造数据格式:把"句子 → 情绪标签"改写成 Gemma 4 习惯的"一问一答"聊天格式(system / user / assistant 三段式)。

  5. 微调前评估 :用未训练的模型先测一次,留作对照基线(pre_finetuning)。

  6. LoRA 微调:核心步骤,冻住原参数,只训练新增的一小撮 LoRA 参数。

  7. 保存成果 :adapter 保存到 ./gemma4-it-emotion-lora-ms-single-gpu

  8. 微调后再评估 :对比 post_finetuningpre_finetuning 的成绩。


三、实操记录:环境准备与避坑

3.1 判断是否需要重建环境

云环境就像"酒店房间":关闭并重启(点击过 Destroy Instance )后,系统会把环境依赖复原到初始状态 ,但 已下载的模型文件(./models./datasets)原封不动还在,无需重新下载。

判断方法:看 Active Instance 面板,如果之前点过 "Destroy Instance",说明环境已被复原,需要先做环境修复。

3.2 环境修复三步走(仅"重启过"才需要)

复制代码
# 第一步:打开新终端
# 第二步:卸载不兼容的旧版组件
pip uninstall torchvision -y

会自动回装:后面的安装脚本(第一个代码单元格)会把配套的最新版自动装回来。

3.3 踩坑实录:torchaudio / torchvision 与 ROCm 的版本冲突

这是本次实操中遇到的最大、也最有代表性的坑,记录下来供以后参考。

坑 1:OSError: ... libtorchaudio.so: undefined symbol

现象:运行第二个代码单元格(导入依赖)时报错,根因链路是:

复制代码
transformers 导入 → 顺手尝试加载 torchaudio
→ torchaudio 的 .so 文件是按旧版 torch 编译的
→ 新版 torch 装好后,旧版 torchaudio 加载失败
→ 整条 import 链路崩溃

解决:这次微调根本不需要音频处理,直接卸载即可:

复制代码
pip uninstall torchaudio -y

然后 Kernel → Restart Kernel,重启内核清掉"半成品"导入状态。

坑 2:"Run All" 又把环境装坏了一次(最容易踩的坑!)

现象 :重启内核后再次 "Run All",同一个错误又出现了 (变成了 RuntimeError: operator torchvision::nms does not exist)。

根因:第一个代码单元格本身就是一条安装命令:

复制代码
!uv pip install -U vllm modelscope transformers accelerate datasets trl peft scikit-learn pandas tqdm torchvision \
  --no-cache -i https://mirrors.cloud.tencent.com/pypi/simple/ --extra-index-url https://wheels.vllm.ai/rocm/

这条命令里含有 torchvision !每次 "Run All" 都会重新执行这一句,把刚刚手动卸载掉的 torchvision 重新装回来 ,而这个新版 torchvision 与当前 ROCm 版 torch2.10.0+git8514f05)不兼容,于是错误死循环。

最终解法

复制代码
pip uninstall torchvision -y
pip uninstall torchaudio -y

然后 Kernel → Restart Kernel这次不要点 "Run All" ,跳过第一个安装单元格,从第二个单元格开始逐个/分段运行(Run → Run All Below,从第二格开始)。

💡 经验总结 :在 ROCm 环境下,如果 transformerstorchvision::nms / libtorchaudio.so 相关错误,第一反应应该是检查"是否被某条安装命令重新装上了不兼容版本的 torchvision/torchaudio",而不是反复重装。

坑 3:自动保存报错 "File Save Error: Invalid response: 502"

现象:代码运行过程中弹出文件保存失败的弹窗。

结论 :这只是 Notebook 自动保存 时网络瞬时抖动导致的提示,与代码运行结果无关,点击 "Close" 关掉即可,不影响训练继续进行。

3.4 释放显存的注意事项

如果之前启动过 vLLM 对话服务,点击运行微调 Notebook 之前 ,必须先回到对应终端按 Ctrl+C 停止它------显存(VRAM)是共享资源,对话服务占满显存会导致微调因"显存不足"报错。

3.5 三个存储区域,避免文件"消失"

区域 路径 特点
左侧文件栏 /workspace(默认工作区) 自动存盘,但无法跨机器移动
系统网络同步盘 /network-workspace(绝对路径) 每人 20GB,支持跨机器同步,重要文件建议放这里
运行环境/已装库 (pip 安装的库等) 断开连接或 Destroy 后瞬间清空 ,需用 requirements.txt/Dockerfile 保存

四、运行结果记录

4.1 设备检查

复制代码
torch version: 2.10.0+git8514f05
torch.cuda.is_available(): True
torch.cuda.device_count(): 1
current device: 0
device name: AMD Radeon Graphics

✅ 在 ROCm 环境下 torch.cuda 接口照常可用,AMD 显卡已就位。

4.2 数据集概览

AI-ModelScope/emotion(= dair-ai/emotion 官方镜像),字段 text(string)/ label(int,对应 6 类):

复制代码
0 → sadness   1 → joy    2 → love
3 → anger     4 → fear   5 → surprise

本次实际使用数据量(出于跑通速度考虑,做了截取):

复制代码
TRAIN_LIMIT      = 4000
VALIDATION_LIMIT = 400
TEST_LIMIT       = 400
EVAL_LIMIT       = 400

4.3 微调前评估(pre_finetuning,基线)

复制代码
{'accuracy': 0.625,
 'macro_f1': 0.4824237718715574,
 'invalid_predictions': 2,
 'evaluated_examples': 400}

4.4 LoRA 微调过程

复制代码
Trainable LoRA parameters: 50,499,584
Total parameters:          7,991,600,416
Trainable ratio:            0.6319%
​
TrainOutput(global_step=250,
            training_loss=0.314503368973732,
            metrics={'train_runtime': 1027.0109,
                     'train_samples_per_second': 3.895,
                     'train_steps_per_second': 0.243,
                     'epoch': 1.0})

4.5 微调后评估(post_finetuning)

复制代码
{'accuracy': 0.915,
 'macro_f1': 0.8644926299207644,
 'invalid_predictions': 0,
 'evaluated_examples': 400}

4.6 前后对比

stage accuracy macro_f1 invalid_predictions evaluated_examples
pre_finetuning 0.625 0.4824 2 400
post_finetuning 0.915 0.8645 0 400

4.7 真实新句子测试(模型从未见过)

微调后的模型对全新句子的预测:

复制代码
I feel completely heartbroken and alone.            => sadness   ✅
This is the best day of my life!                    => joy       ✅
I am really scared about what might happen tomorrow.=> fear      ✅
I can't believe they remembered my birthday!        => surprise/joy
I am so angry that nobody listened to me.           => anger     ✅
I really love spending time with my family.         => love      ✅

说明模型真正学到了"识别情绪"的规律,而不是死记硬背训练集(呼应"epoch=1,避免过拟合"的设计目的)。

4.8 产出文件清单

全部保存在 ./gemma4-it-emotion-lora-ms-single-gpu 目录下:

  • LoRA adapter(核心产物,本次训练新学到的"情绪识别本领")

  • train_metrics.json:训练过程指标

  • gemma4_emotion_before_after_metrics.csv:前后对比表

  • gemma4_emotion_prediction_examples.csv / gemma4_emotion_changed_predictions.csv:逐样本预测明细 / 预测发生变化的样本

  • pre/post_finetuning_predictions.csv:前后各自的逐样本预测

  • pre/post_finetuning_classification_report.csv:分类报告(precision/recall/f1)

  • pre/post_finetuning_confusion_matrix.csv:混淆矩阵


五、这台环境还能用来微调什么?

同一台机器(单卡 + LoRA)+ 同一套流程,换数据集即可复用,适合轻量、专门、目标明确的场景:

  1. 换个"分类"任务(最容易上手):差评/好评、垃圾信息识别、新闻分类、用户意图识别......数据格式同样是"一句话 → 一个标签"。

  2. 教它固定的"输出格式" :把口语句子改写成结构化信息(如 "老王明天下午三点来开会" → {姓名:老王, 时间:明天15:00})。

  3. 调出特定的"语气/人设":某品牌客服口吻、某种文风的文案生成器(例如:用某位作者的作品微调出"风格化创作专家")。

  4. 让它更懂某个垂直领域:某门课程、某产品、某行业术语的问答助手。

复用步骤永远是这三步:

  1. 准备自己的数据集(整理成"一问一答"格式)

  2. 把 Notebook 里加载数据的部分换成自己的数据

  3. 点运行,等待完成

提示:单卡小模型的强项是把"通才"在某件事上调得更专、更稳,不适合从头打造一个无所不能的大助手。建议先从"换个分类数据集"起步,跑顺后再挑战更复杂的玩法(重新加载 adapter、加大数据量、调整 epoch 等)。


六、收尾:释放云资源

AMD Radeon Cloud 免费额度按使用时长扣费,哪怕不跑代码也持续扣费。本次任务定位是"体验流程",无需下载文件,确认打卡截图已保存后:

回到 Profile(个人主页)→ Active Instance → 点击红色 "Destroy Instance" 即可。


附录:完整代码(按 Notebook 执行顺序)

以下代码按 Notebook 中各个 Cell 的执行顺序拼接而成(不含第一个安装依赖的 Cell,因为它会重装 torchvision 导致版本冲突,详见上文"踩坑实录")。后续如需复现/复用,建议:先按需在终端手动安装好依赖,再依次运行下列代码块。

复制代码
import os
import re
import json
import random
import warnings
​
import numpy as np
import pandas as pd
import torch
​
from tqdm.auto import tqdm
from datasets import Dataset, DatasetDict, ClassLabel, load_dataset
from sklearn.metrics import accuracy_score, classification_report, confusion_matrix, f1_score
​
from modelscope import snapshot_download
from modelscope.hub.snapshot_download import dataset_snapshot_download
​
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
from peft import LoraConfig, PeftModel
from trl import SFTConfig, SFTTrainer
​
warnings.filterwarnings("ignore")
​
# -----------------------------
# 基础配置
# -----------------------------
# 魔搭上的模型 ID(Gemma 4 E4B-it 在 ModelScope 上的官方仓库,instruction-tuned 版本,
# 仓库内自带官方 chat_template.jinja,无需手动处理 chat template)。
# 仓库地址: https://www.modelscope.cn/models/google/gemma-4-E4B-it
MODELSCOPE_MODEL_ID = "google/gemma-4-E4B-it"
​
# 魔搭上的数据集 ID(dair-ai/emotion 在 ModelScope 上的官方镜像)。
MODELSCOPE_DATASET_ID = "AI-ModelScope/emotion"
​
# 微调输出目录
OUTPUT_DIR = "./gemma4-it-emotion-lora-ms-single-gpu"
​
# 数据量控制。先用小数据跑通,确认没问题后再加大。
TRAIN_LIMIT = 4000
VALIDATION_LIMIT = 400
TEST_LIMIT = 400
EVAL_LIMIT = 400
​
SEED = 42
MODEL_DTYPE = torch.bfloat16
BF16 = True
FP16 = False
​
SYSTEM_PROMPT = """You are an emotion classification assistant.
Read the user's text and answer with exactly one label.
Only choose from: sadness, joy, love, anger, fear, surprise.
Return only the label and nothing else."""
​
LABEL_PATTERN = re.compile(r"(sadness|joy|love|anger|fear|surprise)", re.IGNORECASE)
​
os.makedirs(OUTPUT_DIR, exist_ok=True)
os.makedirs("./models", exist_ok=True)
os.makedirs("./datasets", exist_ok=True)
​
print("torch version:", torch.__version__)
print("torch.cuda.is_available():", torch.cuda.is_available())
print("torch.cuda.device_count():", torch.cuda.device_count())
if torch.cuda.is_available():
    print("current device:", torch.cuda.current_device())
    print("device name:", torch.cuda.get_device_name(0))
复制代码
# -----------------------------
# 固定随机种子
# -----------------------------
def setup_seed(seed: int = 42):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    set_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)
​
setup_seed(SEED)
复制代码
# -----------------------------
# 从魔搭 ModelScope 下载 Gemma 模型
# -----------------------------
MODELSCOPE_MODEL_ID = "google/gemma-4-E4B-it"
print("Downloading model from ModelScope...")
print("ModelScope model id:", MODELSCOPE_MODEL_ID)
​
model_dir = snapshot_download(
    MODELSCOPE_MODEL_ID,
    cache_dir="./models",
)
​
print("Downloaded model dir:", model_dir)
​
# 后续统一使用本地路径加载
LOCAL_MODEL_DIR = model_dir
复制代码
# -----------------------------
# 从魔搭 ModelScope 加载情绪分类数据集
# -----------------------------
import glob
​
EMOTION_LABEL_NAMES = ["sadness", "joy", "love", "anger", "fear", "surprise"]
​
​
# 直接把魔搭上的数据集仓库(parquet 文件)整体下载到本地,然后用 datasets 库从本地 parquet 加载。
# 不走 MsDataset.load -> datasets.load_dataset 的桥接路径,可以规避 modelscope 与 datasets 之间
# `as_dataset() got an unexpected keyword argument 'verification_mode'` 这类版本错配错误。
print("Downloading dataset from ModelScope...")
print("ModelScope dataset id:", MODELSCOPE_DATASET_ID)
​
dataset_dir = dataset_snapshot_download(
    MODELSCOPE_DATASET_ID,
    cache_dir="./datasets",
)
print("Downloaded dataset dir:", dataset_dir)
​
​
def _parquet_files_for(split_name: str):
    pattern = os.path.join(dataset_dir, "data", f"{split_name}-*.parquet")
    files = sorted(glob.glob(pattern))
    if not files:
        raise FileNotFoundError(
            f"No parquet files matched pattern: {pattern}. "
            f"Please check the dataset repo layout under {dataset_dir}."
        )
    return files
​
​
raw_dataset = load_dataset(
    "parquet",
    data_files={
        "train": _parquet_files_for("train"),
        "validation": _parquet_files_for("validation"),
        "test": _parquet_files_for("test"),
    },
)
​
# 从 parquet 加载时,label 字段类型会退化成普通整数,这里显式 cast 成 ClassLabel,
# 这样后续 `dataset["train"].features["label"].names` 和原始 HF 版接口完全一致。
for split_name in list(raw_dataset.keys()):
    if not isinstance(raw_dataset[split_name].features.get("label"), ClassLabel):
        raw_dataset[split_name] = raw_dataset[split_name].cast_column(
            "label", ClassLabel(names=EMOTION_LABEL_NAMES)
        )
​
print("Raw dataset:", raw_dataset)
​
​
def maybe_limit(split, limit):
    split = split.shuffle(seed=SEED)
    if limit is None:
        return split
    return split.select(range(min(limit, len(split))))
​
​
dataset = DatasetDict({
    "train": maybe_limit(raw_dataset["train"], TRAIN_LIMIT),
    "validation": maybe_limit(raw_dataset["validation"], VALIDATION_LIMIT),
    "test": maybe_limit(raw_dataset["test"], TEST_LIMIT),
})
​
label_names = dataset["train"].features["label"].names
VALID_LABELS = set(label_names)
ALL_EVAL_LABELS = label_names + ["INVALID"]
​
print(dataset)
print("label_names:", label_names)
print("example:", dataset["train"][0])
复制代码
# -----------------------------
# 改写数据格式为"一问一答"聊天格式
# -----------------------------
def to_prompt_completion(example):
    text = example["text"]
    label = label_names[example["label"]]
    user_content = f"Classify the emotion of this text:\n\n{text}"
    return {
        "prompt": [
            {"role": "system", "content": SYSTEM_PROMPT},
            {"role": "user", "content": user_content},
        ],
        "completion": [
            {"role": "assistant", "content": label},
        ],
    }
​
sft_dataset = dataset.map(
    to_prompt_completion,
    remove_columns=dataset["train"].column_names,
)
​
print(sft_dataset)
print(sft_dataset["train"][0])
复制代码
# -----------------------------
# 加载 tokenizer 并确保 chat_template 可用
# -----------------------------
print("Loading tokenizer from:", LOCAL_MODEL_DIR)
​
tokenizer = AutoTokenizer.from_pretrained(
    LOCAL_MODEL_DIR,
    use_fast=True,
    trust_remote_code=True,
)
​
if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token
​
print("pad_token:", tokenizer.pad_token)
print("eos_token:", tokenizer.eos_token)
​
# `google/gemma-4-E4B-it` 的 tokenizer 通常会自带 chat_template。
# 若缺失(缓存不完整等),从同一魔搭仓库拉取官方 chat_template.jinja 注入
# (权重已在上面整仓下载时可跳过额外拉取)。
TEMPLATE_SOURCE_MODEL_ID = "google/gemma-4-E4B-it"
​
def _load_official_gemma_chat_template() -> str:
    """从 gemma-4-E4B-it 仓库下载官方 chat_template.jinja 并返回字符串。
​
    主路径:modelscope.snapshot_download(allow_file_pattern=["chat_template.jinja"])
    兜底:ModelScope raw file API 直接 HTTP GET
    """
    try:
        template_dir = snapshot_download(
            TEMPLATE_SOURCE_MODEL_ID,
            cache_dir="./models",
            allow_file_pattern=["chat_template.jinja"],
        )
        path = os.path.join(template_dir, "chat_template.jinja")
        if os.path.exists(path):
            with open(path, "r", encoding="utf-8") as f:
                return f.read()
    except Exception as e:
        print("snapshot_download(allow_file_pattern) failed, fallback to HTTP. err =", e)
​