news 2026/9/23 3:31:21

PaddleNLP 中的 Gemma 模型精调实战:从 SFT、LoRA 到 DPO/KTO 对齐全流程指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PaddleNLP 中的 Gemma 模型精调实战:从 SFT、LoRA 到 DPO/KTO 对齐全流程指南
  • 人工智能
  • 大模型
  • NLP
  • 深度学习
  • 预训练
  • 微调
  • RLHF
  • 模型量化

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

Gemma 是 Google DeepMind 基于 Gemini 同源研究与技术打造的开源轻量级大语言模型家族。本指南以 docs/en/llm/config/gemma/README.md 为核心骨架,结合 llm/config/gemma 目录下五份开箱即用的精调配置以及 paddlenlp/transformers/gemma 的源码实现,系统讲解在 PaddleNLP 飞桨大模型套件中如何对 Gemma 完成全参 SFT、LoRA 轻量微调以及 DPO、KTO 人类偏好对齐,并介绍相关的 4D 并行分布式策略与 Zero Padding 等优化手段。读完本文,你将能够直接复制配置并一键启动 Gemma-2B/7B 的精调与对齐任务。

1. Gemma 模型概览与支持权重

Gemma 由 Google DeepMind 及 Google 其他团队开发,是一个轻量级、业界领先的开源模型家族,其构建所使用的研究与技术正是用于创建 Gemini 模型的那一套,因此天然具备较强的指令遵循与生成能力。在 PaddleNLP 中,Gemma 已被完整接入预训练权重加载、Tokenizer、模型结构与各类精调流程。

PaddleNLP 目前支持的 Gemma 官方权重如下:

模型权重
google/gemma-7b
google/gemma-7b-it
google/gemma-2b
google/gemma-2b-it

其中带-it后缀的是经过指令微调的对话版本,适合直接推理或作为对齐任务的基座;不带-it的是基座模型,适合继续做 SFT 与对齐训练。这些权重可通过model_name_or_path字段直接指定,PaddleNLP 会从模型中心自动下载并转换为飞桨格式,见 configuration.py 中维护的GEMMA_PRETRAINED_RESOURCE_FILES_MAP

2. 精调前的环境准备与数据格式

按照 llm/README.md 的说明,使用大模型套件前建议安装 PaddleNLP 最新 develop 版本:

pip install --pre --upgrade paddlenlp -f https://www.paddlepaddle.org.cn/whl/paddlenlp.html

2.1 精调数据格式

套件支持的精调数据是每行包含一个字典的 JSON 文件,每个字典包含两个核心字段:

  • srcstrList(str),模型的输入指令(instruction)、提示(prompt),即模型需要执行的任务;
  • tgtstrList(str),模型期望输出的目标内容。

样例数据:

{"src": "Give three tips for staying healthy.", "tgt": "1.Eat a balanced diet and make sure to include plenty of fruits and vegetables. \n2. Exercise regularly to keep your body active and strong. \n3. Get enough sleep and maintain a consistent sleep schedule."}

llm目录下解压官方提供的 alpaca demo 数据集即可快速跑通全流程(数据准备细节见 llm/README.md):

wget https://bj.bcebos.com/paddlenlp/datasets/examples/alpaca_demo.gz tar -xvf alpaca_demo.gz

多轮对话场景下,套件还支持统一的对话模板,相关说明见 多轮对话文档。

3. PaddleNLP 中的 Gemma 模型实现

在进入配置讲解之前,先了解 Gemma 在 PaddleNLP 中的源码实现,有助于理解后续各配置项的底层含义。Gemma 相关代码集中在 paddlenlp/transformers/gemma 目录,由五个文件组成:

  • configuration.pyGemmaConfig配置类与预训练权重资源映射;
  • modeling.py:模型主干实现;
  • modeling_pp.py:流水线并行(Pipeline Parallel)版本实现;
  • tokenizer.py/tokenizer_fast.py:基于 SentencePiece 的慢速/快速分词器。

3.1 GemmaConfig 关键结构参数

在 configuration.py 中,GemmaConfig继承了PretrainedConfig,并通过model_type = "gemma"标识模型类型。GEMMA_PRETRAINED_INIT_CONFIGURATION给出了google/gemma-2b的完整默认结构参数:

