about
The Expressive Power of Transformers with Chain of Thought (Revised) (arxiv.org)
20 points by yawboakye on Mar 24, 2024 | hide | past | pdf | 1 comment on HN

In plain words: Letting a transformer write intermediate steps before answering extends what it can compute; the gain depends on how many it writes. With steps proportional to input length it handles every regular language, the checks a state machine does, which immediate-answer transformers provably cannot.

Abstract · The Expressive Power of Transformers with Chain of Thought

Recent theoretical work has identified surprisingly simple reasoning problems, such as checking if two nodes in a graph are connected or simulating finite-state machines, that are provably unsolvable by standard transformers that answer immediately after reading their input. However, in practice, transformers' reasoning can be improved by allowing them to use a "chain of thought" or "scratchpad", i.e., generate and condition on a sequence of intermediate tokens before answering. Motivated by this, we ask: Does such intermediate generation fundamentally extend the computational power of a decoder-only transformer? We show that the answer is yes, but the amount of increase depends crucially on the amount of intermediate generation. For instance, we find that transformer decoders with a logarithmic number of decoding steps (w.r.t. the input length) push the limits of standard transformers only slightly, while a linear number of decoding steps, assuming projected pre-norm (a slight generalization of standard pre-norm), adds a clear new ability (under standard complexity conjectures): recognizing all regular languages. Our results also imply that linear steps keep transformer decoders within context-sensitive languages, and polynomial steps with generalized pre-norm make them recognize exactly the class of polynomial-time solvable problems -- the first exact characterization of a type of transformers in terms of standard complexity classes. Together, this provides a nuanced framework for understanding how the length of a transformer's chain of thought or scratchpad impacts its reasoning power.

William Merrill, Ashish Sabharwal
arXiv:2310.07923 · cs.LG, cs.CC, cs.CL, cs.LO · submitted Oct 11, 2023 · updated Apr 11, 2024
abstract · pdf · html · 9-page preprint. ICLR camera ready posted April 11

add comment on HN

I struggle heavily with complexity theory, and in particular struggle to follow the definitions, assumptions, and proofs. It can be very hard for me to decide whether to invest in the time and effort needed to decode a paper like this.

Here we have a number of assumptions, and very short proofs, which makes me a little worried that the CoT complexity class might not be useful for understanding the expressive power of actually existing transformers like LLMs - especially ones with hidden weights and architectures.

Can someone with a better understanding remark on this?

Some questions:

1. Does this paper let us better reason about trained hidden weight llms like gpt3.5/gpt4?

2. Does this paper imply that we can always exploit some linear CoT blowup factor to faithfully simulate an automaton, even on a very small transformers, like a 1.5B LLM instance?

3. Does this result attain for untrained or minimally trained transformers? Assuming that's not the case, what assumptions do we make about the network's training?

4. How would we bound the automaton simulation blow up factor for a given llm instance? Can we confirm bounds empirically by benchmarking transformer instances against a corpus of automaton specifications of arbitrary size?

Feedback from a expert on this subject would be really awesome!