One-Layer Transformers Provably Learn In-Context K-Nearest Neighbor Prediction with Chain-of-Thought
Lyumin Wu ⋅ Yuan Cao
Abstract
Chain-of-Thought (CoT) enables transformers to solve complex reasoning tasks by generating intermediate reasoning steps. Despite strong empirical success, the mechanisms by which such reasoning abilities emerge during training remain poorly understood, particularly from a theoretical perspective. In this work, we study this question in a stylized yet fully analyzable setting: in-context $\mathrm{K}$-nearest-neighbor ($\mathrm{K}$-NN) prediction with a one-layer transformer. Earlier work has demonstrated that one-layer transformers can be trained to perform $1$-NN without CoT \citep{Li2024One} by having the softmax attention attend to the nearest neighbor in context. However, this result on $1$-NN cannot be extended to $\mathrm{K}$-NN for $K>1$, as softmax attention is much less naturally suited to attending to the 2nd through $K$-th nearest neighbors. We first give a rigorous negative result showing that this limitation already appears on a well-separated class of binary classification tasks for odd $K>1$. We then give an explicit construction of a one-layer transformer, and show that it can solve $\mathrm{K}$-NN via CoT reasoning. We further prove that, under a stylized training setup, gradient descent can recover this construction, thereby showing that transformers can acquire the capability to solve $\mathrm{K}$-NN through training. In addition, we show that the resulting trained model applies to a broader test class than that assumed in the training analysis. Our results shed light on how CoT expands the ability of transformers to learn and solve tasks that are otherwise hard to solve.
Chat is not available.
Successful Page Load