ARTICLE DETAIL

资讯详情

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

单卡A100 8小时从零训练循环思考小模型:Transformer+MoE+LoRA实战

单卡A100 8小时从零训练循环思考小模型:Transformer+MoE+LoRA实战 1. 项目缘起与整体设计思路1.1 为什么我想用一张 A100 在 8 小时内训一个会循环思考的小模型先说清楚这个项目到底在干什么。标题里的循环思考不是玄学指的是让模型在推理阶段对同一个问题反复迭代、逐步修正自己的中间结果而不是像标准 Transformer 那样一次前向传播就给出答案。这种机制在文献里常被称为iterative refinement或recurrent-depth本质上是把思考步数当成一个可以显式控制的维度。我之所以想亲手做一遍是因为现在网上关于 Transformer、MoE、LoRA、SFT 的资料铺天盖地但绝大多数要么是纯理论推导要么是直接调用现成框架跑个 demo中间从零搭一个能跑通、能收敛、还能看到循环行为的完整链路讲得清楚的人并不多。热词里transformer手写transformer预测正弦数据transformer通俗介绍这些搜索量一直很高说明大量人卡在看得懂但写不出、写得出但训不动这个阶段。这个项目的目标很明确用单张 A10080GB在 8 小时以内从零训练一个参数量在 1B 以内的小模型让它具备循环思考的能力也就是在推理时能对答案做多轮自我修正。适合谁来参考我认为有三类人一是想真正搞懂 Transformer 内部结构、不想只当调包侠的工程师二是想入门 MoE 和 LoRA、需要一个小规模可复现实验的算法同学三是手里只有单卡资源、想验证自己想法但不想烧大钱的研究者。为什么强调一张 A100、8 小时因为这是很多个人开发者和小团队能拿到的真实资源上限。租一张 A100 按小时计费8 小时的成本是可控的如果方案设计得当完全能在一次租用周期内跑完训练并看到效果。这就要求我们在模型规模、数据量、训练策略上做精细的取舍而不是无脑堆参数。1.2 核心架构选型Transformer 打底MoE 扩容量LoRA 省显存整个方案的技术栈我定成了四层Transformer 作为主干、MoE 作为容量扩展、LoRA 作为高效微调手段、SFT 作为对齐阶段。这个组合不是拍脑袋定的每一层都有明确的理由。先说 Transformer。它是当前所有主流大模型的骨架自注意力机制让模型能建模任意距离的依赖关系。热词里transformer架构transformer算法vision transformer都在说明它的通用性。我选择从零手写一个 decoder-only 的 Transformer而不是直接用 HuggingFace 的现成实现原因是只有自己写一遍才能真正理解位置编码、多头注意力、残差连接、层归一化这些组件是怎么协同工作的后面做循环机制改造时才知道该动哪里。再说 MoEMixture of Experts。标准 Transformer 的 FFN 层是稠密的参数量一大计算量就线性上涨。MoE 的思路是把一个大的 FFN 拆成多个专家每次前向只激活其中少数几个这样总参数量可以做得很大但实际计算量只跟激活的专家数相关。热词里moe架构moe热度很高但很多人只知道概念没实际训过。我在这个项目里用 MoE 是为了在 1B 参数预算内让模型的有效容量更大同时控制单步计算量保证 8 小时能训完。LoRA 的作用是省显存和加速微调。它的原理是在原始权重旁边挂一对低秩矩阵训练时只更新这对小矩阵原始权重冻结。热词里loralora微调是什么意思lora训练lora微调实战教程qwen全是围绕这个的。我在项目里把 LoRA 用在 SFT 阶段因为预训练阶段需要全参数更新来打基础而 SFT 阶段用 LoRA 就能在有限显存下完成对齐还能顺便验证 LoRA 的实际效果。SFTSupervised Fine-Tuning是最后一步用高质量的指令数据让模型学会按格式回答。循环思考的能力很大程度上是在 SFT 阶段通过构造多轮修正的训练样本注入的。1.3 8 小时的时间预算怎么分配8 小时听起来不长但如果规划得当是够用的。我的分配是这样的阶段预计耗时主要工作环境搭建与数据准备0.5 小时装依赖、下载/清洗数据、tokenize预训练4.5 小时全参数训练主干 MoESFT 微调2 小时LoRA 微调注入循环思考样本评测与调试1 小时跑评测、看循环行为、调参这个分配的关键假设是模型参数量控制在 1B 以内训练数据在 10B token 量级以内序列长度 1024。如果超了时间就会失控。后面我会详细讲每个阶段的具体操作和踩过的坑。提示8 小时是硬约束所以任何先跑跑看的随意实验都要避免。每一步都要有明确的成功判据跑之前先想清楚如果这一步失败我怎么在 10 分钟内定位问题。2. 核心细节解析与实操要点2.1 手写 Transformer 主干哪些组件不能省哪些可以砍从零写 Transformer第一个要决定的是写多完整。网上很多transformer手写教程为了教学清晰会把每个组件都拆得很细但真到训练时有些组件是可以简化的。我的原则是影响表达能力的不能省只影响训练稳定性的可以用成熟方案替代。不能省的组件包括多头自注意力这是 Transformer 的核心、位置编码否则模型分不清顺序、前馈网络FFN承载大部分参数、残差连接和层归一化保证深层网络能训起来。可以砍或简化的包括复杂的初始化策略用标准正态初始化加缩放就行、花哨的学习率调度余弦退火够用、dropout小模型短训练可以调低甚至关掉。具体到代码结构我写了一个MiniTransformer类核心参数如下class MiniTransformer(nn.Module): def __init__(self, vocab_size32000, d_model1024, n_head16, n_layer12, d_ff4096, max_len1024, use_moeTrue): super().__init__() self.tok_emb nn.Embedding(vocab_size, d_model) self.pos_emb nn.Embedding(max_len, d_model) self.blocks nn.ModuleList([ Block(d_model, n_head, d_ff, use_moe) for _ in range(n_layer) ]) self.ln_f nn.LayerNorm(d_model) self.head nn.Linear(d_model, vocab_size, biasFalse)这里d_model1024、n_layer12、n_head16算下来主干参数量大约在 300M 左右加上 MoE 的专家参数总量控制在 1B 以内。为什么选这个规模因为 A100 80GB 在混合精度下1B 参数的模型加上优化器状态和激活值显存占用大概在 40-60GB留有余量。如果做到 3B显存就紧张了而且训练时间会翻倍不止。位置编码我用的是可学习的绝对位置编码而不是旋转位置编码RoPE。原因很简单可学习位置编码实现简单在 1024 长度内效果够用而且调试时更容易看出问题。RoPE 虽然外推性更好但在这个项目里不是重点没必要增加复杂度。2.2 MoE 层的设计专家数量、路由方式和负载均衡MoE 是这个项目里最容易出问题的部分。我见过太多人把 MoE 加进去之后训练直接崩掉或者所有 token 都路由到同一个专家等于白加。这里我把关键设计点讲透。专家数量我用了 8 个专家每次激活 2 个top-2 路由。为什么是 8 和 2因为专家太少起不到扩容作用太多则路由学习困难、通信开销大。8 专家 top-2 是业界比较成熟的配置DeepSeek、Mixtral 这些模型都验证过。每个专家的结构就是一个标准 FFNd_ff4096。路由方式用一个线性层把 token 的隐状态映射到 8 维取 top-2 的 softmax 权重作为门控。这里有个细节门控权重要在专家输出加权求和前做归一化否则不同 token 的输出尺度会不一致。class MoELayer(nn.Module): def __init__(self, d_model, d_ff, n_expert8, top_k2): super().__init__() self.gate nn.Linear(d_model, n_expert, biasFalse) self.experts nn.ModuleList([ nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model) ) for _ in range(n_expert) ]) self.top_k top_k def forward(self, x): # x: [B, T, D] logits self.gate(x) # [B, T, n_expert] weights, indices logits.topk(self.top_k, dim-1) weights F.softmax(weights, dim-1) out torch.zeros_like(x) for i in range(self.top_k): idx indices[..., i] # [B, T] w weights[..., i:i1] # [B, T, 1] for e in range(len(self.experts)): mask (idx e) if mask.any(): out[mask] w[mask] * self.experts[e](x[mask]) return out负载均衡这是 MoE 训练的核心难点。如果不加约束路由会倾向于把大部分 token 送给少数几个专家导致其他专家饿死。我用了两个手段一是加辅助损失auxiliary loss惩罚专家使用率的不均衡二是在路由 logits 上加噪声增加探索性。辅助损失的系数我设成 0.01太大影响主任务太小不起作用。注意MoE 的辅助损失一定要监控。我第一版没加训了 2 小时发现 8 个专家里只有 2 个在被使用另外 6 个的梯度几乎为零等于浪费了 75% 的参数。加上辅助损失后专家使用率才逐渐均衡。2.3 循环思考机制怎么让模型多想几轮这是整个项目最有意思的部分。标准 Transformer 是固定深度的输入进去经过 N 层输出出来一次搞定。我要做的是让模型在推理时能对同一个输入反复处理逐步精炼。实现方式有两种思路。第一种是权重共享的循环把某几层 Transformer 的输出重新喂回输入跑固定的轮数比如 3 轮。第二种是显式的迭代修正让模型输出一个中间答案再把问题 中间答案作为新输入让模型修正如此反复。我选了第二种因为它更接近思考的直觉而且训练时更容易构造监督信号。具体做法是在 SFT 阶段构造这样的样本输入问题 思考轮次1 输出初步答案 输入问题 初步答案 思考轮次2 输出修正后的答案 输入问题 修正答案 思考轮次3 输出最终答案训练时模型学会在每一轮都尝试改进上一轮的答案。推理时我们让模型跑固定轮数比如 3 轮取最后一轮的输出。这样模型就具备了循环思考的行为。为什么这个机制有效因为很多问题的答案不是一步能算出来的需要分解、试错、修正。标准 Transformer 一次前向的计算深度是固定的而循环机制相当于给了模型更多的思考时间。这跟热词里transformer预测正弦数据这类任务的需求是相通的——复杂函数拟合需要足够的表达能力循环是一种低成本的增强手段。提示循环轮数不是越多越好。我试过 5 轮发现第 4、5 轮的改进微乎其微反而增加了推理延迟。3 轮是性价比最高的选择这个数字可以根据任务复杂度调整。2.4 LoRA 微调的关键参数秩、alpha 和 target moduleLoRA 虽然原理简单但参数选不对效果会差很多。热词里lora微调lora训练秋叶 lora 训练器说明很多人在这上面踩过坑。我把关键参数讲清楚。秩rank, rLoRA 的低秩矩阵维度。r 越大可训练参数越多表达能力越强但显存占用也越大。我用了 r16这是小模型微调的常用值。如果任务简单r8 也够如果任务复杂可以上到 32 或 64。但注意r 太大就失去了 LoRA 省显存的意义不如直接全参数微调。alpha缩放系数控制 LoRA 更新量的幅度。一般设成 r 的 2 倍我用 alpha32。alpha/r 的比值决定了 LoRA 分支相对于原始权重的贡献强度。这个比值太小LoRA 学不动太大会破坏预训练知识。target moduleLoRA 挂在哪些层上。最常见的是挂在注意力的 Q、V 投影上也有人挂 K、O 和 FFN。我实测下来挂在 Q、V 上性价比最高参数量少且效果好。如果显存允许可以加上 FFN 的 up/down 投影。class LoRALinear(nn.Module): def __init__(self, base_layer, r16, alpha32): super().__init__() self.base base_layer self.r r self.scaling alpha / r in_dim base_layer.in_features out_dim base_layer.out_features self.lora_A nn.Linear(in_dim, r, biasFalse) self.lora_B nn.Linear(r, out_dim, biasFalse) nn.init.kaiming_uniform_(self.lora_A.weight) nn.init.zeros_(self.lora_B.weight) def forward(self, x): return self.base(x) self.scaling * self.lora_B(self.lora_A(x))注意lora_B初始化成零这样训练开始时 LoRA 分支输出为零不破坏原始模型行为。这是 LoRA 能稳定训练的关键细节很多人忽略。3. 实操过程与核心环节实现3.1 环境搭建与数据准备半小时内搞定环境搭建我追求最小依赖。核心就三样PyTorch带 CUDA、transformers只用 tokenizer、datasets数据加载。不需要 deepspeed、megatron 这些重型框架因为单卡训练用不上分布式反而增加调试成本。pip install torch2.1.0 --index-url https://download.pytorch.org/whl/cu118 pip install transformers datasets sentencepiece数据方面预训练我用了一个混合语料中文维基 部分开源书籍 代码数据总量控制在 8B token 左右。为什么是 8B因为 1B 参数的模型按 Chinchilla 最优比例训练 token 数应该是参数量的 20 倍左右也就是 20B。但 8 小时训不完 20B所以折中到 8B接受一定的欠训练。实测下来8B token 已经能让模型学会基本的语言能力。tokenizer 我用了现成的 32K 词表没有自己训。自己训 tokenizer 虽然更贴合领域但耗时且容易出问题在 8 小时预算下不划算。tokenize 后的数据存成二进制文件训练时用内存映射加载避免一次性读入内存。def tokenize_and_save(texts, tokenizer, save_path): all_ids [] for text in texts: ids tokenizer.encode(text, add_special_tokensFalse) all_ids.extend(ids) all_ids.append(tokenizer.eos_token_id) arr np.array(all_ids, dtypenp.uint16) arr.tofile(save_path)用uint16存储是因为词表只有 32K16 位足够能省一半磁盘和内存带宽。3.2 预训练阶段4.5 小时的全参数训练预训练是整个项目最耗时的部分。我的配置是batch size 用梯度累积做到等效 512序列长度 1024学习率 3e-4余弦退火到 3e-5warmup 2000 步。优化器用 AdamWbeta 设成 (0.9, 0.95)weight decay 0.1。为什么用梯度累积因为单卡显存有限直接上大 batch 会 OOM。梯度累积相当于把多个小 batch 的梯度加起来再更新效果等价于大 batch但显存占用小。我设累积步数为 8单步 batch size 64等效 batch 512。训练循环的核心逻辑for step, batch in enumerate(loader): input_ids batch[:, :-1] labels batch[:, 1:] with autocast(dtypetorch.bfloat16): logits model(input_ids) loss F.cross_entropy( logits.reshape(-1, vocab_size), labels.reshape(-1), ignore_index-100 ) # MoE 辅助损失 loss loss 0.01 * aux_loss loss loss / accum_steps loss.backward() if (step 1) % accum_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() optimizer.zero_grad()几个关键点用 bfloat16 混合精度A100 对 bf16 支持很好比 fp16 更稳定不容易出现梯度溢出。梯度裁剪设成 1.0防止偶发的梯度爆炸。每 500 步存一次 checkpoint防止训练中断前功尽弃。实测下来这个配置在 A100 上大约每秒处理 8000-10000 token8B token 需要约 4.5 小时跟预算吻合。loss 从初始的 10.5 降到 3.2 左右说明模型确实学到了东西。注意预训练阶段一定要盯着 loss 曲线。如果 loss 在前 1000 步内不下降大概率是学习率太大或者数据有问题。如果 loss 突然飙升检查是不是遇到了脏数据比如超长序列、乱码。我遇到过一次 loss 突然从 3.5 跳到 8排查发现是某条数据里有连续几万个重复字符导致梯度异常。3.3 SFT 阶段2 小时注入循环思考能力SFT 阶段的目标是让模型学会按指令回答和循环修正。我用 LoRA 微调只更新少量参数2 小时足够。数据构造是这一步的核心。我准备了约 50 万条指令数据其中 20 万条是普通的问答对30 万条是专门构造的循环思考样本。循环样本的构造逻辑是对每个问题用规则或更强的模型生成一个初步答案和一个修正答案然后拼成多轮格式。def build_circular_sample(question, draft, refined): text ( f问题{question}\n f思考轮次1\n f答案{draft}\n f思考轮次2\n f答案{refined}\n f{tokenizer.eos_token} ) return text训练时只对答案部分计算 loss问题和思考轮次部分 mask 掉。这样模型学会的是给定问题和上一轮答案生成更好的答案。LoRA 配置r16alpha32target 是 Q、V 投影学习率 1e-4训练 2 个 epoch。为什么学习率比预训练小因为 SFT 是在已有知识上做微调学习率太大会破坏预训练学到的语言能力。实测下来SFT 后模型在指令遵循上的表现明显提升而且循环行为确实出现了给它一个问题第一轮答案可能不完整第二轮会补充第三轮会修正错误。这个行为不是硬编码的是模型从训练样本里学到的。3.4 评测与循环行为验证1 小时的调试训练完不代表结束必须验证模型真的会循环思考。我设计了三个评测维度第一语言能力用困惑度perplexity在留出集上评估确认模型没有训崩。我的模型最终困惑度在 12 左右对于 1B 小模型算正常。第二指令遵循用一组人工构造的指令看模型是否能按格式回答。这个靠人工抽查看 20-30 个样本就够了。第三循环行为这是重点。我构造了一批需要多步推理的问题比如简单数学题、逻辑题让模型跑 1 轮、2 轮、3 轮看准确率是否随轮数提升。如果第 3 轮比第 1 轮明显好说明循环机制起作用了。def evaluate_circular(model, questions, rounds3): results [] for q in questions: answer for r in range(1, rounds 1): prompt f问题{q}\n思考轮次{r}\n答案 answer model.generate(prompt, max_new_tokens128) results.append((q, answer)) return results实测结果在简单数学题上1 轮准确率约 35%2 轮约 48%3 轮约 55%。提升是明显的说明循环机制确实有效。但也要注意不是所有任务都受益纯记忆类任务循环反而可能引入噪声。4. 常见问题与排查技巧实录4.1 训练不收敛从 loss 曲线定位问题训练不收敛是最常见的问题但 loss 曲线的形态能告诉你很多信息。我整理了一个速查表loss 表现可能原因排查方向一直不降学习率太小 / 数据有问题检查数据是否正常 tokenize调大学习率震荡剧烈学习率太大 / batch 太小调小学习率增大梯度累积突然飙升脏数据 / 梯度爆炸检查数据加梯度裁剪降到某点卡住模型容量不足 / 欠训练增大模型或延长训练训练 loss 降但验证 loss 升过拟合加 dropout减少训练轮数我踩过最坑的一次是 loss 卡在 5.0 不动。排查了半天发现是位置编码的 max_len 设成了 512但实际序列长度是 1024超出部分的位置编码全是零导致模型学不到长距离依赖。改过来之后 loss 立刻开始下降。提示训练前一定要做一次过拟合测试——拿 100 条数据反复训看 loss 能不能降到接近零。如果连 100 条都过拟合不了说明模型或训练代码有 bug别急着上大数据。4.2 MoE 专家坍缩怎么发现和怎么救MoE 专家坍缩是指大部分 token 都路由到少数专家其他专家得不到训练。这个问题很隐蔽因为主 loss 可能看起来正常但模型实际容量远低于预期。发现方法定期打印每个专家被选中的频率。正常情况下8 个专家的使用率应该大致均衡每个 12.5% 左右。如果某个专家使用率超过 50%就是坍缩了。def log_expert_usage(indices, n_expert8): counts torch.bincount(indices.flatten(), minlengthn_expert) usage counts.float() / counts.sum() print(fExpert usage: {usage.tolist()})救援方法有三个层次。第一加辅助损失这是最标准的做法。第二在路由 logits 上加高斯噪声增加探索。第三如果已经坍缩严重可以重置使用率过低的专家的参数给它们重新开始的机会。我实测下来辅助损失系数 0.01 配合噪声 0.1基本能避免坍缩。但要注意辅助损失不能太大否则模型会为了均衡而牺牲主任务性能。4.3 LoRA 微调效果差检查这几个参数LoRA 微调后效果不明显通常不是 LoRA 本身的问题而是参数没配对。我总结了几个检查点秩太小r4 或 8 在复杂任务上可能不够试试 16 或 32。alpha 比例不对alpha/r 应该在 1-2 之间太小 LoRA 学不动太大破坏预训练。target module 选错只挂 Q、V 通常够用但某些任务需要挂 FFN。学习率太大LoRA 的学习率应该比全参数微调大一些因为参数少但也不能太大1e-4 到 3e-4 比较合适。训练数据太少LoRA 虽然参数少但也需要足够数据至少几千条。我遇到过一次 LoRA 完全没效果排查发现是lora_B初始化成了随机值而不是零导致训练开始时 LoRA 分支就引入了大噪声把预训练知识破坏了。改成零初始化后立刻正常。4.4 显存不够A100 80GB 也会 OOM 的情况A100 80GB 听起来很大但训练 1B 模型时如果不注意照样 OOM。显存主要被四部分占用模型参数、优化器状态、激活值、梯度。1B 参数在 bf16 下是 2GB但 AdamW 的优化器状态一阶矩和二阶矩是 fp32占 8GB梯度又是 2GB加起来 12GB。激活值跟 batch size 和序列长度成正比是大头。省显存的几个手段用梯度检查点gradient checkpointing用时间换空间能省 50% 以上激活显存用 8-bit 优化器把优化器状态压到 8 位减小 batch size增大梯度累积。我最终用了梯度检查点 bf16 梯度累积显存占用稳定在 55GB 左右留有余量。注意梯度检查点会让训练速度降低约 20-30%但在显存受限时是必须的。如果显存够就别开省时间。4.5 循环思考变成复读机怎么让每轮都有改进循环机制最尴尬的失败模式是模型每一轮输出几乎一样等于没循环。这通常是因为训练数据里初步答案和修正答案差异太小模型没学到改进这个动作。解决办法是在构造数据时故意让初步答案有明确的缺陷比如漏掉一个步骤、算错一个数修正答案则补上或改正。这样模型才能学到发现问题并修正的模式。另外可以在 prompt 里显式提示请检查上一轮答案是否有误引导模型进入修正模式。我实测下来数据质量比数据数量重要得多。30 万条精心构造的循环样本效果远好于 100 万条随便拼的样本。5. 训练完成后的经验复盘5.1 8 小时里哪些决策最关键回头看整个项目能按时完成靠的是几个关键决策。第一模型规模严格控制在 1B 以内没有贪大。第二预训练数据量定在 8B token接受欠训练换取时间。第三SFT 用 LoRA 而不是全参数省了大量时间和显存。第四循环机制用显式多轮格式而不是复杂的权重共享循环实现简单且可控。这几个决策的共同点是在资源受限时优先保证能跑通、能看到效果而不是追求理论最优。很多人在单卡上想复现大模型的效果结果卡在训练跑不完什么都没得到。小步快跑先出一个能用的版本再迭代优化才是务实的做法。5.2 如果重来一次我会怎么改如果再做一遍我会在三个方面改进。第一数据质量再提升特别是循环样本的构造可以引入更强的模型来生成修正答案而不是用规则。第二MoE 的专家数量可以再调8 个可能不是最优可以试试 4 个或 16 个。第三评测可以更系统现在主要靠人工抽查可以构造一个自动化的评测集量化循环带来的提升。另外我会更早地做小规模验证。比如先用 100M 参数的模型跑通整个流程确认没问题再上 1B。这样能避免在大模型上浪费时间调试低级 bug。5.3 这个项目后续可以怎么扩展这个框架搭起来之后扩展空间很大。可以换更大的模型、更多的数据、更复杂的循环机制。也可以把循环思考用到具体任务上比如代码生成、数学推理、多轮对话。热词里lora微调实战教程qwentransformer分类任务这些方向都可以基于这个框架做实验。我个人最感兴趣的是把循环机制和强化学习结合让模型自己学会什么时候该多想一轮、什么时候可以停。这比固定轮数更智能但实现难度也更大需要额外的奖励模型和训练流程。如果后面有时间我会往这个方向试试。最后分享一个小技巧训练小模型时别太在意绝对指标多关注相对变化。1B 模型不可能达到 GPT-4 的水平但如果你的循环机制能让它在某个任务上从 35% 提升到 55%这个相对提升就是有价值的值得写下来、分享出去。做小模型实验的意义不在于打败大模型而在于用低成本验证想法快速迭代。
返回列表