Transformers Learn Shortcuts to Automata

Bingbin Liu, Jordan T. Ash, Surbhi Goel, Akshay Krishnamurthy, Cyril Zhang

Introduction

Modern deep learning systems demonstrate increasing capabilities of algorithmic reasoning. Particularly in modalities such as natural language, math, and code, neural networks can successfully parse and synthesize sequences containing symbolic information and compositional structure. To exhibit these functionalities, these networks are required to learn and execute the relevant discrete algorithms within their internal representations. A core open question in this domain is that of mechanistic understanding: how do neural networks encode the primitives of algorithmic reasoning?

When considering this question, there is an apparent mismatch between classical sequential models of computation (e.g. Turing machines) and the Transformer, the state-of-the-art architecture for neural algorithmic reasoning. If we are to think of algorithms as sequentially-executed computational rules, why should we use a shallowIn practice, a Transformer’s context length (which can be thousands of tokens) is typically far greater than its depth (which can be as small as 66). and non-recurrent architecture to represent them?

We study this question through the lens of semiautomata, which compute state sequences q1,…,qTq_{1},\ldots,q_{T} from inputs σ1,…,σT\sigma_{1},\ldots,\sigma_{T} by application of a recurrent transition function δ\delta (and initial state q0q_{0}):

Semiautomata describe the underlying dynamics of automata, which are simply semiautomata equipped with mappings from states to outputs. With unbounded state spaces, automata can represent all algorithms; however, even bounded automata form a rich class of sequence processing algorithms, containing regular expression parsers and finite-state transducers. In reinforcement learning and control, semiautomata are better known as deterministic Markov models (where σt\sigma_{t} are actions); thus, in addition to algorithmic reasoning, the results in this work also pertain to Transformer dynamics models.

We perform a theoretical and empirical investigation of whether (and how) non-recurrent Transformers perform the computations of semiautomata. We find that Transformers learn shortcut solutions, which correctly and efficiently simulate the sequential transitions of semiautomata using a shallow parallel circuit, rather than naively iterating the single-step recurrence. Shortcuts arise from hierarchical reparameterizations of a semiautomaton’s global transition dynamics.

Our contributions. Our theoretical results provide structural guarantees for the representability of semiautomata (thus, iterative algorithms) by one pass through a shallow, non-recurrent Transformer. In particular, we show that:

Shortcut solutions, with depth logarithmic in the sequence length, always exist (Theorem 1).

Constant-depth shortcuts exist for solvable semiautomata (Theorem 2). These are understood via the Krohn-Rhodes theorem, a landmark result in semigroup theory. Conversely, there do not exist constant-depth shortcuts for non-solvable semiautomata, unless TC0=NC1\mathsf{TC}^{0}=\mathsf{NC}^{1} (Theorem 4).

For a natural class of semiautomata corresponding to path integration in a “gridworld” with boundaries, we show that there are even shorter shortcuts (Theorem 3), beyond those guaranteed by the general structure theorems above.

We accompany these with an extensive set of experimental findings:

Transformers learn shortcuts with standard training (Section 4). Across a wide variety of semiautomaton simulation problems, we find that shallow non-autoregressive Transformers successfully learn shortcut solutions: despite the non-convex optimization problem, gradient-based training works. This suggests that shortcuts are plausible mechanisms for algorithmic reasoning in non-synthetic sequence models, and lies beyond our current theoretical understanding.

Shortcuts are statistically brittle (Section 5). We identify empirical weaknesses of the shortcuts found by Transformers: poor out-of-distribution generalization (including to unseen sequence lengths) and worse performance than RNNs under limited supervision. Toward mitigating these drawbacks and obtaining the best of both worlds, we show that with recency-biased scratchpad training, autoregressive Transformers can easily be guided to learn the iterative RNN-like solutions (chain-of-thought generation).

Emergent reasoning in neural sequence models. Neural sequence models, both recurrent (Wu et al., 2016, Peters et al., 2018, Howard and Ruder, 2018) and non-recurrent (Vaswani et al., 2017, Devlin et al., 2018), have become an era-defining tool for parsing and transducing data with combinatorial structure, such as natural language and code. A nascent frontier is to leverage neural dynamics models, again both recurrent (Hafner et al., 2019) and non-recurrent (Chen et al., 2021a, Janner et al., 2021), for decision making. At the highest level, the present work seeks to understand the mechanisms by which these models perform combinatorial and algorithmic reasoning.

Computational models within neural networks. Despite the preponderance of empirical successes, many mysteries remain, towards understanding the internal mechanisms of neural networks capable of algorithmic reasoning. It is known that self-attention realizes low-complexity circuits (Hahn, 2020, Elhage et al., 2021, Merrill et al., 2021, Edelman et al., 2022), declarative programs (Weiss et al., 2021), and Turing machines (Dehghani et al., 2019, Pérez et al., 2021, Giannou et al., 2023). Interpretable symbolic computations have been extracted from trained models (Clark et al., 2019, Vig, 2019, Tenney et al., 2019, Wang et al., 2022). Our conclusions are closest to the literature on the universal representation on Turing machines (which are automata with unbounded states); however, our work is unique in characterizing the recurrent machines whose execution loops can be efficiently unrolled into a single pass of a shallow Transformer.

Learning elementary algorithms with Transformers. Our work provides a unifying lens on many recent investigations on whether (and how) Transformers represent certain classes of fundamental algorithmic computations. These include bounded-depth Dyck languages (Yao et al., 2021), modular prefix sums (Anil et al., 2022), adders (Nogueira et al., 2021, Nanda and Lieberum, 2022), regular languages (Bhattamishra et al., 2020), and sparse logical predicates (Edelman et al., 2022, Barak et al., 2022), which are all special cases of simulating finite-state automata. Thus, our work provides guarantees of shallow Transformer solutions in all of these settings. Zhang et al. (2022) empirically analyze the behavior and inner workings of Transformers on random-access group operations and note “shortcuts” (which skip over explicit program execution) similar to those we study.

We provide an expanded discussion of related work in Appendix A.5.

Preliminaries

A semiautomaton A:=(Q,Σ,δ){\mathcal{A}}:=(Q,\Sigma,\delta) consists of a set of states QQ, an input alphabet Σ\Sigma, and a transition function δ:Q×Σ→Q\delta:Q\times\Sigma\rightarrow Q. In this work, QQ and Σ\Sigma will always be finite sets. For all positive integers TT and a starting state q0∈Qq_{0}\in Q, A{\mathcal{A}} defines a map from input sequences (σ1,…,σT)∈ΣT(\sigma_{1},\ldots,\sigma_{T})\in\Sigma^{T} to state sequences (q1,…,qT)∈QT(q_{1},\ldots,q_{T})\in Q^{T}: qt:=δ(qt−1,σt)q_{t}:=\delta(q_{t-1},\sigma_{t}) for t=1,…,Tt=1,\ldots,T. This is a deterministic Markov model, in the sense that at time tt, the future states qt+1,…,qTq_{t+1},\ldots,q_{T} only depend on the current state qtq_{t} and the future inputs σt+1,…,σT\sigma_{t+1},\ldots,\sigma_{T}.

We define the task of simulation: given a semiautomaton A\mathcal{A}, starting state q0q_{0}, and input sequence (σ1,…,σT)(\sigma_{1},\ldots,\sigma_{T}), output the state trajectory AT,q0(σ1,…,σT):=(q1,…,qT){\mathcal{A}}_{T,q_{0}}(\sigma_{1},\ldots,\sigma_{T}):=(q_{1},\ldots,q_{T}). Let f:ΣT→QTf:\Sigma^{T}\rightarrow Q^{T} be a function (which in general can depend on A,T,q0\mathcal{A},T,q_{0}). We will say that ff simulates AT,q0\mathcal{A}_{T,q_{0}} if f(σ1:T)=AT,q0(σ1:T)f(\sigma_{1:T})=\mathcal{A}_{T,q_{0}}(\sigma_{1:T}) for all input sequences σ1:T\sigma_{1:T}. Finally, for a positive integer TT, we say that a function class F\mathcal{F} of functions from ΣT→QT\Sigma^{T}\rightarrow Q^{T} simulates A\mathcal{A} at length TT if, for each q0∈Qq_{0}\in Q, there is a function in F\mathcal{F} which simulates AT,q0\mathcal{A}_{T,q_{0}}

Every semiautomaton induces a transformation semigroup T(A){\mathcal{T}}(\mathcal{A}) of functions ρ:Q→Q\rho:Q\rightarrow Q under composition, generated by the per-input-symbol state mappings δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\rightarrow Q. When T(A){\mathcal{T}}(\mathcal{A}) contains the identity function, it is called a transformation monoid. When all of the functions are invertible, T(A){\mathcal{T}}(\mathcal{A}) is a permutation group. See Figure 1 for some examples which appear both in our theory and experiments; additional background (including a self-contained tutorial on the relevant concepts in finite group and semigroup theory) is provided in Appendix A.2. An elementary example is a parity counter (leftmost in Figure 1): the state is a bit, and the inputs are {“toggle the bit”, “do nothing”}; the transformation semigroup is C2C_{2}, the cyclic group of order 2.

2 Recurrent and non-recurrent neural sequence models

An LL-layer (or depth-LL) Transformer is a different sequence-to-sequence network, consisting of alternating self-attention blocks and feedforward MLP blocks

Typical Transformers are shallow, in the sense that L≪TL\ll T. In practice, this makes their inference and gradient computations highly parallelizable, with the number of sequential computation steps scaling linearly in LL, while RNNs require a scaling linear in TT. While there is always a canonical way for RNNs to simulate semiautomata, the answer to the analogous question for shallow Transformers is far less obvious, and forms the main topic of this paper.

Theory: shortcuts abound

To simulate a semiautomaton at length TT, a TT-layer Transformer can implement the same sequential solution as an RNN: let the tt-th layer embed the state transition qt−1↦qtq_{t-1}\mapsto q_{t}. We define shortcuts as solutions which implement the same functionality with a significantly smaller depth.

Let A{\mathcal{A}} be a semiautomaton. For every T≥1T\geq 1, let fTf_{T} be a sequence-to-sequence neural network which simulates A{\mathcal{A}} at length TT. Then, we call this sequence {fT}T≥1\{f_{T}\}_{T\geq 1} a shortcut to A{\mathcal{A}} if the sequence of network depths D:={D(fT)}T≥1D:=\{D(f_{T})\}_{T\geq 1} satisfies D≤o(T)D\leq o(T).

By this definition, shortcuts are quite general, and some are less interesting than others. For example, it is always possible to construct a constant-depth neural network which memorizes all ∣Σ∣T|\Sigma|^{T} values of AT,q0{\mathcal{A}}_{T,q_{0}}, but these networks must be exceptionally wide. There are also solutions which emulate transitions in “chunks”, letting each of (say) T\sqrt{T} layers perform T\sqrt{T} consecutive state transitions; however, without exploiting the structure of the semiautomaton, this would require width Ω(∣Σ∣T)\Omega(|\Sigma|^{\sqrt{T}}). To rule out these cases and focus on interesting shortcuts for Transformers, we want the other size parameters (attention and MLP width) to be small: say, scaling at most polynomially in TT, ∣Q∣|Q|, and ∣Σ∣|\Sigma|. To construct such shortcuts, we need ideas beyond explicit iteration of state transitions.

We begin by noting that polynomial-width shortcuts always exist. This may seem counterintuitive if we restrict ourselves to viewing a Transformer’s intermediate activations as representations of states qtq_{t}, like the RNN solution. Instead, a Transformer can encode and hierarchically compose transformations δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\rightarrow Q (see Figure 2a), leading to far shallower solutions:

Transformers can simulate all semiautomata A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta) at length TT, with depth O(log⁡T)O(\log T), embedding dimension O(∣Q∣)O(|Q|), attention width O(∣Q∣)O(|Q|), and MLP width O(∣Q∣2)O(|Q|^{2}).

This is proven in Appendix C.2, and leverages the ability of a self-attention head to approximate hard attention (i.e. concentrate its mixing weights on a single position). However, self-attention heads can also perform soft attention (i.e. depend on a large number of previous positions), enabling even shallower implementations of certain sequential computations. For example, the parity automaton can be simulated by a single Transformer layer (see Lemma 6): soft attention computes prefix sums in parallel, then the MLP computes “mod 2”. This leads to a significantly more nuanced question: when are there even shallower shortcuts? At first glance, such solutions may seem rare, and specialized to simple cases such as parity.

Our resolution to this question comes from the Krohn-Rhodes decomposition theorem (Krohn and Rhodes, 1965), a landmark result which vastly generalizes the uniqueness of prime integer factorizations, and created the mathematical field of algebraic automata theory (Rhodes et al., 2010). The conclusion is quite unintuitive: allowing for both hard and soft modes of attention, constant-depth shortcuts are surprisingly common!

Transformers can simulate all solvableSee Definition 6. Intuitively, the only obstructions are when the semiautomata contain non-solvable groups such as S5S_{5}, the group of all permutations of 55 elements. semiautomata A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta), with depth O(∣Q∣2log⁡∣Q∣)O(|Q|^{2}\log|Q|), embedding dimension 2O(∣Q∣log⁡∣Q∣)2^{O(|Q|\log|Q|)}, attention width 2O(∣Q∣log⁡∣Q∣)2^{O(|Q|\log|Q|)}, and MLP width ∣Q∣O(2∣Q∣)+2O(∣Q∣log⁡∣Q∣) ⁣ ⋅ ⁣ T|Q|^{O(2^{|Q|})}+2^{O(|Q|\log|Q|)}\!\,\cdot\!\,T.

Much of the appendix is dedicated to providing a user-friendly exposition of the relevant algebraic concepts, culminating in the proof of Theorem 2. We provide a few high-level notes below:

Intuitively (illustrated in Figures 2b and 2c), the Krohn-Rhodes decomposition “factorizes” every solvable semiautomaton into modular counters and memory units, glued together via a feedforward cascade product (Definition 4) whose depth only depends on ∣Q∣|Q|, not TT. These two types of “prime” semiautomata can be efficiently simulated by depth-11 Transformers.

The decomposition depends on the transformation semigroup T(A)\mathcal{T}(\mathcal{A}). It is non-constructive (much like how the existence of prime factorizations doesn’t entail a procedure to find them). Computationally, these solutions still have to be found by a search procedure. Remarkably, we find that gradient-based training succeeds empirically, despite the worst-case computational hardness of related problems.

What makes Transformers special (vs. other universal function approximators)? The same underlying semiautomaton-to-circuit constructions could be applied to any universal function approximator (like a vanilla MLP with the same depth). Transformers embed all of the constructions in Theorems 1 and 2 with exceptional efficiency, in terms of the network complexity measures discussed in Section 2. Most importantly, the constructions leverage Transformers’ positional weight sharing, which removes allThe only part of Theorem 2 requiring TT non-tied neurons is the implementation of mod-pp gates. It disappears entirely if we can add auxiliary MLP neurons with periodic activation functions such as x↦sin⁡(x)x\mapsto\sin(x). suboptimal TT factors from the parameter count. We discuss this further in Appendix A.5.

The proof builds a parallel nearest boundary detector for the two boundary (i.e. leftmost and rightmost) states, and can be found in Appendix C.4. We note that gridworlds are known to be extremal cases for the holonomy decomposition in Krohn-Rhodes theory (Maler (2010) discusses this, calling it the elevator automaton). It would be interesting to generalize our improvement and characterize the class of problems for which self-attention affords O(1)O(1) instead of poly(∣Q∣)\text{poly}(|Q|)-depth solutions.

2 Lower bounds

Can Theorem 2 be improved to handle non-solvable semiautomata? (Equivalently: can Theorem 1 be improved to constant depth?) It turns out that as a consequence of a classic result in circuit complexity (Barrington, 1986), this question is equivalent to the major open question of TC0=?NC1\mathsf{TC}^{0}\stackrel{{\scriptstyle?}}{{=}}\mathsf{NC}^{1} (thus: conjecturally, no). Unless these complexity classes collapse, Theorems 1 and 2 are optimal. In summary, simulating non-solvable semiautomata with constant depth is provably hard:

Let A\mathcal{A} be a non-solvable semiautomaton. Then, for sufficiently large TT, no O(log⁡T)O(\log T)-precision Transformer with depth independent of TT and width polynomial in TT can continuously simulate A\mathcal{A} at length TT, unless TC0=NC1\mathsf{TC}^{0}=\mathsf{NC}^{1}.

This is proven in Appendix C.5. The smallest example of a non-solvable semiautomaton has ∣Q∣=60|Q|=60 states, whose transitions generate A5A_{5} (all of the even permutations).

Finally, we note that although our width bounds might be improvable, an exponential-in-∣Q∣|Q| number of hypotheses (and hence a network with poly(∣Q∣)\text{poly}(|Q|) parameters) is unavoidable if one wishes to learn an arbitrary ∣Q∣|Q|-state semiautomaton from data: there are ∣Q∣∣Q∣⋅∣Σ∣|Q|^{|Q|\cdot|\Sigma|} of them, which generate ∣Q∣Ω(∣Q∣2)|Q|^{\Omega(|Q|^{2})} distinct semigroups (Kleitman et al., 1976). If we wish to study how machine learning models can efficiently identify large algebraic structures, we will need finer-grained inductive biases to specify which semiautomata to prefer, a direction for future work.

Experiments: can SGD find the shortcuts?

The results in Section 3 only provide a precise understanding of representability: they show that shortcut solutions exist within the parameter space of a shallow Transformer. To understand whether Transformers can actually learn these shortcuts from data, we must introduce the additional considerations of generalization and optimization. It is notoriously difficult to derive meaningful analyses which account for all of these factors in deep learning; thus, we do not attempt to do so in this work.Our bounds on the parameter count and weight norms do imply classical generalization bounds for appropriately norm-constrained Transformers (Edelman et al., 2022), but these are too coarse-grained to provide non-vacuous predictions of generalization behavior. Instead, we approach the end-to-end question with an empirical lens: trained on sequences arising from a variety of automata, does a shallow (depth-L≪TL\ll T) Transformer converge to correct simulators of these automata?

For a selection of 19 semiautomata corresponding to various groups and semigroups (detailed descriptions in Appendix B.1.1), we train shallow Transformer (GPT-2-like (Radford et al., 2019)) models to map randomly sampled sequences (σ1,…,σT)(\sigma_{1},\ldots,\sigma_{T}) to their corresponding state sequences (q1,…,qT)(q_{1},\ldots,q_{T}), and evaluate their accuracy on held-out sequences. We vary the depth LL from 11 to 1616, and use freshly-sampled sequences of length T=100T=100. In this setup, the number of sequences encountered during training (≤106\leq 10^{6}) is far smaller than the number of distinct input sequences (∣Σ∣100|\Sigma|^{100}). Thus, brute-force memorization cannot solve this task, and generalization is necessary to achieve nontrivial performance.

Strikingly, we obtain positive results (> ⁣99%>\!99\% in-distribution accuracyOur primary goal is to understand if gradient-based training can find shortcut solutions at all, rather than whether such training is stable. Accordingly, unless otherwise noted, we report the performance of the best model among 20 replicates. See Appendix B for details and sensitivity analyses.) for every finite-state semiautomaton we considered, including ones which generate the non-solvable groups A5A_{5} and S5S_{5}. Figure 3a gives a selection of our full results (in Appendix B.1). We find that more complex semiautomata (corresponding to non-abelian groups) require deeper networks to learn, in agreement with our theoretical constructions.

