首页/新闻资讯/正文详情

LoRA + DeepSpeed 微调 ChatGLM:多GPU训练最佳实践与源码解析

发布时间:2026/9/24 18:04:23 来源:云帆数科 栏目:资讯中心
LoRA + DeepSpeed 微调 ChatGLM:多GPU训练最佳实践与源码解析
简介这是面向深度学习开发者的大模型微调实战项目聚焦LORA与Deepspeed在多GPU环境下对ChatGLM进行高效微调旨在解决大规模参数训练中的显存占用、通信开销与训练效率难题。项目源码完整覆盖数据处理、模型加载、分布式训练、损失函数、优化器与学习率调度、评估验证等关键环节。压缩包共376个文件约170MB以Python脚本为主辅以配置文件、实验数据、模型权重及说明文档既便于直接运行调试也适合拆解学习。通过该项目可深入理解低秩近似压缩原理、ZeRO零冗余优化器、混合精度训练及多卡并行策略为构建大规模对话系统打下扎实基础。已有795人学习下载适合具备PyTorch基础并希望进阶分布式大模型训练的开发者。1. 大模型微调实战LORA DeepSpeed 跑通 ChatGLM这份源码把多 GPU 训练的最短路径走完了聊到 ChatGLM 这类大模型的微调很多人的第一反应是「显存不够、通信太慢、根本跑不动」。而 LORA 加 DeepSpeed 的组合恰恰是把这三座山一起搬走的方案。这份源码项目不是教材里那种跑通即删的 demo而是一套完整的训练闭环——数据预处理、LORA 参数注入、DeepSpeed 多卡配置、训练脚本、评估逻辑都齐了。你拿到手之后改改数据路径和几个超参就能在自己机器上复现一次从加载预训练权重到产出微调模型的全过程。适合正在做对话系统、想把手头 ChatGLM 调成特定风格或特定知识域的工程师也适合被群里的「大模型微调」话题绕得云里雾里的新手——照着跑一遍比看十篇原理文章都管用。2. LORA 低秩微调与 DeepSpeed 选型为什么这两件套是 ChatGLM 微调的默认答案2.1 LORA 到底做了什么四个关键参数直接决定微调成本和效果全量微调一个 ChatGLM-6B光把模型参数加载到显存就需要大约 12GB再加上优化器状态、梯度等等一块 24GB 的卡勉强能塞下但训练时的 batch size 会小得可怜。而 LoRALow-Rank Adaptation的思路完全绕开了这个问题它冻结预训练模型的全部权重只在模型的某些线性层旁边加一组低秩分解矩阵训练时只更新这两张小矩阵。常见实现里这组矩阵的维度由r控制——通常取 8 或 16。r越大可学习的参数量越多模型能学到的任务特征越丰富但显存占用和过拟合风险也跟着涨。代码里一般会看到这样的配置from peft import LoraConfig, get_peft_model lora_config LoraConfig( r8, # 低秩矩阵的秩控制可学习参数量 lora_alpha32, # 缩放系数影响每次参数更新的幅度 lora_dropout0.1, # 防过拟合 target_modules[query_key_value], # ChatGLM 的自注意力层中要注入 LoRA 的模块 biasnone ) model get_peft_model(model, lora_config) model.print_trainable_parameters()这段代码的关键在target_modules。ChatGLM 的 Self-Attention 层里线性变换模块名是query_key_value这套源码里已经把该注的模块名配好了如果你换别的底座模型这一步是第一个要检查的地方——拿model.named_modules()打印一遍看你的目标层叫什么名字否则 LoRA 静默地不生效训练半天 loss 纹丝不动。lora_alpha的取值一般设成r的两倍到四倍它做的事是把 LoRA 的输出再乘一个缩放系数。不是越大越好——太大了微调后的输出分布容易偏离底座模型的语言习惯生成出来的句子开始「发飘」。2.2 DeepSpeed 的三个 Stage这份源码为什么用的是 ZeRO-2 加 offloadDeepSpeed 的核心是 ZeROZero Redundancy Optimizer。通俗地说传统数据并行里每张卡都存一份完整模型副本、梯度副本和优化器状态副本内存被大量冗余占用ZeRO 把这些状态按 GPU 数量分片每张卡只存自己负责的那一部分。Stage 1只分片优化器状态适合显存缺口不大时过渡。Stage 2分片优化器状态加梯度这是多卡微调最常见的配置也是这份源码默认的 stage。Stage 3把模型参数也分片了适合单卡装不下完整模型的场景但通信开销显著上升。这份源码里 DeepSpeed 配置大概长这样{ zero_optimization: { stage: 2, offload_optimizer: { device: cpu }, contiguous_gradients: true, overlap_comm: true }, fp16: { enabled: true, loss_scale: 0, initial_scale_power: 16 }, gradient_accumulation_steps: 4 }offload_optimizer把优化器状态放到 CPU 内存里这一步能省出大量显存代价是训练速度变慢——属于典型的拿时间换空间。如果你显存足够比如单卡 24GB 跑 ChatGLM-6B 微调把offload_optimizer去掉速度会快不少。overlap_comm是让梯度通信和反向传播计算重叠执行,这个开关对训练吞吐影响非常大,建议保持开启。2.3 为什么这套组合是当前微调的最佳性价比ChatGLM 类模型微调的本质是「用最小可训练参数量撬动最大行为改变」。全量微调需要更新全部 60 亿参数而 LoRA 把可训练参数压到千万级别——这份源码跑起来之后model.print_trainable_parameters()会告诉你实际只训练了 0.1% 到 0.5% 的参数。DeepSpeed 的作用则是让这批参数在多张卡上被高效地并行训练两者互补性极强。很多人问用 HuggingFace 的 PEFT 库单独做 LoRA 行不行当然行但对 ChatGLM 这种规模的模型单卡训练效率太低DeepSpeed 的引入让多卡并行、梯度累积、混合精度这些手段一次性到位这也是它成为「默认答案」的原因。3. 环境配置与多 GPU 训练脚本落地从零把训练跑起来3.1 环境准备CUDA、PyTorch 和配套库的版本匹配这份资源的源码在环境要求上不算苛刻,但版本错位会让你第一晚全耗在报错上。我建议的顺序是先把驱动和 CUDA 确认好再建虚拟环境装 PyTorch最后装配套库卡住就查版本号。常规做法是用 conda 先把环境隔离出来conda create -n chatglm_lora python3.10 conda activate chatglm_lora pip install torch2.1.0 torchvision0.16.0 --index-url https://download.pytorch.org/whl/cu121 pip install transformers4.36.2 peft0.7.1 deepspeed0.14.0 pip install datasets accelerate tensorboardPyTorch 版本和 CUDA 版本要严格对照。ChatGLM 的源码在 transformers 4.36 以上版本跑得比较顺太低会缺rotate_half的实现太高又可能踩 ChatGLM 建模文件里 API 变更的坑。如果用的是 Ampere 架构以上的卡A100、RTX 3090/4090PyTorch 2.1 版本的推荐安装方式就是上面这个命令把cu121替换成你的 CUDA 版本号。装完之后跑一行快速验证python -c import torch; print(torch.cuda.is_available(), torch.cuda.device_count())输出True和你预期的卡数说明 GPU 侧通了。接着验证 DeepSpeed 是否正常识别分布式环境可以在多卡机器上跑deepspeed --num_gpus2 --master_port29500 your_train_script.py3.2 训练入口脚本解析main.py 里每一段在干什么源码中的训练脚本结构清晰核心流程是加载 tokenizer 和模型、套 LoRA、准备数据、配优化器、拉起 DeepSpeed 引擎、开始训练。我拆两个关键片段讲。模型加载部分model AutoModel.from_pretrained( config[model_name_or_path], trust_remote_codeTrue, torch_dtypetorch.float16 ) model get_peft_model(model, lora_config)trust_remote_codeTrue是 ChatGLM 这类带领 Python 建模文件的模型必须的它的建模代码直接写在模型的 GitHub 仓库里而不是 transformers 内置不开这个开关模型根本加载不出来。优化器和训练超参配置部分from transformers import get_linear_schedule_with_warmup optimizer torch.optim.AdamW(model.parameters(), lr2e-5, weight_decay0.01) total_steps len(train_dataloader) * config[num_epochs] // config[gradient_accumulation_steps] scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_stepsint(total_steps * 0.03), num_training_stepstotal_steps )注意这里的学习率是 2e-5。LoRA 微调的学习率如果直接用预训练阶段的 5e-5 甚至 1e-4loss 大概率会炸或震荡——因为新的低秩矩阵是从随机初始化开始的学习率太大前面的几步更新就会把输入分布推到预训练参数没见过的地方。通常 1e-4 到 3e-4 是一个经验范围ChatGLM 这种中文生成任务我一般起步用 2e-5 到 5e-5。3.3 关键参数速查与调法参数默认值调参方向r8数据量大、任务难就升到 16数据小就保持 8lora_alpha32与r联动通常是r的 2~4 倍lora_dropout0.1过拟合时升到 0.2正常不用动learning_rate2e-5微调一般不超过 1e-4gradient_accumulation_steps4等效增大 batch size调大不影响单卡显存但增加训练时间per_device_train_batch_size1~2主要受显存限制结合 gradient accumulation 控制总 batch size3.4 评估与验证训练完不测试等于白搞这段看很多人的实现会被省掉但不验证调参就是盲人摸象。源码里给的评估逻辑是留出一部分验证集算 loss 和生成效果。我的习惯是每个 epoch 结束后跑一次验证集 loss同时随机抽几条对话让模型真实生成一版传到 TensorBoard 和模型权重一起保存——只盯 loss 下降不够loss 降到 0.8 但生成的是乱码那说明模型过拟合了。这里贴一个标准的验证循环model.eval() eval_loss 0.0 with torch.no_grad(): for batch in eval_dataloader: outputs model(**batch) eval_loss outputs.loss.item() print(fEval loss: {eval_loss / len(eval_dataloader):.4f})注意验证阶段必须切到model.eval()并关梯度否则 dropout 层的随机行为会让你的验证 loss 出现无意义的波动容易误判模型在收敛还是在发散。4. 数据准备与预处理对话格式怎么喂给 ChatGLM 才能学得会4.1 ChatGLM 微调的数据长什么样ChatGLM 微调的数据不是随便一段文本就能用的它期望的是对话结构的输入——也就是多轮的人机交互历史。通常的格式是 JSON 或 JSONL{ conversations: [ { role: user, content: 请介绍一下杭州的特色美食。 }, { role: assistant, content: 杭州的经典美食有西湖醋鱼、东坡肉、龙井虾仁等其中最出名的是西湖醋鱼…… }, { role: user, content: 有没有适合带回去当伴手礼的 }, { role: assistant, content: 推荐买龙井茶、酥饼和藕粉这些都是方便携带的杭州特产…… } ] }这份源码的数据准备脚本会把上面这种结构拼成 ChatGLM 训练用的 input-target 格式input是完整的上下文包括 role 标记target只计算 assistant 回复部分的 loss。原理上微调就是要让模型学会「看到上面这段对话后下一个 token 应该落在 assistant 回答上」。4.2 数据预处理的两个关键环节第一是 tokenizer 的 max length 设置。ChatGLM 的位置编码是 2048超过这个长度位置编码信息就截断了。对话太长的样本会被截断截断位置如果正好把 assistant 的回答切掉这条样本就废了——没有正确的 target。处理逻辑一般用右侧截断inputs tokenizer( text, max_lengthconfig[max_seq_length], truncationTrue, paddingmax_length, return_tensorspt )第二个环节是 mask 掉 user 部分的 loss。很多新手微调 LLM 时连用户输入部分的 loss 也算进去了这会导致模型学会「复述用户的话」而不是「接住对话继续往下说」。源码里的做法是把 user 部分和 assistant 部分的 label 区别开user 部分的 label 置为 -100在计算交叉熵时自动跳过labels input_ids.clone() labels[user_token_mask] -1004.3 数据量需要多少、数据质量怎么看微调 LORA 的数据量没有绝对答案。一个任务型对话场景几千条高质量样本往往就够用了但如果你想改的是模型的「整体风格」那数据量要上万而且多样性比数量更重要。判断数据质量有个土办法把每条样本的 user 内容按长度和关键词聚类如果大量样本说的是同一个话题、同一个句式模型就会被带偏。数据预处理跑完后源码里通常会有一段把 tokenizer 解码回文本的验证代码我强烈建议你跑一下看到拼出来的训练文本是「人话」再进训练不然模型学到的是拼接 token 的杂音。5. 训练过程避坑与常见问题排查显存、loss、通信三大战场5.1 显存不够明明加了 DeepSpeed 还是 OOM现象训练脚本跑起来几秒钟就报CUDA out of memory有时是RuntimeError: CUDA error: out of memory有时是 DeepSpeed 报AssertionError。原因最常见的不是模型本身装不下而是 batch size 设大了或者没有把输入序列的长度压到合理范围。若序列长度 2048、per-device batch size 设到 4即使 LoRA 参数很少中途激活值一层层积压也会把显存挤爆。还有一种可能offload_optimizer没开优化器状态全部放在 GPU 上。解决先把per_device_train_batch_size降到 1gradient_accumulation_steps提到 8 到 16等效 batch size 不变但峰值显存大幅下降。同时检查max_seq_length如果对话平均长度只有几百 token把序列长度从 2048 砍到 1024显存占用几乎减半。最后确认 DeepSpeed 配置里offload_optimizer已打开。这样三层处理下来4 卡 12GB 的配置也能跑通 ChatGLM-6B。5.2 loss 不降或剧烈震荡现象训练了十几个 steploss 一直维持在 2.0 左右纹丝不动或者从 1.5 突然跳到 3.0 再回到 1.2整个训练曲线像心电图的病例。原因loss 不降的典型原因是学习率过低比如用了 1e-6 这类来自误复制的配置或者 LoRA 根本没注入成功——target_modules填的模块名不对模型在训练时没有任何可训练参数你看到的 loss 全部来自冻结参数的原始表现。loss 震荡则多半是学习率偏大、batch size 太小梯度噪声过大。解决训练前先跑model.print_trainable_parameters()确认可训练参数数量不是 0。然后看学习率LoRA 微调的经验值是 2e-5 到 5e-5低于 1e-5 模型基本学不动高于 1e-4 容易崩。还不行就检查混合精度的 gradient scaling 是否出了问题把 DeepSpeed 配置里fp16: enabled暂时设为false试一轮能排除一半的玄学问题。5.3 多卡训练时某一个 GPU 显存溢出现象nvidia-smi看到的显存占用严重不均衡有的卡吃了 80%有的卡才 30%训练速度也不如预期经常卡住。原因最常见的原因是 batch size 不能被 GPU 数量整除导致数据分配到最后一个 rank 时数量不同另一个原因是deepspeed --num_gpus的数量和配置文件里的实际卡数不一致某个 rank 承担了额外的工作量。解决先确认启动命令里--num_gpus和实际的可见卡数一致然后在训练脚本开头打印一下torch.cuda.device_count()和当前 rank 的local_rank确认 DeepSpeed 正确初始化了。再把per_device_train_batch_size设为 GPU 数量的整数倍保证每个 rank 拿到的数据一样多。最后用export CUDA_VISIBLE_DEVICES0,1,2,3指定卡避免和别的任务抢卡。5.4 训练中断后 resume 恢复现象训练到一半由于断电、OOM 或者其他原因中断了再次启动时所有进度清零得从头再来。如果是几千条数据加上十几个 epoch重新跑一夜又得耗掉。原因训练脚本没有保存 checkpoint 或者没有实现断点续训。很多 demo 脚本只保存最终模型不按 step 或 epoch 周期保存。解决源码里如果有transformers.Trainer则天然支持resume_from_checkpoint参数如果是手写的训练代码需要自己实现把 optimizer 状态、scheduler 状态、epoch、step 全部保存到磁盘的 checkpoint 文件。我的习惯是每 500 步保存一次完整的 checkpoint至少保留最近两个训练中断后的后悔药只有这一个没有其他捷径。记得验证 resume 之后 loss 曲线是连续的如果断裂说明 optimizer 状态没存完整。5.5 生成质量差loss 很低但对话输出是复读机现象验证集 loss 降到 0.8 以下但是实际对话时模型频繁重复同一句话或者答非所问看起来像「强行续写」而不是「回复」。原因这种现象跟数据问题强相关——训练数据中或存在大量短对话、重复句式或者数据里 user 和 assistant 的角色标错模型学到的是「不管别人说什么都接那几个固定的词」。另一个原因是 target 部分的 mask 没做好user 输入也被计入了 loss模型在生成时倾向于复述历史对话。解决回头检查训练数据看 assistant 回复是不是都存在内容分布有没有大量重复。检查 labels 里 user 部分是否已置为 -100。生成时用sampleTrue、temperature0.7、repetition_penalty1.2这几个参数组合跑一批对照量化对比「微调前」和「微调后」的输出差异这是判断 LoRA 是否真正生效的最终手段。6. 微调效果验证与模型合并把 LoRA 权重落回完整模型整套流程跑完后还有一个绕不开的动作LoRA 训练出来的权重默认是增量权重不是独立模型。部署或继续做下游任务时需要把 LoRA 权重合并进底座模型或者至少掌握怎么用 PEFT 单独加载 LoRA 权重做推理。from peft import PeftModel base_model AutoModel.from_pretrained( config[model_name_or_path], trust_remote_codeTrue, torch_dtypetorch.float16 ) model PeftModel.from_pretrained(base_model, ./output/chatglm_lora) merged_model model.merge_and_unload() merged_model.save_pretrained(./merged_model)merge_and_unload()把 LoRA 分支的低秩矩阵乘回原模型的权重矩阵得到一个独立的完整模型文件之后可以拿 vLLM 或原生 transformers 直接加载部署。验证阶段的重点用同一组 Prompt 对比底座模型和 LoRA 微调后的输出。这里给出一个我常用的测试脚本模板prompts [ 你是一个电商客服请回答这件衣服可以七天无理由退货吗, 帮我写一份杭州一日游的行程。 ] for prompt in prompts: inputs tokenizer(prompt, return_tensorspt).to(cuda) outputs model.generate( **inputs, max_new_tokens256, do_sampleTrue, temperature0.7, repetition_penalty1.2 ) print(tokenizer.decode(outputs[0], skip_special_tokensTrue))对比时注意三点第一看微调后的输出是否更贴近目标风格比如客服场景中是否更礼貌、更结构化第二确认模型没有「灾难性遗忘」——底座模型原本擅长的一些通识问答能力如果退化明显说明 LoRA 的r太大或学习率太高第三用同一批问题反复生成几次看看输出多样性是否正常不要出现复读机。最后说说我现在的习惯——拿到任何一套开源微调代码我先检查三个文件数据预处理脚本、训练启动参数、DeepSpeed 配置。先把这三处涉及的参数读懂再跑print_trainable_parameters和一步试跑看到 loss 正常下降才放它去长时间训练。省下的时间足够多跑五轮实验。微调大模型的坑深浅都藏在数据和参数里把这些摸透换任何底座模型你都能快速迁过去。这套源码的价值也正在于此——它把最短路径画出来了剩下的就是你自己的数据工程和调参功课希望帮到你。本文还有配套的精品资源点击获取

相关推荐

红帽 RHCSA/RHCE/RHCA 要不要报机构?2026 真实备考建议,看完少踩坑
红帽 RHCSA/RHCE/RHCA 要不要报机构?2026 真实备考建议,看完少踩坑

红帽 RHCSA/RHCE/RHCA 要不要报机构?2026 真实备考建议,看完少踩坑 很多想入行 Linux、运维、云原生的朋友,都会把红帽认证当成能力标杆。 但一打开论坛,说法两极分化:有人自学一次通关,有人报班花上万踩坑… · 2026/9/24 18:04:23

Java课设超市订单管理系统:从能跑到能讲清楚的完整落地路径
Java课设超市订单管理系统:从能跑到能讲清楚的完整落地路径

简介:这是一套面向高校计算机专业学生与Java初学者的超市订单管理系统课程设计源码,基于MySQL数据库与原生JDBC实现,采用ServletJSP的Web工程结构,适合作为大学实训、课程设计或毕业设计的参考方案。压缩包共231个文件&#xff0c… · 2026/9/24 18:04:16

从一场“大型商综火灾救援”写到毕业论文:消防指挥同学的 AI 搭子怎么选?
从一场“大型商综火灾救援”写到毕业论文:消防指挥同学的 AI 搭子怎么选?

消防指挥专业的论文,难就难在它不是坐在书斋里“谈理论”。 比如你要写《大型商业综合体火灾应急救援指挥流程优化研究》,通常得完成这样一份成果:基于典型火灾案例,梳理接警响应、力量编成、现场侦察、分区管控、内攻搜救、供水保… · 2026/9/24 18:04:16

MCP标准化代码执行:Agent工具调用稳定性的最佳实践
MCP标准化代码执行:Agent工具调用稳定性的最佳实践

坦白说,刚开始给我自己的Agent接入代码执行能力的时候,我踩的坑比收获多。典型的场景是:Agent规划得头头是道,到了真正要跑一段脚本处理日志、批量改文件名、算一组指标的时候,要么输出的函数调用格式不对,… · 2026/9/24 19:18:15

在 Ubuntu 22.04 上安装 Docker Desktop
在 Ubuntu 22.04 上安装 Docker Desktop

在 Ubuntu 22.04 上安装 Docker Desktop 原文链接:How To Install Docker Desktop on Ubuntu 22.04 近期在做一个agent项目,为了方便部署迁移还是使用docker进行吧,先在本机安装docker 一、Docker初识 什么是docker? Docker 是… · 2026/9/24 19:18:15

跨域Cookie写入失败?CORS与SameSite配置全攻略
跨域Cookie写入失败?CORS与SameSite配置全攻略

搞前后端分离最头疼的接口联调阶段,十次里有八次都栽在跨域上。尤其当你辛辛苦苦把登录接口调通,结果发现浏览器控制台报了个“has been blocked by cors policy”的错误,而这次不是因为没配CORS,是因为你配了CORS,但C… · 2026/9/24 19:18:15

CORS跨域设置Cookie实战:登录态保持与常见坑解析
CORS跨域设置Cookie实战:登录态保持与常见坑解析

前后端分离的项目做久了,CORS 是早晚要正面刚的问题。平时最烦的无非两种:一种是接口直接被浏览器拦了,报 No Access-Control-Allow-Origin header is present;另一种更隐蔽,跨域请求能通、数据能回来,就是… · 2026/9/24 19:18:15

2025年Windows装机必备软件清单:21个门类深度配置与避坑指南
2025年Windows装机必备软件清单:21个门类深度配置与避坑指南

1. 为什么还要聊2020年的Windows软件清单2020年到现在,Windows生态其实经历了不少变化,但有一批软件的生命力强得离谱——五年过去了我自己的机器上还在跑,而且装机量只增不减。这份清单最初是我给团队新同事做装机参考整理的,一共… · 2026/9/24 19:18:15

2026用户行为分析工具选型:20款实测与落地避坑指南
2026用户行为分析工具选型:20款实测与落地避坑指南

先说个跟标题相关的背景:我去年接了三个不同的数据分析项目,分别是一个电商独立站、一个SaaS产品和一家做内容社区的公司。三个项目的共同点,就是都需要上一套好用的用户行为分析工具。为了给这三个项目做选型,我把市面上叫得上名… · 2026/9/24 19:18:08

基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程
基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程

简介:这是一套面向计算机、人工智能、自动化等专业学生与教师的毕业设计级项目资源,围绕YOLOv8实现渔船作业监控系统,可用于毕设、课程设计、大作业或项目立项演示。压缩包共97个文件,约24.21MB,以70个Python源码文件为… · 2026/9/24 0:00:13

1D-CNN时间序列建模实战:从Conv1d原理到工业落地
1D-CNN时间序列建模实战:从Conv1d原理到工业落地

简介:面向时间序列数据建模的一维卷积神经网络完整实现,适合深度学习入门者及需要快速验证时序模型的研究者,能够从音频、文本、传感器或股价等序列中挖掘局部特征与时间依赖。压缩包体积很小,只有3KB,内含3个Python脚… · 2026/9/24 0:00:26

柔软的L:汉语语流中被忽视的舌肌张力控制
柔软的L:汉语语流中被忽视的舌肌张力控制

1. 这个“L”不是字母表里的L,而是舌尖上的L最近在几个方言群和语音教学社群里,反复看到有人发一句:“也说字母L:柔软的长舌”。初看以为是英语发音课笔记,点开才发现全是方言爱好者、播音系学生、语言康复师甚至戏曲演… · 2026/9/24 0:00:44

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码