参数gemma-2b 默认值含义
hidden_size2048隐藏层维度
intermediate_size16384MLP 中间层维度
num_hidden_layers28Transformer 编码器层数
num_attention_heads8注意力头数
num_key_value_heads1KV 头数(GQA 分组查询注意力)
num_attention_heads / num_key_value_heads8:1说明采用 GQA 结构
rms_norm_eps1e-6RMS 归一化层 epsilon
vocab_size256000词表大小
max_position_embeddings8192最大序列长度
bos/eos/pad_token_id2/1/0特殊 token 编号
head_dim256每头维度
rope_theta10000.0RoPE 旋转位置编码基频
tie_word_embeddingsTrue是否绑定输入输出词嵌入

GemmaConfig.__init__还暴露了fuse_attention_qkv(融合 QKV 投影)、fuse_attention_ffn(融合 FFN 投影)、alibi(是否使用 ALiBi 位置编码)等优化开关,并在rope属性中定义rope = not self.alibi,即默认启用 RoPE 旋转位置编码。

3.2 模型结构核心模块

在 modeling.py 中,Gemma 的模型结构由以下核心组件构成,从源码结构可以清晰地看到其 Decoder-only 架构:

  • GemmaRMSNorm(L352):RMS 归一化层;
  • GemmaRotaryEmbedding(L378):旋转位置编码嵌入层;
  • GemmaMLP(L418):前馈网络(包含gate_projup_projdown_proj三个投影);
  • GemmaAttention(L462):GQA 分组查询注意力实现;
  • GemmaDecoderLayer(L785):单个 Transformer 解码层(RMSNorm + Attention + MLP 的残差结构);
  • GemmaModel(L1071):整个骨干网络;
  • GemmaForCausalLM(L1443):带语言建模头的因果语言模型入口,支持generate推理。

分词器方面,tokenizer.py 中的GemmaTokenizer基于 SentencePiece(tokenizer.model),model_input_names = ["input_ids", "attention_mask"],特殊 token 为<unk><bos><eos><pad>

4. 全参精调:SFT

llm/config/gemma/sft_argument.json 是 Gemma 全参 SFT 的推荐配置,完整内容如下:

{ "model_name_or_path": "google/gemma-2b", "dataset_name_or_path": "./data", "output_dir": "./checkpoints/sft_ckpts", "per_device_train_batch_size": 2, "gradient_accumulation_steps": 1, "per_device_eval_batch_size": 8, "eval_accumulation_steps": 16, "num_train_epochs": 3, "learning_rate": 3e-05, "warmup_steps": 30, "logging_steps": 1, "evaluation_strategy": "epoch", "save_strategy": "epoch", "src_length": 512, "max_length": 1024, "fp16": true, "fp16_opt_level": "O2", "do_train": true, "do_eval": true, "disable_tqdm": true, "load_best_model_at_end": true, "eval_with_do_generation": false, "metric_for_best_model": "accuracy", "recompute": true, "save_total_limit": 1, "tensor_parallel_degree": 1, "pipeline_parallel_degree": 1, "sharding_parallel_degree": 8, "sharding": "stage2", "zero_padding": false, "unified_checkpoint": true, "use_flash_attention": true }

关键配置项说明:

  • model_name_or_pathgoogle/gemma-2b,指定基座权重;可替换为google/gemma-2b-itgoogle/gemma-7b等上表权重;
  • dataset_name_or_path:指向./data目录,即上文下载并解压的 alpaca 数据集目录;
  • output_dir:精调产出的模型与 checkpoint 保存路径;
  • 训练超参per_device_train_batch_size=2gradient_accumulation_steps=1num_train_epochs=3learning_rate=3e-5warmup_steps=30,配合evaluation_strategy/save_strategy均为epoch,即每个 epoch 末评估与保存一次;load_best_model_at_end=true会在训练结束后加载验证集上metric_for_best_model="accuracy"最优的 checkpoint;
  • 序列长度src_length=512(输入长度)、max_length=1024(输入 + 输出的总长度上限);
  • 混合精度fp16=true搭配fp16_opt_level="O2",开启飞桨 O2 级别的 FP16 混合精度训练;
  • 显存优化recompute=true开启激活重计算(以时间换显存),use_flash_attention=true使用 FlashAttention 加速注意力计算;
  • 分布式并行tensor_parallel_degree=1pipeline_parallel_degree=1sharding_parallel_degree=8sharding="stage2",构成"分组参数切片的数据并行(ZeRO stage2)"策略,即 8 卡场景下对优化器状态和梯度进行切分,显著降低每卡显存占用;unified_checkpoint=true启用统一 checkpoint 格式,便于跨并行策略复用权重;
  • Zero Paddingzero_padding=false,如需进一步降低无效 pad token 占比、提升训练效率,可开启为true(配合 FlashMask 使用效果更佳)。