Our theoretical results show that there are logarithmic-depth (S5,A5S_{5},A_{5}) and constant-depth (all the others) solutions which simulate these semiautomata with exactly 100% accuracy. Of course, with such long sequences, black-box evaluation of whether this accuracy is reached in the population distribution is computationally infeasible. However, we note that without periodic activations (or some other mechanism for extrapolating to unseen count values), our theoretical constructions require MLPs to memorize the mod-nn function. This will be revisited in the out-of-distribution evaluation experiments in Section 5.2, but there is even a corresponding implication for in-distribution mistakes: if a model never sees outlier counts during training, it is expected to make mistakes on those outliers when they appear in evaluation.

Our theoretical results identify shortcut solutions which follow multiple, mutually incompatible paradigms. In general, we do not attempt a full investigation of mechanistic interpretability of the trained models. In particular, we do not claim that the networks discover implementations are isomorphic to those described in the proofs of Theorem 1, 2, and 3. However, as a preliminary exploration, we visualize some of the attention patterns in Figure 3b within successfully-trained models, finding attention heads which perform flat summations (with uniform attention) and conditional resets, agreeing with the construction in Theorem 3.

Although sufficiently deep networks find the solutions with non-negligible probability, the training dynamics are unstable; Figure 3b,c show example training curves, which exhibit high variance, negative progress, or accuracy that decays with continued training. In the same vein as the “synthetic reasoning tasks” introduced by Zhang et al. (2022), we hope that semiautomaton simulation will be useful as a clean, nontrivial testbed (with multiple difficulty knobs) for debugging and improving training algorithms, and perhaps the neural architectures themselves. More details are deferred to Appendix B.1.1.

Further experiments: more challenging settings

The results from Sections 3 and 4 show that Transformers can learn shortcuts end-to-end, unobstructed by depth, generalization, or optimization. However, the experiments in Section 4 are idealized in several ways; a natural question is whether these findings are robust to various challenges that arise in practice. In this section, we investigate the robustness of the shallow Transformer solutions, compared to those found by RNNs (the “natural” architecture for simulating semiautomata). We consider harder forms of supervision (Section 5.1) and evaluation (Section 5.2); details are deferred to Appendix B.2.

2 Out-of-distribution shortcomings of shortcut solutions

The theoretical construction of modular counters (Lemma 6) suggests a possible failure mode: if attention performs prefix addition and the MLP computes the sum modulo nn, the MLP could fail on sums unseen during training. This suggests that if the distribution over σ1:T\sigma_{1:T} shifts between training and testing (but the semiautomaton remains the same), a non-recurrent shortcut solution might map inputs into an intermediate latent variable space (like the sum) which fails to generalize.

More ambitiously, we could try to use these models to extrapolate to longer sequence lengths TT than those seen in the training data. Promoting this difficult desideratum of length generalization is an intricate problem in its own right; see Yao et al. (2021), Anil et al. (2022) for more experiments similar to ours. Figure 7 shows the performance on sequences of various lengths. In contrast to LSTM’s perfect performance on all scenarios, Transformer’s accuracy drops sharply as we move to lengths unseen during training. This is not purely due to unseen values of the positional encoding: randomly shifting the positions during training can cover all the positions seen during testing, which helps improve the length generalization performance but cannot make it perfect; we see similar results for removing positional encodings altogether. Finally, we empirically show that the above flaws are circumventable. Using a combination of scratchpad (a.k.a. “chain-of-thought”) (Nye et al., 2021, Wei et al., 2022) and recency bias (Press et al., 2022), we demonstrate that Transformers can be guided towards learning recurrent (depth-TT) solutions, which generalize out-of-distribution and to longer sequence lengths (Figure 7, yellow curves). Details are deferred to Section B.2.3.

Throughout the deep learning literature, the term shortcut is often used to refer to undesired (i.e., misleading, spurious, or overfitting) statistical properties of learned representations (Geirhos et al., 2020, Robinson et al., 2021). Meanwhile, under our computational (o(T)o(T) circuit depth) definition, shortcut solutions are perfectly valid ways to represent recurrent computations. The results in this section establish a connection between these notions: partially-learned computational shortcuts can be statistical shortcuts. Specifically, a non-recurrent architecture can “hallucinate” intermediate variables other than the state (e.g. the “count” variable for the parity automaton), is thus sensitive to the coverage of these variables in the training data. This leads to out-of-distribution generalization failures (e.g. on rare counts) which are not present in recurrent models.

The experiments in this section highlight a statistical penalty for learning recurrent computations with a non-recurrent architecture. However, the computational advantage of a shallow architecture is extremely appealing: maximally leveraging parallel computation, training and inference can be done much faster (O(log⁡T)O(\log T) or O(1)O(1) time, compared to O(T)O(T)). This highlights a delicate tradeoff between RNNs and Transformers, where neither architecture dominates the other, even when considering this elementary class of algorithmic problems. Attaining the best of both worlds with a practical architecture is an interesting avenue for future work.

Conclusions and future work

We have conducted a theoretical and empirical analysis of how shallow Transformers can learn shortcuts to the recurrent computations of finite-state automata. These shortcuts replace TT sequential iterations of a recurrent unit with a single pass through L≪TL\ll T parallel self-attention layers. Our theoretical results show that shortcuts are ubiquitous, and characterize extremely shallow ones (with LL independent of the context length TT) using algebraic machinery (Krohn-Rhodes theory). Empirically, we have shown that gradient-based optimization successfully finds these shortcuts. While the solutions found in practice generalize near-perfectly in-distribution (Section 4), they lack out-of-distribution robustness (Section 5). We hope that these results shed new light on the internal reasoning mechanisms of Transformers, as well as the design space for architectures capable of algorithmic reasoning.

Future work. This work is an initial foray into the interplay between neural architectures, algebraic automata theory, and circuit complexity. We list some open questions:

Finer-grained circuit complexity of self-attention: For certain automata of interest (e.g. bounded Dyck language parsers (Yao et al., 2021), and the gridworld automata from Theorem 3), there exist extremely shallow Transformer solutions, with depth independent of both TT and ∣Q∣|Q|. Which other natural classes of automata have this property of “beyond Krohn-Rhodes” representability?

When and why does gradient-based optimization work? The precise understanding of hierarchical representation learning in neural networks is an active frontier of research. The empirical tractability of finding “algebraic” shortcut solutions with gradient descent is an especially striking case, as related problems are known to be PSPACE\mathsf{PSPACE}-hard (kozen1977lower, beaudry1992membership, Cho and Huynh, 1991).

Mechanistic interpretability: Automaton simulation problems generate a rich class of challenging test cases for the reverse engineering of trained networks (see Nanda and Lieberum (2022)), and may yield further insights on the inductive biases of Transformers. In our preliminary attempts, we were only able to interpret a small number of simple models.

Acknowledgements

We are very grateful to Abhishek Shetty for helpful discussions about circuit complexity. We also thank Ashwini Pokle for thoughtful comments and suggestions towards improving clarity and readability.

Reproducibility Statement

Complete proofs of the theoretical results are provided in Appendix C, with a self-contained tutorial of relevant algebraic concepts in Appendix A.2. For the empirical results, all our datasets are derived from synthetic distributions, which are clearly described in Appendix B.1 and B.2. The architectures, implementations (with references to popular base repositories), and hyperparameters (including training procedure) are documented in Appendix B.3. Our open-source data-generating code is available from our project page.

References

Appendix

Below, we list our notational conventions for indices, vectors, matrices, and functions.

For a natural number nn, [n][n] denotes the index set {1,2,…,n}\{1,2,\ldots,n\}.

We will sometimes index vectors and matrices with named indices (such as ⊥\bot for padding tokens) instead of integers, for clarity.

eie_{i} denotes the ii-th elementary (one-hot) unit vector. Likewise as above, we sometimes use non-integer indices (e.g. e⊥e_{\bot}).

For a function f:X×Y→Zf:X\times Y\rightarrow Z and all y∈Yy\in Y, we will let f(⋅,y):X→Zf(\cdot,y):X\rightarrow Z denote the restriction of ff to yy (and similarly for other restrictions). This appears in the per-input state transition functions δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\rightarrow Q, as well as the functions represented by neural networks for a particular choice of weights.

For functions f,gf,g, f∘gf\circ g denotes composition: (f∘g)(x):=f(g(x))(f\circ g)(x):=f(g(x)). When we compose neural networks f:X×Θf→Y,g:Y×Θg→Zf:X\times\Theta_{f}\rightarrow Y,g:Y\times\Theta_{g}\rightarrow Z with parameter spaces Θf,Θg\Theta_{f},\Theta_{g}, we will use f∘g:X×(Θf×Θg)→Zf\circ g:X\times(\Theta_{f}\times\Theta_{g})\rightarrow Z to indicate the composition f(g(x;θg)θf)f(g(x;\theta_{g})\theta_{f}).

A.2 Automata, semigroups, and groups

Recall that a semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) has a state space QQ, an input alphabet Σ\Sigma, and a transition function δ:Q×Σ→Q\delta:Q\times\Sigma\to Q. For any natural number TT and a starting state q0∈Qq_{0}\in Q, by repeated composition of the transition function δ\delta, one can use A{\mathcal{A}} to define a map from a sequence of inputs (σ1,…,σT)∈ΣT(\sigma_{1},\ldots,\sigma_{T})\in\Sigma^{T} to a sequence of states (q1,…,qT)∈QT(q_{1},\ldots,q_{T})\in Q^{T} via:

Here and below, it is helpful to use a matrix-vector notation to express the computation of semiautomata. For a given semiautomaton we can always identify the state space QQ with index set {1,…,∣Q∣}\{1,\ldots,|Q|\} and use a one-hot encoding of states into {0,1}∣Q∣\{0,1\}^{|Q|}. For each input symbol σ∈Σ\sigma\in\Sigma, we associate a transition matrix δ(⋅,σ)∈{0,1}∣Q∣×∣Q∣\delta(\cdot,\sigma)\in\{0,1\}^{|Q|\times|Q|} with entries [δ(⋅,σ)]q′,q=1{δ(q,σ)=q′}[\delta(\cdot,\sigma)]_{q^{\prime},q}=\mathbf{1}\{\delta(q,\sigma)=q^{\prime}\}. This implies that for all q,σq,\sigma, we have eδ(q,σ)=δ(⋅,σ)eqe_{\delta(q,\sigma)}=\delta(\cdot,\sigma)e_{q}, so that the computation of the semiautomaton amounts to repeated matrix-vector multiplication.

While semiautomata are remarkably expressive, we discuss a few simple examples throughout this background section to elucidate the key concepts.

Let Q=Σ={0,1}Q=\Sigma=\{0,1\} and let δ(q,0)=q\delta(q,0)=q and δ(q,1)=1−q\delta(q,1)=1-q. Then, starting with q0=0q_{0}=0, the state at time tt, qtq_{t}, is 11 if the binary sequence (σ1,…,σt)(\sigma_{1},\ldots,\sigma_{t}) has an odd number of 11s.

Let Q={1,2},Σ={⊥,1,2}Q=\{1,2\},\Sigma=\{\perp,1,2\} and let δ\delta be given by

As the name suggests, this semiautomaton implements a simple memory operation where the state at time tt is the value of the most recent non-⊥\perp input symbol.

Let SS be a natural number, Q={0,1,…,S}Q=\{0,1,\ldots,S\} and Σ={L,⊥,R}\Sigma=\{L,\perp,R\}. Then the transition matrices are given by:

This semiautomaton describes the movement of an agent along a line segment where actions −1-1 and +1+1 correspond to decrementing and incrementing the state respectively, except that the decrement input has no effect at state and the increment input has no effect at state SS.

Note that we have chosen a convention which differs slightly from the main paper (i.e. Figure 1): we enumerate the indices starting from rather than 11. This is because the proofs are stated more naturally when the boundaries of the gridworld are identified with the indices and SS.

For a semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) each input symbol σ∈Σ\sigma\in\Sigma defines a function δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\to Q. These functions can be composed in the standard way, and we use δ(⋅,σ1:t)\delta(\cdot,\sigma_{1:t}) to denote the tt-fold function composition. Note that δ(q0,σ1:t)\delta(q_{0},\sigma_{1:t}) is precisely the value of the state at time tt on input σ1:t\sigma_{1:t}. Thus, the set of all functions that can be obtained by composition of the transition operator, formally

plays a central role in describing the computation of the semiautomaton. This object is a transformation semigroup. We now turn to describing the necessary algebraic background.

Recall that a group (G,⋅)({\mathcal{G}},\cdot) is a set G{\mathcal{G}} equipped with a binary operation ⋅:G×G→G\cdot:{\mathcal{G}}\times{\mathcal{G}}\to{\mathcal{G}} such that

(identity) There exists an identity element e∈Ge\in{\mathcal{G}} such that e⋅g=g⋅e=ge\cdot g=g\cdot e=g for all g∈Gg\in{\mathcal{G}}.

(invertibility) Every element g∈Gg\in{\mathcal{G}} has an inverse g−1∈Gg^{-1}\in{\mathcal{G}} such that g⋅g−1=g−1⋅g=eg\cdot g^{-1}=g^{-1}\cdot g=e

(associativity) The binary operation is associative: (g1⋅g2)⋅g3=g1⋅(g2⋅g3)(g_{1}\cdot g_{2})\cdot g_{3}=g_{1}\cdot(g_{2}\cdot g_{3}).

A monoid is less structured than a group; there must be an identity element and the binary operation must be associative, but invertibility is relaxed. A semigroup is even less structured: the only requirement is that the binary operation is associative.

It is common to let G{\mathcal{G}} be a subset of functions from Q→QQ\to Q where QQ is some ground set and let the binary operation be function composition. In this case, the structure is called a permutation group or transformation monoid/semigroup depending on which subset of the above properties hold. For transformation groups, since every element has an inverse under function composition, it is immediate that every element is some permutation over the ground set.

In fact, taking G{\mathcal{G}} to be a subset of functions as above is without loss of generality: by Cayley’s theorem every group is isomorphic (equivalent after renaming elements) to a transformation group on some ground set, and we can take the ground set to have the same number of elements as the original group (for finite groups). Analogously, all semigroups are isomorphic to a transformation semigroup, but the ground set may need one additional element (for the identity); this is Cayley’s theorem for semigroups. It is also clear that every transformation semigroup can be realized by some semiautomaton by trivially having the input symbols correspond to the functions in G{\mathcal{G}}.More succinctly, inputs can correspond to a generating set of the group, but this is not relevant for our results. Therefore we have lost no structure when passing from finite semiautomata to finite semigroups.

Before discussing the compositional structure of semigroups, we give one more canonical example.

Let SS be a natural number let Q={0,1,…,S−1}Q=\{0,1,\ldots,S-1\} and let Σ={1}\Sigma=\{1\} have only one element. The dynamics are given by δ(q,1)=(q+1)mod  S\delta(q,1)=(q+1)\mod S. Clearly this semiautomaton implements counting modulo SS. The underlying group is the cyclic group, denoted CSC_{S}, which is isomorphic to the integers mod SS with addition as the binary operation. Note that in this case, the operation is commutative, which makes the group abelian.

The most natural way to compose larger groups from smaller ones is via the direct product. Given two groups GG and HH, we can form a new group with elements {(g,h):g∈G,h∈H}\{(g,h):g\in G,h\in H\} with a binary operation that is applied component-wise (g,h)⋅(g′,h′)=(g⋅g′,h⋅h′)(g,h)\cdot(g^{\prime},h^{\prime})=(g\cdot g^{\prime},h\cdot h^{\prime}) (here, ⋅\cdot is overloaded to be the group operation for all three groups). This direct product group is denoted G×HG\times H. In the context of permutation groups, say GG is a permutation group over ground set QGQ_{G} and HH is over ground set QHQ_{H}. Then G×HG\times H has ground set QG×QHQ_{G}\times Q_{H} and every function in G×HG\times H factorizes component-wise, i.e., every element in G×HG\times H is identified with a permutation (qG,qH)↦(g(qG),h(qH))(q_{G},q_{H})\mapsto(g(q_{G}),h(q_{H})) where g∈G,h∈Hg\in G,h\in H.

Observe that G×HG\times H contains normal subgroups which are isomorphic to both GG and HH. To see this, take N={(eG,h):h∈H}N=\{(e_{G},h):h\in H\} where eGe_{G} is the identity element in GG. Then since geG=eGgge_{G}=e_{G}g and since HH is closed under its group operation, we have (g,h)N=N(g,h)(g,h)N=N(g,h) for all (g,h)∈G×H(g,h)\in G\times H. A symmetric argument shows that GG is also a normal subgroup of the direct product.

Note that we can analogously define direct products in the absence of the group axioms, and thus for monoids and semigroups. This gives a natural construction of the semigroup corresponding to moving around both axes of a 2-dimensional rectangular gridworld, as a concatenation of two non-interacting 1-dimensional gridworlds:

If GSG_{S} is the transformation semigroup of the 1-d grid world with S+1S+1 states, then GS×GSG_{S}\times G_{S} corresponds to a 2-dimensional gridworld. A semiautomaton that yields this transformation semigroup has state space Q={(i,j):i,j∈{0,…,S}}Q=\{(i,j):i,j\in\{0,\ldots,S\}\} and 5 actions: increment or decrement ii or jj, subject to boundary effects, or do nothing.

The definition of direct product extends straightforwardly to more than two terms G1×G2×…×GnG_{1}\times G_{2}\times\ldots\times G_{n}; we identify the items with tuples (g1,g2,…,gn)(g_{1},g_{2},\ldots,g_{n}).

However, it is possible to compose larger groups so that one of the subgroups is not a normal subgroup. This operation is called a semidirect product, with the group law (g,h)⋅(g′,h′)=(g⋅ϕh(g′),h⋅h′)(g,h)\cdot(g^{\prime},h^{\prime})=(g\cdot\phi_{h}(g^{\prime}),h\cdot h^{\prime}) for some ϕh\phi_{h} to be defined later. Observe that in the direct product G×HG\times H, we have constructed the elements from ordered pairs (g∈G,h∈H)(g\in G,h\in H), lifting GG and HH into a shared product space (i.e., the Cartesian product of the underlying sets of GG and HH), defining the group operation as simply applying those of GG and HH separately.

In fact, there are other ways, to define the group operation in the product space, but a difficulty arises: we need to find other nontrivial multiplication rules on pairs (g,h)(g,h), and we cannot take for granted that an arbitrary binary operation satisfies the group axioms. We would like to define other operations (g,h)⋅(g′,h′)(g,h)\cdot(g^{\prime},h^{\prime}) which output an element of gg and an element of hh. An attempt would be to pick two arbitrary injective homomorphisms ϕG,ϕH\phi_{G},\phi_{H} which embed GG and HH into a “shared space,” so that elements of GG and HH can be multiplied together:

However, the middle equality may not hold, because ϕG(g′)\phi_{G}(g^{\prime}) and ϕH(h)\phi_{H}(h) are not guaranteed to commute. (Observe that for the special case of g↦(g,eH),h↦(eG,h)g\mapsto(g,e_{H}),h\mapsto(e_{G},h), these two elements always commute, giving rise to the direct product.)

which is of the form ϕG(⋅)⋅ϕH(⋅)\phi_{G}(\cdot)\cdot\phi_{H}(\cdot) since both GG and HH are themselves closed. This condition is precisely that GG is a normal subgroup.

This object is the semidirect product, and it is denoted G⋊HG\rtimes H. Note that the choice of mapping ϕ\phi is unspecified in the notation, and, in general, different choices of ϕ\phi will yield different structures for the semidirect product.

