#1013·mamba

支持初始/最终循环状态的训练 API (TBPTT/分块序列训练)

作者: jkxdl创建于 2026年8月7日更新于 2026年8月12日

Hi, 感谢您出色的 Mamba 实现。 我正在使用 Mamba 来实现实时机器人控制策略。 在推理过程中,模型逐步运行,并在整个训练过程中持续维护一个递归/卷积状态(隐藏/缓存状态): obs_0 -> state_1 obs_1 + state_1 -> state_2 ... obs_t + state_t -> action_t, state_{t+1} 然而,训练 API 只接受完整序列,并且似乎在内部初始化了递归状态。 它没有暴露与推理缓存/状态 API 兼容的初始状态输入或最终状态输出。 这使得高效的分块训练变得困难: 在完整训练集上进行训练,可以保持与推理相同的连续状态行为,但训练集的长度高度可变,导致多 GPU 加载不平衡。 在独立固定长度的分块上进行训练,会在每个分块边界重置状态,从而导致训练和推理不一致,因为推理状态在一个训练集内从未重置。 截断 BPTT 将是一个很好的解决方案,但它需要在分块之间传递递归状态,同时将其从自动微分图中分离。 是否可以支持以下类似的 API: output, final_state = model( x, initial_state=state, return_final_state=True ) 其中: initial_state 使用与增量推理缓存/状态相同的表示; final_state 可以传递到下一个训练分块; 调用者可以在分块之间使用 final_state = detach(final_state) 进行 TBPTT; 结果与在一次前向传播中处理连接的序列数值上一致,除了浮点数差异之外。 例如: state = None for chunk in episode_chunks: output, state = model( chunk, initial_state=state, return_final_state=True ) loss = compute_loss(output) loss.backward() state = detach_state(state) 这将支持固定长度的分块训练、标记/帧均衡的分布式批量以及更好的 GPU 利用率,同时保持在推理时使用的连续状态行为。 如果此功能已通过现有的低级 API 或缓存对象支持,我可能会错过,那么一个澄清也会非常有帮助。 再次感谢。

内容来源: state-spaces/mamba