Learning Dynamics of Chain-of-Thought State Tracking in a Solvable Transformer Model
Abstract
Chain-of-thought generation can turn a multi-step computation into a sequence of locally checkable state updates, yet the training dynamics of such problems remain poorly understood. We study this question in a solvable setting: a simplified one-block transformer trained by supervised next-token prediction on state sequences generated by composing permutations. Using a statistical-physics mean-field description, we derive dynamics for three order parameters measuring attention retrieval and on/off-target logic overlaps. These equations match simulations for the order parameters, and they predict a sharp transition in final rollout accuracy. The analysis reveals staged learning: the MLP logic first learns a mixed heuristic and then attention locks onto the relevant tokens.