Finally, when G=N⋊HG=N\rtimes H, both NN and HH are subgroups of GG, but NN is also a normal subgroup. To see this, we need to check that hN=NhhN=Nh for any h∈Hh\in H. This is equivalent to hnh−1∈Nhnh^{-1}\in N for each h,nh,n, but we defined the group operation to be hnh−1=ϕh(n)∈Nhnh^{-1}=\phi_{h}(n)\in N, specifically so this would hold. On the other hand, HH may not be a normal subgroup, and in this sense the semidirect product is a generalization of the direct product (for which both subgroups are normal). However, when the mapping ϕ\phi is trivial, that is ϕh(n)=n\phi_{h}(n)=n then both NN and HH are normal subgroups, and one can verify that in this case the semidirect product and direct product coincide.

The transformation semigroup for this semiautomaton is CS⋊C2C_{S}\rtimes C_{2} where CSC_{S} is the cyclic group on SS elements (cf. Example 4). C2C_{2} has two elements, the identity ee and one element hh such that hh=ehh=e. CSC_{S} has SS elements where each element gg is a function that adds some number k∈{0,…,S−1}k\in\{0,\ldots,S-1\} to the input modulo SS. The inverse g−1g^{-1} is naturally to subtract kk to the input, modulo SS. The homomorphism ϕ\phi in the semidirect product is such that ϕe(g)=g\phi_{e}(g)=g and ϕh(g)=g−1\phi_{h}(g)=g^{-1}.

We define one more type of product between groups NN and HH: the wreath product N≀H:=(N×…×N)⋊HN\wr H:=(N\times\ldots\times N)\rtimes H. This is a group containing ∣N∣∣H∣⋅∣H∣|N|^{|H|}\cdot|H| elements (rather than ∣N∣⋅∣H∣|N|\cdot|H|, like the direct and semidirect products). Intuitively, it is defined by creating one copy of NN per element in HH via the direct product, then letting HH specify a way to exchange these copies. Formally, N≀HN\wr H is the unique group generated by

where we have enumerated the elements of HH in arbitrary order, such that each πh:[H]→[H]\pi_{h}:[H]\rightarrow[H] is the permutation defined by right multiplication h′↦h′hh^{\prime}\mapsto h^{\prime}h (by convention).

A naive way to construct the Rubik’s Cube is to assign labels {1,…,54}\{1,\ldots,54\} to the stickers on the cube, and define the Rubik’s Cube group GG via the sticker configurations reachable by the 66 face turns (which each specify a permutation δL,δR,δU,δD,δB,δF:→\delta_{L},\delta_{R},\delta_{U},\delta_{D},\delta_{B},\delta_{F}:\rightarrow of the stickers). This establishes GG as a subgroup of S54S_{54}. First, notice that the 6 central stickers never move (so this is really improvable to S48S_{48}). Next, notice that the 24=8×324=8\times 3 vertex stickers never switch places with the 24=12×224=12\times 2 edge stickers. The vertex stickers form a subset of the wreath product C3≀S8C_{3}\wr S_{8}, while the edge stickers form a subset of the wreath product C2≀S12C_{2}\wr S_{12}. In all, this realizes GG as a subgroup of a direct product of wreath products:

Among other consequences towards solving the Rubik’s Cube, this gives an improved upper bound on the size of GG (which turns out to still be off by a factor of 12, because of nontrivial invariants preserved by the face rotations, a.k.a. unreachable configurations).

When NN is a normal subgroup of GG, the quotient group G/NG/N is defined as {gN:g∈G}\{gN:g\in G\} with binary operation (gN)(g′N)=(gg′)N(gN)(g^{\prime}N)=(gg^{\prime})N. The fact that NN is a normal subgroup implies that this is a well defined group. We can also check that if G=N⋊HG=N\rtimes H then the quotient group G/NG/N is isomorphic to HH, which matches the intuition for multiplication and division.

A GG group is simple if it has no non-trivial normal subgroups. Intuitively, a simple group cannot be factorized into components; this generalizes the fact that a prime number admits no non-trivial factorization. When GG is not simple then it has a non-trivial normal subgroup, say NN. We call NN proper if N≠GN\neq G. We call a proper subgroup NN maximal if there is no other proper normal subgroup N′◃GN^{\prime}\triangleleft G such that N◃N′N\triangleleft N^{\prime}. Equivalently, NN is a maximal proper normal subgroup if and only if G/NG/N is simple. This is akin to extracting a prime factor from a number, since the quotient group G/NG/N cannot be further factorized. We will revisit this idea of factorization when defining composition series and solvable groups in Section C.3.2.

Finally, we provide some additional terminology related to these different notions of products, which provide a cleaner unifying language in which to state our constructions. Let N,HN,H be arbitrary groups. Which groups GG contain a normal subgroup isomorphic NN, such that the quotient G/NG/N is isomorphic to HH? Such a group GG is said to be an extension of NN over HH. The direct product G=N×HG=N\times H is known as the trivial extension. A semidirect product G=N⋊HG=N\rtimes H is known as a split extension. However, not all extensions are split extensions; the smallest example is the quaternion group Q8Q_{8}, the group of unit quaternions {±1,±i,±j,±k}\{\pm 1,\pm i,\pm j,\pm k\} under multiplication (i2=j2=k2=ijk=−1i^{2}=j^{2}=k^{2}=ijk=-1), which cannot be realized as a semidirect product of smaller groups. In general, it is very hard to derive interesting properties of a group extension based on the properties of NN and HH. Fortunately, there is a characterization of general extensions. The Krasner-Kaloujnine universal embedding theorem (Krasner and Kaloujnine, 1951) states that all extensions GG can be found as subgroups of the wreath product N≀HN\wr H. The proof of Theorem 2 essentially shows how to implement the different kinds of group extensions, given constructions which implement the substructures N,HN,H. In the worst case, we will have to implement a wreath product.

A.3 Shallow circuit complexity classes

We provide an extremely abridged selection of relevant concepts in circuit complexity. For a systematic introduction, refer to (Arora and Barak, 2009). In particular, we discuss each circuit complexity class and inclusion below:

NC0\mathsf{NC}^{0} is the class of constant-depth, constant-fan-in, polynomial-sized AND\mathsf{AND}/OR\mathsf{OR}/NOT\mathsf{NOT} circuits. If a constant-depth Transformer only uses the constant-degree sparse selection constructions in (Edelman et al., 2022), it can be viewed as representing functions in this class. However, the representational power of these circuits is extremely limited: they cannot express any function which depend on a number of inputs growing with TT.

AC0\mathsf{AC}^{0} is the class of constant-depth, unbounded-fan-in, polynomial-sized AND\mathsf{AND}/OR\mathsf{OR} circuits, allowing NOT\mathsf{NOT} gates only at the inputs. A classic result is that the parity of TT bits is not in AC0\mathsf{AC}^{0} (Furst et al., 1984); Hahn (2020) concludes the same for bounded-norm (and thus bounded-Lipschitz-constant) constant-depth Transformers.

ACC0\mathsf{ACC}^{0} extends AC0\mathsf{AC}^{0} with an additional type of unbounded-fan-in gate known as MODm\mathsf{MOD}_{m} for arbitrary number mm, which checks if the sum of the input bits is a multiple of mm. Theorem 2 comes from the fact that the semigroup word problem (which is essentially identical to semiautomaton simulation) is in this class; see (Barrington and Thérien, 1988).

TC0\mathsf{TC}^{0} extends AC0\mathsf{AC}^{0} with an additional type of unbounded-fan-in gate called MAJ\mathsf{MAJ}, which computes the majority of an odd number of input bits (a threshold gate). It is straightforward to simulate modular counters using a polynomial number of parallel thresholds (i.e. ACC0⊆TC0\mathsf{ACC}^{0}\subseteq\mathsf{TC}^{0}). Whether this inclusion is strict (can you simulate a threshold in constant depth with modular counters?) is a salient open problem in circuit complexity. Threshold circuits are a very natural model for objects of interest in machine learning like decision trees and neural networks (Merrill et al., 2021).

NC1\mathsf{NC}^{1} is the class of O(log⁡T)O(\log T)-depth, constant-fan-in, polynomial-sized AND\mathsf{AND}/OR\mathsf{OR}/NOT\mathsf{NOT} circuits. It is an extremely popular and natural complexity class capturing efficiently parallelizable algorithms. It is unknown whether any of the inclusions in the “larger” classes TC0⊆NC1⊆L⊆P\mathsf{TC}^{0}\subseteq\mathsf{NC}^{1}\subseteq\mathsf{L}\subseteq\mathsf{P} are strict.

A.4 The Transformer architecture

In this section, we define the Transformer function class used in our theoretical results, and discuss remaining discrepancies with the true architecture.

the causally-masked softmax at row tt is defined to be softmax(z1:t)\mathsf{softmax}(z_{1:t}) on the first tt coordinates, and 0 on the rest. To implement the causal masking operation, it is customary to set the entries above the diagonal of the attention score matrix XWQWK⊤X⊤XW_{Q}W_{K}^{\top}X^{\top} to −∞-\infty, then obtaining CausalAttn(XWQWK⊤X⊤)\mathsf{CausalAttn}(XW_{Q}W_{K}^{\top}X^{\top}) via a row-wise softmax (letting e−∞e^{-\infty} evaluate to ).

In general, for any positive integer HH, a multi-headed self-attention block consists of a sum of HH copies of the above construction, each with its own parameters.

This component is often called soft attention: the softmax performs continuous selection, taking a convex combination of its inputs. In contrast, hard attention refers to attention heads which perform truly sparse selection (putting weight 11 on the position with the highest score, and on all others).

At the level of granularity of the results in this paper (up to negligible changes in the width and weight norms), this changes very little from the viewpoint of representation. A residual connection can be implemented (or negated) by appending two ReLU activations to a non-residual network:

Similarly, a residual connection can be implemented with one attention head (with internal embedding dimension k=dk=d), as long as it is able to select its own position (which will be true in all of our constructions).

For simplicity of presentation, we omit the normalization layers which are usually present after each attention and MLP block. It would be straightforward (but an unnecessary complication) to modify the function approximation gadgets in our constructions to operate with unit-norm embeddings.

Finally, it will greatly simplify the constructions to add padding tokens: to simulate a semiautomaton at length TT, we will choose to prepend τ\tau tokens, with explicitly chosen embeddings, which do not depend on the input σ1:T\sigma_{1:T}. Theorem 1 uses τ=Θ(T)\tau=\Theta(T) padding, and Theorem 2 uses τ=1\tau=1. In both cases, padding is not strictly necessary (the same functionality could be implemented by the MLPs without substantially changing our results), but we find that it leads to the most intuitive and concise constructions.

We define the following quantities associated with a Transformer network, and briefly outline their connection to familiar concepts in circuit complexity:

The dimensions according to the definition of a sequence-to-sequence network: sequence length TT and embedding dimension dd. Up to a factor of bit precision, this corresponds to the number of inputs in a classical Boolean circuit. We will exclusively define architectures where dd is independent of TT.

A bound on its ∞\infty-weight norms: the largest absolute value of any trainable parameter. These can be converted into norm-based generalization bounds via the results in Edelman et al. (2022). Note that the results in this paper go beyond the sparse variable creation constructions of bounded-norm attention heads; in general, the norms scale with TT. The attention heads express meaningful non-sparse functions. Aside from the positive experimental results, we do not directly investigate generalization in this paper.

A.5 Additional discussion of related work

We first provide references for the “reasoning-like” applications of neural networks mentioned in the main paper.

Program synthesis: (Chen et al., 2021b, Schuster et al., 2021, Li et al., 2022).

Mathematical reasoning: (Lample and Charton, 2019, Polu and Sutskever, 2020, Drori et al., 2022).

Neural dynamics models for decision-making: recurrent (Hafner et al., 2019, Ye et al., 2021, Micheli et al., 2022) and non-recurrent (Chen et al., 2021a, Janner et al., 2021).

We provide an expanded discussion of empirical analyses of neural networks trained on synthetic combinatorial tasks.

Pointer Value Retrieval: Zhang et al. (2021a) propose a benchmark of tasks based on pointer value retrieval (PVR) to study the generalization ability (in-distribution as well as distribution shift) of different neural network architectures. Their key idea behind the task is “indirection through a pointer rule”, that is, a specific position of the input acts as a pointer to the relevant position (window) of the input which contains the answer. Using our results, we can implement a certain sub-class of PVR tasks: (1) we use the first attention layer to identify the pointer, and (2) we use the second attention layer to select the window between the pointer value and the width. (2) is doable with O(1)O(1) attention heads if we are computing a function that is based on the sum (for example,mod  n\mod n). Otherwise it would require the window size number of attention heads similar to our grid-world construction.

LEGO: Zhang et al. (2022) propose a task based on solving a simple chain-of-reasoning problem based on group-based equality constraints. They study the ability of transformers to generalize the entire chain of reasoning given only part of the chain while training. A direct comparison to our setting is not clear since this task is not modelled as a sequence-to-sequence task, however it serves as another example of the emergence of “shortcut” solutions: transformers solve certain variables without resolving the chain of reasoning.

Dyck: Several works (Hahn, 2020, Ebrahimi et al., 2020, Newman et al., 2020, Yao et al., 2021) have studied the ability of Transformers to represent Dyck languages, both for generation and closing bracket prediction. The most closely related to our work is Yao et al. (2021), which constructs a clever depth-2 as well as a depth-DD solution for bounded-depth DD Dyck languages. Bounded-depth Dyck can be captured by our semiautomata formalism and our main construction would recover the depth-2D2^{D} solution by default. Their depth-2 construction bears semblance to the constructions we use in Theorem 3: they implement a counter in the first layer similar to our mod  n\mod n construction, and implement a proximity-based depth matching in the second layer. Our grid-world construction generalizes their construction to a significantly more complex problem. We view our work as a generalization of their results to a wider class of semiautomata.

Parity: Another commonly studied synthetic setup is the task of learning parities. Edelman et al. (2022), Barak et al. (2022) perform a theoretical and empirical study of the ability of Transformers (and other architectures) to learn sparse parities where the support size k≪Tk\ll T. Bhattamishra et al. (2020), Schwarzschild et al. (2021) study the task of computing prefix sum in the binary basis (which is essentially parity of the prefix sum) for Transformers and recurrent models, repsectively. Anil et al. (2022), Wei et al. (2022) study essentially the same problem however they model the task as a natural language task and use pretrained Transformers.

Modular addition: In the pursuit of understanding grokking, Nanda and Lieberum (2022) focus on the task of adding two 5 digit numbers modulo a large prime (113113 in their setting). They take the viewpoint of mechanistic interpretability and attempt to reverse engineer what the Transformer is learning on this task in the low sample regime. They claim that the trained model learns sinusoidal encodings that we also use in our theoretical constructions. Note that our setting of modular counters performs a TT-way summation, while their setting involves only a 2-way summation (with carryover). Inspired by their work, we do some preliminary investigation into interpreting the trained Transformer on the grid world (see Figure 10).

Dyck languages are particularly interesting for their completeness property: the Chomsky-Schützenberger representation theorem (Chomsky and Schützenberger, 1959) states that all context-free languages can be (homomorphically) represented by the intersection of a Dyck language and a regular language. For more on this topic, see the discussion in Yao et al. (2021). In the context of regular languages (which in general induce finite-state automata), our findings imply that O(log⁡T)O(\log T)-depth networks can simulate all context-free languages (Theorem 1), and O(1)O(1)-depth networks can represent some of them. The obstructing regular languages are the ones whose associated syntactic monoids are non-solvable. We further note that the gridworld semigroups are aperiodic and thus simulable by star-free regular expressions (Schützenberger, 1965) and AC0\mathsf{AC}^{0} circuits (Chandra et al., 1983, Barrington and Thérien, 1988). We did not see a way for this to generically entail O(1)O(1)-depth shortcuts with self-attention. For the relation between the Chomsky hierarchy and various neural networks in practice, Delétang et al. (2022) provide an extensive empirical study for memory-augmented RNNs and Transformers on tasks spanning all 4 levels of the hierarchy, and conclude the Transformers lack the ability to even recognize regular languages. Their results do not contradict with ours, since they measure performance on “inductive inference”, which is similar to our length generalization setup where we also see the failure of Transformer.

There has been much recent interest in quantifying out of distribution generalization of trained models under distribution shifts that maintain some notion of “logical” invariance. Wei et al. (2022), Anil et al. (2022) empirically investigate the ability of pre-trained Transformers to generalize to longer sequence length for parity-like problems modelled as language tasks. Xu et al. (2020) study size generalization in graph neural networks where they train on small graphs and evaluate on larger sized graphs with similar structural properties. Schwarzschild et al. (2021), Bansal et al. (2022) focus on length and algorithmic generalization for recurrent models where they train on simple/easy instances of the underlying problem and evaluate on harder/complex instances using the power of recurrence to simulate extra computational steps, inspired by the ideas of Neural Turing Machines (Graves et al., 2014) and Adaptive Computation Time (Graves, 2016). We view our results as complementing those of Yao et al. (2021), Anil et al. (2022) for a richer class of problems. Our use of scratchpad is inspired by Nye et al. (2021), Wei et al. (2022), Anil et al. (2022).

Our work is not the first to notice that Transformer architectures make brittle predictions out-of-distribution. Indeed, even the seminal paper introducing the architecture (Vaswani et al., 2017) notes that length generalization is promoted by a subtle hyperparameter choice (namely, the positional encoding scheme). Furthermore, there have been several attempts to reconcile this gap by modifying Transformers to behave more like RNNs; (Dehghani et al., 2019, Nye et al., 2021, Wei et al., 2022, Anil et al., 2022, Hutchins et al., 2022). Kasai et al. (2021) consider training a non-recurrent Transformer, and finetuning it into an RNN. All of these works have some element of natural language experiments: either the task is end-to-end language modeling, or the synthetic reasoning task is framed as a natural language problem, for a pretrain-finetune pipeline. We view our work as strengthening the foundations of these lines of inquiry. Theoretically, we provide structural guarantees for how shallow non-recurrent models can (perhaps deceptively) fit recurrent dynamics over long sequences. Empirically, we perform a pure (no confounds arising from the influence of a natural langauge corpus) analogue of the experiments seeking to help neural networks follow long chains of reasoning.

As mentioned briefly towards the end of Section 5, the setting of indirectly-supervised semiautomata matches that of autoregressive generative modeling (a.k.a. next-token prediction), if the continuations of the sequence depend on the state of a latent semiautomaton. This is the case in (for example) generating Dyck languages (Yao et al., 2021), where the possible continuations are {\{all possible open brackets, if the stack qtq_{t} is not full}∪{\}\cup\{close bracket which pairs with the top of the stack qt}q_{t}\}. We note that when an autoregressive model is used for sequence generation via a token-by-token inference procedure, this amounts to a special case of scratchpad inference (with a naive 11-step training procedure): the constant-depth network is used as a single iteration of a recurrent network, whose state is the completed prefix of the current generated sequence. Non-autoregressive natural language generation and transduction are an exciting area of research (Gu et al., 2017); for a recent survey, see Xiao et al. (2022). Our results are relevant to this line of work, suggesting that there may not be an expressivity barrier to expressing deep recurrent linguistic primitives, but there may be issues with out-of-distribution robustness.

Pérez et al. (2021) show that an infinite-precision Transformer achieves Turing completeness, as a single forward pass through a 3-layer decoder can simulate one transition step of a Turing machine. Giannou et al. (2023) exhibit a 13-layer Transformer whose weights are hard-coded to a universal Turing machine, and can be looped to perform any computation. These works show that one pass through a Transformer can implement a single computational step of a Turing machine. In contrast, our results show how a shallow Transformer can sometimes execute a computational loop over the entire context in a single non-recurrent pass. This requires a significantly more refined analysis, which depends on the global algebraic structure induced by the automata in question.

Another area where tools from abstract algebra are used to reason about neural networks is geometric deep learning, a research program which seeks to understand how to specify inductive biases stemming from algebraic invariances. For a recent survey, see Bronstein et al. (2021). In contrast, this work studies the ability of a fixed architecture to learn a wide variety of algebraic operations, in the absence of special priors (but a large amount of data). There are certainly possible connections (e.g. “how do you bias an architecture to perform operations in a known group, when there is limited data?”) to explore in future work.

