首页
/ FastChat 微调实践指南:FSDP、QLoRA 与多硬件全参数训练

FastChat 微调实践指南:FSDP、QLoRA 与多硬件全参数训练

2026-09-05 17:37:45作者:蔡怀权

FastChat 仓库为 Vicuna 系模型与 FastChat-T5 提供了完整的监督微调(SFT)工具链,docs/training.md 是这套工具链的官方操作手册,涵盖基于 FSDP 的 T5 全参数微调、基于 DeepSpeed ZeRO 的 LoRA/QLoRA 低秩微调,以及面向本地 NPU 集群的 Vicuna-7B 全参数训练。读完本文,你将掌握这四种训练路径的完整命令、参数含义与适用场景,并能结合 fastchat/train 下的源码理解数据预处理、loss 掩码与检查点保存的底层机制。

训练数据格式:所有脚本共用的对话 JSON

四条训练路径共享同一种数据格式。仓库自带示例数据 data/dummy_conversation.json,每条样本是一个对象,conversations 字段交替存放 humangpt 两个角色的消息:

[
  {
    "id": "identity_0",
    "conversations": [
      { "from": "human", "value": "Who are you?" },
      { "from": "gpt", "value": "I am Vicuna, a language model trained by researchers from LMSYS." }
    ]
  }
]

所有训练脚本通过 --data_path 指向该 JSON 文件。从源码结构看,fastchat/train/train.pymake_supervised_data_module 直接 json.load 该文件并取 example["conversations"] 作为训练输入,因此自定义数据集时只需保持这一结构即可。

全参数微调:数据预处理与 loss 掩码机制

