多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

ChatGLM3-Base有监督微调实战:从LoRA配置到私有化部署

ChatGLM3-Base有监督微调实战:从LoRA配置到私有化部署 简介针对大模型微调需求的实战型资源面向算法工程师与人工智能学习者围绕ChatGLM3-Base模型提供可直接运行的有监督微调完整实现覆盖数据准备、模型加载、参数训练到结果推理的闭环帮助使用者解决微调落地过程中的数据准备、训练配置与评估改进等常见难题。压缩包共十一个文件大小仅为七百九十七KB包含六个脚本分别承担数据预处理、模型引擎封装、训练器构建、工具函数与结果发送等核心模块、一个推理演示笔记本、一份已标注样本数据、一篇教程文档以及两张训练曲线截图结构紧凑便于按需复用。已有二百七十二人学习下载。项目不局限于跑通代码而是按照有监督微调的典型流程拆解出样本构造、模型引擎封装、训练执行与推理验证等关键步骤并穿插超参数调优、学习率策略、正则化技术等进阶话题帮助学习者建立系统方法论。通过实际动手操作学习者能独立完成ChatGLM3-Base模型的监督微调任务并可将这套方法迁移至其他预训练语言模型沉淀可复用的工程代码与排错经验适合希望快速进入大模型微调实战的开发者。1. 先跑通一条 SFT 最小链路再谈调优一个反直觉的结论很多人微调 ChatGLM3 第一反应是拿 Chat 版继续训练但真正做行业私有化模型正确的起点是 Base 版。Chat 版里一切会聊天的能力都是别人用某套 SFT 数据配比训练出来的你拿它当底座再灌自己的数据等于在一个不可见的分布上叠 buff效果会卡在一个说不清的天花板上。反过来把 ChatGLM3-Base 当白纸用整理好的指令-回答对从头做有监督微调你能完全控制对话风格、回复格式和知识边界。这套链路的核心资产是数据和配置。模型加载参数错一个loss 直接不降对话模板错一个标记模型学到的全是格式噪声超参不合适12GB 卡和 24GB 卡完全是两种体验。这就是为什么实战项目习惯把源码和流程一起打包——缺经验时一份能跑的代码比读十篇文档有用得多。这篇沿着「环境基座→数据处理→训练配置→推理评估」展开最后落到几条能直接带进生产的排错经验上。适合有一定 transformers 基础、正准备在业务数据上跑第一个 SFT 模型的人。2. 环境与基座怎么把 ChatGLM3-Base 加载成可训练状态2.1 最小环境清单与加载代码ChatGLM3 系列依赖自定义模型代码transformers 在加载时必须显式允许远程代码执行否则会直接报错。依赖版本上transformers4.37.0、torch2.1、peft0.9、accelerate0.28是一套经过验证的组合。先把环境装好再跑下面这段最小加载。import torch from transformers import AutoModel, AutoTokenizer model_id THUDM/chatglm3-6b-base tokenizer AutoTokenizer.from_pretrained( model_id, trust_remote_codeTrue, padding_sideright, ) model AutoModel.from_pretrained( model_id, trust_remote_codeTrue, torch_dtypetorch.bfloat16, # 30 系及以上显卡用 bf16V100 换 fp16 device_mapauto, # 让 accelerate 自动分配显存 ) model.gradient_checkpointing_enable() model.train() print(model.dtype, model.device)trust_remote_codeTrue是 ChatGLM 系列绕不开的参数它会执行仓库里的自定义 Python 代码来构建模型结构。padding_sideright设置为右侧填充SFT 阶段如果混用左右填充attention mask 的位置偏置会干扰训练。torch_dtype建议优先 bf16它的数值范围比 fp16 大训练更稳V100 这类老卡不支持 bf16退到 fp16。device_mapauto会按显存把层分到多张卡单卡场景下它等价于直接把模型放到cuda:0。2.2 显存策略全参、LoRA 与 4bit 量化三挡6B 模型 fp16 权重约 12GB看起来 24GB 卡能装下但训练还要算梯度和优化器状态直接全参微调需要大约 4 倍权重内存。方案显存需求适用场景全参微调40GB 以上有多卡 A100/GH200 的团队LoRA约 16-18GB单卡 24GB 的标准选择最推荐QLoRA4bit 量化约 9-12GB单卡 16GB 或以下体验型项目LoRA 只训练注入的低秩矩阵冻结原权重QLoRA 在此基础上把基座量化到 4bit进一步压显存。对大多数业务场景LoRA 的效果已经足够量化主要解决「卡不够」的硬约束。QLoRA 的加载代码需要在from_pretrained里增加量化配置from transformers import BitsAndBytesConfig bnb_config BitsAndBytesConfig( load_in_4bitTrue, bnb_4bit_quant_typenf4, bnb_4bit_compute_dtypetorch.bfloat16, bnb_4bit_use_double_quantTrue, ) model AutoModel.from_pretrained( model_id, trust_remote_codeTrue, quantization_configbnb_config, device_mapauto, )bnb_4bit_quant_type用 nf4它对正态分布权重有更好的量化精度bnb_4bit_compute_dtype决定反量化后的计算精度保持 bf16 即可use_double_quant再省一点显存但会稍微增加耗时。量化后训练前要调用prepare_model_for_kbit_training(model)它会把需要训练的层转为 fp32否则 LoRA 层可能出现精度问题。2.3 加载后先做一次冒烟验证训练跑十几个小时才发现模型加载有问题成本太高。加载后立刻做一次 forward 验证确认 tokenizer 和模型能正确协作。input_text tokenizer.apply_chat_template( [{role: user, content: 你好}], tokenizeTrue, return_tensorspt, add_generation_promptFalse, ) with torch.no_grad(): out model(input_idsinput_text.to(model.device)) print(logits shape:, out.logits.shape)apply_chat_template是这条链路里最重要的一个 API它把对话结构按照模型预设的模板转成 token 序列。ChatGLM3-Base 本身没有对话能力但它的 tokenizer 依然带有模板能力SFT 训练的数据必须经过它处理。add_generation_promptFalse表示只编码对话本身不追加模型回复的引导标记这个参数到了推理阶段会反过来用。如果 logits shape 第一维是序列长度、最后一维是词表大小就说明模型链路通了。3. 数据工程不解决好对话格式与 token 截断SFT 就是白练3.1 指令数据集的三种常见形态SFT 训练数据本质上是一条条「指令/上下文 标准回答」的配对。项目里常见的格式有两种Alpaca 式的单轮instruction/input/output和 ShareGPT 式的多轮conversation数组。实际业务数据以多轮居多下面以 JSON Lines 为例{system: 你是某保险公司的客服助手回答须包含条款编号。, conversation: [ {role: user, content: 重疾险等待期是多久}, {role: assistant, content: 根据条款 2.3 条等待期为 90 天。} ]}system字段放角色设定和输出约束ChatGLM3 模板支持把它放在对话最前面。conversation里的 role 只用user和assistant两种多轮就继续往后追加。这部分工作的重点不是写解析代码而是确保每条对话的语义闭环上下文完整、回答不依赖外部记忆。数据里的system如果每一条都不同模型会把它当成对话内容的一部分去拟合所以相同业务域建议统一系统指令。3.2 tokenize、标签掩码与截断策略数据解析成统一结构后需要转换成模型接受的input_ids和labels。最直接的做法是把整条对话通过apply_chat_template转成 token然后让模型对整段序列做语言建模。这里有一个关键选择labels是否要掩码掉用户输入部分。import json from torch.utils.data import Dataset import torch class SFTDataset(Dataset): def __init__(self, path, tokenizer, max_len2048): self.tokenizer tokenizer self.examples [] with open(path, encodingutf-8) as f: for line in f: row json.loads(line) sys_prompt row.get(system, 你是一个可靠的助手。) conversation [{role: system, content: sys_prompt}] conversation.extend(row[conversation]) encoded tokenizer.apply_chat_template( conversation, tokenizeTrue, return_tensorspt, truncationTrue, max_lengthmax_len, add_generation_promptFalse, ) input_ids encoded[0] # 简化方案全部 token 参与 loss 计算 labels input_ids.clone() self.examples.append({input_ids: input_ids, labels: labels}) def __len__(self): return len(self.examples) def __getitem__(self, i): return self.examples[i]这段代码的labels等于input_ids是最省事的做法模型对整个对话序列做预测。更严格的做法是把labels中用户输入、系统提示和模板标记替换为-100只让 assistant 轮次的回答参与 loss 计算这样模型不会浪费建模能力去「背」用户的问法。ChatGLM3 官方微调脚本里有现成的掩码实现原理是扫描模板中 assistant 轮次的起止 token 位置把区间之外的位置全部置-100。如果数据里用户输入长度占比高掩码方案收敛更稳如果只做格式对齐全体计算 loss 也能用。截断策略是另一个坑。max_length截断是硬切多轮对话一旦超过长度后面最新的问答会被直接丢掉。实际项目中更稳妥的做法是保留 system 和最后两轮对话中间部分舍得删。SFT 数据里「最新指令」的权重远高于历史上下文开头两轮没被模型记住不致命但最后一轮被截掉这条样本就废了。3.3 数据量级与配比参考不是所有任务都需要几万条数据。SFT 的下限比多数人想象的低关键是任务类型和数据质量。任务类型建议规模数据侧重点回复风格/格式对齐1k-5k 条覆盖全部要求格式正例即可垂直领域问答10k-50k 条知识覆盖面、难例比例工具调用/结构化输出2k-10k 条严格校验输出 schema 的样本通用能力保持业务数据的 10%-20%混合通用指令防止灾难性遗忘loss不降时先别调超参回头检查数据是不是 system 指令每一条都不同、是不是 assistant 回答里夹杂了大量噪音、是不是截断把答案切没了。数据里若有 5% 的坏样本模型会用 20% 的容量去拟合这 5% 的噪声。4. LoRA 与超参ChatGLM3-Base 有监督微调的核心配置4.1 用 peft 配置 LoRA 目标模块数据准备好后进入训练配置。LoRA 的做法是冻结原模型在目标线性层旁路插入低秩矩阵。ChatGLM3-6B 的自注意力和前馈网络里都有可注入的线性层target_modules直接决定哪些层被训练。from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training model prepare_model_for_kbit_training(model) # 4bit 量化场景必加 lora_config LoraConfig( r16, lora_alpha32, target_modules[query_key_value, dense], lora_dropout0.05, biasnone, task_typeCAUSAL_LM, ) model get_peft_model(model, lora_config) model.print_trainable_parameters()target_modules里的模块名不是拍脑袋写的。ChatGLM3-6B 的 SelfAttention 把 Q、K、V 合并在一个query_key_value矩阵里输出投影叫dense。如果你用的包版本有命名差异先打印模块名核对print([n for n, _ in model.named_modules() if query in n or dense in n])r是低秩矩阵的秩业务任务 8 够用跨领域泛化需求高可以加到 32lora_alpha是放缩系数通常设成r的 2 倍过大等于把 LoRA 权重放大基座能力容易漂移lora_dropout防过拟合数据量少于 5k 时可加到 0.1。biasnone表示不训练任何偏置项这是 LoRA 的默认推荐全参微调才需要额外考虑。4.2 用 Trainer 拉起一个有监督微调训练数据已经是 token 化后的Dataset用 Hugging Face Trainer 最顺。不用 SFTTrainer 的原因是后者期待原始文本字段拿 token 化好的数据进去还要再过一遍打包逻辑容易重复截断。训练前必须实现一个 collator把 batch 内样本垫到相同长度。def sft_collator(features): input_ids [f[input_ids] for f in features] labels [f[labels] for f in features] max_len max(len(x) for x in input_ids) batch_ids, batch_labels [], [] for ids, lbs in zip(input_ids, labels): pad_len max_len - len(ids) batch_ids.append(ids [tokenizer.pad_token_id] * pad_len) batch_labels.append(lbs [-100] * pad_len) return { input_ids: torch.tensor(batch_ids), attention_mask: (torch.tensor(batch_ids) ! tokenizer.pad_token_id).long(), labels: torch.tensor(batch_labels), }attention_mask让模型忽略 padding 位置labels里的-100是 PyTorch 交叉熵的忽略索引自动排除掉这些位置。padding 一律放右侧配合加载时设置的padding_sideright。训练参数按 LoRA 的常见配置来from transformers import Trainer, TrainingArguments training_args TrainingArguments( output_dirsft_ckpt, per_device_train_batch_size2, gradient_accumulation_steps8, # 等效 batch size 2*8 16 learning_rate2e-4, num_train_epochs3, lr_scheduler_typecosine, warmup_ratio0.03, logging_steps10, save_strategyepoch, save_total_limit2, bf16True, gradient_checkpointingTrue, max_grad_norm1.0, report_tonone, ) trainer Trainer( modelmodel, argstraining_args, train_datasettrain_dataset, data_collatorsft_collator, ) trainer.train()per_device_train_batch_size2配gradient_accumulation_steps8是 24GB 单卡下比较稳的组合。等效 batch size 16-32 是 SFT 的常见区间太小收敛不稳太大容易过拟合。学习率 2e-4 是 LoRA 的常见起点这个值比全参微调的 1e-5 高一个数量级因为只更新低秩矩阵。gradient_checkpointingTrue用计算换显存如果卡足够大可以关掉提速。save_total_limit2防止 checkpoints 把磁盘写满。如果不习惯手写训练循环LLaMA-Factory 这类高效微调平台把 ChatGLM3 的 LoRA 配置模板化界面里选 Base 模型、填数据路径就能训练但超参逻辑和 Trainer 是一样的理解上面的参数含义后再去用它排查问题会顺手得多。4.3 断点续训与损失异常判定训练中断是常态不要每次从零开始。Trainer 天然支持断点续训python train.py --resume_from_checkpoint sft_ckpt/checkpoint-1000对应的代码里需要判断 checkpoint 目录是否存在或者直接作为命令行参数传入trainer.train(resume_from_checkpoint...)。续训时会自动恢复优化器状态和学习率调度器位置只有数据顺序会重置。现象原因处理loss 恒定为常数模板输出为空、labels 全 -100打印一条 tokenize 结果人工核对loss 前几步直接下降但立刻停滞数据里有大量重复样本检查去重、调低 lrloss 正常下降但 eval loss 上升过拟合减少 epoch、增大 dropout、加通用数据5. 推理与评估验证 SFT 生效的三个关键信号5.1 加载 adapter 并做生成推理训练产物不是完整模型而是 LoRA adapter。推理时先加载 Base 模型再挂 adapter。注意生成阶段的模板处理和训练不同add_generation_promptTrue让模板在最后补上模型回复的引导标记。from peft import PeftModel base_model AutoModel.from_pretrained( model_id, trust_remote_codeTrue, torch_dtypetorch.float16, device_mapauto, ) model PeftModel.from_pretrained(base_model, ./sft_ckpt/final) model.eval() conversation [ {role: system, content: 你是保险客服助手回答必须引用条款编号。}, {role: user, content: 重疾险等待期多久}, ] prompt tokenizer.apply_chat_template( conversation, tokenizeFalse, add_generation_promptTrue ) inputs tokenizer(prompt, return_tensorspt).to(model.device) out model.generate( **inputs, max_new_tokens512, do_sampleTrue, temperature0.7, top_p0.9, repetition_penalty1.05, ) answer tokenizer.decode(out[0][inputs[input_ids].shape[1]:], skip_special_tokensTrue) print(answer)max_new_tokens限制新增生成长度而不是总长度temperature0.7保留多样性又不至于发散repetition_penalty1.05对长回复场景很有用。decode 时通过inputs[input_ids].shape[1]切掉输入部分只保留模型新生成的 token。5.2 用三个信号判断 SFT 是否真的生效第一个信号是格式约束力。SFT 最直接的效果是「按格式说话」评估时看模型是否稳定输出训练集定义的格式比如是否带条款编号、是否遵守角色设定。第二个信号是知识边界。Base 模型对垂直领域的问题只会泛泛而谈SFT 后应该能引用训练数据里的实体和规则。第三个信号是反事实能力。把训练样本里的实体换掉比如把「重疾险」换成「医疗险」模型应该基于规则重新作答而不是复读训练数据原文。如果换实体后仍然一字不差输出训练答案说明过拟合严重需要加正则或减 epoch。SFT 和 RL 在这里有一个明确分工SFT 负责教会模型格式和能力RLHF 负责对齐人类偏好。项目标题写明 SFT评估时就不应该期待模型学会「拒绝回答」或「承认不知道」那是偏好对齐的范畴。5.3 合并权重与交付部署上线服务时把 adapter 合并回 Base 模型得到一个完整的独立权重省去每次加载时的 PeftModel 包装。merged model.merge_and_unload() merged.save_pretrained(./chatglm3-sft-merged) tokenizer.save_pretrained(./chatglm3-sft-merged)合并后的权重可以用 vLLM 这类推理框架直接加载部署也可以转成 GGUF 格式跑在更低配置的 CPU 环境上。私有化部署的常见坑是合并前用merge_and_unload()而不是model model.merge()后者会把 LoRA 权重残留在模型里后续推理状态不对。6. 从训练到私有化部署值得记住的 5 条实战排错经验6.1 高频失败模式的处置现象根因处置训练 loss 第 1 步就是 nanbf16 在旧卡上不可用换 fp16 重跑loss 不降模板或数据问题抽 50 条数据冒烟训练 1 步打印 loss 和样本生成结果重复repetition_penalty 太小或温度过低调高到 1.1 左右复读训练原文过拟合或数据泄漏减 epoch、降 lora_alpha、检查验证集是否混入训练数据部署后效果与验证时不一致加载流程少了 add_generation_prompt检查推理模板是否和训练一致每一类都可以在 30 分钟内定位。loss 异常先跑冒烟训练把max_steps50配logging_steps1观察第一步 loss 是否接近log(vocab_size)附近偏差过大说明模板把序列拼错了。6.2 写一个可复用的验证剧本评估 SFT 效果不建议人肉敲 prompt 一个个试。维护一份独立的验证集每轮训练后批量生成结果再人工抽审#!/bin/bash DECODED_CKPTsft_ckpt/final python eval_generate.py \ --base_model THUDM/chatglm3-6b-base \ --adapter_path $DECODED_CKPT \ --eval_file data/valid.jsonl \ --max_new_tokens 512 \ --output_file eval_results/$(date %Y%m%d_%H%M).jsonleval_generate.py里循环读取valid.jsonl逐条走 5.1 的生成流程把输入、标准答案、模型输出写进结果文件。评估指标不需要上复杂框架先统计三点格式违规率、回答平均长度、以及「输出与标准答案完全相同」的比例——最后一点恰好对应过拟合。演示场景建议挑一个格式约束明显的任务比如 JSON 输出或固定话术模板视觉差异比闲聊任务更直观学生或业务方一眼能看出模型变化。跑完验证剧本把效果稳定的 checkpoint 合并、导出、记录训练参数这一轮 SFT 才算真正收尾。下一轮迭代时用同样的剧本对比新旧模型输出比任何 loss 曲线都可信。本文还有配套的精品资源点击获取
返回列表