Our theoretical results can be interpreted as a depth separation result: contingent on TC0≠NC1\mathsf{TC}^{0}\neq\mathsf{NC}^{1}, it takes strictly more layers to simulate non-solvable semiautomata, compared to their solvable counterparts. In a similar spirit, there have been several works establishing depth separation for feed-forward neural networks (mostly using ReLU activations) (Telgarsky, 2016, Eldan and Shamir, 2016, Daniely, 2017, Lee et al., 2017, Safran et al., 2019). These results are usually constructive in nature, that is, they show the existence of functions that can be represented by depth LL but would require exponential width for depth L−1L-1 (or L\sqrt{L}, depending on the result).

More elementary neural architectures, such as MLPs, have the ability to represent arbitrary functions, given sufficiently many neurons (Hornik et al., 1989, Cybenko, 1989). The ACC0\mathsf{ACC}^{0} circuit construction described in (Barrington and Thérien, 1988) can be implemented by any such architecture, not just the Transformer– there is a naive black-box way to “compile” each gate in the circuit into a network with the same depth as the ACC0\mathsf{ACC}^{0} circuit. A natural question is: why, then, should we prefer Transformers? The primary advantage of Transformers comes from the position-wise weight sharing of the attention layers and the casual structure from causal attention maps. Unlike MLPs, the shared parameters in Transformers allow for significant reduction in the total parameter count for representing the two main operations across positions: modular counters, and resets. In particular, these functions can be represented with O(1)O(1) trainable parameters in the attention and MLP weight matrices (i.e. independent of TT), as opposed to the Θ(T)\Theta(T) parameters in a vanilla MLP, where position-wise parameter sharing is not available. In a sense, the Transformer architecture is naturally suited for implementing this construction.

Appendix B Experiments

This section contains a full description and discussion of the in-distribution simulation experiments from Section 4.

This defines a supervised learning problem over sequences.

Note that without intermediate states in the input these problems exhibit long-range dependencies: for example, in the parity semiautomaton (and for any semiautomaton whose transformation semigroup is a group), every qtq_{t} depends on every preceding input {σt′:t′<t}\{\sigma_{t^{\prime}}:t^{\prime}<t\}. Indeed, this is why previous studies have used group operations as a benchmark for reasoning (Anil et al., 2022, Zhang et al., 2022).

We proceed to enumerate the semiautomata considered in these simulation experiments.

Cyclic groups C2,C3,…,C8C_{2},C_{3},\ldots,C_{8}. For each cyclic group CnC_{n} (realized as Q:={0,1,…,n−1}Q:=\{0,1,\ldots,n-1\} under mod-nn addition), we choose the generator set Σ\Sigma to be the full set of group elements {0,…,n−1}\{0,\ldots,n-1\}. An alternative could be to let Σ\Sigma be a minimalIn the sense that it induces a non-trivial learning problem on this group.If we only pick the generator {1}\{1\}, the output sequence is deterministic, and there is no learning problem. set {0,1}\{0,1\}, which we do not use in the experiments.

Direct products of cyclic groups C2×C2,C2×C2×C2C_{2}\times C_{2},C_{2}\times C_{2}\times C_{2}, realized as concatenated copies of the component semiautomata. Note that C6C_{6} (which is isomorphic to C2×C3C_{2}\times C_{3}), included in the above set, is another example.

Dihedral groups D6,D8D_{6},D_{8}. Our realization of D2nD_{2n} chooses Q={0,1,…,n−1}×{0,1}Q=\{0,1,\ldots,n-1\}\times\{0,1\} and Σ={(1,0),(0,1)}\Sigma=\{(1,0),(0,1)\}. Since these groups are non-abelian, it is already not so straightforward (compared to parity) to see why constant-depth shortcuts should exist.

Permutation groups A4,S4,A5,S5A_{4},S_{4},A_{5},S_{5}. We choose QQ to be the set of n!n! permutations for SnS_{n} (symmetric group), and QQ to be the set of n!2\frac{n!}{2} even permutations for AnA_{n} (alternating group on nn elements). The generator set for SnS_{n} consists of the minimal generators, a transposition and an nn-cycle, as well as 6 other permutations. These other permutations are chosen following the ordering given by the sympy.combinatorics\mathsf{sympy.combinatorics} package. They are not necessary for covering the state space (since the minimal set of 2 permutations already suffice to cover QQ), but can help speed up the mixing of the states. For AnA_{n}, we choose the 33-cycles of the form (12i)(12i) for i∈{3,4,⋯ ,n}i\in\{3,4,\cdots,n\}. Note that A4,S4A_{4},S_{4} are solvable (leading to constant-depth shortcuts), while A5,S5A_{5},S_{5} are not. Also, note that to learn a constant-depth shortcut for A4A_{4}, a model needs to discover the wondrous fact that A4A_{4} has a nontrivial normal subgroup, that of its double transpositions.

The quaternion group Q8Q_{8}. This is the smallest example of a non-abelian solvable group which is not realizable as a semidirect product of smaller groups, thus requiring the full wreath product construction (Lemma 10) in our theory.

We focus on the online learning setting for all experiments in this paper: at training iteration ii, draw a fresh minibatch of samples from DA{\mathcal{D}}_{\mathcal{A}}, compute the network’s loss and gradients on this minibatch, and update the model’s weights using a standard first-order optimizer (we use AdamW (Loshchilov and Hutter, 2017)). This is to mitigate the orthogonal challenge of overfitting; note that the purpose of these experiments is to determine whether standard gradient-based training finds shortcut solutions in these combinatorial settings (in a reasonable amount of time), not how efficiently. We do not investigate how to improve sample efficiency in this paper. The results in the paper are based on sinusoidal positional encodings (Vaswani et al., 2017) unless otherwise specified.

We report our main results with sequence length T=100T=100, which is large enough to rule out memorization: for this choice of TT, the inputs come from a uniform distribution over ∣Σ∣100>1030|\Sigma|^{100}>10^{30} sequences, rendering it overwhelmingly unlikely for a sample to appear twice between training and evaluation. We observed positive results in most of the settings for larger TT, but training became prohibitively unstable and computationally expensive; mitigating this is an interesting direction for future empirically-focused studies.

We seek to investigate the sufficient depth for learning to simulate each semiautomaton. Thus, for each problem setting, we vary the number of layers LL in the Transformer between 11 and 1616. Note that we do not attempt in this work to distinguish between depths O(log⁡T)O(\log T) and O(1)O(1), nor do we attempt to tackle the problem of exhaustively enumerating and characterizing the shortcut solutions for any particular semiautomaton.

The minimum number of layers required to achieve 99%+ performance reflects our beliefs on the difficulty of the task: a high-level trend is that the semigroups which don’t contain groups (which only require memory lookups) are the easiest to learn, and among the groups, the larger non-abelian groups require more layers to learn, with the non-solvable group S5S_{5} requiring the largest depth.However, we stress that these experiments do not control for the fact that larger groups have richer supervision (for example, A5A_{5} has more informative labels than A4A_{4}), possibly accounting for the counterintuitive result that the latter requires more layers, despite being a subgroup of the former. Between the non-abelian groups, the difficulty of learning Q8Q_{8} compared to D8D_{8} (which has the same cardinality) agrees with our theoretical characterizations of the respective constant-depth shortcuts for these groups: D8D_{8} can be written as a semidirect product of smaller groups, while Q8Q_{8} cannot, so our theoretical construction of a constant-depth shortcut must embed Q8Q_{8} in a larger structure (i.e. the wreath product).

Throughout these experiments, we observe the following forms of training instability: high variance in training curves (based on initialization and random seeds for the gradient-based optimization algorithm), and negative progress (i.e. non-monotonic loss curves), even for training runs which eventually converge successfully. This is evident in Figure 3(b),(c) and in the significant difference between the maximum accuracies in Figure 8 and the median in Figure 9.

To stabilize training, we experiment with dropout and exponential moving average (EMA)We use the EMA implementation from https://github.com/fadel/pytorch_ema.. The effectiveness of dropout varies across datasets; for example, we find using a dropout of 0.1 (the best among {0,0.1,0.2,0.3}\{0,0.1,0.2,0.3\}) to be helpful for Dihedral and Quaternion, while such dropout hurts the training of Dyck and Gridworld. We find EMA to be generally useful, and fix the decay parameter γ=0.9\gamma=0.9 in the experiments since the performance of the EMA model does not seem to be sensitive to the choice of γ∈{0.85,0.9,0.95}\gamma\in\{0.85,0.9,0.95\}. Further, increasing the patience of the learning rate scheduler can be helpful.

B.2 Section 5: Failures of shortcuts in more challenging settings

Our theoretical and main empirical findings have shown that not only do shallow non-recurrent networks subsume deeper finite-state recurrent models in theory, these shallow solutions can also be found empirically via standard gradient-based training. However, experiments in Section 4 and Appendix B are in an idealized setting, with full state supervision during training and in-distribution evaluation at test time. This section studies more challenging settings where these assumptions are relaxed. We consider training under indirect (Section B.2.1) or incomplete (Section B.2.2) state supervision, and evaluation on sequences that is out-of-distribution (Section B.2.3) or of longer lengths (Section B.2.4).

Permutations with single-element observations: We take the permutation group S5S_{5} with QQ is the set of 5!5! operations. The observation function φ:Q→{1,2,3,4,5}\varphi:Q\rightarrow\{1,2,3,4,5\} returns the first value of the permutation. For example, φ((2,1,4,3,5)=2\varphi((2,1,4,3,5)=2. We use a set of 5 generators for the experiments.

We consider two distributions on the input sequences: (1) the input is always a sequence of the form abababa⋯abababa\cdots (i.e. the process is never in the absorbing state), which is the setup in Bhattamishra et al. (2020); and (2) the input is of the form abababa⋯abababa\cdots with probability 0.5, and is some randomly drawn string of a,ba,b otherwise. Note that case (1) can be solved purely based on the positional encoding, since the label is 1 when the position is a multiple of 4 and 0 otherwise, while case (2) is more difficult since the model needs to take into account the input tokens.

We train GPT-2-like models on sequences of length 40. We use 16 layers for S5S_{5} and 8 layers for other tasks, with embedding dimension d=512d=512 and H=8H=8 attention heads. As shown in Figure 11, the model is able to achieve near-perfect in-distribution accuracies for all tasks. An interesting side finding is that the choice of positional encoding turns out to be important for both cases of (abab)∗(abab)^{*}: learning is challenging for linear encoding (i.e. pi∝ip_{i}\propto i) but is easy when using sinusoidal positional encoding, which is likely because the sinusoidal encoding naturally matches the periodicity in (abab)∗(abab)^{*}. In all other experiments, we use sinusoidal positional encodings unless otherwise noted.

Another challenge of limited supervision is that the observation sequence may be incomplete, that is, we may not be able to get supervision on the states at every time step. We consider the task of learning length 100 sequences, where the state at each position is revealed with some probability preveal∈(0,1]p_{\mathsf{reveal}}\in(0,1].

The previous subsections show positive results on learning shallow non-recurrent shortcuts with limited supervision during training, either in the form of indirect observations or incomplete observation sequences. In this section, we study challenges at test time, and evaluate Transformers on their out-of-distribution generalization performance. For this and the next subsection, the models are trained in the standard way with full state supervision. The training sequences are of length 40, where each position has an equal probability of being 0 or 1, i.e. Pr\/[σ=1]=0.5\mathop{\bf Pr\/}[\sigma=1]=0.5. At test time, the sequences of the same length as training, but the Bernoulli parameter Pr\/[σ=1]\mathop{\bf Pr\/}[\sigma=1] varies in the range {0.05,0.1,0.15,…,0.9,0.95}\{0.05,0.1,0.15,\ldots,0.9,0.95\}.

Figure 13 (left) shows the accuracy as Pr\/[σ=1]\mathop{\bf Pr\/}[\sigma=1] varies. The performance of the Transformer degrades sharply as the test distribution changes away from training, failing at out-of-distribution generalization. Given the theoretical construction of modular counters (Lemma 6), our hypothesis is that Transformer may be learning a shortcut solution that computes the parity by counting the number of 1s, and that counts less frequently seen during training will cause the model to fail. The experimental results agree with the hypothesis: as Pr\/[σ=1]\mathop{\bf Pr\/}[\sigma=1] deviates from 0.5, it is less likely for the value of the count (which concentrates around T×Pr\/[σ=1]T\times\mathop{\bf Pr\/}[\sigma=1]) to be seen during training, hence the performance degrades. In contrast, an LSTM recurrent network maintains perfect accuracy when evaluated on all values of Pr\/[σ=1]\mathop{\bf Pr\/}[\sigma=1].

We further test this hypothesis by checking how the accuracy changes as we vary the count (i.e. the number of 1s) in the input sequence. As shown in Figure 13 (right), Transformer’s performance degrades as the count moves away from the expected number during training, agreeing with the hypothesis. It might appear strange that GPT fails at a lower count more than a higher count. However, this may be because the shortcut learns a correlation between the count and the position: during training, a lower count is more likely to appear early in an input sequence, as opposed to the testing scenario where a lower count is equally likely to appear at a later part of an sequence. This is further supported by the observation that training the model with randomly shifted positions significantly improves the performance at lower counts.

We investigate one established mitigation for the out-of-distribution brittleness of non-recurrent Transformers: scratchpad training and inference. Given a sequence of inputs (σ1,…,σT)(\sigma_{1},\ldots,\sigma_{T}) and states (q1,…,qT)(q_{1},\ldots,q_{T}), in the standard (non-recurrent) sequence-to-sequence learning pipeline, the network receives σ1:T\sigma_{1:T} as input, and outputs the sequence of predictions for qtq_{t}. In scratchpad training (Nye et al., 2021, Wei et al., 2022), we instead feed the network an interleaved sequence of inputs and states (σ1,q1,σ2,q2,σ3,q3,…,qT−1,σT)(\sigma_{1},q_{1},\sigma_{2},q_{2},\sigma_{3},q_{3},\ldots,q_{T-1},\sigma_{T}) (with an appropriately expanded token vocabulary), and define the network’s state predictions to be those at the appropriately aligned positions: (q^1,⊥,q^2,⊥,…,⊥,q^T)(\hat{q}_{1},\bot,\hat{q}_{2},\bot,\ldots,\bot,\hat{q}_{T}) (where ⊥\bot denotes a position where the prediction is ignored by the loss function). During inference, we iteratively fill in the state predictions. This removes the need for the network to learn long-range dependencies in a single non-recurrent pass, by splitting it into TT sequential state prediction problems which can depend on previous predicted state q^t−1\hat{q}_{t-1}; one can think of this as a way to guide a shallow Transformer to learn the recurrent solution (i.e. explicit depth-Θ(T)\Theta(T) iteration of the state transition function), rather than a shortcut.

We note that introducing the scratchpad itself is not sufficient to remove the parallel solution, since the model can simply ignore the scratchpad positions and find the same parallel shortcut as before. The good news is that we can couple scratchpad with an explicit recency bias in the attention mechanism (Press et al., 2022) which biases the model towards putting more attention weights on closer input. Intuitively, if the model is only allowed to put attention on the current input token and the current scratchpad (which is simply the current state), then the model is forced to be recurrent; recency bias can be considered as a soft relaxation of the same idea. Combining scratchpad and recency bias, we are able to train a Transformer to learn the recurrent solution, which is resilient to distribution shift; see Figure 13 (left). Notice that this mitigation completely foregoes the computational advantage of a shallow shortcut; we leave it to future work to obtain shortcuts which are resilient to distribution shift. Towards this, the constructions used in the proof of Theorem 1 may be helpful. Finally as a side note, even though the state transitions are Markov, the dependency in the input sequence can still be long range, so we do not expect recency bias to help without scratchpad, since in this case the output can depend uniformly on each input positions (e.g. consider parity).

Figure 14 shows the performance on sequences of various lengths. In contrast to LSTM’s perfect performance on all scenarios, Transformer’s accuracy drops sharply as we move to lengths unseen during training. This is not purely due to unseen values of the positional encoding: randomly shifting the positions during training can cover all the positions seen during testing, which helps improve the length generalization performance but cannot make it perfect; we see similar results for removing positional encodings altogether. However, similar to the OOD setup in the previous subsection, we empirically show that the above flaws are circumventable. Using a combination of scratchpad (a.k.a. “chain-of-thought”) (Nye et al., 2021, Wei et al., 2022) and recency bias (Press et al., 2022), we demonstrate that Transformers can be guided towards learning recurrent (depth-TT) solutions, which generalize out-of-distribution and to longer sequence lengths (Figure 14, yellow curves). The results also confirm that the inclusion of recency bias is necessary: without it, scratchpad training shows no improvement on length generalization.

Figure 14 also shows some interesting findings related to positional encoding, which is believed to be a key component for Transformers and a topic with active research (Ke et al., 2020, Chu et al., 2021). While this work does not aim to improve positional encoding, some of our results may be of interest for future research.

Training with shifted positions: In general, unseen positions appear to be a major contributor to Transformer’s failure of length generalization. This is evidenced by the comparison between Transformer trained with absolute positional encoding, and Transformers trained with random shifts added to the positional encoding: for each batch, we sample a random positive integer in and add it to the position indices before calculating the positional encoding; this random integer is the same for each batch and varies across batches. Figure 14 shows that adding such random shifts gives a significant boost to Transformer’s length generalization performance, for both Dyck and C2C_{2}. This suggests that a main challenge to length generalization is the distribution shifts due to positions unseen during training, and finding better positional encoding could be a potential remedy for poor length generalization.

As a side note, we also find that removing positional encoding altogether helps improve generalization for both parity and Dyck. For the former, removing positional encodings makes sense since parity is a symmetric function where the ordering of the arguments does not matter, Empirically, we are able to achieve non-trivial accuracy (even when evaluated at the sequence level) without positional encoding, whereas Bhattamishra et al. (2020) reports 0 accuracy. The discrepancy may be due to different model size: Bhattamishra et al. considers Transformers with up to 4 layers, 4 heads and dimension up to 32, whereas for the parity experiments we consider Transformers with 8 layers, 8 heads, and dimension 512. though the positive result for Dyck is less clearly understood. Note that removing positional encoding does not mean having no position information, since the use of the causal mask implicitly encodes the position, which is also noted in Bhattamishra et al. (2020) and concurrent work by Haviv et al. (2022). Understanding this phenomenon is tangential to the current work and is left to future work.

B.3 Additional details

For GPT-2 models, we fix the embedding dimension and MLP width to 512 and the number of heads to 8 in all experiments in Section 4, and vary the number of layers from 1 to 16. For LSTM, we fix the embedding dimension to 64, the hidden dimension to 128, and the number of layers to 1. We use the AdamW optimizer (Loshchilov and Hutter, 2017), with learning rate in {3e-5, 1e-4, 3e-4} for GPT-2 or {1e-3, 3e-3} for LSTM, weight decay 1e-4 for GPT-2 or 1e-9 for LSTM, and batch size 16 for GPT-2 or 64 for LSTM. As detailed in Section B.1.1, the models are trained in an online fashion with freshly drawn samples in each batch. The number of freshly drawn samples ranges from 600k to 5000k for different datasets, which is much fewer than the number of possible strings of length 100.

Our experiments are implemented with PyTorch (Paszke et al., 2019). The Transformers architectures are taken from the HuggingFace Transformers library (Wolf et al., 2019), using the GPT-2 configuration as a base. The LSTM architecture is the default one provided by the PyTorch library.

The experiments were performed on an internal cluster with NVIDIA Tesla P40, P100, V100, and A100 GPUs. For the experiments in Section 4, each training run took up to 1010 hours on a single GPU, for a total of ≈104\approx 10^{4} GPU hours. The remaining experiments in Section 5 amount to less than 1%1\% of this expenditure.

Appendix C Proofs

We first recall the notions of simulation introduced in Section 2:

A function can simulate an automaton for particular choices of T,q0T,q_{0}. For a semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta), a function f:ΣT→QTf:\Sigma^{T}\rightarrow Q^{T} simulates AT,q0{\mathcal{A}}_{T,q_{0}} if f(σ1:T)=AT,q0(σ1:T)f(\sigma_{1:T})={\mathcal{A}}_{T,q_{0}}(\sigma_{1:T}) for all input sequences σ1:T\sigma_{1:T}. Here, the right-hand side denotes the sequence of states q1:Tq_{1:T} induced by the input sequence σ1:T\sigma_{1:T} under the transitions δ\delta starting from state q0q_{0}.

A function class can simulate multiple functions associated with a semiautomaton. For a semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) and a positive integer TT, a function class F{\mathcal{F}} (a set of functions f:ΣT→QTf:\Sigma^{T}\rightarrow Q^{T}) simulates A{\mathcal{A}} at length TT if, for every q0∈Qq_{0}\in Q, there is function fq0∈Ff_{q_{0}}\in{\mathcal{F}} which simulates AT,q0{\mathcal{A}}_{T,q_{0}}.