在讲解具体命令之前,先看全参数训练脚本 fastchat/train/train.py 的预处理逻辑,它解释了为什么微调只对模型回复计算损失:

  • preprocess 函数使用 Vicuna 对话模板(get_conversation_template("vicuna"))把 human/gpt 消息渲染成带 USER:/ASSISTANT: 标记的完整 prompt,再按 tokenizer.model_max_length 填充截断(见 train.py#L92-L177)。
  • 随后将 input_ids 克隆为 targets,用分隔符 conv.sep + conv.roles[1] + ": " 逐轮切分,把用户指令片段和 padding 位置的标签全部替换为 IGNORE_TOKEN_ID(即 LabelSmoother.ignore_index),最终只对 ASSISTANT 回复的 token 计算交叉熵损失。
  • 数据集有两种加载方式:SupervisedDataset 在构造时一次性完成全部 tokenization;LazySupervisedDataset 则按索引惰性预处理并做单样本缓存。训练脚本中的 --lazy_preprocess True 就是切换二者的开关(见 train.py#L235-L253),适合数据量大、不想预先把全部序列驻留显存的场景。

另外,train() 入口处有一个值得注意的细节:当 model_max_length 超过模型配置的 max_position_embeddings 时,会自动设置 rope_scaling = {"type": "linear", "factor": ceil(model_max_length / orig_ctx_len)}(见 train.py#L265-L275)。这意味着如果你想把 Vicuna-7B 的微调长度扩展到 16k,直接把 --model_max_length 设为 16384 即可,RoPE 缩放因子会按整倍数自动推导。

保存环节同样针对 FSDP 做了适配:trainer_save_model_safe 使用 FullStateDictConfig(offload_to_cpu=True, rank0_only=True) 把分片参数聚合到 CPU 后由 rank 0 统一落盘(见 train.py#L81-L89),避免多卡保存时的显存峰值。

使用 FSDP 微调 FastChat-T5(4 x A100 40GB)

docs/training.md 给出的 T5 全参数微调命令如下,官方标注为 4 x A100(40GB)配置:

torchrun --nproc_per_node=4 --master_port=9778 fastchat/train/train_flant5.py \
    --model_name_or_path google/flan-t5-xl \
    --data_path ./data/dummy_conversation.json \
    --bf16 True \
    --output_dir ./checkpoints_flant5_3b \
    --num_train_epochs 3 \
    --per_device_train_batch_size 1 \
    --per_device_eval_batch_size 1 \
    --gradient_accumulation_steps 4 \
    --evaluation_strategy "no" \
    --save_strategy "steps" \
    --save_steps 300 \
    --save_total_limit 1 \
    --learning_rate 2e-5 \
    --weight_decay 0. \
    --warmup_ratio 0.03 \
    --lr_scheduler_type "cosine" \
    --logging_steps 1 \
    --fsdp "full_shard auto_wrap" \
    --fsdp_transformer_layer_cls_to_wrap T5Block \
    --tf32 True \
    --model_max_length 2048 \
    --preprocessed_path ./preprocessed_data/processed.json \
    --gradient_checkpointing True

关键参数说明:

  • --fsdp "full_shard auto_wrap":启用 FSDP 全分片并自动包装,配合 --fsdp_transformer_layer_cls_to_wrap T5Block 指定以 T5Block 为分片粒度,使编码器/解码器的每层块成为独立 FSDP unit;
  • --preprocessed_path:T5 训练脚本特有参数,指定预处理缓存路径(见 train_flant5.py#L57-L60),与 --data_path 配合使用可跳过重复预处理;
  • --model_max_length 2048train_flant5.py 中该参数默认即为 2048(train_flant5.py#L67-L72);
  • 其余为 HF Trainer 标准超参:3 个 epoch、每卡 batch size 1、梯度累积 4(等效全局 batch 16)、余弦学习率调度、3% 步数 warmup、按步保存并只保留最近 1 个 checkpoint。

T5 脚本还有两个与因果 LM 训练不同的实现细节:

  1. 必须使用非 fast 的 T5Tokenizer。源码注释明确指出,fast tokenizer 会在特殊 token 前错误地补空格(见 train_flant5.py#L405-L413)。
  2. 词表扩展smart_tokenizer_and_embedding_resize 会向 T5 词表追加 [PAD]<{\n 等 T5 特殊字符 token,并把新增 embedding 初始化为旧 embedding 的均值(见 train_flant5.py#L84-L112)。数据侧则以 ### user:\n / ### assistant:\n 信号切分多轮对话,只把回答部分(外加 EOS)构造为 labels,问题部分被 mask(见 train_flant5.py#L142-L176)。

训练后必须做权重修复。原文档特别提醒:用 HF + FSDP 训练 Flan-T5 保存的 checkpoint 中,共享嵌入(shared embeddings)权重会损坏,训练完成后要调用仓库自带工具函数修复才能正常加载。该函数就是 fastchat/utils.py 中的 clean_flant5_ckpt:它读取 checkpoint 目录下的 pytorch_model.bin.index.json,取出 shared.weight,再将其回写到 decoder.embed_tokens.weightencoder.embed_tokens.weight 两个分片中。使用方式:

from fastchat.utils import clean_flant5_ckpt
clean_flant5_ckpt("./checkpoints_flant5_3b")

使用 (Q)LoRA 微调 Vicuna-7B(DeepSpeed ZeRO2)

原文档给出的 Vicuna-7B QLoRA 命令基于 ZeRO2,并明确说明两个前提约束:QLoRA 与 ZeRO3 当前不兼容(LoRA 支持 ZeRO3,参考配置见 playground/deepspeed_config_s3.json);以及依赖版本要求 bitsandbytes>=0.39.0、transformers>=4.30.0

deepspeed fastchat/train/train_lora.py \
    --model_name_or_path ~/model_weights/llama-7b \
    --lora_r 8 \
    --lora_alpha 16 \
    --lora_dropout 0.05 \
    --data_path ./data/dummy_conversation.json \
    --bf16 True \
    --output_dir ./checkpoints \
    --num_train_epochs 3 \
    --per_device_train_batch_size 1 \
    --per_device_eval_batch_size 1 \
    --gradient_accumulation_steps 1 \
    --evaluation_strategy "no" \
    --save_strategy "steps" \
    --save_steps 1200 \
    --save_total_limit 100 \
    --learning_rate 2e-5 \
    --weight_decay 0. \
    --warmup_ratio 0.03 \
    --lr_scheduler_type "cosine" \
    --logging_steps 1 \
    --tf32 True \
    --model_max_length 2048 \
    --q_lora True \
    --deepspeed playground/deepspeed_config_s2.json

LoRA 相关参数在 fastchat/train/train_lora.pyLoraArguments 中有完整定义(train_lora.py#L55-L65),除命令中显式给出的三项外,还有两个可调项:

参数 默认值 说明
--lora_r 8 低秩分解的秩 r
--lora_alpha 16 缩放系数 alpha,实际缩放为 alpha/r
--lora_dropout 0.05 LoRA 层 dropout
--lora_target_modules ["q_proj", "v_proj"] 注入 LoRA 的线性层,默认只改 attention 的 Q/V 投影
--lora_bias "none" 支持 "none" / "all" / "lora_only",控制是否保存偏置项
--q_lora False 置 True 启用 4bit 量化底模(QLoRA)

--deepspeed 指向的 playground/deepspeed_config_s2.json 是仓库自带的 ZeRO2 配置:zero_optimization.stage=2,优化器状态 offload 到 CPU(offload_optimizer.device=cpu),并开启 contiguous_gradientsoverlap_comm;batch size、梯度累积步数、fp16 均设为 "auto",由命令行参数接管。若改用 LoRA + ZeRO3,则切换到 playground/deepspeed_config_s3.json,该配置在 ZeRO3 基础上同时 offload 优化器与参数到 CPU,并开启 stage3_gather_16bit_weights_on_model_save,保证存盘时聚合出完整 16bit 权重。

从源码可以印证文档中的兼容性警告:train_lora.py 在检测到 q_lora=True 且启用了 FSDP 或 ZeRO3 时会打印 "FSDP and ZeRO3 are both currently incompatible with QLoRA"(train_lora.py#L121-L126)。QLoRA 路径下底模通过 BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", ...) 以 NF4 + 双重量化加载(train_lora.py#L128-L146),计算 dtype 跟随 --bf16/--fp16 选择。训练结束后 rank 0 只保存 LoRA 适配器权重(model.save_pretrained 保存 PEFT state dict),而非完整模型(train_lora.py#L211-L218),因此 checkpoint 目录很小,后续需要与底模合并或加载适配器推理。

仓库还提供了一个可直接套用的 LoRA 训练脚本 scripts/train_lora.sh:以 lmsys/vicuna-7b-v1.5 为底模,--q_lora False(纯 LoRA)、--fp16 True--num_train_epochs 150,并新增了 --gradient_checkpointing True--flash_attn False 两个选项——flash_attn 为 True 时会调用 fastchat/train/llama_flash_attn_monkey_patch.py 中的 replace_llama_attn_with_flash_attn 把 Llama 的 attention 替换为 FlashAttention 实现。

T5-XL / XXL 的 QLoRA 微调

对 encoder-decoder 架构的 Flan-T5,使用独立的 fastchat/train/train_lora_t5.py

deepspeed fastchat/train/train_lora_t5.py \
        --model_name_or_path google/flan-t5-xl    \
        --data_path ./data/dummy_conversation.json \
        --bf16 True \
        --output_dir ./checkpoints_flant5_3b \
        --num_train_epochs 3 \
        --per_device_train_batch_size 1 \
        --per_device_eval_batch_size 1  \
        --gradient_accumulation_steps 4  \
        --evaluation_strategy "no"  \
        --save_strategy "steps"  \
        --save_steps 300 \
        --save_total_limit 1 \
        --learning_rate 2e-5 \
        --weight_decay 0.     \
        --warmup_ratio 0.03    \
        --lr_scheduler_type "cosine"   \
        --logging_steps 1 \
        --model_max_length 2048    \
        --preprocessed_path ./preprocessed_data/processed.json \
        --gradient_checkpointing True \
        --q_lora True     \
        --deepspeed playground/deepspeed_config_s2.json

该脚本复用 train_flant5.py 的数据模块与词表扩展逻辑(smart_tokenizer_and_embedding_resizemake_supervised_data_module),但 LoRA 的 task_type 设为 SEQ_2_SEQ_LM,且由于 QLoRA 同样不兼容 ZeRO3,这里依然使用 deepspeed_config_s2.json

在本地 NPU 上全参数微调 Vicuna-7B

原文档最后给出了面向昇腾等本地 NPU 集群的 Vicuna-7B 全参数训练命令,以 8 x NPU 为例,通过 --nproc_per_node 指定 NPU 数量:

torchrun --nproc_per_node=8 --master_port=20001 fastchat/train/train.py \
    --model_name_or_path ~/vicuna-7b-v1.5-16k  \
    --data_path data/dummy_conversation.json \
    --fp16 True \
    --output_dir output_vicuna \
    --num_train_epochs 3 \
    --per_device_train_batch_size 8 \
    --per_device_eval_batch_size 1 \
    --gradient_accumulation_steps 1 \
    --evaluation_strategy "no" \
    --save_strategy "steps" \
    --save_steps 1200 \
    --save_total_limit 10 \
    --learning_rate 2e-5 \
    --weight_decay 0. \
    --warmup_ratio 0.03 \
    --lr_scheduler_type "cosine" \
    --logging_steps 1 \
    --fsdp "full_shard auto_wrap" \
    --fsdp_transformer_layer_cls_to_wrap 'LlamaDecoderLayer' \
    --model_max_length 2048 \
    --gradient_checkpointing True \
    --lazy_preprocess True

与 T5 版本对比,该命令有三个针对性调整:

  1. FSDP 包装单元改为 'LlamaDecoderLayer',与 Llama 的层结构对应(T5 版本是 T5Block);
  2. 底模选择 vicuna-7b-v1.5-16k,即 16k 上下文的 Vicuna;--model_max_length 2048 低于 16384,不会触发 RoPE 重缩放;若需要按 16k 长度训练,可按前述机制把该参数调大;
  3. 显式开启 --lazy_preprocess True,利用 LazySupervisedDataset 惰性 tokenization 控制 CPU 内存。

该命令走的是 fastchat/train/train.py 的完整流程:加载数据 → 构造掩码后的 input_ids/labels/attention_mask → HF Trainer 训练 → 检测 output_dir 下是否已有 checkpoint-*,有则自动 resume_from_checkpoint 续训,否则从零开始(train.py#L303-L314)。

训练参数速查

四个脚本共享的 HF Trainer 核心参数在原文档中的取值基本一致,可作为默认起点:

参数 取值 说明
--learning_rate 2e-5 三条路径统一
--lr_scheduler_type cosine 余弦退火
--warmup_ratio 0.03 3% 步数线性预热
--weight_decay 0. 文档命令未启用权重衰减
--num_train_epochs 3 配合 --save_steps 300/1200 与 --save_total_limit 控制磁盘占用
--model_max_length 2048 T5 脚本的默认值;Llama 脚本默认为 512,命令中显式覆盖
精度 --bf16 True--fp16 True NPU 场景用 fp16,GPU 场景用 bf16

小结

docs/training.md 给出的四条命令覆盖了 FastChat 微调的典型场景:T5 系列全参数微调(FSDP,训练后需用 clean_flant5_ckpt 修复共享嵌入)、Vicuna-7B 的 QLoRA(ZeRO2 + NF4 量化底模,注意与 ZeRO3 的不兼容)、T5 系列 QLoRA,以及 NPU 集群上的 Vicuna 全参数训练。实现层面的三个关键点值得记住:loss 只计算在 assistant 回复 token 上;LoRA checkpoint 只保存适配器权重;FSDP 保存走 CPU 聚合 + rank0 落盘。更深入的模型加载与适配细节可继续参考 fastchat/model/model_adapter.py,数据准备与清洗工具位于 fastchat/data 目录。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
528
588
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
906
1.83 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
docsdocs
暂无描述
Markdown
891
5.78 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.53 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.34 K
1.45 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
987
506
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384