llm目录下,单卡启动 Gemma SFT 精调:

python -u run_finetune.py ./config/gemma/sft_argument.json

多卡(8 卡)启动时使用paddle.distributed.launch

python -u -m paddle.distributed.launch --devices "0,1,2,3,4,5,6,7" run_finetune.py ./config/gemma/sft_argument.json

其中run_finetune.py是套件的统一精调入口(见 llm/run_finetune.py),它根据配置文件自动组装 Trainer 并应用 4D 并行策略。

5. 偏好对齐:DPO 与 LoRA DPO

在 SFT 之后,通常需要借助人类偏好数据对模型进行对齐。PaddleNLP 在 llm/alignment/dpo 目录提供了 DPO 训练入口,Gemma 的 DPO 配置位于 dpo_argument.json 与 dpo_lora_argument.json。

5.1 全参 DPO

dpo_argument.json 的完整内容:

{ "model_name_or_path": "google/gemma-2b", "train_dataset_path": "./data/train.jsonl", "dev_dataset_path": "./data/dev.jsonl", "output_dir": "./checkpoints/dpo_ckpts", "per_device_train_batch_size": 1, "gradient_accumulation_steps": 1, "per_device_eval_batch_size": 1, "num_train_epochs": 1, "max_steps": 100, "learning_rate": 1e-06, "warmup_steps": 10, "logging_steps": 1, "evaluation_strategy": "steps", "save_strategy": "steps", "eval_steps": 100, "save_steps": 500, "max_seq_len": 4096, "max_prompt_len": 2048, "bf16": true, "fp16_opt_level": "O2", "do_train": true, "do_eval": true, "disable_tqdm": true, "load_best_model_at_end": true, "tensor_parallel_degree": 2, "sharding": "stage1", "use_flash_attention": true, "recompute": false, "recompute_granularity": "full", "beta": 0.1, "benchmark": false, "loss_type": "sigmoid", "label_smoothing": 0.0, "unified_checkpoint": true, "autotuner_benchmark": false, "lazy": false, "seed": 42, "sft_loss_ratio": 0.0 }

与 SFT 配置相比,DPO 特有的关键项:

  • 数据:改为train_dataset_path/dev_dataset_path,指向包含偏好对(chosen/rejected)的train.jsonldev.jsonl
  • 序列长度max_seq_len=4096max_prompt_len=2048,DPO 需要同时容纳 prompt 与答案对,序列更长;
  • DPO 超参beta=0.1为 DPO 正则化系数(控制对参考策略的偏离程度),loss_type="sigmoid"为 DPO 原始 sigmoid 损失形式,label_smoothing=0.0关闭标签平滑,sft_loss_ratio=0.0表示不混入 SFT 损失;
  • 精度:使用bf16=true(BF16 混合精度),在大规模对齐训练中 BF16 相比 FP16 拥有更大的动态范围;
  • 分布式tensor_parallel_degree=2配合sharding="stage1",采用 2 路张量并行 + ZeRO stage1(切分优化器状态)的组合。

启动命令(参照 llm/README.md 中 DPO 的用法,将配置替换为 Gemma 的):

python -u ./alignment/dpo/run_dpo.py ./config/gemma/dpo_argument.json

多卡启动:

python -u -m paddle.distributed.launch --devices "0,1,2,3,4,5,6,7" ./alignment/dpo/run_dpo.py ./config/gemma/dpo_argument.json

5.2 LoRA DPO(轻量对齐)

当显存受限或希望快速迭代时,可使用 dpo_lora_argument.json 中的 LoRA DPO 配置。它在全参 DPO 的基础上,仅额外改动以下几点:

"learning_rate": 1e-05, "gradient_accumulation_steps": 8, "tensor_parallel_degree": 1, "lora": true, "lora_rank": 64, "rslora_plus": true
  • lora=true:启用 LoRA 低秩适配,只训练注入的低秩矩阵,冻结其余全部参数;
  • lora_rank=64:LoRA 低秩矩阵的秩,秩越大适配能力越强、参数量也越多;
  • rslora_plus=true:开启 RsLoRA+ 变体(基于秩缩放改进的 LoRA,对较大 rank 更稳定);
  • 同时将tensor_parallel_degree降为 1(LoRA 冻结主干后对张量并行的依赖降低),并通过gradient_accumulation_steps=8模拟更大的等效 batch。

LoRA DPO 的启动命令与全参 DPO 相同,仅配置文件替换为dpo_lora_argument.json

python -u ./alignment/dpo/run_dpo.py ./config/gemma/dpo_lora_argument.json

6. 偏好对齐:KTO 与 LoRA KTO

KTO(Kahneman-Tversky Optimization)是另一种不需要成对偏好数据的人类对齐方法,它只需要对单个输出打上"可取/不可取"标签。PaddleNLP 提供 KTO 训练入口于 llm/alignment/kto,Gemma 对应配置为 kto_argument.json 与 kto_lora_argument.json。

6.1 全参 KTO

kto_argument.json 的核心内容:

{ "model_name_or_path": "google/gemma-2b", "train_dataset_path": "./data/train.jsonl", "dev_dataset_path": "./data/dev.jsonl", "output_dir": "./checkpoints/kto_ckpts", "per_device_train_batch_size": 1, "gradient_accumulation_steps": 8, "per_device_eval_batch_size": 1, "num_train_epochs": 1, "max_steps": 100, "learning_rate": 2e-06, "warmup_steps": 10, "max_seq_len": 4096, "max_prompt_len": 2048, "bf16": true, "fp16_opt_level": "O2", "tensor_parallel_degree": 8, "sharding": "stage1", "use_flash_attention": true, "recompute": false, "recompute_granularity": "full", "beta": 0.1, "unified_checkpoint": true, "seed": 42 }

特点说明:

  • 与 DPO 共用beta=0.1正则化系数,但损失函数采用 KTO 形式,无需成对样本;
  • learning_rate=2e-6,相比 DPO 更低,KTO 对学习率更为敏感;
  • tensor_parallel_degree=8配合sharding="stage1",即 8 路张量并行,适合单机多卡显存吃紧时对大模型做对齐;
  • recompute_granularity="full"定义了重计算的粒度,开启recompute时可按 full 粒度重算整层激活。

启动命令(参照 llm/README.md 的 KTO 用法):

python -u -m paddle.distributed.launch --devices "0,1,2,3,4,5,6,7" ./alignment/kto/run_kto.py ./config/gemma/kto_argument.json

6.2 LoRA KTO

kto_lora_argument.json 在 KTO 基础上仅追加 LoRA 相关开关:

"learning_rate": 2e-05, "lora": true

即使用默认 rank 的 LoRA 对 Gemma 做轻量 KTO 对齐,学习率由全参的2e-6提升至2e-5(LoRA 训练通常需要更高的学习率)。启动命令:

python -u -m paddle.distributed.launch --devices "0,1,2,3,4,5,6,7" ./alignment/kto/run_kto.py ./config/gemma/kto_lora_argument.json

7. 4D 并行与显存优化策略小结

Gemma 各配置文件中的分布式字段统一遵循飞桨大模型套件的 4D 并行设计(详见 llm/README.md),用户只需修改 Trainer 配置即可组合不同策略:

  • 数据并行(DP):默认维度,每卡持有完整模型副本、切分数据;
  • 分组参数切片(Sharding / ZeRO)sharding="stage1"只切分优化器状态,sharding="stage2"进一步切分梯度,sharding="stage3"再切分模型参数,配合sharding_parallel_degree指定切分卡数;
  • 张量并行(TP)tensor_parallel_degree将单层权重按头/列切分到多卡;
  • 流水线并行(PP)pipeline_parallel_degree将模型按层切分到多卡,modeling_pp.py即为 Gemma 的流水线并行实现。

在显存优化层面,套件还提供zero_padding(零填充,减少 pad token 无效计算)、recompute(激活重计算)、use_flash_attention(FlashAttention 内核加速)、unified_checkpoint(统一 checkpoint 便于并行策略间迁移)等选项,均已在上述 Gemma 配置中给出推荐取值,可根据硬件显存规模灵活调整。