When WW is a linear threshold function z↦arg max⁡q[Wz]qz\mapsto\operatorname*{arg\,max}_{q}[Wz]_{q}, this corresponds to a standard classification head. However, our constructions may leverage other encodings of discrete objects.

We provide some simple function approximation results below.

The inner dimension is d′=4∣X∣d^{\prime}=4|{\mathcal{X}}|, and the weights satisfy

For each x0∈Xx_{0}\in{\mathcal{X}}, we construct an indicator ψx0(x)\psi_{x_{0}}(x) for x0x_{0}, out of 4 ReLU units. Letting Δ′:=Δ/4\Delta^{\prime}:=\Delta/4, the construction is

The second layer simply sums these indicators, weighted by each f(x0)f(x_{0}). ∎

Letting Xi{\mathcal{X}}_{i} denote the set of unique values in coordinate ii, the inner MLP dimensions are as follows:

When we apply Lemmas 1 and 2 in recursive constructions, and Bx/Δ≥1B_{x}/\Delta\geq 1, we will opt to use the bound ∥b1∥∞≤6Bx/Δ\left\|b_{1}\right\|_{\infty}\leq 6B_{x}/\Delta, to reduce the clutter of propagating the 22 term without resorting to asymptotic notation.

The inner dimension is d′=2d^{\prime}=2, and the weights satisfy

We construct the threshold using 2 ReLU units. The construction is

We record some useful lemmas pertaining to approximating hard coordinate selection with soft attention. The following is a simplified version of Lemma B.7 from (Edelman et al., 2022) (which generalizes this to multi-index selection):

Let t∗:=arg max⁡tztt^{*}:=\operatorname*{arg\,max}_{t}z_{t}. Suppose that for all t′≠t∗t^{\prime}\neq t^{*}, zt′≤zt∗−γz_{t^{\prime}}\leq z_{t^{*}}-\gamma. Then,

Without loss of generality, max⁡z=γ\max z=\gamma (since the softmax function is invariant under shifting all inputs by the same value), so that all other coordinates are non-positive. Also, assume T≥2T\geq 2 (the T=1T=1 case is trivial). We have

Thus, the 1-norm of the difference is bounded by

We note the following elementary fact about 2-dimensional circular embeddings.

Consider p1,…,pTp_{1},\ldots,p_{T}, the TT equally-spaced points on the 22-dimensional circle:

C.2 Proof of Theorem 1: Logarithmic-depth shortcuts via parallel prefix sum

In this section, we give the full statement and proof of the universal existence of logarithmic-depth shortcuts.

Let A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta) be a semiautomaton, q0∈Qq_{0}\in Q, and T≥1T\geq 1. Then, there is a depth-⌈log⁡2T⌉\lceil\log_{2}T\rceil Transformer which continuously simulates AT,q0{\mathcal{A}}_{T,q_{0}}, with embedding dimension 2∣Q∣+22|Q|+2, MLP width ∣Q∣2+∣Q∣|Q|^{2}+|Q|, and ∞\infty-weight norms at most max⁡{4∣Q∣+2,10Tlog⁡∣Q∣+log⁡T}\max\{4|Q|+2,10T\sqrt{\log|Q|+\log T}\}. It has H=2H=2 heads with embedding dimension ∣Q∣|Q| implying 2∣Q∣+22|Q|+2 attention width, and a 33-layer MLP.

The basic idea is that all prefix compositions δ(⋅,σt)∘…∘δ(⋅,σ1)\delta(\cdot,\sigma_{t})\circ\ldots\circ\delta(\cdot,\sigma_{1}) can be evaluated in logarithmic depth using a binary tree whose leaves are the per-input transition functions δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\rightarrow Q. The attention heads select the pairs of functions that need to be composed, while the feedforward networks implement function composition. The network will manipulate functions in terms of their transition maps: for example, the encoding of f:=(1↦1,2↦1,3↦2)f:=(1\mapsto 1,2\mapsto 1,3\mapsto 2) is

We will produce a construction for the case where TT is a power of 2; general TT can be handled via padding. To simplify the construction, we also introduce TT padding positions −(T−1),…,0-(T-1),\ldots,0 at the beginning; while this greatly simplifies the positional selection construction, this padding construction could be replaced with a slightly more complicated MLP. Also, in this construction, we do not need to use residual connections; the parallel prefix sum algorithm we use can be executed “in place”, saving a logarithmic factor in the width. We do assume access to the 2 positional embeddings at each layer; in the absence of residual connections, the identity function restricted to these 2 dimensions can be implemented by the MLP and attention heads.

Let L=log⁡2TL=\log_{2}T be the depth of the binary tree. We choose d:=2∣Q∣+2d:=2|Q|+2. Instead of indexing the dimensions by [d][d], we give them names:

Left function encoding dimensions (q,L)(q,\mathsf{L}) for each q∈Qq\in Q.

Right function encoding dimensions (q,R)(q,\mathsf{R}) for each q∈Qq\in Q.

Positional encoding dimensions P1,P2\mathsf{P}_{1},\mathsf{P}_{2}.

Without loss of generality, let Q=[∣Q∣]={1,…,Q}Q=[|Q|]=\{1,\ldots,Q\} (selecting an arbitrary enumeration of the state space). Also, assume ∣Q∣≥2|Q|\geq 2 (if not, add a dummy state). We choose E(σt):=∑q∈Qδ(q,σt)⋅e(q,R)E(\sigma_{t}):=\sum_{q\in Q}\delta(q,\sigma_{t})\cdot e_{(q,\mathsf{R})}, mapping each input symbol to the “transition map” of its transitions. At the padding positions −(T−1),…,0-(T-1),\ldots,0, we will encode the “go to q0q_{0}” function: ∑q∈Qq0⋅e(q,R)\sum_{q\in Q}q_{0}\cdot e_{(q,\mathsf{R})}.

We first introduce the construction for function composition with a 3-layer ReLU MLP, which will be used by all layers. It gives an exponential improvement over the generic universal function approximation gadget from Lemma 2.

The intermediate dimensions are d1=∣Q∣2+∣Q∣d_{1}=|Q|^{2}+|Q| and d2=∣Q∣2d_{2}=|Q|^{2}, and weight norms are bounded by 4∣Q∣+24|Q|+2.

We also add QQ more weights which let the inputs pass through along the e(q,R)e_{(q,\mathsf{R})} directions (add QQ more rows e(q,R)⊤e_{(q,\mathsf{R})}^{\top} to W1,W2′W_{1},W_{2}^{\prime}, calling these indices ∙q\bullet q for all q∈Qq\in Q; set biases to 0), for a total of 4∣Q∣2+∣Q∣4|Q|^{2}+|Q| hidden units and ∣Q∣2+∣Q∣|Q|^{2}+|Q| output dimensions of W2′W_{2}^{\prime}.

The rest of the construction uses a standard parallel algorithm for computing all prefix function compositions: at layer l∈[L]l\in[L], compose the function at position tt with the function at position t−2l−1t-2^{l-1}. This is a standard algorithm for computing all prefix compositions of associative binary operations with a logarithmic-depth circuit (Hillis and Steele Jr., 1986). We choose the position embeddings to enable implementing these “look-backs” with rotation matrices. For each t∈{−T+1,…,0,1,…,T}t\in\{-T+1,\ldots,0,1,\ldots,T\}, we use the circle embeddings

Let θ:=−π2l−1T\theta:=-\frac{\pi 2^{l-1}}{T}, γ:=100T2(log⁡∣Q∣+log⁡T)\gamma:=100T^{2}(\log|Q|+\log T).

Select WQ[L]=WQ[R]=WK[R]:=γ⋅(eP1e1⊤+eP2e2⊤)W_{Q}^{[\mathsf{L}]}=W_{Q}^{[\mathsf{R}]}=W_{K}^{[\mathsf{R}]}:=\sqrt{\gamma}\cdot(e_{\mathsf{P}_{1}}e_{1}^{\top}+e_{\mathsf{P}_{2}}e_{2}^{\top}).

Select WK[L]:=γ⋅(eP1e1⊤+eP2e2⊤)ρθW_{K}^{[\mathsf{L}]}:=\sqrt{\gamma}\cdot(e_{\mathsf{P}_{1}}e_{1}^{\top}+e_{\mathsf{P}_{2}}e_{2}^{\top})\rho_{\theta}, where ρθ\rho_{\theta} is the rotation matrix

Select WV[L]=WV[R]:=∑q∈Qe(q,L)eq⊤W_{V}^{[\mathsf{L}]}=W_{V}^{[\mathsf{R}]}:=\sum_{q\in Q}e_{(q,\mathsf{L})}e_{q}^{\top}.

Select WC[L]=∑q∈Qeqe(q,L)⊤W_{C}^{[\mathsf{L}]}=\sum_{q\in Q}e_{q}e_{(q,\mathsf{L})}^{\top}, WC[R]=∑q∈Qeqe(q,R)⊤W_{C}^{[\mathsf{R}]}=\sum_{q\in Q}e_{q}e_{(q,\mathsf{R})}^{\top}.

Thus, at the final layer, the (q,R)(q,\mathsf{R}) dimensions at position tt contains the transition map of the prefix composition

It suffices to choose WW to be z↦e(q,R)⊤zz\mapsto e_{(q,\mathsf{R})}^{\top}z for an arbitrary qq to read out the sequence of states as scalar outputs in [∣Q∣][|Q|]. To output a one-hot encoding, an additional MLP (appended to the end of the final layer) would be required. ∎

C.3 Proof of Theorem 2: Constant-depth shortcuts via Krohn-Rhodes decomposition

We begin with the full statement of the theorem:

Let A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta) be a solvable semiautomaton (see Definition 6), q0∈Qq_{0}\in Q, and T≥1T\geq 1. Then, there is a depth-O(∣Q∣2log⁡∣Q∣)O(|Q|^{2}\log|Q|) Transformer which continuously simulates AT,q0{\mathcal{A}}_{T,q_{0}}, with embedding dimension O(2∣Q∣∣T(A)∣)O(2^{|Q|}|{\mathcal{T}}({\mathcal{A}})|), MLP width ∣Q∣O(2∣Q∣)+O(2∣Q∣∣Q∣ ∣T(A)∣ T)|Q|^{O(2^{|Q|})}+O(2^{|Q|}|Q|\,|{\mathcal{T}}({\mathcal{A}})|\,T), attention width O(∣Q∣2∣Q∣∣T(A)∣)O(|Q|2^{|Q|}|{\mathcal{T}}({\mathcal{A}})|) heads, and weight norms are bounded by 6∣Q∣ Tlog⁡T+6max⁡{∣Q∣,∣Σ∣}6|Q|\,T\log T+6\max\{|Q|,|\Sigma|\}.

We will begin by presenting self-contained constructions for the two atoms in the Krohn-Rhodes decomposition: a modular counter and a memory unit. In Appendix C.3.2, we will introduce necessary background from Krohn-Rhodes theory (including the definition of a solvable semiautomaton). In Appendix C.3.3 and C.3.4, we will complete the proof of Theorem 2.

We will start with a construction of a tiny network which lets us simulate any semiautomaton whose transformation semigroup is a cyclic group. Later on, we will use copies of this unit to handle all solvable groups. The construction simply uses attention to perform a flat prefix sum, and an MLP to compute the modular sum.

For any positive integer nn, define the mod-nn modular counter semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta):

Let A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta) be the mod-nn modular counter semiautomaton. Let q0∈Qq_{0}\in Q, and T≥1T\geq 1. Then, there is a depth-11 Transformer which continuously simulates AT,q0{\mathcal{A}}_{T,q_{0}}, with embedding dimension 33, width 4nT4nT, and ∞\infty-weight norms at most 4nT+24nT+2. It has H=1H=1 head with embedding dimension k=1k=1, and a 2-layer ReLU MLP.

The intuition is simply that the lower triangular matrix causal mask can implement simulation in this cyclic group by performing unweighted prefix sums. The only subtlety is that selecting WQ=WK=0W_{Q}=W_{K}=0 does not quite give us prefix sums: the attention mixture weights at position tt are 1t∑t′∈[t]et′\frac{1}{t}\sum_{t^{\prime}\in[t]}e_{t^{\prime}}, while we would like the normalizing factor to be uniform across positions (1/T1/T rather than 1/t1/t). It is possible to undo this normalization using the MLP; however, a particulaly simple solution is to use an additional padding input ⊥\bot and 1-dimensional position embeddings to “absorb” a fraction of the attention proportional to 1−t/T1-t/T.

We proceed to formalize this construction, beginning with the input embedding and attention block:

Select d:=3,k:=1,H:=1d:=3,k:=1,H:=1. Intuitively, the 3 dimensions implement {\{input/output, padding, position}\} “channels”.

Include an extra position ⊥\bot, with embedding E(⊥):=e2E(\bot):=e_{2} and position encoding P⊥,::=0P_{\bot,:}:=0. Think of this as padding at position ; it is not masked out by the causal attention mask at any position t≥1t\geq 1.

For t∈[T]t\in[T], select Pt,::=γte3P_{t,:}:=\gamma_{t}e_{3}, where γt:=log⁡(2T−t)\gamma_{t}:=\log(2T-t) is such that 1eγt+t=12T\frac{1}{e^{\gamma_{t}}+t}=\frac{1}{2T}.

Select WQ:=e3,WK:=e2,WV:=e1,WC⊤:=e1W_{Q}:=e_{3},W_{K}:=e_{2},W_{V}:=e_{1},W_{C}^{\top}:=e_{1}.

In the output of this attention module, for any input sequence σ1:T\sigma_{1:T}, the 1st channel of the output at position tt is then

where zt∈{0,…,n−1}z_{t}\in\{0,\ldots,n-1\} is such that δ(⋅,σt)=gzt\delta(\cdot,\sigma_{t})=g^{z_{t}}. The MLP simply needs to memorize the function

We invoke Lemma 1, with Δ=12T,Bx=n−12<n2,By=n\Delta=\frac{1}{2T},B_{x}=\frac{n-1}{2}<\frac{n}{2},B_{y}=n. The number of possible values of SS (the cardinality of X{\mathcal{X}} in Lemma 1) is at most nTnT. ∎

It turns out that to simulate semigroups instead of groups, the only additional ingredient is a memory unit, a semiautomaton for which there are “read” and “write” operations. The minimal example of this is a flip-flop (Example 2), a semiautomaton which can sequentially remember and retrieve a single bit ∈{0,1}\in\{0,1\}, and whose transformation semigroup is the flip-flop monoid. It will be convenient to generalize this object to QQ states:

For a given state set QQ, define the memory semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta):

Let A=(Q,Σ,δ)\mathcal{A}=(Q,\Sigma,\delta) be the memory semiautomaton. Let q0∈Qq_{0}\in Q, and T≥1T\geq 1. Then, there is a depth-11 Transformer which continuously simulates AT,q0{\mathcal{A}}_{T,q_{0}}, with embedding dimension 44, width 4∣Q∣4|Q|, and ∞\infty-weight norms at most 2Tlog⁡(∣Q∣T)2T\log(|Q|T). It has H=1H=1 head with embedding dimension k=2k=2, and a 2-layer ReLU MLP.

We start in state q0∈Qq_{0}\in Q. Our goal is to identify the closest non-⊥\bot token and output the corresponding state. The attention construction is:

where the first coordinate denotes the action that sets the stateTechnically σ=⊥\sigma=\bot does not reset the state. We will see that when q0q_{0} is selected, it must be that the semiautomaton is always in state q0q_{0}., the second coordinate denotes whether the input is the no-op action ⊥\bot, and the fourth coordinate is padding.

We use positional encoding Pt,::=(t/T)⋅e3P_{t,:}:=(t/T)\cdot e_{3}.

Denote this max position as jmax⁡j_{\max}. In the setting of hard attention, the output for the ithi_{th} token after the attention module is E(σjmax⁡)⊤e1E(\sigma_{j_{\max}})^{\top}e_{1}. In particular, this value is q0q_{0} if and only if σj=⊥,∀j≤i\sigma_{j}=\bot,\forall j\leq i, i.e. the semiautomaton never leaves the starting state. Otherwise, the value is the value of the nearest non-⊥\bot state (including the current state).

The key idea behind the proof of Theorem 2 is that all semigroups (and thus, all transformation semigroups of semiautomata) admit a “prime factorization” into elementary components, which turn out to be simple groups and copies of the flip-flop monoid, which have both been discussed in Appendix A.2. This is somewhat counterintuitive: the only constraint on the algebraic structure of a semigroup is associativity (and indeed, there are many more semigroups than groups), but all of these structures can be built using these two types of “atoms”. These components, as well as the cascade product under which this notion of “factorization” is defined, are naturally and efficiently implementable by constant-depth self-attention networks.

We begin by discussing the analogous decomposition for groups, which generalizes the fact that integers have unique prime factorizations. Let GG be a finite group, and let

be a composition series: each HiH_{i} is a maximal proper normal subgroup of Hi+1H_{i+1}; 11 denotes the trivial group with 11 element. Then the quotient group Hi+1/HiH_{i+1}/H_{i} is called a composition factor. The Jordan-Hölder theorem tells us that one can think about the set of composition factors as an invariant of GG.

Any two composition series of GG are equivalent: they have the same length nn, and the sequences of compositions factors Hi+1/HiH_{i+1}/H_{i} are equivalent under permutation and isomorphism.

When each Hi+1/HiH_{i+1}/H_{i} is abelian, GG is called a solvable group. It turns out that each Hi+1/HiH_{i+1}/H_{i} is a simple group, so the composition factors of solvable groups can only be cyclic groups of prime order (because every finitely generated abelian group is a direct product of cyclic groups, and, of these, only those of prime order are simple). The smallest non-solvable group is A5A_{5}, realizable as the group of even permutations of 55 elements. As a part of Theorem 2, we will use the composition series to iteratively build neural networks which simulates solvable group operations, requiring intricate constructions to do this with depth independent of the sequence length TT.

Now, we move on to semigroups. When not all of the input symbols to a semiautomaton induce permutations, we no longer have the group axiom of invertibility (also, if there is no explicit identity symbol, we are not guaranteed to have the monoid axiom of an identity element either). Intuitively, this would seem to induce a much larger family of algebraic structures; an analogy, which is formalizable by representation theory, is that we are now considering a collection of general matrices under multiplication, instead of only invertible ones. The non-invertible transitions collapse the rank of the transformations, reducing the set of reachable transformations whenever they are included in an input sequence.

