Continuous Thought Machine源码级拆解:从原理到PyTorch最小实现

Continuous Thought Machine源码级拆解:从原理到PyTorch最小实现 最近关于“模型到底会不会思考”的讨论越来越多尤其是在推理模型热度上升之后。Sakana AI 提出的 Continuous Thought Machine 是这一方向上比较有代表性的概念之一。它想做的不只是“让模型多输出几个隐含的推理 token”而是希望在推理过程中维护一个连续、可更新的思维状态让模型真正“边想边解”。如果你去看这个方向的源码会发现它并不是一个单一的模型文件而是由多个模块组合而成。本文从源码级视角拆解 Continuous Thought Machine 的核心机制并给出一个可以在本地 CPU 上运行的 PyTorch 最小实现同时梳理在实际大模型上落地时需要关注的工程问题。无论你是对推理模型感兴趣的研究者还是希望在业务中使用高质量推理能力的大模型开发这篇文章都会有一些可以参考的内容。1. 背景为什么“连续思维”会成为研究热点1.1 从“生成答案”到“生成思考”传统的大语言模型本质上是下一个 token 预测器。给定一段上下文模型一次前向传播预测下一个 token然后把这个 token 接在输入后面继续预测下一个 token。这种方法在对话、翻译、摘要等任务上表现很好但面对复杂数学题、逻辑推理、代码 Debug 时直接生成答案很容易出错。后来出现了 Chain-of-Thought 方法。核心思路是让模型在输出最终答案之前先输出一段“思考过程”。比如User小明有 12 个苹果送给小红一半又买来 4 个现在有几个 Assistant小明原来有 12 个苹果。 送出一半后剩下12 / 2 6 个。 又买来 4 个6 4 10 个。 所以答案是 10。这种方法有效但它有一个天然局限思考过程被“token 化”了。也就是说模型只能用语言模型本身的离散 token 来表达思考。很多直觉判断、中间状态、候选方案的比较很难用几句话写清楚。1.2 什么是 Continuous Thought MachineContinuous Thought Machine 的核心思想是把“思考状态”从 token 流中抽离出来用一组连续向量continuous vector来表示。这个向量可以跨时间步长携带也可以被多次更新然后再反过来影响每一步的 token 预测。你可以把它理解成一个内嵌在语言模型里的“工作记忆”。传统 Transformer 的工作记忆是 KV Cache它记录的是过去所有 token 的信息。而 Continuous Thought Machine 额外维护了一个“当前在想什么”的状态这个状态不直接输出为文字而是作为隐变量参与每一层计算。对比一下两种路径维度Chain-of-ThoughtContinuous Thought Machine思考的载体离散 token连续向量推理时开销增加输出长度增加每个步骤内部迭代次数可解释性较好人能读到思考过程较弱思维状态不可直接阅读计算效率随着思考长度线性增长可以通过限制迭代次数来控制扩展性容易接入现有模型需要改造模型结构1.3 它解决了什么问题从源码实现的角度看Continuous Thought Machine 重点解决三个问题思考的自由度不强迫模型把思考过程写成 token允许它在连续空间中试探多种可能。推理的深度通过循环更新思维状态模型可以反复“追问”已有信息而不必生成冗长的文字。计算的灵活性可以设计成“简单问题少迭代复杂问题多迭代”的动态机制而不是所有问题都用同样的计算量。当然这种设计也会带来新问题比如训练不稳定、推理延迟增加、思维状态不可解释等。后面会专门展开。2. 源码级拆解核心模块与数据流如果你去阅读一个 Continuous Thought Machine 风格的开源实现通常会发现它由四个核心模块组成输入编码器、思维状态模块、token 解码器、更新控制器。下面按模块拆开讲。2.1 输入编码器把文本压缩成初始思维状态输入编码器的作用是把当前问题转换成模型可以处理的向量表示。在很多实现中它并不是一个独立的 BERT 或 T5 Encoder而是直接使用 Transformer 第一层或前几层的输出。关键点在于编码器不仅要输出每个 token 对应的隐藏状态还要生成一个“初始思维向量”。这个向量可以被看作模型对当前问题的第一印象。伪代码如下class InputEncoder(nn.Module): def __init__(self, d_model): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.pos PositionalEncoding(d_model) def forward(self, tokens): x self.embed(tokens) self.pos(tokens) return x # [batch, seq_len, d_model]初始思维向量通常是所有 token 隐藏状态的平均池化或注意力池化结果thought x.mean(dim1) # [batch, d_model]这一步的源码实现通常很简单但重要。2.2 思维状态模块持续更新的“工作记忆”这是整个 Continuous Thought Machine 的核心。它接收上一步的思维状态结合当前 token 序列的表示输出更新后的思维状态。常见实现有两种风格循环神经网络风格用一个 GRU / LSTM 单元来更新思想向量。Transformer 风格把思想向量作为一个额外的 token拼在序列前面经过若干层 Transformer 后取出该位置的向量作为新状态。第二种方式更容易利用现有 Transformer 优化成熟度也是很多开源实现偏好的写法。示意代码如下class ThoughtUpdater(nn.Module): def __init__(self, d_model, n_iters4): super().__init__() self.n_iters n_iters self.transformer nn.TransformerEncoderLayer( d_modeld_model, nhead8, batch_firstTrue ) def forward(self, x, thought): for _ in range(self.n_iters): # 把 thought 作为首 token 拼到序列前 seq torch.cat([thought.unsqueeze(1), x], dim1) out self.transformer(seq) thought out[:, 0] # 取首位置作为新的思想向量 x out[:, 1:] # 剩余位置仍然对应 token 序列 return x, thought这里每一步迭代都让思维向量“重新观察”一次当前序列。迭代次数 n_iters 就相当于思考的步数。n_iters 越大模型越有机会修正自己的中间状态但计算量也越大。2.3 token 解码器让思维影响每一步输出核心问题是我们有了思维向量怎么让它影响当前时刻的 token 预测最简单的做法是在每层 Transformer 输出后将思维向量通过一个线性投影然后加到 token 隐藏状态上class ThoughtInjectedLayer(nn.Module): def __init__(self, d_model): super().__init__() self.gate nn.Linear(d_model, d_model) self.act nn.SiLU() self.norm nn.LayerNorm(d_model) def forward(self, token_hidden, thought): influence self.act(self.gate(thought)) return self.norm(token_hidden influence.unsqueeze(1))更精细的方案是采用门控融合也就是让模型学习“当前 token 应该受到思维状态多大影响”。这类似于 GRU 中的更新门。2.4 更新控制器动态决定思考步数真实业务中不可能所有问题都固定思考 4 步。简单问题一步就好复杂问题可能要 16 步甚至更多。于是需要一个控制器预测每一层还需要迭代多少次。常见实现方式分类器方案在思维状态上接一个线性层输出“继续思考 / 停止思考”的概率。不确定性方案计算思维向量变化量||thought_new - thought_old||如果变化很小就提前停止。这个设计直接影响工程落地成本因为思考步数是推理延迟的主要来源。3. 最小实现PyTorch 版 Continuous Thought 模块为了让你更直观地理解我整理了一个可以在 CPU 上直接跑通的最小实现。完整代码会包含数据构造、模型定义、训练与推理目标是让读者看到 Continuous Thought 的核心循环如何工作。3.1 环境准备本文示例只需要 Python 3.9 以上版本和 PyTorch 2.0 以上版本。pip install torch不需要额外的大型依赖。3.2 定义数据集为了演示训练流程我构造一个简单的模式学习任务。给定一串随机整数序列模型要学会预测下一位。这里我故意生成一些具有一定规律的数据方便观察 loss 下降。import torch import torch.nn as nn import torch.optim as optim vocab_size 32 seq_len 16 batch_size 64 def make_batch(batch_size, seq_len, vocab_size): data torch.randint(0, vocab_size, (batch_size, seq_len 1)) x data[:, :-1] y data[:, 1:] return x, y3.3 定义 ContinuousThoughtLayer这是最核心的部分。我实现一个简单的 Continuous Thought 层内部迭代 n_iters 次每次都会用注意力机制更新 token 表示并且把思维向量一并更新。class ContinuousThoughtLayer(nn.Module): def __init__(self, d_model, nhead, d_ff, n_iters4): super().__init__() self.n_iters n_iters self.self_attn nn.MultiheadAttention( d_model, nhead, batch_firstTrue ) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model), ) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.thought_linear nn.Linear(d_model, d_model) self.gate nn.Linear(d_model * 2, d_model) def forward(self, x, thoughtNone): x: [batch, seq_len, d_model] thought: [batch, d_model] 或 None b, s, d x.shape if thought is None: thought x.mean(dim1) for _ in range(self.n_iters): # 1. token 序列自注意力 attn_out, _ self.self_attn(x, x, x) x self.norm1(x attn_out) # 2. 思维向量转换后与 token 交互 thought_t self.thought_linear(thought).unsqueeze(1) gate_input torch.cat([x, thought_t.expand(b, s, d)], dim-1) gate torch.sigmoid(self.gate(gate_input)) x x gate * thought_t # 3. FFN x self.norm2(x self.ffn(x)) # 4. 更新思维向量基于 token 池化 thought x.mean(dim1) thought thought torch.tanh(thought) return x, thought这里我使用了两步设计用 gate 控制思维向量对 token 的影响。用“残差更新 tanh”控制思维向量的变化范围避免状态爆炸。3.4 定义完整语言模型在 ContinuousThoughtLayer 外层加上词嵌入层和输出层就是一个可以训练的简易语言模型。class ContinuousThoughtLM(nn.Module): def __init__(self, vocab_size, d_model128, nhead8, d_ff256, n_iters4): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.pos_embed nn.Parameter(torch.zeros(1, seq_len, d_model)) nn.init.normal_(self.pos_embed, std0.02) self.thought_layer ContinuousThoughtLayer( d_modeld_model, nheadnhead, d_ffd_ff, n_itersn_iters, ) self.out_proj nn.Linear(d_model, vocab_size) def forward(self, x): h self.embed(x) self.pos_embed h, thought self.thought_layer(h) return self.out_proj(h) def generate(self, x, max_new_tokens8): self.eval() with torch.no_grad(): for _ in range(max_new_tokens): logits self(x) next_token logits[:, -1, :].argmax(dim-1, keepdimTrue) x torch.cat([x, next_token], dim1) if x.shape[1] seq_len: x x[:, -seq_len:] return xgenerate方法里我使用贪心解码仅用于演示。生产环境通常会用采样、top-p、beam search 等策略。3.5 训练循环训练过程与普通语言模型完全一致用交叉熵损失优化下一个 token 预测任务。model ContinuousThoughtLM(vocab_size) optimizer optim.AdamW(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() model.train() for step in range(200): x, y make_batch(batch_size, seq_len, vocab_size) logits model(x) loss criterion(logits.reshape(-1, vocab_size), y.reshape(-1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() if step % 20 0: print(fstep {step}, loss {loss.item():.4f})运行 200 步后loss 会明显下降。虽然这个玩具任务没有展示出 Continuous Thought 的绝对优势但它已经形成了一个可运行的最小闭环。3.6 推理验证随便给一个长度为 8 的输入序列让模型继续补全x, _ make_batch(1, 8, vocab_size) out model.generate(x, max_new_tokens4) print(原始输入:, x[0].tolist()) print(生成结果:, out[0].tolist())输出结果会是词汇表范围内的整数序列表示模型学会了在局部模式下续写。4. 在大模型上落地的三种低成本方案上面是最小实现但它离实际生产还有距离。真实业务中我们通常不会直接从零训练一个 Continuous Thought Machine而是基于现有开源大模型来做改造。下面给出三种成本从低到高的方案。4.1 方案一Self-Consistency 多次采样投票这是最接近“持续思考”精神但完全不改模型结构的方法。核心思想是同一个问题多次采样每次给模型不同的随机性最后对答案做多数投票。from transformers import AutoTokenizer, AutoModelForCausalLM model_name Qwen/Qwen2.5-7B-Instruct tokenizer AutoTokenizer.from_pretrained(model_name) model AutoModelForCausalLM.from_pretrained(model_name) prompt 一个边长为 5 的正方形对角线长度是多少请直接给出结果。 inputs tokenizer(prompt, return_tensorspt) outputs model.generate( **inputs, do_sampleTrue, temperature0.7, num_return_sequences8, max_new_tokens128, ) for i, output in enumerate(outputs): text tokenizer.decode(output, skip_special_tokensTrue) print(f--- 第 {i1} 次采样 ---) print(text)这里的重点是num_return_sequences。多次采样相当于让模型从不同“思维轨迹”中逼近答案然后用投票或规则筛选出最终结果。优点是改动小缺点是计算成本成倍增加。4.2 方案二两阶段思考把中间结论写回 Prompt这个方案更接近 Chain-of-Thought但比简单的“输出思考过程”多一步让模型先输出思考摘要再把摘要追加到原问题后面作为第二轮推理的上下文。import json def think_then_answer(question): thinking_input question \n\n请先不要给最终答案而是列出你的推理思路和需要确认的条件。 thinking_text generate_text(thinking_input) final_input f 原始问题{question} 你已经完成的思考 {thinking_text} 请基于上面的思考给出最终答案并检查是否有计算错误。 return generate_text(final_input)这种方式非常容易实现。它相当于强制模型进行“反思”。在很多推理任务上效果比直接生成答案要好。缺点是需要手动设计 Prompt 模板并且会消耗两倍以上的 token。4.3 方案三在开源模型上做轻量微调如果业务场景相对固定比如只做代码 Bug 修复或数学应用题可以考虑在开源模型上加入特殊的“思考 token”。例如在训练数据中加入think ... /think标记让模型生成完整推理过程。training_example s用户请修复下面代码中的问题 def add(a, b): return a - b 助手think 观察函数名是 add但返回值是 a - b。 这会导致调用 add(1, 2) 返回 -1而预期是 3。 修复方案把 return a - b 改为 return a b。 /think answerdef add(a, b): return a b /answer/s微调的目标是让模型学会在think和/think之间展开推理。这一步不需要重新设计模型结构只需要构造合适的训练数据用 LoRA 或 QLoRA 低成本微调即可。但在生产项目中有一点需要特别注意思考 token 会显著增加推理延迟。如果产品对首 token 时间要求很高必须做延迟测试。5. 常见问题与排查思路我在实现和调研过程中整理了一些高频问题很多都和训练稳定性、推理性能有关。问题现象常见原因解决思路训练时 loss 震荡严重思维向量更新幅度过大给思维向量加 tanh 或 LayerNorm必要时降低学习率推理比普通模型慢很多固定 n_iters 太大改成动态提前停止简单问题降低 n_iters思维向量对输出没有影响gate 初始化为接近 0将 gate 的偏置初始化为 0或去掉 gate 做残差融合长序列显存爆掉每轮迭代都保存了完整梯度检查 n_iters 是否相当于计算图展开层数可考虑 gradient checkpoint答案不稳定多次采样后缺少聚合规则引入 self-consistency 投票或规则校验模型“一本正经胡说八道”思考状态没有监督信号结合奖励模型或规则校验给思考过程增加额外 loss5.1 排查流程如果遇到训练不收敛建议按下面顺序排查先关闭 Continuous Thought 模块只训练普通的 Transformer 基线确认数据和 loss 计算没有问题。固定 n_iters1确认单次迭代能收敛。逐步增大 n_iters观察 loss。如果增大后反而变差说明思维向量的更新路径存在问题优先检查梯度范数。加入梯度裁剪把 max_norm 设为 1.0。如果思维向量仍然爆炸就在每次更新后加 LayerNorm。这套流程能快速定位是数据问题、模型结构问题还是优化超参数问题。6. 工程落地与最佳实践从“跑通 demo”到“上生产”中间还有很多路要走。下面是我认为最重要的几项工程建议。6.1 把思考步数做成可配置项不同业务的时效要求完全不同。智能客服场景可能要求 500ms 内返回而数学解题场景可以接受 10 秒。因此Continuous Thought Machine 的 n_iters 不应该写死在模型里而应该做成服务配置{ model: continuous-thought-7b, max_thought_iters: 8, stop_threshold: 0.001, temperature: 0.6, top_p: 0.95 }同时可以在网关层根据问题的难度路由到不同参数。简单问题走低迭代数复杂问题走高迭代数。6.2 对思维状态做监控连续思维向量虽然不可直接阅读但可以通过它的变化量来判断模型是否稳定。比如记录每一步迭代后thought向量的范数范数异常增大可能意味着模型进入不稳定状态。范数变化长期接近 0说明思维状态已经“饱和”后续迭代收益很低。将这些指标接入 Prometheus 或 Grafana比只看最终准确率更能提前发现问题。6.3 设置结果校验器生产环境中不能完全依赖模型自我修正。建议针对业务场景做一个轻量级校验器。例如数学题解析模型输出中的数字做公式验证。代码修复运行单元测试。数据抽取用正则或 schema 校验字段。校验器输出的分值可以反馈到推理链路中如果低分自动触发“再多想一次”的循环。6.4 控制安全边界给模型更多“思考”时间不等于给模型更多权利。在接入企业数据时仍然要遵循最小权限原则。尤其是在涉及数据库操作、文件删除、支付回调等高风险动作时不要直接信任模型生成的“最终答案”而是必须通过权限校验、人工审批或沙箱执行。Continuous Thought 的每一轮迭代都可能引入新的上下文这也意味着攻击者可能通过 prompt 注入把恶意指令“藏在”思考过程中。上线前建议做针对性的红队测试至少覆盖越狱提示、间接注入、角色切换三类场景。6.5 离线评测指标推理模型的评测不能只看 top-1 准确率。建议增加三个指标Passk采样 k 个结果只要有 1 个正确就算通过。适合评测模型的潜力。Self-Consistency 准确率采样多次后做投票看投票结果是否正确。适合评测实际可用性。平均思考步数记录每个正确/错误答案消耗的 n_iters观察模型是否把计算量花在了合适的地方。没有评测反馈Continuous Thought 的迭代优化就是盲人摸象。7. 总结与下一步Continuous Thought Machine 给 LLM 推理提供了一种新的设计思路不再把“思考”局限成一段可见的文字而是把它当作一组在连续空间中不断更新的隐状态。从这个角度看它与循环神经网络、状态空间模型有很强的思想延续性但又和现代 Transformer 做了结合。本文从源码级拆解了核心模块输入编码器如何生成初始思维状态思维状态如何参与 token 计算又如何被循环更新同时给出了一个完整的 PyTorch 最小实现方便你直接运行。最后也讨论了基于现有大模型实现类似效果的三种低成本方案以及训练、推理和工程落地中的关键坑点。如果你之后继续深入建议按这个顺序学习阅读并运行文中的最小实现理解“连续状态”和“token 生成”是如何交织的。尝试把 n_iters 改成动态停止观察推理速度和生成质量的变化。在开源模型上复现 Self-Consistency 方案用标准评测集测量准确率提升幅度。如果对训练感兴趣可以进一步研究 GRPO 等强化学习算法把“思考过程”和“最终奖励”联系起来。在做生产落地时不要一开始就追求复现一个完整的 Continuous Thought Machine。更稳妥的做法是先用低成本方案验证效果再逐步增加模型的“持续思考”能力。把成本和收益量化清楚比追新概念本身更有价值。