
在实际的大模型应用中In-Context LearningICL是最常用也最容易误读的一种能力模型不更新任何参数只靠 prompt 中给出的少量示例就能在新输入上完成任务。普通实现通常就是把示例拼进上下文然后让 Transformer 做一次前向传播最后自回归生成答案。这个过程简单但局限也很明显如果任务需要多步推理需要反复对照上下文中的不同示例一次固定深度的前向传播并不一定能完成足够的信息交换。于是In-Context Learning with Recurrent Latent Reasoning 这一方向开始被关注BDH-CQ 就是其中一个值得拆解和实验的思路。从命名上看BDH 可以理解为 Block-Diagonal Hidden StateCQ 可以理解为 Context Query。组合起来的思想是模型先把上下文中的每条示例编码成一块块隐藏状态然后在生成每个 token 前用 query 在这些块之间进行多轮循环读取和更新。这样做的价值在于把“看一遍上下文”变成“在潜在空间里反复思考几轮”。需要注意如果后续找到原始论文应以论文原始定义为准本文只从工程复现角度把 BDH-CQ 当作一个可运行的循环潜在推理框架来理解。下面的最小可复现 PyTorch 原型会涉及合成 few-shot 任务构造、BDH-CQ 模块实现、训练评估、参数调优和常见问题排查。最终目标不是复现任何官方 benchmark而是搭建一套可以在本机运行、能验证“循环步数是否带来增益”的方法。适合正在研究 prompt 内部机制、设计记忆增强模型或者准备把循环推理模块接入自己项目的开发者参考。1. 普通 ICL 的局限与循环潜在推理的出现1.1 普通 ICL 为什么不足以完成复杂任务先定义一个普通 ICL 的抽象过程。给定上下文 C 包含若干条示例和一条查询 q模型输出 y。在 Transformer 中这一步通常执行y Decoder(Transformer_Encoder(Embed([C; q])))所有信息交流都发生在若干层 self-attention 中。层数是固定的注意力权重只通过一次前向传播计算没有额外的“迭代计算”机制。对于简单任务比如让模型按照上下文格式输出“姓名张三”一次前向传播完全够用。但对复杂任务比如从多条数学示例中总结规则或者从多次状态转移中推断下一步模型需要把不同示例中的关键信息相互比较、消歧、再和当前查询结合。普通 Transformer 是一次性完成这些交互缺少显式中间状态。从另一个角度看ICL 的能力上限受上下文长度、注意力模式和层数共同影响。上下文越长token 之间的距离越远注意力要覆盖的信息越多层数越少可以完成的非线性变换越少。一旦一次前向传播无法完成信息综合模型就退化成“模仿 prompt 格式”而不是真正执行推理。这也解释了很多场景下 ICL 表现不稳定的原因模型可能记住了格式但没有形成对任务规则的可靠内部表示。1.2 循环潜在推理把前向传播看成可迭代计算“循环潜在推理”可以理解为模型在内部维护一个潜在状态 h先在隐空间中执行 T 次迭代更新最后通过解码头输出。更新过程用公式表示h^(0) Q(x) for t 1..T: h^(t) g(h^(t-1), Memory(C)) y Decode(h^(T))这里的 Memory(C) 是上下文编码后的记忆Q(x) 是 query 的编码g 可以是注意力、MLP 或更复杂的模块T 是循环步数。相比普通前向传播循环潜在推理的核心区别是计算量不再只由层数决定还受 T 控制。因此模型可以在相同参数下对同一个 query 做更多轮推理。这也是很多“深度思考”类方法的设计动机。T 越大潜在状态可以和上下文反复交互但也会带来梯度路径变长、训练不稳定、推理变慢等问题。所以不能盲目增加 T需要结合具体任务验证哪个步数性价比最高。1.3 BDH-CQ 想解决的问题BDH-CQ 的核心假设是不同上下文示例之间会产生记忆干扰。如果所有示例都放在同一个向量空间里无差别交互模型很难区分哪些信息属于任务规则、哪些只是干扰项。于是它使用 Block-Diagonal Hidden State把隐状态切成若干块每个块负责相对独立的上下文记忆再通过 Context Query 控制当前 query 如何读取这些块。这样可以降低块间干扰同时让 query 在多次循环中聚焦到与当前问题最相关的块。在工程实现上BDH-CQ 并不一定指某个唯一的模型结构而是一类“块状记忆 循环查询”的设计。下面用简化版本把这个机制完整实现出来并放到一个能控制难度的合成任务上验证。2. 环境准备与最小项目结构2.1 实验目标本实验的目标是验证两个问题循环潜在推理能不能比单次前向传播获得更高的 few-shot 准确率。在上下文示例数量增加时循环模块是否能更充分利用新增信息。为了让结果可解释不使用标准 NLP 数据集而是构造合成 few-shot 任务。合成任务可以精确控制规则和噪声也方便对比不同循环步数、不同块数的效果。2.2 Python 环境与依赖推荐使用 Anaconda 或 venv 创建独立环境。建议 Python 3.10 或更高版本PyTorch 2.1 或更高版本。以下命令创建一个名为 bdh-cq 的环境conda create -n bdh-cq python3.10 -y conda activate bdh-cq pip install torch2.1.2 numpy scikit-learn安装完成后验证 PyTorch 是否可用python -c import torch; print(torch.__version__)如果输出类似2.1.2cu121说明环境正常。如果使用 CPU 环境也可以安装 CPU 版本本实验数据量小CPU 训练完全够用。依赖清单如下组件建议版本用途Python3.10运行环境PyTorch2.1模型训练与张量运算NumPy1.24数据生成辅助scikit-learn1.3可选用于计算指标实际项目中如果已有虚拟环境可以不额外创建。但建议保证版本一致避免因 API 变化导致代码无法运行。2.3 项目文件结构用以下结构组织代码bdh-cq/ config.py # 超参数配置 dataset.py # 合成 few-shot 数据构造 model.py # BDH-CQ 层和 few-shot 模型 train.py # 训练与保存 checkpoint eval.py # 评估不同 shots 和 steps 的效果config.py 中的默认配置class Config: vocab_size 64 d_model 64 num_blocks 4 steps 3 batch_size 64 lr 1e-3 epochs 30 max_grad_norm 1.0这些参数会在后续各节详细解释。现在先进入数据集和模型实现。3. 实现 BDH-CQ 最小原型3.1 构造一个能验证 ICL 规律的合成任务为了让模型必须依赖上下文而不是记住固定映射任务设计成每个 batch 随机生成一个偏移量 offset上下文由若干键值对组成查询键是上下文中没有出现过的 key正确答案是(query_key offset) % vocab_size。模型需要从上下文样例中推出当前 batch 的 offset才能正确预测。下面是一个数据构造函数import random import torch def make_batch(vocab_size, num_shots, batch_size, seedNone): if seed is not None: random.seed(seed) ctx_tokens [] query_tokens [] labels [] masks [] for _ in range(batch_size): offset random.randint(1, vocab_size // 2) keys random.sample(range(1, vocab_size - 1), num_shots) context [] valid [] for k in keys: v (k offset) % vocab_size if v 0: v 1 context [k, v] valid [1, 1] qk random.choice([x for x in range(1, vocab_size - 1) if x not in keys]) qv (qk offset) % vocab_size if qv 0: qv 1 context.append(qk) valid.append(1) ctx_tokens.append(context) query_tokens.append(qk) labels.append(qv) masks.append(valid) max_len max(len(c) for c in ctx_tokens) for i in range(len(ctx_tokens)): pad_len max_len - len(ctx_tokens[i]) ctx_tokens[i] [0] * pad_len masks[i] [0] * pad_len return ( torch.tensor(ctx_tokens), torch.tensor(query_tokens), torch.tensor(labels), torch.tensor(masks, dtypetorch.bool), )这里把真实 token 从 1 开始编号0 保留给 padding。query key 刻意不出现在上下文中模型无法通过简单复制答案完成预测必须从k - (koffset)的对应关系中总结规律。3.2 BDH-CQ 层实现BDH-CQ 层负责执行循环潜在推理。输入是上下文编码ctx和查询编码q输出是更新后的查询表示。关键操作包括使用 query 与上下文做点积注意力读取当前最相关的上下文内容。把隐藏状态切成多个块每个块用独立 MLP 更新。重复 T 次上述过程形成循环潜在推理。import torch import torch.nn as nn class BDH_CQLayer(nn.Module): def __init__(self, d_model, num_blocks, steps): super().__init__() self.d_model d_model self.num_blocks num_blocks self.steps steps assert d_model % num_blocks 0 self.block_dim d_model // num_blocks self.block_nets nn.ModuleList([ nn.Sequential( nn.Linear(self.block_dim * 2, self.block_dim * 2), nn.ReLU(), nn.Linear(self.block_dim * 2, self.block_dim), ) for _ in range(num_blocks) ]) self.norm nn.LayerNorm(d_model) self.cq_proj nn.Linear(d_model, d_model) def forward(self, ctx, q, maskNone): # ctx: [batch, seq_len, d_model] # q: [batch, d_model] h q for _ in range(self.steps): # 1. 用 query 读取上下文 attn_logits torch.matmul(ctx, h.unsqueeze(-1)).squeeze(-1) if mask is not None: attn_logits attn_logits.masked_fill(~mask, -1e9) attn torch.softmax(attn_logits / (self.d_model ** 0.5), dim-1) ctx_vec torch.matmul(attn.unsqueeze(1), ctx).squeeze(1) # 2. 按块更新潜在状态 h_blocks h.view(-1, self.num_blocks, self.block_dim) c_blocks ctx_vec.view(-1, self.num_blocks, self.block_dim) inp torch.cat([h_blocks, c_blocks], dim-1) outs [] for i, net in enumerate(self.block_nets): outs.append(net(inp[:, i])) h_new torch.stack(outs, dim1).view_as(h) # 3. 残差 LayerNorm h self.norm(h h_new) return self.cq_proj(h)这段代码有几个关键点点积注意力把 query 作为查询向量从上下文中检索相关信息。块更新通过view(-1, num_blocks, block_dim)完成每个块只使用自己的 MLP模拟 block-diagonal 的权重大结构。残差连接和 LayerNorm 用于稳定循环训练。steps控制内部迭代次数是循环潜在推理的核心。3.3 完整模型、训练和评估把 embedding、BDH-CQ 层和分类头组合起来import torch.nn.functional as F class FewShotICLModel(nn.Module): def __init__(self, config): super().__init__() self.embed nn.Embedding(config.vocab_size, config.d_model) self.bdh_cq BDH_CQLayer(config.d_model, config.num_blocks, config.steps) self.head nn.Linear(config.d_model, config.vocab_size) def forward(self, ctx, q, maskNone): ctx_emb self.embed(ctx) q_emb self.embed(q) h self.bdh_cq(ctx_emb, q_emb, mask) return self.head(h)训练循环使用交叉熵损失和 AdamW 优化器。为了处理循环展开带来的梯度波动增加梯度裁剪def train_step(model, optimizer, batch): ctx, q, label, mask batch logits model(ctx, q, mask) loss F.cross_entropy(logits, label) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() return loss.item()评估时统计预测准确率torch.no_grad() def evaluate(model, batch): ctx, q, label, mask batch logits model(ctx, q, mask) pred logits.argmax(dim-1) return (pred label).float().mean().item()这里模型使用了全局 embedding可能也会学到一些词表统计信息。但因为 offset 每个 batch 随机变化固定映射无法解决所有情况模型只能依赖上下文。4. 运行验证循环步数是否带来增益4.1 训练脚本入口训练命令可以直接传入关键参数。示例python train.py --epochs 30 --steps 3 --num_blocks 4 --d_model 64训练过程会打印每个 epoch 的 loss 和验证准确率。建议在实验开始前用固定随机种子固定数据顺序保证可复现。4.2 验证不同 shots 和 steps为了回答“循环步数有没有用”需要固定其他条件只改变steps。分别训练steps1和steps3的模型然后在不同 shot 数量下评估。评估代码思路如下shots_list [1, 2, 4, 8] for steps in [1, 3]: model FewShotICLModel(config) train(model, config, stepssteps) for shots in shots_list: acc evaluate_with_shots(model, shots) print(fsteps{steps}, shots{shots}, acc{acc:.4f})实验时最好在每个配置上使用 3 到 5 个随机种子报告均值与标准差。否则单次结果波动较大容易误判。4.3 典型结果解读下面是一张示意结果表不是任何正式 benchmark 数据只用于说明常见趋势循环步数shots1shots2shots4shots8steps10.550.630.700.74steps30.600.710.800.86在这个合成任务上常见的观察是随着 shots 增加准确率整体上升说明模型确实在利用更多上下文。steps3 在 shots 更多时优势更明显。shots1 时循环推理优势有限因为上下文只包含一个样例信息本身不足。如果你在自己的实验中看到 steps1 和 steps3 差别很小可以检查任务是否太简单。需要把 offset 的随机范围加大或者减少 embedding 的编码能力才更容易体现循环推理的价值。4.4 结果观察点验证时需要同时观察训练 loss、验证准确率和梯度范数。如果 loss 持续下降但验证准确率不涨大概率是模型记住了训练分布没有真正利用上下文。此时可以增加新 batch 的 offset 随机性或降低词表大小减少全局记忆空间。5. 核心超参数与训练稳定性5.1 参数速查表超参数含义实验建议调大影响调小影响d_model隐层维度32 到 128表达更强显存增加可能欠拟合num_blocks块数4 到 8减少块间干扰每块维度变小块间干扰增大steps循环步数2 到 4更多迭代训练更慢推理能力不足lr学习率5e-4 到 1e-3训练不稳定收敛变慢batch_size批大小32 到 128梯度更稳定显存增加梯度噪声大max_grad_norm梯度裁剪1.0无可能影响收敛5.2 循环步数不是越大越好循环步数 T 是 BDH-CQ 最重要的超参数之一。T 越大潜在状态可以迭代更多次理论上推理能力更强。但实际中T 过大会导致梯度路径过长容易出现梯度消失或梯度爆炸。训练耗时线性增加。模型可能对训练数据过拟合在测试集上反而下降。推理阶段延迟增加生产环境不可接受。建议从steps2或steps3开始观察验证集趋势再逐步增加。如果steps3比steps1没有明显提升不需要继续调大问题更可能出在任务设计或数据质量上。5.3 块数与块维度num_blocks控制隐藏状态被切成多少块。块数越多每块维度越小块与块之间的信息隔离越强。这种设计有利于减少上下文示例之间的干扰但也会限制每个块单独的表达能力。对于d_model64块数选择 4 或 8 比较合适。若num_blocks8每块维度为 8MLP 参数量会明显减少可能需要配合更深层或更宽的 block 网络。反过来如果num_blocks1就退化成普通全局 MLP 更新失去了 block-diagonal 的意义。5.4 训练稳定性措施循环展开模型训练时稳定性比普通前向模型更关键。推荐以下几点使用残差连接并保证每个循环内至少有一个归一化层。使用梯度裁剪裁剪阈值常设为 1.0。学习率不要一次性设太大建议配合 warmup。观察梯度范数日志如果从正常范围突然变成NaN优先检查数据和 mask。如果显存受限可以使用梯度检查点降低显存占用。示例代码torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)这段代码放在loss.backward()之后、optimizer.step()之前。6. 常见问题排查6.1 典型报错与解决问题现象可能原因检查方式处理建议loss 为 NaN学习率过大或数据里有 padding token 被当作真实标签打印 loss、检查 label 是否为 0降低 lr调整数据生成逻辑避免 0 token加入循环后准确率反而下降steps 过大、过拟合、梯度不稳定