A landmark result of Krohn and Rhodes (1965) tames the seemingly vast and unorderly universe of general finite semigroups. It extends the Jordan-Hölder theorem to the case of semigroups, for a more sophisticated notion of decomposition. Since that work, many variations have arisen, in terms of its precise statement, construction of the decomposition, and proof of correctness. Out of these, an important development is the holonomy decomposition method (Zeiger, 1967, Eilenberg, 1974), which forms the basis of our results. We extract the definitions and theorems from Maler and Pnueli (1994), whose exposition emphasizes explicitly tracking the construction of the semiautomaton. We also refer to Maler (2010), Egri-Nagy and Nehaniv (2015), Zimmermann (2020) as recent expositions, containing historical context.

Let nn be a positive integer. For each i∈[n]i\in[n], let A(i)=(Q(i),Σ(i),δ(i)){\mathcal{A}}^{(i)}=(Q^{(i)},\Sigma^{(i)},\delta^{(i)}) be a semiautomaton. For i∈{2,…,n}i\in\{2,\ldots,n\}, let ϕ(i):Q(1)×⋯×Q(i−1)×Σ→Σ(i)\phi^{(i)}:Q^{(1)}\times\cdots\times Q^{(i-1)}\times\Sigma\rightarrow\Sigma^{(i)} denote a dependency function. This object ({A(i)};{ϕ(i)}\{{\mathcal{A}}^{(i)}\};\{\phi^{(i)}\}) is called a transformation cascade, and defines a cascade product semiautomaton A=(Q(1)×⋯×Q(n),Σ(1),δ){\mathcal{A}}=(Q^{(1)}\times\cdots\times Q^{(n)},\Sigma^{(1)},\delta) by “feedforward simulation” under the dependency function. We define δ\delta by the ii-th component of its output (which we call δ(≤i):Q(1)×⋯×Q(n)×Σ(1):Q(i)\delta^{(\leq i)}:Q^{(1)}\times\cdots\times Q^{(n)}\times\Sigma^{(1)}:Q^{(i)}):

The corresponding transformation semigroup T(A){\mathcal{T}}({\mathcal{A}}) is known as a cascade product semigroup.

Intuitively, the cascade specifies a way to compose semiautomata hierarchically: the first layer i=1i=1 maps input sequences to its state sequence, and each internal layer receives an input which depends on the states of all of the preceding layers. Algebraically, the cascade product semigroup is a subsemigroup of the larger wreath product of semigroups (the straightforward analogue of the wreath product of groups, discussed in Section A.2). Although this is useful from an algebraic point of view, we will not use this perspective; the cascade product is a smaller substructure of the wreath product which is sufficient for semiautomaton simulation.

Finally, it will be convenient to define permutation-reset semiautomata, which are a useful intermediate step in the Krohn-Rhodes decomposition. To obtain our final result, we will further break these semiautomata down into flip-flops and simple groups.

A semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) is a permutation-reset semiautomaton if, for each σ∈Σ\sigma\in\Sigma, the transition function δ(⋅,σ):Q→Q\delta(\cdot,\sigma):Q\rightarrow Q is either a bijection (i.e. a permutation over the states of QQ) or constant (i.e. maps every state to some q(σ)q(\sigma)). Associated with each permutation-reset semiautomaton is its permutation group, generated by only the bijections.

Now we can state the Krohn-Rhodes theorem, which decomposes every finite semiautomaton into a transformation cascade.

Let A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) be a semiautomaton. Then, there exists a transformation cascade {A(1),…,A(n);ϕ(2),…,ϕ(n)}\{{\mathcal{A}}^{(1)},\ldots,{\mathcal{A}}^{(n)};\phi^{(2)},\ldots,\phi^{(n)}\}, defining a cascade product semiautomaton A′{\mathcal{A}}^{\prime}, such that:

The input symbol space of A(1){\mathcal{A}}^{(1)} (and thus, that of A′{\mathcal{A}}^{\prime}) is Σ\Sigma, the same as that of A{\mathcal{A}}.

Letting Q(i)Q^{(i)} denote the state space of A(i){\mathcal{A}}^{(i)}, there exists a function W:Q(1)×⋯×Q(n)→Q{\mathcal{W}}:Q^{(1)}\times\cdots\times Q^{(n)}\rightarrow Q such that W∘AT,q0′{\mathcal{W}}\circ{\mathcal{A}}^{\prime}_{T,q_{0}} simulates AT,q0{\mathcal{A}}_{T,q_{0}} for all T≥1,q0∈QT\geq 1,q_{0}\in Q. For each i∈[n]i\in[n], the transformation semigroup T(A(i)){\mathcal{T}}({\mathcal{A}}^{(i)}) is a permutation-reset semiautomaton with at most ∣Q∣|Q| states, whose permutation group is a (possibly trivial) subgroup of T(A){\mathcal{T}}({\mathcal{A}}) (Maler and Pnueli (1994), Theorem 4).

The number of semiautomata in the cascade is n≤2∣Q∣n\leq 2^{|Q|}. Furthermore, the cascade has at most ∣Q∣|Q| levels: the indices can be partitioned into at most L≤∣Q∣L\leq|Q| contiguous subsets N(1)={1,…,n1},N(2),{n1+1,…,n1+n2},…,N(L)={n−nL+1,…,n}N^{(1)}=\{1,\ldots,n_{1}\},N^{(2)},\{n_{1}+1,\ldots,n_{1}+n_{2}\},\ldots,N^{(L)}=\{n-n_{L}+1,\ldots,n\} such that ϕ(i)\phi^{(i)} only depends on input indices from previous partitions (Maler and Pnueli (1994), Claim 11 & Corollary 12).

With this decomposition, we are now able to define a solvable semiautomaton.

Let A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) be a semiautomaton. We call A{\mathcal{A}} solvable if all the permutation groups associated with all of the permutation-reset automata from Theorem 7 are solvable groups.

The remainder of this section will build our construction from the bottom up:

Appendix C.3.3 will build up from the base case of cyclic groups (Lemma 6), using increasingly sophisticated notions of group products, culminating in a recursive construction which simulates all stages of the Jordan-Hölder composition series. The crucial step is a construction for simulating the semidirect product of groups, given networks which simulate the individual components; this allows us to handle the solvable non-abelian groups.

Appendix C.3.4 will build networks which simulate permutation-reset semiautomata. A new base case arises: the memory unit (Lemma 7), a semiautomaton whose transformation semigroup is a generalization of the flip-flop monoid. Combining the constructions for solvable groups and memory units, we obtain simulators for solvable permutation-reset semiautomata. Finally, the cascade product guaranteed by Krohn-Rhodes (Theorem 7) glues all of these pieces together, giving us the final result.

We begin by handling groups. Now, we are ready to specify the recursive constructions which “glue” these components together to form solvable groups. We will proceed in a “bottom-up” order:

Define a canonical semiautomaton AG{\mathcal{A}}^{G} corresponding to each group GG (Definition 7), such that if a network can simulate AG{\mathcal{A}}^{G}, it can simulate any other semiautomaton whose transformation semigroup T(AG){\mathcal{T}}({\mathcal{A}}^{G}) is GG. This lets us talk about simulating groups, rather than particular semiautomata. We will show how to turn simulators for groups NN and HH into simulators for extensions of NN by HH, for increasingly sophisticated extensions, until all cases have been captured.

Show how to build the trivial extension: given networks which simulate the groups NN and HH, simulate the direct product G≅N×HG\cong N\times H, by simply running the individual simulators in parallel (Lemma 8). Combined with Lemma 6, this immediately allows us to simulate arbitrary abelian groups with depth 11, since every abelian group is isomorphic to a direct product of cyclic groups.

Show how to build a split extension: given networks which simulate a normal subgroup NN and quotient HH, construct a network which simulates any semidirect product G≅N⋊HG\cong N\rtimes H (Lemma 9). This is the first place where we will require a sequential cascade of layers. It will allow us to handle certain families of non-abelian groups (including S3,D2n,A4,S4S_{3},D_{2n},A_{4},S_{4})

Show how to build arbitrary extensions (any GG which contains NN as a normal subgroup, and for which the quotient group G/NG/N is isomorphic to HH), using the wreath product (Lemma 10), which contains all of the group extensions. The wreath product is itself the semidirect product between a ∣H∣|H|-way direct product and HH, so this can be done in a constant number of layers, by the above. This finally lets us implement any step of a composition series. In particular, using cyclic groups as a simulable base case, this shows that we can simulate all solvable groups.

It will be convenient to associate with each group GG a canonical “complete” semiautomaton for the class of all semiautomata A{\mathcal{A}} for which T(A)=G{\mathcal{T}}({\mathcal{A}})=G. It is simply the one whose input symbol space Σ\Sigma is every transformation reachable by some sequence of inputs (i.e. every element of T(A){\mathcal{T}}({\mathcal{A}})). (For a semigroup, we would also want to adjoin the identity element if it is missing, however, we will only find it useful to define this for groups.)

Let GG be a finite group. Then, we define the canonical group semiautomaton for GG as the semiautomaton (Q,Σ,δ)(Q,\Sigma,\delta) defined by:

Q:=GQ:=G, the set of elements of GG. Note that if (for example) G=SnG=S_{n}, we are setting the state space to be the set of n!n! permutations, not the ground set [n][n].

Σ:=G\Sigma:=G. That is, we include all functions in the input symbol space.

δ(g,h):=h⋅g\delta(g,h):=h\cdot g, for all ∀g∈Q,h∈Σ\forall g\in Q,h\in\Sigma. (In algebraic terms, we are embedding the GG into its left regular representation, a.k.a. left multiplication action.) Thus, if we take q0q_{0} to be the identity element, the sequence of states q1,q2,…,qTq_{1},q_{2},\ldots,q_{T} corresponds to qt=σtσt−1…σ1q_{t}=\sigma_{t}\sigma_{t-1}\ldots\sigma_{1}.

When we simulate the canonical group semiautomaton, we will always choose q0q_{0} to be the identity element eGe_{G}.

A sequence-to-sequence network is said to continuously simulate GG at length TT if it continuously simulates the canonical group semiautomaton of GG at length TT.

To reduce notational clutter, we will access the shape attributes of an implementation via “object-oriented” notation, defining

to respectively denote the complexity-parameterizing quantities

We also enforce that throughout our constructions of networks which simulate groups, we will maintain that all networks and their submodules manipulate encodings via integer vectors in a consecutive range {0,…,n−1}\{0,\ldots,n-1\}. Furthermore, the identity element will always map to the zero vector. We will keep track of the dimensionality of these vectors sim(G,T).repDim≤d\mathsf{sim}(G,T).\mathsf{repDim}\leq d, and their maximum entries sim(G,T).repSize−1\mathsf{sim}(G,T).\mathsf{repSize}-1. All encoders EE and decoders WW will map all group elements to and from this kind of representation, and we will choose W=E−1W=E^{-1}. In all, the networks will keep a repDim\mathsf{repDim}-dimensional “workspace” of integer vectors, with entries bounded by repSize−1\mathsf{repSize}-1. When combining groups via the various products constructions, we will combine the components’ individual workspaces to create a larger workspace for the product group’s elements.

We make some additional remarks on implementations:

Note that the canonical semiautomaton “forgets” about the semiautomaton abstraction, and never assumes that GG is a permutation group on the original state space QQ of the semiautomaton we would like to simulate. Indeed, when N◃GN\triangleleft G are permutation groups on QQ, there is no natural permutation group on QQ associated with the quotient H≅G/NH\cong G/N; it turns out that will consider simulators for NN and HH.There is nothing in general preventing quotient groups from being extremely large groups which are not realizable as smaller permutation groups. For concrete examples, see (Kovács and Praeger, 1989). When we ultimately specialize to simulating the composition series of solvable groups, the largest groups we will handle will be the cyclic groups of prime order, so we will in the end be guaranteed that the groups we want to simulate are realizable with ≤∣Q∣\leq|Q| states, but not directly or canonically.

To return to solving the simulation problem for some semiautomaton A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) whose transformation semigroup is isomorphic to GG (at length TT and initial state q0q_{0}), let μ:G→SQ\mu:G\rightarrow S_{Q} denote this isomorphism. We use AT,eGG(σ1:T){\mathcal{A}}^{G}_{T,e_{G}}(\sigma_{1:T}) as the network, with an encoding layer E∘μ−1E\circ\mu^{-1}, and decoding layer (π↦π(q0))∘μ∘W(\pi\mapsto\pi(q_{0}))\circ\mu\circ W, which can be memorized by an MLP of width O(∣G∣)O(|G|) via Lemma 2.

The modular counter semiautomaton, for which we constructed a simulator in Lemma 6, is the canonical group semiautomaton for the corresponding cyclic group CnC_{n}. Calling this construction simCn\mathsf{sim}_{C_{n}}, we can easily verify that it satisfies the canonical simulator’s conditions, and:

simCn.headDim=1\mathsf{sim}_{C_{n}}.\mathsf{headDim}=1.

simCn.mlpWidth=4∣G∣⋅T\mathsf{sim}_{C_{n}}.\mathsf{mlpWidth}=4|G|\cdot T.

simCn.normBound≤4∣G∣⋅T+2≤6∣G∣⋅T\mathsf{sim}_{C_{n}}.\mathsf{normBound}\leq 4|G|\cdot T+2\leq 6|G|\cdot T.

simCn.repDim=1\mathsf{sim}_{C_{n}}.\mathsf{repDim}=1.

simCn.repSize=∣G∣\mathsf{sim}_{C_{n}}.\mathsf{repSize}=|G|.

As a precursor to the more sophisticated products, we formalize the obvious fact that two non-interacting parallel semiautomata can be simulated without increasing the depth. First, we define the direct product semiautomaton:

Let A=(Q,Σ,δ),A′=(Q′,Σ′,δ′){\mathcal{A}}=(Q,\Sigma,\delta),{\mathcal{A}}^{\prime}=(Q^{\prime},\Sigma^{\prime},\delta^{\prime}) be two semiautomata. Then, A×A′=(Q×Q′,Σ∪{e}×Σ′∪{e},δ×δ′){\mathcal{A}}\times{\mathcal{A}}^{\prime}=(Q\times Q^{\prime},\Sigma\cup\{e\}\times\Sigma^{\prime}\cup\{e\},\delta\times\delta^{\prime}) denotes the natural direct product semiautomaton. Its states are ordered pairs (q∈Q,q′∈Q′)(q\in Q,q^{\prime}\in Q^{\prime}). Its input symbols are defined similarly, adjoining identity inputs (so that δ(q,e)=q,δ′(q′,e)=q′\delta(q,e)=q,\delta^{\prime}(q^{\prime},e)=q^{\prime}). The transitions δ×δ′\delta\times\delta^{\prime} are defined such that

Note that under this definition, we have T(A×A′)=T(A)×T(A′){\mathcal{T}}({\mathcal{A}}\times{\mathcal{A}}^{\prime})={\mathcal{T}}({\mathcal{A}})\times{\mathcal{T}}({\mathcal{A}}^{\prime}). In particular, for two groups G,HG,H, we have G×H=T(AG)×T(AH)=T(AG×H)=G×HG\times H={\mathcal{T}}({\mathcal{A}}^{G})\times{\mathcal{T}}({\mathcal{A}}^{H})={\mathcal{T}}({\mathcal{A}}^{G\times H})=G\times H.

Let G(1),…,G(n)G^{(1)},\ldots,G^{(n)} be a collection of finite groups, and let T≥1T\geq 1. Suppose each group admits a simulation simi:=sim(G(i),T)\mathsf{sim}_{i}:=\mathsf{sim}(G^{(i)},T). Then, there is a simulation of the direct product group sim×:=sim(G(1)×…×G(n),T)\mathsf{sim}_{\times}:=\mathsf{sim}(G^{(1)}\times\ldots\times G^{(n)},T), whose sizes satisfy:

sim×.depth=max⁡i{simi.depth}\mathsf{sim}_{\times}.\mathsf{depth}=\max_{i}\{\mathsf{sim}_{i}.\mathsf{depth}\}.

sim×.dim=∑i{simi.dim}\mathsf{sim}_{\times}.\mathsf{dim}=\sum_{i}\{\mathsf{sim}_{i}.\mathsf{dim}\}.

sim×.heads=∑i{simi.heads}\mathsf{sim}_{\times}.\mathsf{heads}=\sum_{i}\{\mathsf{sim}_{i}.\mathsf{heads}\}.

sim×.headDim=max⁡i{simi.headDim}\mathsf{sim}_{\times}.\mathsf{headDim}=\max_{i}\{\mathsf{sim}_{i}.\mathsf{headDim}\}.

sim×.mlpWidth=∑i{simi.mlpWidth}\mathsf{sim}_{\times}.\mathsf{mlpWidth}=\sum_{i}\{\mathsf{sim}_{i}.\mathsf{mlpWidth}\}.

sim×.normBound≤max⁡i{simi.normBound}\mathsf{sim}_{\times}.\mathsf{normBound}\leq\max_{i}\{\mathsf{sim}_{i}.\mathsf{normBound}\}.

sim×.repDim=∑i{simi.repDim}\mathsf{sim}_{\times}.\mathsf{repDim}=\sum_{i}\{\mathsf{sim}_{i}.\mathsf{repDim}\}.

sim×.repSize=max⁡i{simi.repSize}\mathsf{sim}_{\times}.\mathsf{repSize}=\max_{i}\{\mathsf{sim}_{i}.\mathsf{repSize}\}.

First, we pad all of the individual simi\mathsf{sim}_{i} with layers implementing identity (add residual connections, and set attention WVW_{V} and all MLP weight matrices to 0), so that all of them have depth max⁡i{simi.depth}\max_{i}\{\mathsf{sim}_{i}.\mathsf{depth}\}.

Then, the intuition is to construct the direct product semiautomaton by concatenating the “workspaces” of each G(i)G^{(i)}. In other words, we set the canonical encoding sim×.E\mathsf{sim}_{\times}.E of (g(1),…,g(n))(g^{(1)},\ldots,g^{(n)}) to be the concatenation of each simi\mathsf{sim}_{i}’s encodings.

The direct product simply lets each simi\mathsf{sim}_{i} take inputs and outputs in its individual workspace. To enable this, we need enough parallel dimensions. We set an embedding space of dimension ∑i{simi.dim}\sum_{i}\{\mathsf{sim}_{i}.\mathsf{dim}\} (and similarly within the heads and MLPs), partitioning the coordinates such that in the product construction, each simi.E\mathsf{sim}_{i}.E and simi.θ\mathsf{sim}_{i}.\theta only reads and writes to its own dimensions.

This clearly simulates the direct product group. Figure 16 provides a sketch of this construction. ∎

Note that the direct product construction already allows us to simulate all finite abelian groups in constant depth, since each such group is isomorphic to the direct product of a collection of abelian groups of prime power order.

Now, as a harder (and conceptually crucial) case, we show how to simulate a group which is a semidirect product of two groups we already know how to simulate. This encompasses the direct product as a special case, but can now handle some non-abelian groups which admit such decompositions (like the dihedral group D2nD_{2n}). The catch is that we will have to simulate these groups using a sequential cascade of the individual simulators. This is the key lemma which lets us simulate non-abelian groups:

Let GG be a finite group which is isomorphic to a semidirect product: G≅N⋊HG\cong N\rtimes H, where NN is a normal subgroup of GG. Let T≥1T\geq 1. Suppose N,HN,H admit simulations simN:=sim(N,T),simH:=sim(H,T)\mathsf{sim}_{N}:=\mathsf{sim}(N,T),\mathsf{sim}_{H}:=\mathsf{sim}(H,T). Then, there is a simulation of GG, sim⋊:=sim(G,T)\mathsf{sim}_{\rtimes}:=\mathsf{sim}(G,T), whose sizes satisfy:

