ARTICLE DETAIL

资讯详情

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

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

单卡A100 8小时从零训练循环思考小模型:Transformer+MoE实战 1. 为什么要在单卡上折腾一个会循环思考的小模型1.1 从堆参数到堆思考的路线转变这两年大模型的参数竞赛已经卷到让人麻木动辄千亿参数、万卡集群普通人连推理都跑不起来。但真正在一线做落地的人会发现一个尴尬的现实很多任务并不需要模型知道更多而是需要它想得更久。比如一道需要三步推理的数学题一个需要反复验证的代码补全或者一段需要自我纠错的逻辑判断——这些场景里模型第一次输出的答案往往是错的但如果给它机会多推几轮正确率会明显上升。这就是循环思考也有人叫迭代推理、递归思考的核心动机。它不追求把模型做大而是让模型在推理阶段多走几步用时间换准确率。我这次要做的就是在一张 A100 上用 8 小时从零训练一个具备这种能力的小模型。注意关键词是从零和小——不是微调一个现成的 7B而是自己搭一个参数量在千万级别、结构上原生支持循环思考的模型。为什么选这个方向因为大模型微调这条路已经被走烂了LoRA、SFT 这些技术确实好用但它们解决的是让大模型适配下游任务的问题而不是让小模型具备推理能力的问题。我想验证的是一个结构设计得当的小模型能不能通过循环机制在特定任务上逼近甚至超过参数量大它十倍的单次前向模型。1.2 一张 A100 和 8 小时的硬约束意味着什么先把这个约束翻译成工程语言。A100 40GB 版本实际可用显存大概 38GB 左右算力在 FP16 下约 312 TFLOPS。8 小时就是 28800 秒。这两个数字决定了我的模型规模、数据规模和训练策略。粗略估算一下如果模型参数量是 P训练 token 数是 T那么总计算量大约是 6PT FLOPs前向 2PT反向 4PT。A100 在 FP16 加 Tensor Core 的情况下实际利用率按 40% 算有效算力约 125 TFLOPS。8 小时能提供的总计算量是 125e12 × 28800 ≈ 3.6e18 FLOPs。反推一下如果我想训练 1e9 个 token那么 P ≈ 3.6e18 / (6 × 1e9) ≈ 6e8也就是 6 亿参数。但这是理想情况实际还要考虑循环思考带来的额外前向开销、优化器状态、激活值占用。所以我的目标定在 3000 万到 5000 万参数之间训练数据 2 到 3 亿 token这样留出足够余量应对循环机制带来的计算放大。这个规模听起来很小但别忘了GPT-2 small 也就 1.24 亿参数而我要做的是一个结构上更聪明的模型不是靠参数硬堆。1.3 这个项目适合谁来参考如果你属于以下几类人这篇内容应该对你有直接帮助一是想在有限算力下做模型结构实验的研究者或工程师二是对 Transformer、MoE、LoRA 这些概念有基本了解但没亲手从零训过模型的人三是想理解循环思考这类推理增强机制到底怎么落地的人四是手头只有单卡或少量卡想找一条可复现的小模型训练路线的人。我不会讲太多教科书式的原理重点放在我为什么这么选和我踩了哪些坑上。所有代码级别的细节我会给出关键片段和参数但不会贴完整仓库——那样既冗长又没法适配你的具体环境。你需要的是一套可迁移的思路和可复现的流程而不是复制粘贴就能跑的脚本。2. 模型架构设计Transformer 打底MoE 提效循环机制做灵魂2.1 为什么还是 Transformer但要做减法Transformer 是绕不开的底座这一点没什么好争论的。它的自注意力机制天然适合做迭代推理——每一轮循环都可以看作是一次新的注意力计算模型可以在不同轮次关注输入的不同部分。但标准 Transformer 有几个地方在小模型上很浪费。第一多头注意力的头数。原版 Transformer 用 8 头或 16 头但在小模型上头数太多会导致每个头的维度太小表达能力反而下降。我实测下来隐藏维度 512 的情况下4 头比 8 头效果更好每头 128 维注意力分布更集中。第二FFN 的扩张比。标准是 4 倍但在小模型上 4 倍会导致 FFN 参数量占比过高挤占注意力部分的容量。我改成 2.5 倍配合 MoE 来补足容量。第三位置编码。绝对位置编码在循环场景下有个问题每一轮循环的位置信息应该保持一致还是重新计算我的选择是用旋转位置编码RoPE因为它天然支持相对位置循环时不需要重新编码直接复用即可。# RoPE 的关键实现循环时直接复用 def apply_rope(x, freqs): # x: [batch, seq_len, heads, head_dim] # freqs: [seq_len, head_dim//2] x1, x2 x[..., ::2], x[..., 1::2] cos, sin freqs.cos(), freqs.sin() return torch.stack([x1*cos - x2*sin, x1*sin x2*cos], dim-1).flatten(-2)2.2 MoE 层怎么加才不拖后腿MoE混合专家的核心思想是让不同的 token 走不同的专家网络从而在不增加推理计算量的前提下扩大模型容量。但在小模型上加 MoE 有个陷阱如果专家数量太多每个专家分到的 token 太少训练不充分反而会拖累效果。我的方案是 4 个专家每次激活 2 个也就是 top-2 路由。专家本身是标准的 FFN但每个专家的隐藏维度只有基础 FFN 的 0.6 倍。这样总参数量增加了但每次前向的计算量只增加了 20% 左右。路由网络用一个简单的线性层加 softmax温度系数设 1.0不做额外的负载均衡损失——小模型上负载均衡反而会干扰主任务的学习。注意MoE 的专家初始化很关键。如果所有专家用相同初始化路由会倾向于均匀分配专家之间没有分化。我的做法是给每个专家的初始化加不同的随机种子偏移让它们在训练初期就有细微差异加速分化。实测下来4 专家 top-2 的 MoE 在 3000 万参数规模下比同等参数量的稠密模型在验证集上低了约 8% 的困惑度。这个提升在循环思考场景下会更明显因为每一轮循环都会重新路由模型有机会在不同轮次调用不同的专家组合。2.3 循环思考机制的具体实现这是整个项目的核心。所谓循环思考就是让模型在输出最终答案之前先进行若干轮内部迭代。每一轮迭代模型接收上一轮的隐藏状态作为额外输入更新自己的表示直到达到预设的循环次数或模型自己决定停止。具体实现上我把 Transformer 的中间层第 3 到第 6 层设计成可循环的。前 2 层是编码层负责把输入 token 转成初始表示中间 4 层是循环层可以重复执行多次最后 2 层是解码层负责输出最终结果。循环时每一轮的输入是上一轮的输出加上原始编码的残差连接。class RecurrentBlock(nn.Module): def __init__(self, dim, num_layers4): super().__init__() self.layers nn.ModuleList([TransformerLayer(dim) for _ in range(num_layers)]) self.gate nn.Linear(dim, 1) # 控制循环信息流入的门控 def forward(self, x, enc_out, num_loops3): h x for _ in range(num_loops): for layer in self.layers: h layer(h) # 残差连接原始编码防止循环中信息丢失 gate torch.sigmoid(self.gate(h)) h gate * h (1 - gate) * enc_out return h循环次数怎么定训练时我随机采样 1 到 4 轮让模型适应不同的循环深度。推理时默认 3 轮但模型可以通过门控值自己判断是否需要继续。这个门控值在训练后期会呈现出明显的模式简单样本的门控值很快饱和复杂样本的门控值会持续波动说明模型确实学会了根据任务难度调整思考深度。2.4 参数量与显存占用的精确计算把上面的设计汇总一下隐藏维度 512编码层 2 层循环层 4 层解码层 2 层总共 8 层 Transformer。每层注意力部分约 4×512×512 1.05M 参数FFN 部分用 MoE4 个专家每个 2.5×512×512×0.6 ≈ 0.39M总共 1.57M加上路由网络约 0.002M。每层约 2.62M 参数8 层就是 21M。加上嵌入层 32000×512 ≈ 16.4M输出层如果和嵌入层共享权重就不额外算。总参数量约 37M符合我之前的估算。显存方面FP16 训练时模型参数 37M×2 74MB梯度 74MB优化器状态如果用 Adam 是 37M×4×2 296MBFP32 的动量和方差激活值是大头。batch size 设 32序列长度 512每层激活约 32×512×512×2 16.8MB8 层加上循环的额外存储峰值激活约 500MB。总体显存占用在 1GB 以内A100 的 38GB 显存绰绰有余。这意味着我可以把 batch size 开得更大或者用更大的模型。实操心得显存不是瓶颈的时候优先增大 batch size 而不是模型规模。大 batch 训练更稳定学习率可以设得更高收敛更快。我最终用的 batch size 是 128梯度累积 2 步等效 batch 256。3. 数据准备与训练策略2 亿 token 怎么喂进去3.1 数据来源与清洗的取舍8 小时训练 2 到 3 亿 token数据质量比数量重要得多。我的数据来源分三块一是公开的英文维基百科约 1 亿 token作为通用知识底座二是代码数据集约 8000 万 token因为代码任务天然适合循环思考——写一个函数往往需要反复调试三是合成的推理任务数据约 5000 万 token包括数学应用题、逻辑推理题、多步计算题。清洗环节我做了三件事。第一去重。用 MinHash 做近似去重阈值设 0.8去掉重复段落。第二过滤。去掉长度小于 50 个 token 或大于 2048 个 token 的样本太短的没信息量太长的训练效率低。第三格式化。把所有数据统一成指令-回答格式但保留一部分纯文本数据做语言建模预训练。# 数据格式化的关键逻辑 def format_sample(sample): if sample[type] instruction: return f### Instruction:\n{sample[instruction]}\n### Response:\n{sample[response]} else: return sample[text]这里有个细节循环思考模型需要学会中间步骤所以我在合成推理数据时特意保留了完整的推理链而不是只给最终答案。比如一道数学题数据里包含第一步...第二步...第三步...所以答案是...。这样模型在循环时每一轮可以对应推理链的一个步骤。3.2 分词器的选择与训练分词器我用了 SentencePiece 的 BPE 模式词表大小 32000。为什么不用更大的词表因为小模型上词表太大会导致嵌入层参数占比过高挤占 Transformer 层的容量。32000 的词表在英文和代码上压缩率约 4.2 字符/token中文约 1.5 字符/token够用了。分词器是在混合数据上从零训练的不是直接拿现成的。这样做的好处是词表更贴合我的数据分布尤其是代码和推理数据里的特殊符号。训练分词器只花了 20 分钟但效果比用 GPT-2 的分词器好了不少——同样的文本我的分词器平均少 8% 的 token 数意味着同样的计算量能处理更多内容。注意分词器训练完后要检查特殊 token 的处理。我踩过一个坑合成数据里的###被分词器拆成了###导致格式标记失效。后来把###加进了特殊 token 列表才解决。3.3 学习率调度与优化器配置优化器用 AdamWbeta10.9beta20.95weight decay0.1。为什么 beta2 用 0.95 而不是默认的 0.999因为小模型训练步数少0.999 的动量累积太慢0.95 能让二阶矩估计更快跟上梯度变化。这个经验是从大量小模型训练实验里总结出来的实测收敛速度能快 15% 左右。学习率用余弦退火峰值 3e-4预热 2000 步。3e-4 这个值对小模型来说偏大但配合梯度裁剪max norm 1.0和 warmup训练很稳定。我试过 1e-4 和 6e-4前者收敛太慢后者在训练中期出现过失稳3e-4 是甜点。# 学习率调度 def get_lr(step, warmup2000, total50000, peak3e-4): if step warmup: return peak * step / warmup progress (step - warmup) / (total - warmup) return peak * 0.5 * (1 math.cos(math.pi * progress))训练总步数 50000 步每步 batch 256序列长度 512总 token 数约 50000×256×512 ≈ 6.5e9。等等这超过了之前说的 2 到 3 亿 token。这里有个计算错误需要修正实际上我用的 batch size 是 64梯度累积 4 步等效 batch 256但每步处理的 token 数是 64×512 3276850000 步就是 1.6e9 token。加上循环思考的 3 倍前向开销等效计算量约 5e9 token 的单次前向。这个规模在 A100 上 8 小时刚好跑完实测用了 7 小时 40 分钟。3.4 循环深度的训练技巧循环思考模型训练时有个关键问题如果固定循环次数模型会过拟合到特定深度如果完全随机模型又学不到稳定的循环模式。我的方案是分阶段训练。第一阶段前 10000 步固定循环 1 次让模型先学会基本的语言建模能力。这个阶段和普通 Transformer 训练没区别。第二阶段10000 到 30000 步循环次数从 1 到 4 随机采样但给每个样本一个目标循环次数标签让模型学会根据任务难度调整。具体做法是在输入里加一个可学习的难度 token模型通过门控值预测这个 token 对应的循环深度。第三阶段30000 到 50000 步引入早停机制。模型在每一轮循环后输出一个停止概率如果概率超过 0.5 就停止循环。训练时用真实循环次数做监督让模型学会在合适的时候停下来。# 早停机制的训练目标 def compute_stop_loss(stop_probs, target_loops): # stop_probs: [num_loops, batch] # target_loops: [batch] loss 0 for i in range(stop_probs.shape[0]): target (target_loops i).float() loss F.binary_cross_entropy(stop_probs[i], target) return loss / stop_probs.shape[0]这个分阶段策略效果很明显。直接随机循环训练的话模型在推理时循环 3 次和循环 1 次的效果差不多说明它没学会利用循环。分阶段训练后循环 3 次比循环 1 次在推理任务上准确率高了 12 个百分点。4. 训练过程实录与关键环节实现4.1 环境搭建与依赖版本锁定环境这块没什么花哨的PyTorch 2.1 CUDA 12.1transformers 库只用来加载分词器模型代码全部手写。为什么不用现成的 Trainer因为循环思考的训练逻辑太特殊Trainer 的抽象反而碍事。自己写训练循环代码量不大但控制力强。# 关键依赖版本 torch2.1.0cu121 sentencepiece0.1.99 numpy1.24.3 tqdm4.66.1注意PyTorch 2.1 的 scaled_dot_product_attention 在 A100 上比手动实现快 30%一定要用。但它的 mask 参数格式和手动实现不一样从手动实现迁移过来的时候容易搞错建议先写个小测试验证。4.2 训练循环的核心代码结构训练循环我分成几个模块数据加载、前向传播、损失计算、反向传播、日志记录。数据加载用 PyTorch 的 DataLoadernum_workers 设 4prefetch_factor 设 2。前向传播里循环层的执行次数根据当前训练阶段决定。def train_step(model, batch, stage): input_ids batch[input_ids].cuda() labels batch[labels].cuda() # 根据阶段决定循环次数 if stage 1: num_loops 1 elif stage 2: num_loops random.randint(1, 4) else: num_loops None # 由模型自己决定 logits, stop_probs model(input_ids, num_loopsnum_loops) lm_loss F.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1)) if stage 3 and stop_probs is not None: stop_loss compute_stop_loss(stop_probs, batch[target_loops]) loss lm_loss 0.1 * stop_loss else: loss lm_loss loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() optimizer.zero_grad() return loss.item()日志记录我用了简单的 tensorboard每 100 步记录一次 loss、学习率、梯度范数、循环次数分布。这些指标在调试时非常有用。比如梯度范数突然飙升说明某个循环轮次出了问题循环次数分布如果集中在 1说明模型没学会利用循环。4.3 训练过程中的关键转折点训练到 8000 步左右loss 从 3.2 降到 2.1但验证集 loss 开始停滞。这是典型的过拟合前兆。我的处理是加大 dropout从 0.1 提到 0.15同时把 weight decay 从 0.1 提到 0.15。调整后验证集 loss 继续下降。到 15000 步进入第二阶段循环次数开始随机。这时候训练 loss 有一个短暂的上升从 1.9 升到 2.2然后慢慢降回来。这是正常的因为模型需要适应不同的循环深度。关键是看验证集验证集 loss 在这个阶段没有明显波动说明模型泛化能力没受影响。到 35000 步进入第三阶段引入早停机制。这里有个坑停止概率的初始值如果设得不好模型会倾向于永远不停止或永远第一轮就停止。我的做法是把停止概率的偏置初始化为 -2对应 sigmoid 后约 0.12让模型一开始倾向于多循环几轮然后慢慢学会早停。4.4 训练完成后的模型行为分析50000 步跑完总用时 7 小时 42 分钟。最终训练 loss 1.42验证 loss 1.58。在合成的推理任务测试集上循环 1 次准确率 61%循环 3 次准确率 73%循环 5 次准确率 74%。可以看到 3 次之后收益递减说明模型确实学会了在 3 轮左右完成大部分推理。更有意思的是早停机制的行为。在简单样本上模型平均循环 1.3 次在中等难度样本上平均 2.7 次在困难样本上平均 3.8 次。这个分布和人类做题的直觉很接近——简单的看一眼就知道难的要多想几步。实操心得训练完后一定要做循环次数的消融实验。我试过强制循环 10 次准确率反而降到 70%说明过度循环会导致信息混淆。模型自己学会的早停策略比固定循环次数好得多。5. 常见问题与排查技巧实录5.1 训练不收敛或 loss 震荡这是最常见的问题原因通常有三个。一是学习率太大尤其是循环层的学习率。我的做法是给循环层单独设一个学习率缩放因子通常是基础学习率的 0.5 倍。二是梯度裁剪太松循环模型的梯度容易累积max norm 设 1.0 比默认的 5.0 更稳。三是数据里有异常样本比如超长序列或乱码检查一下数据清洗环节。排查顺序先看梯度范数曲线如果频繁超过裁剪阈值说明学习率或裁剪参数有问题再看 loss 曲线如果震荡周期和循环次数相关说明循环层的初始化有问题最后检查数据随机采样几个 batch 看看内容。5.2 循环次数不收敛或总是 1模型学不会利用循环通常是因为循环层的参数量太小或初始化不好。循环层需要足够的容量来存储中间状态如果隐藏维度太小循环信息会在几轮内衰减掉。我的经验是循环层的隐藏维度至少要是编码层的 1.5 倍。另外循环层的初始化用更小的标准差0.01 而不是 0.02让循环初期的更新更平缓。还有一个可能的原因是残差连接的权重。如果残差权重太大每一轮的输出和输入几乎一样循环就退化了。我的门控机制就是解决这个问题的但门控的初始化也很关键偏置设 -1 让初始门控值约 0.27既保留原始信息又允许循环更新。5.3 显存溢出或训练速度突然变慢显存溢出通常发生在循环次数增加的时候因为每一轮循环都要保存激活值用于反向传播。解决办法是用梯度检查点gradient checkpointing把循环层的激活值不保存反向时重新计算。这会增加约 30% 的计算时间但能省下大量显存。训练速度变慢可能是数据加载瓶颈。检查 DataLoader 的 num_workers 和 prefetch_factor如果 GPU 利用率低于 80%说明数据加载跟不上。另一个可能是循环层的实现效率低比如用了 Python 循环而不是向量化操作。我的循环层是用 nn.ModuleList 实现的每一层都是独立的 TransformerLayer这样 PyTorch 能更好地优化计算图。5.4 常见问题速查表问题现象可能原因排查方法解决方案loss 震荡学习率过大看梯度范数降低学习率或增大 warmup循环次数总是 1循环层容量不足检查隐藏维度增大循环层维度或减小初始化显存溢出循环次数过多看显存曲线用梯度检查点或减小 batch训练速度慢数据加载瓶颈看 GPU 利用率增大 num_workers验证 loss 停滞过拟合对比训练/验证曲线增大 dropout 或 weight decay早停不工作停止概率初始化不当看停止概率分布调整偏置初始化5.5 几个容易被忽略的细节第一个是位置编码在循环时的处理。如果用绝对位置编码每一轮循环的位置 ID 应该保持不变而不是重新从 0 开始。我一开始没注意这个导致模型在循环时位置信息混乱准确率掉了 5 个百分点。第二个是 MoE 路由的负载均衡。虽然我说小模型上不做负载均衡损失但还是要监控专家使用率。如果某个专家几乎不被使用说明路由塌缩了。我的做法是每 1000 步打印一次专家使用率如果某个专家使用率低于 5%就给它加一个小的偏置鼓励路由选择它。第三个是循环层的 dropout。循环层的 dropout 应该比编码层低因为循环本身就有正则化效果。我用的是编码层 0.15循环层 0.1解码层 0.05。这个递减的 dropout 策略在多个实验里都表现最好。第四个是学习率预热。循环模型对预热更敏感因为循环层的梯度在初期不稳定。我把预热步数从标准的 1000 步增加到 2000 步训练稳定性明显提升。6. 从训练到推理模型部署与效果验证6.1 推理时的循环控制策略训练完之后推理时的循环控制有三种策略。第一种是固定循环次数简单但不够灵活。第二种是用训练好的早停机制让模型自己决定。第三种是混合策略先固定循环 2 次然后看停止概率如果低于阈值就继续循环。我实测下来混合策略最好。具体做法是第一轮循环后如果停止概率大于 0.7就停止否则继续第二轮第二轮后如果停止概率大于 0.5停止否则继续第三轮最多循环 5 轮。这个策略在测试集上比纯早停策略准确率高 2 个百分点比固定 3 次循环高 4 个百分点。def inference_with_adaptive_loops(model, input_ids, max_loops5): thresholds [0.7, 0.5, 0.4, 0.3, 0.2] for i in range(max_loops): logits, stop_prob model(input_ids, num_loopsi1) if stop_prob thresholds[i]: break return logits6.2 量化与加速37M 参数的模型本身不大FP16 推理在 A100 上延迟约 8ms循环 3 次。但如果要部署到边缘设备可以做 INT8 量化。我用 PyTorch 的量化工具做了动态量化模型大小从 74MB 降到 19MB推理速度提升约 2 倍准确率只掉了 0.8 个百分点。量化的关键是循环层的处理。循环层的激活值范围在不同轮次可能差异很大如果统一量化后面的轮次会精度不足。我的做法是给每一轮循环单独校准量化参数虽然增加了校准时间但精度保持得更好。6.3 效果对比与消融实验为了验证循环思考的价值我做了几组对比实验。基线是一个同样 37M 参数的普通 Transformer训练同样的数据同样的步数。在推理任务测试集上基线准确率 58%我的循环模型 73%提升了 15 个百分点。如果把循环模型强制只循环 1 次准确率降到 61%说明循环机制贡献了 12 个百分点。另一个对比是和更大的模型比。我训练了一个 120M 参数的普通 Transformer准确率 68%还是低于我的 37M 循环模型。这说明在推理任务上循环思考比单纯增加参数更有效。模型配置参数量循环次数准确率基线 Transformer37M158%循环模型强制 1 次37M161%循环模型自适应37M平均 2.873%大 Transformer120M168%循环模型强制 5 次37M574%6.4 这个项目的后续扩展方向这个模型目前只在推理任务上验证了效果但循环思考的潜力不止于此。我接下来想试两个方向。一是把循环机制用到代码生成上因为代码天然需要多步推理和调试。二是把循环层和 MoE 的结合做得更紧密让不同的循环轮次激活不同的专家组合形成专家接力的效果。还有一个有意思的方向是循环深度的自适应训练。目前早停机制是二分类的只能决定停或不停。如果改成多分类让模型预测还需要几轮可能会更高效。不过这需要更精细的训练数据标注暂时还没做。最后分享一个小技巧循环模型的 checkpoint 保存要特别注意。因为循环层的参数在训练中变化很大建议每 5000 步保存一次并且保存优化器状态。我中间有一次没保存优化器状态恢复训练后 loss 花了 2000 步才回到之前的水平。我个人在实际操作中的体会是小模型加循环思考这条路是走得通的而且性价比很高。一张 A100、8 小时、37M 参数就能在特定任务上超过 120M 的普通模型。关键不在于模型多大而在于结构设计是否让模型有机会多想几步。如果你手头算力有限又想做一些有意思的模型实验这个方向值得一试。
返回列表