8. 小结

PaddleNLP 为 Gemma 提供了从权重加载、模型结构到精调对齐的完整链路支持:通过 llm/config/gemma 下的sft_argument.jsondpo_argument.jsondpo_lora_argument.jsonkto_argument.jsonkto_lora_argument.json五份配置,分别覆盖全参 SFT、全参/LoRA DPO、全参/LoRA KTO 五种主流训练范式;底层由 GemmaConfig 与 modeling.py 提供与官方一致的 GQA 注意力、RMSNorm、RoPE 等架构实现;训练侧则由 llm/run_finetune.py、llm/alignment/dpo/run_dpo.py、llm/alignment/kto/run_kto.py 统一承载,并集成了 4D 并行、混合精度、FlashAttention 与统一 checkpoint 等工程化能力。开发者只需准备好src/tgt(或偏好对/带标签)格式的数据,修改配置文件中的模型与路径字段,即可在单卡或多卡环境下完成 Gemma 的定制化训练。

  • 人工智能
  • 大模型
  • NLP
  • 深度学习
  • 预训练
  • 微调
  • RLHF
  • 模型量化

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 3:30:56

告别呼吸的痛:从入门到精通的调试心法

告别呼吸的痛:从入门到精通的调试心法 复制来的代码跑不通,看着满屏红色的报错信息,是不是感觉胸口发闷,像得了呼吸的痛?别慌,这是每个开发者从入门到精通必经的“渡劫”时刻。很多新手遇到这种情况,第一反应是删掉重写,或者在Stack Overflow上疯狂搜索,结果越改越乱。…

作者头像 李华
网站建设 2026/9/23 3:30:37

告别文档迷宫:快乐大本营官网手写实现对比与选型

告别文档迷宫:快乐大本营官网手写实现对比与选型 官方文档太长抓不住重点,这是每个开发者在接手新项目或探索新技术栈时的第一痛点。面对【快乐大本营官网】这种高并发、重交互的页面结构,直接照抄文档里的示例代码往往只能解决表面问题,无法应对真实的业务复杂性。要想真正吃透其背后的逻辑,必须动手【手写实现】核心…

作者头像 李华
网站建设 2026/9/23 3:30:25

我是歌手梁博入门到精通性能优化避坑指南

我是歌手梁博入门到精通性能优化避坑指南 看了一堆教程还是不会写项目,这是不是你的真实写照? 别慌,你不是一个人。 很多开发者在从【入门到精通】的进阶路上,都卡在了“原理懂但手废”的死胡同里。…

作者头像 李华
网站建设 2026/9/23 3:30:18

2026最新:别只背邹忌讽齐王纳谏原文,用代码重构讽谏逻辑

2026最新:别只背邹忌讽齐王纳谏原文,用代码重构讽谏逻辑 你背得滚瓜烂熟的《邹忌讽齐王纳谏》,是不是在考场上让你拿了满分,但在实际业务里却让你束手无策?很多开发者陷入同一个死胡同: 学会语法却不知怎么搭项目…

作者头像 李华
网站建设 2026/9/23 3:30:04

3个高频面试题拆解:GUI界面选型避坑指南

3个高频面试题拆解:GUI界面选型避坑指南 官方文档厚得像砖头,翻半天还是不知道哪个框架适合你的项目?别急,GUI界面开发里的坑,我踩了十年,今天直接给你掏心窝子讲透。 这不只是技术选型,更是面试桌上的 高频面试题 。面试官问“为什么选这个框架”,你要是只会背“性能好”,基本就凉一半。Stack…

作者头像 李华
网站建设 2026/9/23 3:29:36

19e数字便民图解原理:3步搞定项目落地难题

19e数字便民图解原理:3步搞定项目落地难题 是不是刷了上百篇技术博客,收藏了无数“保姆级教程”,结果真上手写个像样的项目,脑子还是空的?那种“懂了但不会”的无力感,真的能把人逼疯。很多开发者卡在从“看代码”到“写代码”的鸿沟上,根本原因不是智商不够,而是缺乏对底层逻辑的直观感知。单纯看文字描述太抽…

作者头像 李华