sim⋊.depth=simN.depth+simH.depth+2\mathsf{sim}_{\rtimes}.\mathsf{depth}=\mathsf{sim}_{N}.\mathsf{depth}+\mathsf{sim}_{H}.\mathsf{depth}+2.

sim⋊.dim=simN.dim+simH.dim\mathsf{sim}_{\rtimes}.\mathsf{dim}=\mathsf{sim}_{N}.\mathsf{dim}+\mathsf{sim}_{H}.\mathsf{dim}.

sim⋊.heads=max⁡{simN.heads,simH.heads}\mathsf{sim}_{\rtimes}.\mathsf{heads}=\max\{\mathsf{sim}_{N}.\mathsf{heads},\mathsf{sim}_{H}.\mathsf{heads}\}.

sim⋊.headDim=max⁡{simN.headDim,simH.headDim}\mathsf{sim}_{\rtimes}.\mathsf{headDim}=\max\{\mathsf{sim}_{N}.\mathsf{headDim},\mathsf{sim}_{H}.\mathsf{headDim}\}.

sim⋊.mlpWidth=max⁡{sim{N,H}.mlpWidth,4∣G∣}\mathsf{sim}_{\rtimes}.\mathsf{mlpWidth}=\max\{\mathsf{sim}_{\{N,H\}}.\mathsf{mlpWidth},4|G|\}.

sim⋊.normBound≤max⁡{sim{N,H}.normBound,6 sim{N,H}.repSize,simN.repDim+simH.repDim}\mathsf{sim}_{\rtimes}.\mathsf{normBound}\leq\max\{\mathsf{sim}_{\{N,H\}}.\mathsf{normBound},6\,\mathsf{sim}_{\{N,H\}}.\mathsf{repSize},\mathsf{sim}_{N}.\mathsf{repDim}+\mathsf{sim}_{H}.\mathsf{repDim}\}.

sim⋊.repDim=simN.repDim+simH.repDim\mathsf{sim}_{\rtimes}.\mathsf{repDim}=\mathsf{sim}_{N}.\mathsf{repDim}+\mathsf{sim}_{H}.\mathsf{repDim}.

sim⋊.repSize=max⁡{simN.repSize,simH.repSize}\mathsf{sim}_{\rtimes}.\mathsf{repSize}=\max\{\mathsf{sim}_{N}.\mathsf{repSize},\mathsf{sim}_{H}.\mathsf{repSize}\}.

The intuition is as follows, using the dihedral group D2n≅Cn⋊C2D_{2n}\cong C_{n}\rtimes C_{2} as an example:

For simplicity, let us think of the “reversible car on a circular world” semiautomaton, whose transformation semigroup is D2nD_{2n}. Its state consists of a direction ∈{+1,−1}\in\{+1,-1\}, and a position ∈{0,1,…,n−1}\in\{0,1,\ldots,n-1\}. It has two types of inputs: “advance by ii” (increment the position by ii in the current direction, modulo nn), and “reverse” (flip the sign of the direction). Our simulation task is to track the car’s state sequence, given a sequence of inputs (in constant depth, of course).

It is intuitively clear that we can (and should) compute the sequence corresponding to “direction at time tt”, which is equivalent to simulating the parity semiautomaton.

We will convert the “advance” moves via a “basis transformation”: whenever the current direction is −1-1, an “advance by ii” should be converted into −i-i. Then, we have reduced the problem to the prefix sum.

Let us write down the properties of ϕ\phi:

ϕ\phi is a homomorphism. That is, ϕh⋅h′=ϕh(ϕh′(⋅))=ϕh∘ϕh′\phi_{h\cdot h^{\prime}}=\phi_{h}(\phi_{h^{\prime}}(\cdot))=\phi_{h}\circ\phi_{h^{\prime}} as permutations on NN.

The output of that homomorphism, ϕh\phi_{h}, is also a homomorphism. That is, ϕh(gg′)=ϕh(g)⋅ϕh(g′)\phi_{h}(gg^{\prime})=\phi_{h}(g)\cdot\phi_{h}(g^{\prime}).

Let us roll out the definition of the semidirect product, given a sequence of inputs (gt,ht)(g_{t},h_{t}):

In general, by induction, letting (g≤t,h≤t)(g_{\leq t},h_{\leq t}) denote (gt,ht)⋯(g1,h1)(g_{t},h_{t})\cdots(g_{1},h_{1}), we have

Applying ϕh≤t−1\phi_{h_{\leq t}}^{-1} on both sides, we notice that

Thus, it suffices to compute each h≤t=htht−1…h1h_{\leq t}=h_{t}h_{t-1}\ldots h_{1}, map each gt↦ϕh≤t−1(gt)g_{t}\mapsto\phi_{h_{\leq t}}^{-1}(g_{t}), compute the prefix products in these “coordinates”, then invert the mapping to get back g≤tg_{\leq t}.

Like before, we partition the embedding dimension in our construction sim⋊\mathsf{sim}_{\rtimes} into blocks, one for each component simulator. Let us index the dimensions by the dN:=simN.dimd_{N}:=\mathsf{sim}_{N}.\mathsf{dim} indices in the “NN channel” and analogously for the dHd_{H}-dimensional “HH channel”. We choose the canonical encoding EE to map elements to their individual channels:

We proceed to specify the construction layer-by-layer. Let L{N,H}L_{\{N,H\}} denote sim{N,H}.depth\mathsf{sim}_{\{N,H\}}.\mathsf{depth}.

As suggested by the intuitive sketch, we begin with LHL_{H} Transformer layers, which are just a copy of simH.θ\mathsf{sim}_{H}.\theta, reading and writing in the HH channel, with a parallel residual layer in the NN channel. So far, after these LHL_{H} layers, the output at each position tt is an integer vector, whose HH channel contains h≤th_{\leq t}, and whose NN channel contains simN.E(gt)\mathsf{sim}_{N}.E(g_{t}).

To do this, we invoke Lemma 2 (choosing the output to be in the same representation as that used by simN.E\mathsf{sim}_{N}.E, in the NN channel), with

We also add residual connections in the HH channel. In summary, after this layer, the output at each position tt is an integer vector, whose HH channel contains h≤th_{\leq t}, and whose NN channel contains ϕh≤t−1(gt)\phi_{h\leq t}^{-1}(g_{t}).

This uses Lemma 2, with exactly the same bounds.

At the end of this final “unmixing” layer, the output at each position tt is an integer vector, whose HH channel contains h≤th_{\leq t}, and whose NN channel contains g≤tg_{\leq t}; thus, this is a valid simulation of the semidirect product.

This construction is sketched in Figure 17. ∎

Note that H≅G/NH\cong G/N does not imply that GG is a semidirect product of NN and HH. Thus, although simulating semidirect products allows us to handle some families of non-abelian groups, this does not yet allow us to handle general solvable groups (i.e. general steps of a composition series, even with a cyclic quotient group). The smallest example is the non-abelian quaternion group Q8Q_{8}, the group of unit quaternions under multiplication, which cannot be realized as a semidirect product of subgroups. Instead, we need to appeal to the Krasner–Kaloujnine universal embedding theorem (Krasner and Kaloujnine, 1951): a characterization of all of the groups GG which are extensions of NN by HH, as subgroups of the wreath product N≀HN\wr H.

Let GG be a finite group which is isomorphic to a wreath product: G≅N≀HG\cong N\wr H. Let T≥1T\geq 1. Suppose N,HN,H admit simulations simN:=sim(N,T),simH:=sim(H,T)\mathsf{sim}_{N}:=\mathsf{sim}(N,T),\mathsf{sim}_{H}:=\mathsf{sim}(H,T). Then, there is a simulation of GG, sim≀:=sim(G,T)\mathsf{sim}_{\wr}:=\mathsf{sim}(G,T). In the case where sim≀.repDim=1\mathsf{sim}_{\wr}.\mathsf{repDim}=1, the sizes satisfy:

sim≀.depth=simN.depth+simH.depth+2\mathsf{sim}_{\wr}.\mathsf{depth}=\mathsf{sim}_{N}.\mathsf{depth}+\mathsf{sim}_{H}.\mathsf{depth}+2.

sim≀.dim=∣H∣⋅simN.dim+simH.dim\mathsf{sim}_{\wr}.\mathsf{dim}=|H|\cdot\mathsf{sim}_{N}.\mathsf{dim}+\mathsf{sim}_{H}.\mathsf{dim}.

sim≀.heads=max⁡{∣H∣⋅simN.heads,simH.heads}\mathsf{sim}_{\wr}.\mathsf{heads}=\max\{|H|\cdot\mathsf{sim}_{N}.\mathsf{heads},\mathsf{sim}_{H}.\mathsf{heads}\}.

sim≀.headDim=max⁡{simN.headDim,simH.headDim}\mathsf{sim}_{\wr}.\mathsf{headDim}=\max\{\mathsf{sim}_{N}.\mathsf{headDim},\mathsf{sim}_{H}.\mathsf{headDim}\}.

sim≀.mlpWidth=max⁡{∣H∣⋅simN.mlpWidth,simH.mlpWidth,5∣H∣2∣N∣}\mathsf{sim}_{\wr}.\mathsf{mlpWidth}=\max\{|H|\cdot\mathsf{sim}_{N}.\mathsf{mlpWidth},\mathsf{sim}_{H}.\mathsf{mlpWidth},5|H|^{2}|N|\}.

sim≀.normBound≤max⁡{sim{N,H}.normBound,6 ∣H∣}\mathsf{sim}_{\wr}.\mathsf{normBound}\leq\max\{\mathsf{sim}_{\{N,H\}}.\mathsf{normBound},6\,|H|\}.

sim≀.repDim=∣H∣⋅simN.repDim+1\mathsf{sim}_{\wr}.\mathsf{repDim}=|H|\cdot\mathsf{sim}_{N}.\mathsf{repDim}+1.

sim≀.repSize=max⁡{simN.repSize,simH.repSize}\mathsf{sim}_{\wr}.\mathsf{repSize}=\max\{\mathsf{sim}_{N}.\mathsf{repSize},\mathsf{sim}_{H}.\mathsf{repSize}\}.

Even though the wreath product’s algebraic structure can be very complex, the construction just requires us to implement its relatively simple description. Applying Lemma 8, we have a network sim×\mathsf{sim}_{\times} which simulates N×…×NN\times\ldots\times N. Then, we simply apply Lemma 9, using simH\mathsf{sim}_{H} to “re-map” inputs to sim×\mathsf{sim}_{\times} for the normal subgroup. This construction is sketched in Figure 18.

We can make one interesting improvement over a generic application of Lemmas 8 and 9: the structure of the mixing function ϕ\phi, which specifies the semidirect product, is extremely regular. Very fortunately, the structure of ϕ\phi allows us to avoid any dependence on the size of the wreath product group (∣N∣∣H∣⋅∣H∣|N|^{|H|}\cdot|H|) in the size measures of the implementation. A general automorphism on N×⋯×NN\times\cdots\times N is specified by its ∣N∣∣H∣|N|^{|H|} values. However, in this case, ϕ\phi is just a permutation, specified by how each of the ∣H∣|H| channels should switch places. Thus, much like the function composition gadget in Theorem 1, we can construct a simpler MLP than the generic one from Lemma 2.

Specifically, we would like to approximate the function ϕ:H×(N×⋯×N)→(N×⋯×N)\phi:H\times(N\times\cdots\times N)\rightarrow(N\times\cdots\times N), which simply applies πh\pi_{h} to the indices:

In the component neural networks’ representation space, we need the MLP to implement

recalling that the elements of g,hg,h are represented by integer vectors with ∞\infty-norm at most sim{N,H}.repBound\mathsf{sim}_{\{N,H\}}.\mathsf{repBound}. Notice that when the representation of ∣H∣|H| is a single integer, restricting to any particular coordinate in the representation of an element gg, this is the same composition problem of function transition maps solved by Lemma 5 in the proof of Theorem 1, which uses its left inputs to permute its right inputs (modulo converting the representations from {0,…,∣H∣−1}\{0,\ldots,|H|-1\} to {1,…,∣H∣}\{1,\ldots,|H|\}, which we can do by shifting the indicators at the input and final-layer output weights). Thus, ∣N∣⋅simN.dim|N|\cdot\mathsf{sim}_{N}.\mathsf{dim} parallel copies of the 3-layer function composition MLP suffice, yielding

When the information about group elements in HH is encoded by multiple integers, it is straightforward to extend this construction, by replacing the one-dimensional indicator with the multidimensional indicator from Lemma 2. We will skip the details of this case, since our final results are only about solvable groups; when we want to simulate a general group extension, it will always come from the composition series, so that HH is always a cyclic group of prime order. ∎

Thus, for general group extensions GG, we can construct sim≀\mathsf{sim}_{\wr}, the wreath product simulator for N≀HN\wr H, and combine the individual simulators. Note that we can throw away the excess group elements from the simulator: only include in sim≀.E,sim≀.W\mathsf{sim}_{\wr}.E,\mathsf{sim}_{\wr}.W the group elements which correspond to the subgroup isomorphic to GG. Then, no part of this construction needs to maintain a width or matrix entry scaling with ∣N≀H∣|N\wr H|.

Putting all of this together, we state an intermediate theorem, which is our most general result for groups:

Let GG be a solvable group which is isomorphic to a permutation group on nn elements. Let T≥1T\geq 1. Then, there is a Transformer network sim:=sim(G,T)\mathsf{sim}:=\mathsf{sim}(G,T) which simulates GG at length TT, for which we have the following size bounds:

sim.depth≤3log⁡2∣G∣\mathsf{sim}.\mathsf{depth}\leq 3\log_{2}|G|.

sim.mlpWidth≤20nT∣G∣\mathsf{sim}.\mathsf{mlpWidth}\leq 20nT|G|.

sim.normBound≤6nT\mathsf{sim}.\mathsf{normBound}\leq 6nT.

We start with a simulation of H1H_{1}, which must be a cyclic group, and build the sequence of group extensions recursively until we obtain GG. In the worst case (in the sense that the implementation size bounds from Lemma 10 are maximized), each step in the composition series must be manifested by a wreath products with K:=CnK:=C_{n}. Recall that we have:

simK.mlpWidth=4nT\mathsf{sim}_{K}.\mathsf{mlpWidth}=4nT.

simK.normBound≤6nT\mathsf{sim}_{K}.\mathsf{normBound}\leq 6nT.

simK.repSize≤n\mathsf{sim}_{K}.\mathsf{repSize}\leq n.

simHi+1.depth≤simKi.depth+3\mathsf{sim}_{H_{i+1}}.\mathsf{depth}\leq\mathsf{sim}_{K_{i}}.\mathsf{depth}+3 (1 more layer to simulate the cyclic group KiK_{i}, and 2 from the wreath product’s mixing operations).

simHi+1.dim≤∣Ki∣⋅simHi.dim+1\mathsf{sim}_{H_{i+1}}.\mathsf{dim}\leq|K_{i}|\cdot\mathsf{sim}_{H_{i}}.\mathsf{dim}+1 (noting that all of the components can reuse the same ⊥\bot and positional encoding dimensions).

simHi+1.heads≤∣Ki∣⋅simHi.heads+1\mathsf{sim}_{H_{i+1}}.\mathsf{heads}\leq|K_{i}|\cdot\mathsf{sim}_{H_{i}}.\mathsf{heads}+1.

simHi+1.headDim≤max⁡{1,1,…,1}=1\mathsf{sim}_{H_{i+1}}.\mathsf{headDim}\leq\max\{1,1,\ldots,1\}=1.

simHi+1.mlpWidth≤max⁡{∣Ki∣⋅simHi.mlpWidth,4nT,5∣Ki∣2⋅∣Hi∣}\mathsf{sim}_{H_{i+1}}.\mathsf{mlpWidth}\leq\max\{|K_{i}|\cdot\mathsf{sim}_{H_{i}}.\mathsf{mlpWidth},4nT,5|K_{i}|^{2}\cdot|H_{i}|\}.

simHi+1.normBound≤max⁡{6nT,6 ∣Ki∣}\mathsf{sim}_{H_{i+1}}.\mathsf{normBound}\leq\max\{6nT,6\,|K_{i}|\}.

simHi+1.repDim=∣Ki∣⋅simHi.repDim+1\mathsf{sim}_{H_{i+1}}.\mathsf{repDim}=|K_{i}|\cdot\mathsf{sim}_{H_{i}}.\mathsf{repDim}+1.

simHi+1.repSize≤n\mathsf{sim}_{H_{i+1}}.\mathsf{repSize}\leq n.

Now, using this construction and the results developed in the previous section for groups, we complete the construction for semigroups:

We combine the memory gate construction (Lemma 7) and any network simulating a group to implement the corresponding permutation-reset semiautomaton (Definition 5), the elements of the cascade in Theorem 7.

To finish, we implement the cascade product (Definition 4) of these permutation-reset semiautomata, guaranteed to exist by Theorem 7. This gives the full result.

simM.mlpWidth=4∣Q∣\mathsf{sim}_{M}.\mathsf{mlpWidth}=4|Q|.

simM.normBound≤2Tlog⁡(∣Q∣ T)\mathsf{sim}_{M}.\mathsf{normBound}\leq 2T\log(|Q|\,T).

Let A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) be a permutation-reset semiautomaton (see Definition 5), and let GG denote its permutation group. Let T≥1,q0∈QT\geq 1,q_{0}\in Q. Let simG:=sim(G,T)\mathsf{sim}_{G}:=\mathsf{sim}(G,T) be a Transformer network which continuously simulates GG at length TT. Then, there is a Transformer network simG′\mathsf{sim}_{G}^{\prime} which continuously simulates AT,q0{\mathcal{A}}_{T,q_{0}}, with size bounds:

simG′.depth=simG.depth+simM.depth+1≤3log⁡2∣G∣+2\mathsf{sim}_{G}^{\prime}.\mathsf{depth}=\mathsf{sim}_{G}.\mathsf{depth}+\mathsf{sim}_{M}.\mathsf{depth}+1\leq 3\log_{2}|G|+2.

simG′.dim=simG.dim+simG.repDim+simM.dim≤∣G∣+∣Q∣+4\mathsf{sim}_{G}^{\prime}.\mathsf{dim}=\mathsf{sim}_{G}.\mathsf{dim}+\mathsf{sim}_{G}.\mathsf{repDim}+\mathsf{sim}_{M}.\mathsf{dim}\leq|G|+|Q|+4.

simG′.heads=simG.heads+simM.heads≤2∣G∣+1\mathsf{sim}_{G}^{\prime}.\mathsf{heads}=\mathsf{sim}_{G}.\mathsf{heads}+\mathsf{sim}_{M}.\mathsf{heads}\leq 2|G|+1.

simG′.headDim=simG.headDim+simM.headDim+simG.repDim≤∣Q∣+3\mathsf{sim}_{G}^{\prime}.\mathsf{headDim}=\mathsf{sim}_{G}.\mathsf{headDim}+\mathsf{sim}_{M}.\mathsf{headDim}+\mathsf{sim}_{G}.\mathsf{repDim}\leq|Q|+3.

simG′.mlpWidth=simG.mlpWidth+simM.mlpWidth+∣G∣2∣Q∣≤20nT∣G∣+4∣Q∣+∣G∣2∣Q∣\mathsf{sim}_{G}^{\prime}.\mathsf{mlpWidth}=\mathsf{sim}_{G}.\mathsf{mlpWidth}+\mathsf{sim}_{M}.\mathsf{mlpWidth}+|G|^{2}|Q|\leq 20nT|G|+4|Q|+|G|^{2}|Q|.

