ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从零开始训练LLM:数据、模型、训练与推理的完整工程实践

从零开始训练LLM:数据、模型、训练与推理的完整工程实践 1. 项目定位从零开始到底意味着什么1.1 是抄代码还是搞懂每一层很多人看到“ai-engineering-from-scratch”第一反应是我知道就是照着书把一个大语言模型写出来。我最初也是这么想的但真正动手之后才发现如果只是把代码敲出来跑通那叫“复现”不叫“从零开始”。真正意义上的 from scratch是你拿到一堆原始文本、一张显卡、一个空白的框架项目最后得到一个能对话、能写代码、能推理的模型。这个过程里没有现成的模型权重没有huggingface上直接加载的checkpoint所有东西都要自己造。我在这个项目里给自己定了几条硬性标准也正是这几条标准让整个项目没有沦为“调包侠练习”不用任何预训练权重哪怕是开源的也不行。Tokenizer、数据集构造、模型结构、训练循环、推理逻辑全部自己实现。每一个模块都要能回答出“为什么这样设计”而不是“别人都这么搞所以我也这么搞”。1.2 项目范围怎么切——先定小边界刚开始我犯了一个很典型的错误野心太大。想着直接冲一个上百B参数的MoE模型做完整套对齐流程最后发现自己连数据清洗都搞不定。后来我把项目缩成了三个阶段每个阶段都有一个可验收的交付物阶段一训练一个约30M参数的Transformer在单一数据集上跑通完整的预训练流程。阶段二扩展到100M250M参数引入RoPE、GQA、Flash Attention训练数据规模提升到10B token左右。阶段三在阶段二的模型上做SFT和基础推理优化形成一个小型但完整的“对话模型”。这个切法很关键。很多人失败不是因为智商不够而是因为第一步迈太大。你不可能在第一次就复现一个GPT-4但你可以先用小模型把全链路跑通再逐步放大。1.3 为什么非要自己搭一遍直接下载一个Llama的权重、写几行inference代码五分钟就能出结果。但那种做法解决不了任何工程问题因为你不知道一个模型从数据到部署之间有多少决策点。自己搭一遍的价值在于你会被迫面对这些平时根本不会注意到的问题词表大小选多少直接决定embedding层的参数量和训练吞吐。学习率峰值设多少跟你的batch size、数据量、warmup步数全都耦合在一起。序列长度和batch size怎么配比才能让你的显存利用率不打折。数据配比不对时loss曲线会呈现什么样子如何提前发现。这些问题在“下载模型、调API”的流程里永远碰不到但它们恰恰是AI工程的核心。我的体会是如果你能独立训练一个100M级别的模型你对大模型的理解深度会远超那些只会调接口的人。2. 数据工程比模型更值得花时间的部分2.1 数据从哪来、怎么筛做GPT系列模型的人常说一句话模型架构决定上限数据质量决定你实际能达到多少。这句话真不是客套。我第一阶段用的数据是公开的英文语料比如OpenWebText类的开源数据集但直接用原始版本根本不行里面各种乱七八糟的噪声能把训练结果拉垮几个百分点。我筛数据经历了下面几步先做语言过滤把非英文内容按比例去掉保留目标语言的文本占比在95%以上。再做质量过滤用启发式规则删掉全是重复字符的文本、过短的段落少于50字符、HTML标签残留、以及“乱码”特征明显的片段。最后做去重用MinHash LSH对文本做近似去重这一步能显著减少模型背诵训练集的倾向也直接提升下游评估得分。我踩过最深的坑是“元组重复”。就是整篇文章不重复但里面某些段落反复出现例如新闻网站的文章互相转载、版权页和导航栏。这种东西光靠全文去重是抓不到的必须做“段落级”去重即把每篇文章拆成段落然后对段落做MinHash签名比较。实操建议对于10B token以下的数据规模一套基础的规则过滤MinHash去重完全够用不要一上来就上基于模型的分类器成本太高且收益不明显。2.2 Tokenizer训练与词表大小选择Tokenizer是很多人容易忽略的模块但它对你的训练效率影响巨大。我训练了一个BPE词表用的工具是sentencepiece或者tokenizers库核心训练参数如下词表大小8K阶段一、32K阶段二。字符覆盖度0.9999保证几乎所有字符都能被编码。特殊token加上pad、bos、eos、unk训练时还要额外加对话模板用的分隔符。为什么词表大小对模型很重要因为词表大小直接决定embedding层的维度。假设隐藏层维度是1024词表8K对应的embedding矩阵参数量是8M如果词表扩到32Kembedding参数就变成32M。对小模型来说这个占比非常可观。你不能盲目追求大词表也不能选太小导致每个token的信息密度太低。我实际测试过一个30M参数的模型词表8K换成32K同样训练步数下loss反而更高——为什么因为embedding参数暴增之后在相同的总参数量预算下Transformer层的参数被挤占模型容量反而降低了。训练Tokenizer时的另一个关键点必须保证训练语料与预训练语料分布一致。如果你的分词器只在通用英文上训练后面又想在上面做代码生成你会发现代码里的空格缩进、特殊符号被拆得稀碎生成效果一塌糊涂。2.3 样本配比与Dataloader实现细节多数据集混合时比如通用文本、代码、数学混合data sampling的比例分配是一个核心工程问题。最简单的做法是按token数比例混合但实际效果通常不好。我参考了一些大模型训练的经验最终采用了“按字节预算配比”的方式先给每个数据集设定一个目标token占比比如通用文本60%、代码25%、数学15%然后按比例进行采样并且在每个epoch里主动做“数据集轮换”。具体实现上有个细节不要让模型按照固定顺序吃遍整个数据集这样会让它在不同域名之间突然切换导致loss剧烈波动。正确做法是提供多个数据流每个数据流内部保持数据分布稳定然后在每个step随机选择数据流。Dataloader还有一个容易犯错的点——padding策略。训练时如果你的单条样本长度参差不齐直接拼batch会浪费大量显存。我的做法是按长度分桶bucket不同桶内做padding。这样短样本不会被长样本拖累训练吞吐能提升30%~40%。3. 模型搭建现代LLM的关键组件逐个拆3.1 RMSNorm与Pre-Norm结构现代大语言模型几乎都不再用原始的LayerNorm而是用RMSNorm。两者的差别可以简单理解成LayerNorm还要计算均值并做中心化RMSNorm只做缩放不做平移省掉了一组参数的均值归约操作计算更轻量而且在深度学习任务上效果不输LayerNorm。RMSNorm的公式RMSNorm对输入x的每个特征维度算均方根RMS(x) sqrt(mean(x^2) eps)然后用x / RMS(x)做归一化再乘一个可学习的权重gamma。在我的实现里这个模块的代码只有十几行import torch import torch.nn as nn class RMSNorm(nn.Module): def __init__(self, dim, eps1e-6): super().__init__() self.eps eps self.weight nn.Parameter(torch.ones(dim)) def forward(self, x): rms torch.sqrt(x.pow(2).mean(-1, keepdimTrue) self.eps) return x / rms * self.weight另一个关键结构是Pre-Norm。简单说在残差连接之前先做Normalization再做子层计算。即output x sublayer(norm(x))。Pre-Norm的好处是训练更稳定因为每个残差分支的输入都被归一化过梯度回传时不会因为深层网络导致爆炸或消失。几乎所有现代LLM包括GPT、LLaMA系列都采用Pre-Norm RMSNorm的组合。3.2 RoPE旋转位置编码到底在编码什么传统Transformer用绝对位置编码Sinusoidal或可学习的Positional Embedding但这类编码是“加到token向量里”的不能直接刻画相对位置关系。RoPE的出发点很不一样它对query和key向量注入位置信息的方式是“旋转”。以一个二维向量为例RoPE会根据位置给它乘一个旋转矩阵。旋转矩阵 R(m) [[cos(mθ), -sin(mθ)], [sin(mθ), cos(mθ)]]其中m是位置下标θ是预设的频率。当你把query和key同时旋转之后两者做内积时结果只跟它们的相对距离有关。这个性质在长文本任务上非常重要。实现上RoPE是在attention计算之前对Q和K做变换def apply_rope(x, cos, sin): # x: (batch, seq, num_heads, head_dim) d x.shape[-1] x1 x[..., : d // 2] x2 x[..., d // 2 :] return torch.cat([x1 * cos - x2 * sin, x1 * sin x2 * cos], dim-1)你不需要自己手写旋转矩阵的每个角度PyTorch有现成的三角函数计算重点是理解“Q和K必须使用相同的位置编码逻辑”不然后面做推理时长度外推会直接出问题。关于长度外推模型没见过更长的序列但你希望它能处理更长的文本RoPE虽然本身具备一定的外推能力但如果训练时序列固定为2048直接推到8192效果会很差。工程上常用的办法是NTK-aware scaling或者对高频分量做插值这属于训练后的优化可以在推理阶段实施。3.3 GQA注意力省显存的折中方案传统Multi-Head AttentionMHA每个注意力头都有自己的K和V投影矩阵。Grouped Query AttentionGQA的思路是让多个查询头共享同一组键值头从而大幅减少KV缓存和参数。具体来说假设你有32个Query头、4个KV头每个KV头服务8个Query头。这样KV投影的参数直接变成原来的1/8推理时的KV cache也缩小到原来的1/8。为什么这个设计可行因为研究发现多个Query头关注的信息模式有很多重叠KV头不需要完全独立。训练阶段GQA还可以省显存但对训练的收益不如推理阶段明显。我的实际经验是显存瓶颈往往不在模型参数而在KV cache。把KV cache减到1/8意味着同样的显存能把推理batch size放大好几倍。实现GQA时有个绕不开的细节从MHA“升级”到GQA之后已经训练好的模型权重是没法直接用的需要将原来每个KV头的参数做平均或者复制到新的共享KV头上。所以如果你计划最终用GQA最好一开始就按GQA设计而不是训练完之后再转换。3.4 从Dense到MoE的选型依据很多人做到第二阶段就会忍不住想上Mixture of ExpertsMoE。MoE的核心思想是每一层不再是单个前馈网络而是多个专家网络由Router路由网络根据token决定激活哪些专家。MoE的价值在于总参数量可以非常大但每个token只激活一小部分参数推理成本不会线性增长。但MoE的工程复杂度也不是玩笑。负载不均衡是最典型的问题Router很容易“偷懒”把大多数token都路由到同一个专家上导致其他专家变成摆设。我当时的处理策略是加辅助负载均衡损失def load_balancing_loss(router_logits, num_tokens, num_experts): # router_probs: (num_tokens, num_experts) router_probs router_logits.softmax(dim-1) # 专家被路由的平均概率 expert_load router_probs.mean(dim0) # 辅助损失 专家数量 * 平均负载向量的平方和 aux_loss num_experts * (expert_load.pow(2).sum()) return aux_loss辅助损失越小说明路由越均匀。实际中我会把它乘以一个系数0.01级别加到主loss上。我的建议如果你在单个GPU上训练不要上MoE。MoE更适合多机多卡训练和推理场景单卡上频繁的通信开销会让你被内存带宽卡死。项目阶段二直接做Dense Transformer足够学到该学的东西。4. 训练工程把模型真正跑起来4.1 混合精度与梯度累积的参数计算训练阶段的第一个大坑就是显存。一个100M参数的模型全精度FP32占400MB看起来不大但加上优化器状态AdamW要保存两个动量、梯度、激活值整体轻松翻几倍。工程上第一件事就是上混合精度。最常见的做法是bf16训练。为什么用bf16而不是fp16因为bf16的指数范围和fp32一致梯度在反向传播时不容易下溢训练稳定性好得多。NVIDIA从Ampere架构开始支持bf16。混合精度的核心是权重用FP32保存一份主副本训练过程中的前向和反向在bf16下计算优化器更新在FP32下完成。这样做既享受低精度带来的显存和速度优势又避免精度损失导致模型不收敛。显存估算可以按一个经验公式来模型参数每个参数2字节bf16。梯度每个参数2字节。AdamW优化器状态每个参数8字节两个动量各4字节按FP32算。所以实际的显存下限大约是模型参数量 × 12 字节 激活值/中间变量一个250M参数的模型仅参数、梯度、优化器就要约3GB。加上激活值取决于序列长度和batch size单卡训练250M模型至少需要10GB以上的显存这个数字很容易超所以我会提前算好。先看你一共需要多少有效batch size比如256条样本而单卡上一个前向反向只能放下16条样本那就需要梯度累积。梯度累积就是把多个mini-batch的梯度累加之后再统一做一次优化器更新accumulation_steps 256 // 16 # 16步累积 optimizer.zero_grad() for micro_step in range(accumulation_steps): loss model(input_batch) loss loss / accumulation_steps # 除以步数保持梯度均值稳定 loss.backward() optimizer.step()注意loss一定要除以累积步数否则等价于你把学习率放大了accumulation_steps倍训练直接发散。4.2 学习率方案怎么定学习率是训练里最敏感的超参数之一。对于GPT类的自回归模型我用的是一套被广泛验证过的方案线性warmup比如500步内从0升到峰值目的是让模型在训练初期梯度方向不稳定的阶段慢慢起步防止早期震荡。到达峰值后按余弦退火衰减到最低值一般是峰值的十分之一。峰值学习率本身和batch size、模型规模相关。经验上1e-4到3e-4是一个常见的区间。具体峰值怎么选我试过两组对比一组是固定1e-4另一组是学到了“线性缩放规则”——batch size翻倍时学习率也适当调大但不是严格线性实际操作中我会保守一点按sqrt缩放。观察loss曲线时如果训练初期loss下降非常缓慢可能是学习率太小但如果loss在warmup结束后立刻开始大幅震荡说明峰值学习率太高了。我推荐一个省事的办法先用一个小规模实验比如模型缩小10倍确定学习率范围再放大到完整模型上。小模型训练速度快一次能跑十几种超参组合省下来的时间远大于那点训练成本。warmup步数的计算如果你的总训练步数是50000步我通常设warmup为500~1000步即1%~2%。数据量越大、batch size越大warmup可以适当增加比例。def get_lr(step, total_steps, peak_lr, warmup_steps): if step warmup_steps: # 线性上升 return peak_lr * (step 1) / warmup_steps # 余弦衰减从峰值衰减到 peak_lr / 10 progress (step - warmup_steps) / (total_steps - warmup_steps) return peak_lr / 10 0.9 * peak_lr / 2 * (1 math.cos(math.pi * progress))4.3 训练监控与断点续训训练跑起来之后最忌讳“关进小黑屋不看”。我自己至少经历过三次loss悄悄涨上去但快照已经覆盖的惨案。因此监控体系必须在训练开始前就建好。我的做法是每个step打印一次loss、lr、token吞吐量。每100步记录一次训练集上的loss。每500步做一次小型验证集上的困惑度评估。训练日志直接用wandb或者本地tensorboard都行关键是必须记录可对比的曲线。如果你只有最后一个数字你永远不知道模型在哪一步开始崩溃的。断点续训这件事说得容易做起来难。每次保存checkpoint时至少包含模型权重、优化器状态、学习率调度器状态、随机数生成器状态、当前step。有一步没保存恢复训练时就可能出问题。我当时踩过一个经典坑只保存了模型权重和优化器没保存RNG状态恢复训练后数据顺序发生变化因为Dataloader的shuffle随机种子丢了训练曲线出现一跳一跳的异常。加回来之后问题消失。保存频率上我是每1000步保存一次全量checkpoint每5000步保留一个长期版本磁盘够就多存几份不够就只保留最近三个最佳验证loss版本。4.4 单机多卡与分布式策略当你开始训练250M以上模型时单卡可能已经不够了。最简单的多卡方案是DistributedDataParallelDDP。DDP的原理是把模型复制到每张卡上各自计算梯度然后做梯度同步再更新。DDP最需要注意的地方是batch size的全局语义。如果你用8张卡每张卡batch size 16那么全局有效batch是128。所有依赖batch size的超参数学习率、warmup、梯度累积步数都要按照128来算。另一个细节是数据分片。Pytorch DDP要求每个进程的数据不重叠。我的做法是# 每个进程拿到唯一一份数据切片 dataset MyDataset(data_path) sampler DistributedSampler(dataset, num_replicasworld_size, rankrank) dataloader DataLoader(dataset, samplersampler, batch_sizeper_gpu_batch)每个epoch开始时要记得调用sampler.set_epoch(epoch)否则每个epoch的数据顺序都一样模型容易过拟合数据顺序而不是学习数据分布。如果显存还是不够下一层优化是FSDPFully Sharded Data Parallel它会把模型参数、梯度、优化器状态分片到多张卡上。但FSDP的通信开销更大需要调sharding_strategy和cpu_offload等参数。我的建议是先跑通DDP再用FSDP优化显存千万不要一上来就上FSDP。5. 推理与评估不只是能生成就行5.1 KV Cache的原理与显存估算自回归生成是逐个token进行的。每生成一个新token所有早先token的Key和Value其实不需要重新计算。KV Cache就是把这些中间结果存下来避免重复计算。没有KV Cache生成第N个token时要重新计算前面N-1个token的所有attention复杂度是O(N^2)。 有KV Cache每个token只需要算一次K、V后续直接查缓存复杂度降为O(N)。让我给一个具体的显存估算。假设模型配置是层数L24头数16head_dim64KV头数4序列长度S2048batch size B4。每层每个KV cache元素占用KV头数×head_dim×2K和V×序列长度×batch size。KV cache per layer 4 × 64 × 2 × 2048 × 4 4,194,304 个元素用bf16存储每个元素2字节单层就是约8MB。24层就是约192MB。这个数字看起来不大但如果batch size从4提到32序列长度推到8192KV cache直接就奔着几个GB去了。这也是为什么GQA重要的原因——KV头数减半KV cache直接减半。实现KV Cache的时候我推荐预分配内存而不是动态扩展。比如直接分配(batch, max_seq_len, num_kv_heads, head_dim)的tensor用位置索引不断往里写避免频繁resize带来的开销。5.2 采样策略温度、top-p、重复惩罚训练完成的模型生成的文本默认是贪婪解码每次都选概率最高的token。但贪婪解码有两个问题一是一样的输入永远输出一样的结果二是容易出现重复循环。工程上最常用的解码参数组合是temperature控制概率分布的尖锐程度。温度越低生成越保守温度越高越随机。典型值0.7~0.9用于对话0.1~0.3用于代码或逻辑推理。top_p核采样只从累积概率达到p的最小集合里采样避免从大量低概率token里选出无关内容。典型值0.9~0.95。repeat_penalty对已出现的token概率做惩罚减少重复。典型值1.05~1.2。一个容易忽略的细节是温度缩放发生在softmax之前还是之后。正确的做法是对logits先除以temperature再做softmax。如果你对已经softmax之后的概率做温度操作结果是完全错误的。def sample(logits, temperature0.8, top_p0.9): logits logits / temperature sorted_logits, sorted_indices torch.sort(logits, descendingTrue) cum_probs torch.cumsum(sorted_logits.softmax(-1), dim-1) valid cum_probs top_p # 至少保留一个token valid[..., 0] True logits[~valid] float(-inf) probs logits.softmax(-1) return torch.multinomial(probs, num_samples1)5.3 评估指标的选择与误区训练阶段的评估最常用的指标是困惑度Perplexity, PPL。PPL实际是交叉熵的指数形式PPL exp(loss)。PPL越低越好。但PPL有它自己的局限它衡量的是“模型对训练分布预测的准确度”不代表生成的文本就自然流畅。我在项目里遇到过一个模型PPL降到12但生成的文本还是前言不搭后语。原因是评测集和训练集的分布差异太大PPL在overfit情况下根本没有参考价值。因此训练完成之后还需要额外的评估维度在领域外的通用benchmark上跑测试比如常识问答、数学题。对于对话场景做人工或LLM辅助的偏好评估A/B对比。检查生成样例的长度分布、重复率、中文标点使用是否规范。不要只盯PPL曲线。至少选几个固定prompt在训练过程中定期抽样生成文本看一眼。你肉眼看到的效果往往比任何单指标都更能暴露问题。6. 常见问题排查实录6.1 Loss不降或平台期训练开始后如果loss从头到尾几乎没有下降我建议按下面顺序排查数据问题检查tokenizer是不是把所有文本都编码成了同一个token比如词表太小导致 出现频率过高。打印几条训练样本出来看一眼就知道。模型问题检查是否有残差连接漏接、注意力mask是否错误。一个排查技巧是在固定batch上做一次过拟合测试小数据上训练loss应该能降到接近0做不到说明模型本身有问题。优化器参数AdamW的beta2参数对梯度稀疏场景很敏感。如果设置成0.999但数据里大量填充token更新会变得非常慢尝试beta20.95。学习率峰值太低会导致loss下降缓慢先用一个高于常规的学习率试一把确认模型能学再调回去。平台期loss卡住不动则更棘手常见原因是数据多样性不足。如果所有训练数据都是单一风格模型很容易提前进入饱和。增大数据量、提升数据混合中的多样性 往往比调参有用得多。6.2 Loss突刺与训练发散Loss突刺sudden spike是指训练很久之后loss突然暴涨好几倍然后又慢慢落下来。这是我最怕的问题因为它意味着训练稳定性被打破了。最常见的触发原因数据里混进了异常样本例如超长文本撑爆了上下文或包含非UTF-8字符。学习率在warmup之后的衰减阶段出现数值震荡。bf16的数值精度在某些层的值域上不够导致梯度溢出。我的应对策略在Dataloader里做异常样本过滤比如过滤掉序列长度超过阈值、特殊字符比例过高的样本。把clip_grad_norm加上通常设1.0。就算出现梯度峰值也能在极端情况下保护模型不会彻底发散。如果突刺频繁出现把学习率峰值降低30%再看。始终保留最近一个checkpoint一旦出现发散趋势回滚到突刺之前并降低学习率重启。6.3 显存不够怎么办显存爆掉是训练过程中最令人头疼的问题。我自己的排查顺序减少batch size这个最简单但效率损失也直接。开启gradient_checkpointing。用“训练时重新计算激活值”换取显存速度会慢20%~30%但显存占用能降一半以上。检查是否有不必要的缓存。例如PyTorch的pinned_memory、cuda.MemPool、未释放的中间张量。用torch.cuda.memory_summary()看完整的内存分配情况定位是哪一层的tensor占了大头。对于注意力层务必用Flash Attention或SDPA而不是手写标准attention。标准attention会生成(batch, heads, seq_len, seq_len)的注意力矩阵序列长度2048时占用极其恐怖Flash Attention通过分块计算显存占用从O(n²)降到O(n)。6.4 生成重复与幻觉问题推理阶段最常见的两个质量问题一个是重复一个是幻觉。重复的根源通常是解码策略导致。单纯的temperature再低也会产生“陷入循环”的问题因为模型发现高概率的输出路径只有某几个token转来转去就在那打转。实践经验是重复惩罚repeat penalty比降低temperature更有效。幻觉问题则是模型没有真正学会“区分已知和不知道”。我的经验是幻觉不可能被训练单一解决纯粹靠解码策略无法根除。工程上可以做两个缓解在SFT数据里刻意加入“我不知道”的回答样本教会模型在知识不确定时拒绝回答。在推理时接入检索RAG把模型生成的内容和外部知识源做比对后再返回。这里补充一个很关键的认知幻觉不是bug而是语言模型的固有倾向。大模型本质上是概率性的文本续写器它的目标是生成“合理”的文本而不是生成“正确”的文本。想要减少幻觉必须在训练目标和推理约束上同时下功夫。个人体会这个项目从动手到基本跑通前后花了大概两个多月。最大的收获不是最终那个能生成文本的小模型而是我终于对模型训练各个环节的“手感”有了真实的认知。很多参数看文档是一回事自己调又是一回事。比如学习率从3e-4调到1e-4看起来只是数字变了实际上loss曲线的形态完全不一样训练稳定性也天差地别。这种经验不亲手跑一遍光靠看书是体会不到的。如果你也想做类似的项目我最后分享一个小技巧准备一份“参数改动记录表”。每次改动超参数、数据结构或模型结构都把改动前后的实验结果记录在案。你会惊讶地发现很多训练问题其实是因为某次不经意的改动造成的没有记录时你只能从头慢慢排查。从零开始做AI工程进度不快但每一步都算数。
返回列表