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 ). and non-recurrent architecture to represent them?
We study this question through the lens of semiautomata, which compute state sequences from inputs by application of a recurrent transition function (and initial state ):
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 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 (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 consists of a set of states , an input alphabet , and a transition function . In this work, and will always be finite sets. For all positive integers and a starting state , defines a map from input sequences to state sequences : for . This is a deterministic Markov model, in the sense that at time , the future states only depend on the current state and the future inputs .
We define the task of simulation: given a semiautomaton , starting state , and input sequence , output the state trajectory . Let be a function (which in general can depend on ). We will say that simulates if for all input sequences . Finally, for a positive integer , we say that a function class of functions from simulates at length if, for each , there is a function in which simulates
Every semiautomaton induces a transformation semigroup of functions under composition, generated by the per-input-symbol state mappings . When contains the identity function, it is called a transformation monoid. When all of the functions are invertible, 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 , the cyclic group of order 2.
2 Recurrent and non-recurrent neural sequence models
An -layer (or depth-) 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 . In practice, this makes their inference and gradient computations highly parallelizable, with the number of sequential computation steps scaling linearly in , while RNNs require a scaling linear in . 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 , a -layer Transformer can implement the same sequential solution as an RNN: let the -th layer embed the state transition . We define shortcuts as solutions which implement the same functionality with a significantly smaller depth.
Let be a semiautomaton. For every , let be a sequence-to-sequence neural network which simulates at length . Then, we call this sequence a shortcut to if the sequence of network depths satisfies .
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 values of , but these networks must be exceptionally wide. There are also solutions which emulate transitions in “chunks”, letting each of (say) layers perform consecutive state transitions; however, without exploiting the structure of the semiautomaton, this would require width . 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 , , and . 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 , like the RNN solution. Instead, a Transformer can encode and hierarchically compose transformations (see Figure 2a), leading to far shallower solutions:
Transformers can simulate all semiautomata at length , with depth , embedding dimension , attention width , and MLP width .
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 , the group of all permutations of elements. semiautomata , with depth , embedding dimension , attention width , and MLP width .
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 , not . These two types of “prime” semiautomata can be efficiently simulated by depth- Transformers.
The decomposition depends on the transformation semigroup . 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 non-tied neurons is the implementation of mod- gates. It disappears entirely if we can add auxiliary MLP neurons with periodic activation functions such as . suboptimal 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 instead of -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 (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 be a non-solvable semiautomaton. Then, for sufficiently large , no -precision Transformer with depth independent of and width polynomial in can continuously simulate at length , unless .
This is proven in Appendix C.5. The smallest example of a non-solvable semiautomaton has states, whose transitions generate (all of the even permutations).
Finally, we note that although our width bounds might be improvable, an exponential-in- number of hypotheses (and hence a network with parameters) is unavoidable if one wishes to learn an arbitrary -state semiautomaton from data: there are of them, which generate 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-) 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 to their corresponding state sequences , and evaluate their accuracy on held-out sequences. We vary the depth from to , and use freshly-sampled sequences of length . In this setup, the number of sequences encountered during training () is far smaller than the number of distinct input sequences (). Thus, brute-force memorization cannot solve this task, and generalization is necessary to achieve nontrivial performance.
Strikingly, we obtain positive results ( 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 and . 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 () 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- 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 , the MLP could fail on sums unseen during training. This suggests that if the distribution over 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 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-) 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 ( 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 ( or time, compared to ). 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 sequential iterations of a recurrent unit with a single pass through parallel self-attention layers. Our theoretical results show that shortcuts are ubiquitous, and characterize extremely shallow ones (with independent of the context length ) 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 and . 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 -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 , denotes the index set .
We will sometimes index vectors and matrices with named indices (such as for padding tokens) instead of integers, for clarity.
denotes the -th elementary (one-hot) unit vector. Likewise as above, we sometimes use non-integer indices (e.g. ).
For a function and all , we will let denote the restriction of to (and similarly for other restrictions). This appears in the per-input state transition functions , as well as the functions represented by neural networks for a particular choice of weights.
For functions , denotes composition: . When we compose neural networks with parameter spaces , we will use to indicate the composition .
A.2 Automata, semigroups, and groups
Recall that a semiautomaton has a state space , an input alphabet , and a transition function . For any natural number and a starting state , by repeated composition of the transition function , one can use to define a map from a sequence of inputs to a sequence of states 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 with index set and use a one-hot encoding of states into . For each input symbol , we associate a transition matrix with entries . This implies that for all , we have , 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 and let and . Then, starting with , the state at time , , is if the binary sequence has an odd number of s.
Let and let be given by
As the name suggests, this semiautomaton implements a simple memory operation where the state at time is the value of the most recent non- input symbol.
Let be a natural number, and . Then the transition matrices are given by:
This semiautomaton describes the movement of an agent along a line segment where actions and 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 .
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 . This is because the proofs are stated more naturally when the boundaries of the gridworld are identified with the indices and .
For a semiautomaton each input symbol defines a function . These functions can be composed in the standard way, and we use to denote the -fold function composition. Note that is precisely the value of the state at time on input . 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 is a set equipped with a binary operation such that
(identity) There exists an identity element such that for all .
(invertibility) Every element has an inverse such that
(associativity) The binary operation is associative: .
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 be a subset of functions from where 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 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 .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 be a natural number let and let have only one element. The dynamics are given by . Clearly this semiautomaton implements counting modulo . The underlying group is the cyclic group, denoted , which is isomorphic to the integers mod 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 and , we can form a new group with elements with a binary operation that is applied component-wise (here, is overloaded to be the group operation for all three groups). This direct product group is denoted . In the context of permutation groups, say is a permutation group over ground set and is over ground set . Then has ground set and every function in factorizes component-wise, i.e., every element in is identified with a permutation where .
Observe that contains normal subgroups which are isomorphic to both and . To see this, take where is the identity element in . Then since and since is closed under its group operation, we have for all . A symmetric argument shows that 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 is the transformation semigroup of the 1-d grid world with states, then corresponds to a 2-dimensional gridworld. A semiautomaton that yields this transformation semigroup has state space and 5 actions: increment or decrement or , subject to boundary effects, or do nothing.
The definition of direct product extends straightforwardly to more than two terms ; we identify the items with tuples .
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 for some to be defined later. Observe that in the direct product , we have constructed the elements from ordered pairs , lifting and into a shared product space (i.e., the Cartesian product of the underlying sets of and ), defining the group operation as simply applying those of and 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 , and we cannot take for granted that an arbitrary binary operation satisfies the group axioms. We would like to define other operations which output an element of and an element of . An attempt would be to pick two arbitrary injective homomorphisms which embed and into a “shared space,” so that elements of and can be multiplied together:
However, the middle equality may not hold, because and are not guaranteed to commute. (Observe that for the special case of , these two elements always commute, giving rise to the direct product.)
which is of the form since both and are themselves closed. This condition is precisely that is a normal subgroup.
This object is the semidirect product, and it is denoted . Note that the choice of mapping is unspecified in the notation, and, in general, different choices of will yield different structures for the semidirect product.
Finally, when , both and are subgroups of , but is also a normal subgroup. To see this, we need to check that for any . This is equivalent to for each , but we defined the group operation to be , specifically so this would hold. On the other hand, 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 is trivial, that is then both and 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 where is the cyclic group on elements (cf. Example 4). has two elements, the identity and one element such that . has elements where each element is a function that adds some number to the input modulo . The inverse is naturally to subtract to the input, modulo . The homomorphism in the semidirect product is such that and .
We define one more type of product between groups and : the wreath product . This is a group containing elements (rather than , like the direct and semidirect products). Intuitively, it is defined by creating one copy of per element in via the direct product, then letting specify a way to exchange these copies. Formally, is the unique group generated by
where we have enumerated the elements of in arbitrary order, such that each is the permutation defined by right multiplication (by convention).
A naive way to construct the Rubik’s Cube is to assign labels to the stickers on the cube, and define the Rubik’s Cube group via the sticker configurations reachable by the face turns (which each specify a permutation of the stickers). This establishes as a subgroup of . First, notice that the 6 central stickers never move (so this is really improvable to ). Next, notice that the vertex stickers never switch places with the edge stickers. The vertex stickers form a subset of the wreath product , while the edge stickers form a subset of the wreath product . In all, this realizes 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 (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 is a normal subgroup of , the quotient group is defined as with binary operation . The fact that is a normal subgroup implies that this is a well defined group. We can also check that if then the quotient group is isomorphic to , which matches the intuition for multiplication and division.
A 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 is not simple then it has a non-trivial normal subgroup, say . We call proper if . We call a proper subgroup maximal if there is no other proper normal subgroup such that . Equivalently, is a maximal proper normal subgroup if and only if is simple. This is akin to extracting a prime factor from a number, since the quotient group 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 be arbitrary groups. Which groups contain a normal subgroup isomorphic , such that the quotient is isomorphic to ? Such a group is said to be an extension of over . The direct product is known as the trivial extension. A semidirect product is known as a split extension. However, not all extensions are split extensions; the smallest example is the quaternion group , the group of unit quaternions under multiplication (), 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 and . Fortunately, there is a characterization of general extensions. The Krasner-Kaloujnine universal embedding theorem (Krasner and Kaloujnine, 1951) states that all extensions can be found as subgroups of the wreath product . The proof of Theorem 2 essentially shows how to implement the different kinds of group extensions, given constructions which implement the substructures . 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:
is the class of constant-depth, constant-fan-in, polynomial-sized // 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 .
is the class of constant-depth, unbounded-fan-in, polynomial-sized / circuits, allowing gates only at the inputs. A classic result is that the parity of bits is not in (Furst et al., 1984); Hahn (2020) concludes the same for bounded-norm (and thus bounded-Lipschitz-constant) constant-depth Transformers.
extends with an additional type of unbounded-fan-in gate known as for arbitrary number , which checks if the sum of the input bits is a multiple of . 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).
extends with an additional type of unbounded-fan-in gate called , 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. ). 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).
is the class of -depth, constant-fan-in, polynomial-sized // 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 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 is defined to be on the first 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 to , then obtaining via a row-wise softmax (letting evaluate to ).
In general, for any positive integer , a multi-headed self-attention block consists of a sum of 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 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 ), 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 , we will choose to prepend tokens, with explicitly chosen embeddings, which do not depend on the input . Theorem 1 uses padding, and Theorem 2 uses . 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 and embedding dimension . 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 is independent of .
A bound on its -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 . 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 attention heads if we are computing a function that is based on the sum (for example,). 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- solution for bounded-depth Dyck languages. Bounded-depth Dyck can be captured by our semiautomata formalism and our main construction would recover the depth- 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 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 . 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 ( 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 -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 -depth networks can simulate all context-free languages (Theorem 1), and -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 circuits (Chandra et al., 1983, Barrington and Thérien, 1988). We did not see a way for this to generically entail -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 is not fullclose bracket which pairs with the top of the stack . 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 -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 , 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 but would require exponential width for depth (or , 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 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 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 trainable parameters in the attention and MLP weight matrices (i.e. independent of ), as opposed to the 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 depends on every preceding input . 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 . For each cyclic group (realized as under mod- addition), we choose the generator set to be the full set of group elements . An alternative could be to let be a minimalIn the sense that it induces a non-trivial learning problem on this group.If we only pick the generator , the output sequence is deterministic, and there is no learning problem. set , which we do not use in the experiments.
Direct products of cyclic groups , realized as concatenated copies of the component semiautomata. Note that (which is isomorphic to ), included in the above set, is another example.
Dihedral groups . Our realization of chooses and . 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 . We choose to be the set of permutations for (symmetric group), and to be the set of even permutations for (alternating group on elements). The generator set for consists of the minimal generators, a transposition and an -cycle, as well as 6 other permutations. These other permutations are chosen following the ordering given by the package. They are not necessary for covering the state space (since the minimal set of 2 permutations already suffice to cover ), but can help speed up the mixing of the states. For , we choose the -cycles of the form for . Note that are solvable (leading to constant-depth shortcuts), while are not. Also, note that to learn a constant-depth shortcut for , a model needs to discover the wondrous fact that has a nontrivial normal subgroup, that of its double transpositions.
The quaternion group . 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 , draw a fresh minibatch of samples from , 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 , which is large enough to rule out memorization: for this choice of , the inputs come from a uniform distribution over 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 , 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 in the Transformer between and . Note that we do not attempt in this work to distinguish between depths and , 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 requiring the largest depth.However, we stress that these experiments do not control for the fact that larger groups have richer supervision (for example, has more informative labels than ), 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 compared to (which has the same cardinality) agrees with our theoretical characterizations of the respective constant-depth shortcuts for these groups: can be written as a semidirect product of smaller groups, while cannot, so our theoretical construction of a constant-depth shortcut must embed 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 ) 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 in the experiments since the performance of the EMA model does not seem to be sensitive to the choice of . 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 with is the set of operations. The observation function returns the first value of the permutation. For example, . 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 (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 with probability 0.5, and is some randomly drawn string of 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 and 8 layers for other tasks, with embedding dimension and 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 : learning is challenging for linear encoding (i.e. ) but is easy when using sinusoidal positional encoding, which is likely because the sinusoidal encoding naturally matches the periodicity in . 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 .
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. . At test time, the sequences of the same length as training, but the Bernoulli parameter varies in the range .
Figure 13 (left) shows the accuracy as 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 deviates from 0.5, it is less likely for the value of the count (which concentrates around ) to be seen during training, hence the performance degrades. In contrast, an LSTM recurrent network maintains perfect accuracy when evaluated on all values of .
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 and states , in the standard (non-recurrent) sequence-to-sequence learning pipeline, the network receives as input, and outputs the sequence of predictions for . In scratchpad training (Nye et al., 2021, Wei et al., 2022), we instead feed the network an interleaved sequence of inputs and states (with an appropriately expanded token vocabulary), and define the network’s state predictions to be those at the appropriately aligned positions: (where 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 sequential state prediction problems which can depend on previous predicted state ; one can think of this as a way to guide a shallow Transformer to learn the recurrent solution (i.e. explicit depth- 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-) 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 . 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 hours on a single GPU, for a total of GPU hours. The remaining experiments in Section 5 amount to less than 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 . For a semiautomaton , a function simulates if for all input sequences . Here, the right-hand side denotes the sequence of states induced by the input sequence under the transitions starting from state .
A function class can simulate multiple functions associated with a semiautomaton. For a semiautomaton and a positive integer , a function class (a set of functions ) simulates at length if, for every , there is function which simulates .
When is a linear threshold function , 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 , and the weights satisfy
For each , we construct an indicator for , out of 4 ReLU units. Letting , the construction is
The second layer simply sums these indicators, weighted by each . ∎
Letting denote the set of unique values in coordinate , the inner MLP dimensions are as follows:
When we apply Lemmas 1 and 2 in recursive constructions, and , we will opt to use the bound , to reduce the clutter of propagating the term without resorting to asymptotic notation.
The inner dimension is , 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 . Suppose that for all , . Then,
Without loss of generality, (since the softmax function is invariant under shifting all inputs by the same value), so that all other coordinates are non-positive. Also, assume (the 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 , the equally-spaced points on the -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 be a semiautomaton, , and . Then, there is a depth- Transformer which continuously simulates , with embedding dimension , MLP width , and -weight norms at most . It has heads with embedding dimension implying attention width, and a -layer MLP.
The basic idea is that all prefix compositions can be evaluated in logarithmic depth using a binary tree whose leaves are the per-input transition functions . 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 is
We will produce a construction for the case where is a power of 2; general can be handled via padding. To simplify the construction, we also introduce padding positions 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 be the depth of the binary tree. We choose . Instead of indexing the dimensions by , we give them names:
Left function encoding dimensions for each .
Right function encoding dimensions for each .
Positional encoding dimensions .
Without loss of generality, let (selecting an arbitrary enumeration of the state space). Also, assume (if not, add a dummy state). We choose , mapping each input symbol to the “transition map” of its transitions. At the padding positions , we will encode the “go to ” function: .
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 and , and weight norms are bounded by .
We also add more weights which let the inputs pass through along the directions (add more rows to , calling these indices for all ; set biases to 0), for a total of hidden units and output dimensions of .
The rest of the construction uses a standard parallel algorithm for computing all prefix function compositions: at layer , compose the function at position with the function at position . 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 , we use the circle embeddings
Let , .
Select .
Select , where is the rotation matrix
Select .
Select , .
Thus, at the final layer, the dimensions at position contains the transition map of the prefix composition
It suffices to choose to be for an arbitrary to read out the sequence of states as scalar outputs in . 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 be a solvable semiautomaton (see Definition 6), , and . Then, there is a depth- Transformer which continuously simulates , with embedding dimension , MLP width , attention width heads, and weight norms are bounded by .
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 , define the mod- modular counter semiautomaton :
Let be the mod- modular counter semiautomaton. Let , and . Then, there is a depth- Transformer which continuously simulates , with embedding dimension , width , and -weight norms at most . It has head with embedding dimension , 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 does not quite give us prefix sums: the attention mixture weights at position are , while we would like the normalizing factor to be uniform across positions ( rather than ). It is possible to undo this normalization using the MLP; however, a particulaly simple solution is to use an additional padding input and 1-dimensional position embeddings to “absorb” a fraction of the attention proportional to .
We proceed to formalize this construction, beginning with the input embedding and attention block:
Select . Intuitively, the 3 dimensions implement input/output, padding, position “channels”.
Include an extra position , with embedding and position encoding . Think of this as padding at position ; it is not masked out by the causal attention mask at any position .
For , select , where is such that .
Select .
In the output of this attention module, for any input sequence , the 1st channel of the output at position is then
where is such that . The MLP simply needs to memorize the function
We invoke Lemma 1, with . The number of possible values of (the cardinality of in Lemma 1) is at most . ∎
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 , and whose transformation semigroup is the flip-flop monoid. It will be convenient to generalize this object to states:
For a given state set , define the memory semiautomaton :
Let be the memory semiautomaton. Let , and . Then, there is a depth- Transformer which continuously simulates , with embedding dimension , width , and -weight norms at most . It has head with embedding dimension , and a 2-layer ReLU MLP.
We start in state . Our goal is to identify the closest non- token and output the corresponding state. The attention construction is:
where the first coordinate denotes the action that sets the stateTechnically does not reset the state. We will see that when is selected, it must be that the semiautomaton is always in state ., the second coordinate denotes whether the input is the no-op action , and the fourth coordinate is padding.
We use positional encoding .
Denote this max position as . In the setting of hard attention, the output for the token after the attention module is . In particular, this value is if and only if , i.e. the semiautomaton never leaves the starting state. Otherwise, the value is the value of the nearest non- 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 be a finite group, and let
be a composition series: each is a maximal proper normal subgroup of ; denotes the trivial group with element. Then the quotient group 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 .
Any two composition series of are equivalent: they have the same length , and the sequences of compositions factors are equivalent under permutation and isomorphism.
When each is abelian, is called a solvable group. It turns out that each 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 , realizable as the group of even permutations of 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 .
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 be a positive integer. For each , let be a semiautomaton. For , let denote a dependency function. This object () is called a transformation cascade, and defines a cascade product semiautomaton by “feedforward simulation” under the dependency function. We define by the -th component of its output (which we call ):
The corresponding transformation semigroup is known as a cascade product semigroup.
Intuitively, the cascade specifies a way to compose semiautomata hierarchically: the first layer 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 is a permutation-reset semiautomaton if, for each , the transition function is either a bijection (i.e. a permutation over the states of ) or constant (i.e. maps every state to some ). 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 be a semiautomaton. Then, there exists a transformation cascade , defining a cascade product semiautomaton , such that:
The input symbol space of (and thus, that of ) is , the same as that of .
Letting denote the state space of , there exists a function such that simulates for all . For each , the transformation semigroup is a permutation-reset semiautomaton with at most states, whose permutation group is a (possibly trivial) subgroup of (Maler and Pnueli (1994), Theorem 4).
The number of semiautomata in the cascade is . Furthermore, the cascade has at most levels: the indices can be partitioned into at most contiguous subsets such that 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 be a semiautomaton. We call 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 corresponding to each group (Definition 7), such that if a network can simulate , it can simulate any other semiautomaton whose transformation semigroup is . This lets us talk about simulating groups, rather than particular semiautomata. We will show how to turn simulators for groups and into simulators for extensions of by , for increasingly sophisticated extensions, until all cases have been captured.
Show how to build the trivial extension: given networks which simulate the groups and , simulate the direct product , 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 , 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 and quotient , construct a network which simulates any semidirect product (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 )
Show how to build arbitrary extensions (any which contains as a normal subgroup, and for which the quotient group is isomorphic to ), using the wreath product (Lemma 10), which contains all of the group extensions. The wreath product is itself the semidirect product between a -way direct product and , 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 a canonical “complete” semiautomaton for the class of all semiautomata for which . It is simply the one whose input symbol space is every transformation reachable by some sequence of inputs (i.e. every element of ). (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 be a finite group. Then, we define the canonical group semiautomaton for as the semiautomaton defined by:
, the set of elements of . Note that if (for example) , we are setting the state space to be the set of permutations, not the ground set .
. That is, we include all functions in the input symbol space.
, for all . (In algebraic terms, we are embedding the into its left regular representation, a.k.a. left multiplication action.) Thus, if we take to be the identity element, the sequence of states corresponds to .
When we simulate the canonical group semiautomaton, we will always choose to be the identity element .
A sequence-to-sequence network is said to continuously simulate at length if it continuously simulates the canonical group semiautomaton of at length .
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 . Furthermore, the identity element will always map to the zero vector. We will keep track of the dimensionality of these vectors , and their maximum entries . All encoders and decoders will map all group elements to and from this kind of representation, and we will choose . In all, the networks will keep a -dimensional “workspace” of integer vectors, with entries bounded by . 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 is a permutation group on the original state space of the semiautomaton we would like to simulate. Indeed, when are permutation groups on , there is no natural permutation group on associated with the quotient ; it turns out that will consider simulators for and .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 states, but not directly or canonically.
To return to solving the simulation problem for some semiautomaton whose transformation semigroup is isomorphic to (at length and initial state ), let denote this isomorphism. We use as the network, with an encoding layer , and decoding layer , which can be memorized by an MLP of width 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 . Calling this construction , we can easily verify that it satisfies the canonical simulator’s conditions, and:
.
.
.
.
.
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 be two semiautomata. Then, denotes the natural direct product semiautomaton. Its states are ordered pairs . Its input symbols are defined similarly, adjoining identity inputs (so that ). The transitions are defined such that
Note that under this definition, we have . In particular, for two groups , we have .
Let be a collection of finite groups, and let . Suppose each group admits a simulation . Then, there is a simulation of the direct product group , whose sizes satisfy:
.
.
.
.
.
.
.
.
First, we pad all of the individual with layers implementing identity (add residual connections, and set attention and all MLP weight matrices to 0), so that all of them have depth .
Then, the intuition is to construct the direct product semiautomaton by concatenating the “workspaces” of each . In other words, we set the canonical encoding of to be the concatenation of each ’s encodings.
The direct product simply lets each take inputs and outputs in its individual workspace. To enable this, we need enough parallel dimensions. We set an embedding space of dimension (and similarly within the heads and MLPs), partitioning the coordinates such that in the product construction, each and 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 ). 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 be a finite group which is isomorphic to a semidirect product: , where is a normal subgroup of . Let . Suppose admit simulations . Then, there is a simulation of , , whose sizes satisfy:
.
.
.
.
.
.
.
.
The intuition is as follows, using the dihedral group as an example:
For simplicity, let us think of the “reversible car on a circular world” semiautomaton, whose transformation semigroup is . Its state consists of a direction , and a position . It has two types of inputs: “advance by ” (increment the position by in the current direction, modulo ), 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 ”, which is equivalent to simulating the parity semiautomaton.
We will convert the “advance” moves via a “basis transformation”: whenever the current direction is , an “advance by ” should be converted into . Then, we have reduced the problem to the prefix sum.
Let us write down the properties of :
is a homomorphism. That is, as permutations on .
The output of that homomorphism, , is also a homomorphism. That is, .
Let us roll out the definition of the semidirect product, given a sequence of inputs :
In general, by induction, letting denote , we have
Applying on both sides, we notice that
Thus, it suffices to compute each , map each , compute the prefix products in these “coordinates”, then invert the mapping to get back .
Like before, we partition the embedding dimension in our construction into blocks, one for each component simulator. Let us index the dimensions by the indices in the “ channel” and analogously for the -dimensional “ channel”. We choose the canonical encoding to map elements to their individual channels:
We proceed to specify the construction layer-by-layer. Let denote .
As suggested by the intuitive sketch, we begin with Transformer layers, which are just a copy of , reading and writing in the channel, with a parallel residual layer in the channel. So far, after these layers, the output at each position is an integer vector, whose channel contains , and whose channel contains .
To do this, we invoke Lemma 2 (choosing the output to be in the same representation as that used by , in the channel), with
We also add residual connections in the channel. In summary, after this layer, the output at each position is an integer vector, whose channel contains , and whose channel contains .
This uses Lemma 2, with exactly the same bounds.
At the end of this final “unmixing” layer, the output at each position is an integer vector, whose channel contains , and whose channel contains ; thus, this is a valid simulation of the semidirect product.
This construction is sketched in Figure 17. ∎
Note that does not imply that is a semidirect product of and . 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 , 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 which are extensions of by , as subgroups of the wreath product .
Let be a finite group which is isomorphic to a wreath product: . Let . Suppose admit simulations . Then, there is a simulation of , . In the case where , the sizes satisfy:
.
.
.
.
.
.
.
.
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 which simulates . Then, we simply apply Lemma 9, using to “re-map” inputs to 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 , which specifies the semidirect product, is extremely regular. Very fortunately, the structure of allows us to avoid any dependence on the size of the wreath product group () in the size measures of the implementation. A general automorphism on is specified by its values. However, in this case, is just a permutation, specified by how each of the 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 , which simply applies to the indices:
In the component neural networks’ representation space, we need the MLP to implement
recalling that the elements of are represented by integer vectors with -norm at most . Notice that when the representation of is a single integer, restricting to any particular coordinate in the representation of an element , 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 to , which we can do by shifting the indicators at the input and final-layer output weights). Thus, parallel copies of the 3-layer function composition MLP suffice, yielding
When the information about group elements in 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 is always a cyclic group of prime order. ∎
Thus, for general group extensions , we can construct , the wreath product simulator for , and combine the individual simulators. Note that we can throw away the excess group elements from the simulator: only include in the group elements which correspond to the subgroup isomorphic to . Then, no part of this construction needs to maintain a width or matrix entry scaling with .
Putting all of this together, we state an intermediate theorem, which is our most general result for groups:
Let be a solvable group which is isomorphic to a permutation group on elements. Let . Then, there is a Transformer network which simulates at length , for which we have the following size bounds:
.
.
.
We start with a simulation of , which must be a cyclic group, and build the sequence of group extensions recursively until we obtain . 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 . Recall that we have:
.
.
.
(1 more layer to simulate the cyclic group , and 2 from the wreath product’s mixing operations).
(noting that all of the components can reuse the same and positional encoding dimensions).
.
.
.
.
.
.
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.
.
.
Let be a permutation-reset semiautomaton (see Definition 5), and let denote its permutation group. Let . Let be a Transformer network which continuously simulates at length . Then, there is a Transformer network which continuously simulates , with size bounds:
.
.
.
.
.
.
Without loss of generality, we will let .
We split the embedding space in our construction into two channels: the dimensions used by , and a channel consisting of 4 additional dimensions, to be used by a copy of the memory semiautomaton, whose symbol set is . Let us call these the and channels. For the reset symbols, let denote the 4-dimensional encoding of from the memory semiautomaton.
Since we defined to be isomorphic to the permutation group associated with , there is a bijection between group elements and permutations on . We choose the embedding as follows:
Let denote .
The first layers are chosen to be a copy of in the channel, and only residual connections in the channel. At the end of this, given any inputs which map via to (letting the group operation be identity when is a reset symbol), the outputs in the channel will be -dimensional encodings of the prefix group products . Now, letting denote the most recent reset ( such that 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 , we define to be . We treat like a reset symbol at the beginning of the sequence. Also, note that our canonical group semiautomaton simulator always uses as its initial state.
From this output, simply decodes the correct from dimension 1. ∎
Let be a semiautomaton, and let . Let be the transformation cascade (Definition 4) which simulates , as guaranteed by Theorem 7. For each , let be a Transformer network which continuously simulates the permutation-reset semiautomaton at length . Then, there is a Transformer network which simulates at length . Its size bounds are:
.
.
.
.
.
.
At this point, most of the work has been done for us.
We create a separate channel for each component permutation-reset semiautomaton . This requires a total of embedding dimensions. In addition to these channels, we keep one dimension (with residual connections throughout the network) to represent the input . Let denotes the unit vector along this coordinate. Choosing an arbitrary enumeration to identify with , we select the embeddings to be .
Namely, we invoke Lemma 2 with , giving us for each pre-final-layer an MLP which represents the function
where the inputs are stored in the respective and channels, and the output is written to the channel. Here, the number of input dimensions is
Between 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 .
The final construction concatenates these blocks, so that at the output of the last layer, every channel contains a representation of its corresponding component’s semiautomaton . The guaranteed by Theorem 7 suffices for the overall choice of . ∎
C.4 Proof of Theorem 3: Even shorter shortcuts for gridworld
Recall the gridworld semiautomaton in Example 3, where the state () either move to the adjacent state based upon seeing input token or (modulo boundary effects), or stay unmoved upon seeing . More formally, the transition function is defined as:
In this section, we will show how to implement gridworld simulation using only Transformer layers. Here we restate the theorem in full generality:
The depth in (i) can be reduced to if we allow max pooling, and the dependence on in the width can be removed with sinusoidal activation. We discuss this in detail after the proof along with generalization to the -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 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 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 (rather than the entire state sequence).
We map actions to , i.e. , , and . Let 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 ). 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 attention heads in two attention layers, and therefore do not require a recursive computation from the start state (with depth ).
If , then state at is , otherwise state at is .
The distinct values correspond to 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, and .
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 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 closest distinct values to the current value (suppose we are considering position ) by identifying the positions for closest values in the set , i.e. closest distinct values smaller than , and values larger than . 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 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 .
There exist 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 .
We will show two constructions for implementing this, one of which will use depth and width, and the other will use depth and 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.
-depth construction: The idea is that we can first use layers to construct “features” that contain all the information needed to determine the state, then a 3-layer MLP with 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 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- index gives us the last boundary state (see Algorithm 1 for why this suffices).
Therefore, we can compute this function using features each taking value in and the output having values. These features themselves can be constructed using Lemma 3 with 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 and norms . After this, the rest of the function can be constructed using a 3-layer ReLU network with width and norms bounded by using Lemma 2.
-depth construction: An alternative solution to the above is to pay depth, but reduce the width to be . 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 numbers. We describe the corresponding MLP by components (Fig 20, right):
Find the min value of , denoted as : 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 depth, 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 , 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 -depth and -width for the MLP.
Avoiding width in the MLP 1 using periodic activations. As in the modular addition (Lemma 6) construction, we can use 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 for these activations since we are embedding values only as close as .
Avoiding depth in the MLP 2 using max-pooling. The -depth in MLP 2 is incurred by calculating the min of 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 and embedding with non-uniform angles. This could potentially alleviate the width 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 be a non-solvable semiautomaton. Then, for sufficiently large , no fixed-precision Transformer with depth independent of and width polynomial in can simulate at length , unless .
This follows straightforwardly from the fact that simulating at length is -complete under reductions: given any -depth bounded-fan-in circuit , and a depth- circuit which simulates a semiautomaton whose transformation monoid contains a non-solvable subgroup, there is a procedure which generates a depth- circuit to simulate ; see (Barrington and Thérien, 1988). This in turn comes from the construction used in Barrington’s theorem (Barrington, 1986), which characterizes 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 . 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 bit precision, can be represented with a 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 -way summation over -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 -bit numbers in . ∎