simG′.normBound≤max⁡{simG.normBound,simM.normBound,6∣Q∣}≤6∣Q∣ Tlog⁡T\mathsf{sim}_{G}^{\prime}.\mathsf{normBound}\leq\max\{\mathsf{sim}_{G}.\mathsf{normBound},\mathsf{sim}_{M}.\mathsf{normBound},6|Q|\}\leq 6|Q|\,T\log T.

Without loss of generality, we will let Q:=[∣Q∣]Q:=[|Q|].

We split the embedding space in our construction into two channels: the simG.dim\mathsf{sim}_{G}.\mathsf{dim} dimensions used by GG, and a channel consisting of 4 additional dimensions, to be used by a copy of the memory semiautomaton, whose symbol set is QQ. Let us call these the GG and MM channels. For the reset symbols, let EM(σ)E_{M}(\sigma) denote the 4-dimensional encoding of σ\sigma from the memory semiautomaton.

Since we defined GG to be isomorphic to the permutation group associated with A{\mathcal{A}}, there is a bijection Φ:G→SQ\Phi:G\rightarrow S_{Q} between group elements and permutations on QQ. We choose the embedding EE as follows:

Let LGL_{G} denote simG.depth\mathsf{sim}_{G}.\mathsf{depth}.

The first LGL_{G} layers are chosen to be a copy of simG.θ\mathsf{sim}_{G}.\theta in the GG channel, and only residual connections in the MM channel. At the end of this, given any inputs σ1:T\sigma_{1:T} which map via Φ−1\Phi^{-1} to gtg_{t} (letting the group operation be identity when σt\sigma_{t} is a reset symbol), the outputs in the GG channel will be dG:=simG.repDimd_{G}:=\mathsf{sim}_{G}.\mathsf{repDim}-dimensional encodings of the prefix group products g≤t=gtgt−1⋯g1g_{\leq t}=g_{t}g_{t-1}\cdots g_{1}. Now, letting r(t)r(t) denote the most recent reset (τ≤t\tau\leq t such that στ\sigma_{\tau} is a reset token), we notice that the state we want can be derived from this sequence:

Here, if there have been no resets up to time tt, we define r(t)r(t) to be . We treat q0q_{0} like a reset symbol at the beginning of the sequence. Also, note that our canonical group semiautomaton simulator always uses g0=eGg_{0}=e_{G} as its initial state.

From this output, WW simply decodes the correct qtq_{t} from dimension 1. ∎

Let A=(Q,Σ,δ){\mathcal{A}}=(Q,\Sigma,\delta) be a semiautomaton, and let T≥1T\geq 1. Let {A(1),…,A(n);ϕ(2),…,ϕ(n)}\{{\mathcal{A}}^{(1)},\ldots,{\mathcal{A}}^{(n)};\phi^{(2)},\ldots,\phi^{(n)}\} be the transformation cascade (Definition 4) which simulates A{\mathcal{A}}, as guaranteed by Theorem 7. For each ii, let simi\mathsf{sim}_{i} be a Transformer network which continuously simulates the permutation-reset semiautomaton A(i){\mathcal{A}}^{(i)} at length TT. Then, there is a Transformer network simA\mathsf{sim}_{\mathcal{A}} which simulates A{\mathcal{A}} at length TT. Its size bounds are:

simA.depth=∣Q∣⋅(max⁡i{simi.depth}+1)−1≤3∣Q∣2log⁡∣Q∣+7∣Q∣\mathsf{sim}_{\mathcal{A}}.\mathsf{depth}=|Q|\cdot(\max_{i}\{\mathsf{sim}_{i}.\mathsf{depth}\}+1)-1\leq 3|Q|^{2}\log|Q|+7|Q|.

simA.dim=∑i=1nsimi.dim+1≤2∣Q∣(∣T(A)∣+∣Q∣+4)+1\mathsf{sim}_{\mathcal{A}}.\mathsf{dim}=\sum_{i=1}^{n}\mathsf{sim}_{i}.\mathsf{dim}+1\leq 2^{|Q|}(|{\mathcal{T}}({\mathcal{A}})|+|Q|+4)+1.

simA.heads=∑i=1nsimi.heads≤2∣Q∣+1(∣T(A)∣+1)\mathsf{sim}_{\mathcal{A}}.\mathsf{heads}=\sum_{i=1}^{n}\mathsf{sim}_{i}.\mathsf{heads}\leq 2^{|Q|+1}(|{\mathcal{T}}({\mathcal{A}})|+1).

simA.headDim=max⁡i=1n{simi.headDim}≤∣Q∣+3\mathsf{sim}_{\mathcal{A}}.\mathsf{headDim}=\max_{i=1}^{n}\{\mathsf{sim}_{i}.\mathsf{headDim}\}\leq|Q|+3.

simA.mlpWidth=∑i=1nsimi.mlpWidth+2∣Q∣ ∣Q∣2∣Q∣ ∣Σ∣≤2∣Q∣(20∣Q∣ ∣T(A)∣ T+4∣Q∣+∣T(A)∣2∣Q∣+∣Q∣2∣Q∣ ∣Σ∣)\mathsf{sim}_{\mathcal{A}}.\mathsf{mlpWidth}=\sum_{i=1}^{n}\mathsf{sim}_{i}.\mathsf{mlpWidth}+2^{|Q|}\,|Q|^{2^{|Q|}}\,|\Sigma|\leq 2^{|Q|}(20|Q|\,|{\mathcal{T}}({\mathcal{A}})|\,T+4|Q|+|{\mathcal{T}}({\mathcal{A}})|^{2}|Q|+|Q|^{2^{|Q|}}\,|\Sigma|).

simA.normBound≤max⁡i=1n{simi.normBound}∪{2∣Q∣(∣T(A)∣+5∣Q∣)}+6max⁡{∣Q∣,∣Σ∣}\mathsf{sim}_{\mathcal{A}}.\mathsf{normBound}\leq\max_{i=1}^{n}\{\mathsf{sim}_{i}.\mathsf{normBound}\}\cup\{2^{|Q|}(|{\mathcal{T}}({\mathcal{A}})|+5|Q|)\}+6\max\{|Q|,|\Sigma|\} ≤max⁡{6∣Q∣ Tlog⁡T,2∣Q∣(∣T(A)∣+5∣Q∣)}+6max⁡{∣Q∣,∣Σ∣}\leq\max\{6|Q|\,T\log T,2^{|Q|}(|{\mathcal{T}}({\mathcal{A}})|+5|Q|)\}+6\max\{|Q|,|\Sigma|\}.

At this point, most of the work has been done for us.

We create a separate channel ii for each component permutation-reset semiautomaton A(i){\mathcal{A}}^{(i)}. This requires a total of ∑i=1nsimi.dim\sum_{i=1}^{n}\mathsf{sim}_{i}.\mathsf{dim} embedding dimensions. In addition to these channels, we keep one dimension (with residual connections throughout the network) to represent the input σt\sigma_{t}. Let eΣe_{\Sigma} denotes the unit vector along this coordinate. Choosing an arbitrary enumeration to identify Σ\Sigma with [Σ][\Sigma], we select the embeddings to be E(σ):=σ⋅eΣE(\sigma):=\sigma\cdot e_{\Sigma}.

Namely, we invoke Lemma 2 with Δ=1,Bx=max⁡{∣Q∣,∣Σ∣},By=∣Σ∣\Delta=1,B_{x}=\max\{|Q|,|\Sigma|\},B_{y}=|\Sigma|, giving us for each pre-final-layer ii an MLP which represents the function

where the inputs are stored in the respective i′<ii^{\prime}<i and Σ\Sigma channels, and the output is written to the ii channel. Here, the number of input dimensions is

Between ii in the same layer, these routing constructions need to be executed in parallel, so this incurs another multiplicative factor in the width, bounded conservatively by 2∣Q∣2^{|Q|}.

The final construction concatenates these blocks, so that at the output of the last layer, every channel ii contains a representation of its corresponding component’s semiautomaton QiQ_{i}. The W{\mathcal{W}} guaranteed by Theorem 7 suffices for the overall choice of WW. ∎

C.4 Proof of Theorem 3: Even shorter shortcuts for gridworld

Recall the gridworld semiautomaton in Example 3, where the state (Q={0,1,…,S}Q=\{0,1,\ldots,S\}) either move to the adjacent state based upon seeing input token LL or RR (modulo boundary effects), or stay unmoved upon seeing ⊥\perp. More formally, the transition function is defined as:

In this section, we will show how to implement gridworld simulation using only 22 Transformer layers. Here we restate the theorem in full generality:

The depth in (i) can be reduced to O(S)O(S) if we allow max pooling, and the dependence on TT in the width can be removed with sinusoidal activation. We discuss this in detail after the proof along with generalization to the kk-dimensional gridworld case.

Note that, in order to find the current state, we need to only know the most recent time at which the semiautomata was at a boundary. It is not immediately obvious how to compute the most recent boundary, if one is not allowed to use the trivial sequential simulation algorithm. Our key insight is that this boundary detector can be computed without needing to parse the entire sequence, using the most recent S+1S+1 distinct values of the prefix sums in the sequence.

This algorithm is especially well-suited to the Transformer architecture since: (i) the prefix sum can be computed using one attention layer as in Lemma 6, and (ii) the identification of distinct values can be implemented by a sparse value-dependent lookup similar to the memory lookup in Lemma 7 with the help of the self-attention (context-dependent retrieval, as opposed to a static lookup), and (iii) the positional weight sharing and causal masking enable all of these computations to be performed in parallel. Overall, Theorem 3 consists of a concise implementation which executes all of these most-recent-boundary detectors in parallel.

In what follows, we first describe the algorithm (Algorithm 1) for computing the state of the semiautomata using the S+1S+1 distinct prefix sum values, and give a proof of its correctness. Subsequently, we formalize the Transformer construction that implements the algorithm. A consolidated list of notations used in the algorithm as well as the proofs is provided in Table 1 for the reader.

To convey the essence of the full construction, we first provide pseudocode (rather than Transformer weights) for computing the final state qTq_{T} (rather than the entire state sequence).

We map actions σ∈{L,R,⊥}\sigma\in\{L,R,\bot\} to σ~∈{−1,1,0}\widetilde{\sigma}\in\{-1,1,0\}, i.e. L↦−1L\mapsto-1, R↦1R\mapsto 1, and ⊥↦0\bot\mapsto 0. Let σ~(:)\widetilde{\sigma}^{(:)} denote the sequence of mapped actions, and let 0 be the initial state. The algorithm (Algorithm 1) has two steps: first, we identify the last time the agent is at a boundary (wall) and the type of the boundary (i.e. state 0 or state SS). The final state is then simply the sum of all actions in the sub-sequence, shifted by the last boundary, which is easily computable with 1 attention layer (Lemma 6). Our key insight is that we can identify the boundary using O(S)O(S) attention heads in two attention layers, and therefore do not require a recursive computation from the start state (with depth TT).

If tmin⁡>tmax⁡t_{\min}>t_{\max}, then state at tmin⁡t_{\min} is , otherwise state at tmax⁡t_{\max} is SS.

The S+1S+1 distinct values correspond to S+1S+1 distinct states (covering both boundaries). This implies that the minimum and maximum out of these distinct prefix sums must correspond to the boundaries, that is, qtmax⁡=Sq_{t_{\max}}=S and qtmin⁡=0q_{t_{\min}}=0.

Given the above, our algorithm identifies the boundary correctly and then can just use the prefix sum to evaluate the current state.

In this section we will show how to simulate Algorithm 1 using a 2-layer Transformer with 2S2S attention heads.

In our construction, the first attention layer will compute the prefix sums. This can be mapped to a cyclic group from Lemma 6, however for completeness, we will restate the main construction. The MLP in the first layer will map this prefix sum to a circular embedding (see Proposition 5). The second layer attention will use the circular embedding structure to find S+1S+1 closest distinct values to the current value ztz_{t} (suppose we are considering position t∈[T]t\in[T]) by identifying the positions for closest values in the set {zt−S,zt−S+1,…,zt−1,zt+1,zt+2,…zt+S}\{z_{t}-S,z_{t}-S+1,\ldots,z_{t}-1,z_{t}+1,z_{t}+2,\ldots z_{t}+S\}, i.e. SS closest distinct values smaller than ztz_{t}, and SS values larger than ztz_{t}. This closest distinct value construction can be viewed as a position dependent flip-flop monoid construction, where we need to identify the closest position with a particular action. Note that this set of values would contain the distinct S+1S+1 values needed by the Algorithm 1, hence the second layer MLP can implement the state computation using these values.

The attention construction for the first layer, in full detail:

Select WQ:=e3,WK:=e2,WV:=e1,WC⊤:=e4W_{Q}:=e_{3},W_{K}:=e_{2},W_{V}:=e_{1},W_{C}^{\top}:=e_{4}.

There exist a<b∈{0,1,…,2S}a<b\in\{0,1,\ldots,2S\} such that

Step 2: identify the boundary state in the selected window by comparing the indices of the endpoints of the window;

Step 3: output the final state based on the position of the boundary states and its value relative to the current position tt.

We will show two constructions for implementing this, one of which will use O(1)O(1) depth and 2O(S)2^{O(S)} width, and the other will use O(log⁡S)O(\log S) depth and O(S)O(S) width. The trade-off essentially lies in how a min function is implemented and can be resolved if we allow a min-pooling layer, which we will discuss after the proof.

O(1)O(1)-depth construction: The idea is that we can first use O(1)O(1) layers to construct “features” that contain all the information needed to determine the state, then a 3-layer MLP with 2O(S)2^{O(S)} width can compute the state as a function of these features by Lemma 2. The features we need are the following (the labels underneath are to be consistent with Figure 20, left):

Here the feature in C.3 compares the end points of the S+1S+1 windows, the two features in C.4 compare the window with its adjacent windows on each side, and the last feature in C.5 will be used to eliminate the irrelevant window. Features in C.3 and C.4 can each be computed as a threshold function (at 0) on the difference between the two elements to be compared, which can be implemented using 2 layers by Lemma 3.

Combining these inequalities together concludes the proof.

Given the optimal window, we can use feature C.3 for the relevant window to identify the boundary, since the closer-to-tt index gives us the last boundary state (see Algorithm 1 for why this suffices).

Therefore, we can compute this function using 4S+34S+3 features each taking value in {0,1}\{0,1\} and the output having S+1S+1 values. These features themselves can be constructed using Lemma 3 with Δ=1/4T\Delta=1/4T since the indices are separated by at least this gap. For the indicator index, we can compose two such constructions similar to Lemma 1. This gives us the first layer of MLP with width O(S)O(S) and norms O(T)O(T). After this, the rest of the function can be constructed using a 3-layer ReLU network with width 2O(S)2^{O(S)} and norms bounded by O(S)O(S) using Lemma 2.

O(log⁡S)O(\log S)-depth construction: An alternative solution to the above is to pay O(log⁡(S))O(\log(S)) depth, but reduce the width to be O(S)O(S). We will borrow features in equation C.3-C.5, but construct the MLP explicitly rather than calling Lemma 2 as a black box: the width and depth trade-off essentially correspond to two ways of implementing the min of SS numbers. We describe the corresponding MLP by components (Fig 20, right):

Find the min value of f1(s)f_{1}(s), denoted as f1,min⁡:=min⁡sf1(s)f_{1,\min}:=\min_{s}f_{1}(s): This can be achieved using 1 min-pooling layer. If we allow ReLU only, then this can be implemented with pairwise comparison using a network with ⌈log⁡S+1⌉\lceil\log S+1\rceil depth, 3S3S width and and constant weight norm. The log depth is conjectured to be unimprovable; see discussion after the proof.

Finally, the state is computed as ∑sf2(s)(Bs,Ms)\sum_{s}f_{2}^{(s)}(B_{s},M_{s}), which can be implemented with 1 layer of width 1.

Using standard architectural tools, such as max-pooling, we can improve our construction to get O(1)O(1)-depth and O(S)O(S)-width for the MLP.

Avoiding width TT in the MLP 1 using periodic activations. As in the modular addition (Lemma 6) construction, we can use sin⁡\sin activations in the MLP to directly compute the circular embeddings that are used as input to the second attention layer. This would require only two hidden nodes in the MLP. Note that we do not need precision greater than O(log⁡T)O(\log T) for these activations since we are embedding values only as close as 1/poly(T)1/\mathsf{poly}(T).

Avoiding log⁡(S)\log(S) depth in the MLP 2 using max-pooling. The O(log⁡S)O(\log S)-depth in MLP 2 is incurred by calculating the min of SS numbers and is conjectured to be necessary for ReLU networks (Goel et al., 2017, Mukherjee and Basu, 2017, Hertrich et al., 2021). However, the depth can be reduced to 1 if we allow max-pooling layers, which are commonly used in both theory and practice (Zhang et al., 2021b, He et al., 2016, Vaswani et al., 2021).

Remark: Yao et al. (2021) use layer-norm to compute cos⁡\cos and sin⁡\sin embedding with non-uniform angles. This could potentially alleviate the width TT concern; we leave this exploration to future work.

Since a 2-dimensional gridworld is just the direct product of 1-dimensional gridworlds (by the construction in Lemma 8), we can implement both dimensions in parallel by concatenating the network for each dimension. This can be done by doubling the dimensions, parallel attention heads, and parallel hidden units in the MLP. The attention head parameters for each dimension can be chosen to only focus on the relevant dimension and similarly the MLP can zero out dependence on the other dimension. We can extend this to higher dimension with a multiplicative increase in the size of the parameters.

C.5 Proof of Theorem 4: Depth lower bound for non-solvable semiautomata

Let A\mathcal{A} be a non-solvable semiautomaton. Then, for sufficiently large TT, no fixed-precision Transformer with depth independent of TT and width polynomial in TT can simulate A\mathcal{A} at length TT, unless TC0=NC1\mathsf{TC}^{0}=\mathsf{NC}^{1}.

This follows straightforwardly from the fact that simulating A\mathcal{A} at length TT is NC1\mathsf{NC}^{1}-complete under NC0\mathsf{NC}^{0} reductions: given any O(log⁡T)O(\log T)-depth bounded-fan-in AND/OR/NOT\mathsf{AND}/\mathsf{OR}/\mathsf{NOT} circuit C{\mathcal{C}}, and a depth-DD circuit C′{\mathcal{C}}^{\prime} which simulates a semiautomaton whose transformation monoid contains a non-solvable subgroup, there is a procedure which generates a depth-O(D)O(D) circuit to simulate C{\mathcal{C}}; see (Barrington and Thérien, 1988). This in turn comes from the construction used in Barrington’s theorem (Barrington, 1986), which characterizes NC1\mathsf{NC}^{1} as exactly the set of languages recognizable by bounded-width branching programs. For a closely related reference which follows almost exactly the same argument, see (Mereghetti and Palano, 2000).

Thus, it suffices to show that a constant-depth Transformer is in TC0\mathsf{TC}^{0}. The details of manipulating floating-point numbers with discrete circuits are peripheral to the main results in this paper, so we provide a brief proof sketch. A similar argument is used by Merrill et al. (2021) to establish that “saturated” Transformers (a multi-index analogue of hard-attention Transformers), with O(log⁡T)O(\log T) bit precision, can be represented with a TC0\mathsf{TC}^{0} circuit. We outline a proof (which applies to the formal setting considered by Merrill et al. (2021)) for the notion of Transformers defined in this paper.

The only subtlety arises when there is a TT-way summation over O(log⁡T)O(\log T)-bit numbers, which occur in the softmax and attention mixture layers. For this operation, we can use the construction from (Reif and Tate, 1992), which can even add TT poly(T)\text{poly}(T)-bit numbers in TC0\mathsf{TC}^{0}. ∎