Representational Strengths and Limitations of Transformers

Clayton Sanford, Daniel Hsu, Matus Telgarsky

Introduction

In recent years, transformer networks (Vaswani et al., 2017) have been established as a fundamental neural architecture powering state-of-the-art results in many applications, including language modeling (OpenAI, 2023), computer vision (Dosovitskiy et al., 2021), and protein folding (Jumper et al., 2021). The key building block of transformer models is the self-attention unit, a primitive that represents interactions among input elements as inner-products between low-dimensional embeddings of these elements.

The success of transformer models is linked to their ability to scale their training and generalization performance to larger datasets and sequence lengths. Their representational capacity, however, underlies this scaling power, and is tied to the inductive biases of their learning algorithms. Empirically, transformer models trained with gradient-based learning algorithms exhibit biases towards certain algorithmic primitives (Edelman et al., 2022; Liu et al., 2022) and learn representations that may encode domain-specific information in the self-attention units (Clark et al., 2019; Hewitt and Manning, 2019; Rogers et al., 2020; Chen et al., 2022). These examples indicate that transformer architectures not only provide computational benefits, but also have representational capabilities that are particularly well-matched to practical tasks.

In this paper, we investigate these inductive biases by identifying “natural” computational tasks for which transformers are well-suited, especially compared to other neural network architectures, as well as tasks that highlight the limitations of transformers. The tasks—sparse averaging, pair-matching, and triples-matching—represent primitive operations that aggregate structural information encoded in embeddings. We use these tasks to elucidate the relationship between the embedding dimension mm of a self-attention unit and its expressivity, and to showcase the fundamental representational limitations of self-attention layers.

In our model, the primary computational bottleneck faced by a transformer in computing a “sequence-to-sequence”Note, however, that attention units are permutation equivariant, so the order of elements in the input “sequence” X∈XNX\in\mathcal{X}^{N} is irrelevant. In practice, positional encodings are used when the sequence order is relevant. function f ⁣:XN→YNf\colon\mathcal{X}^{N}\to\mathcal{Y}^{N} is the constrained processing of pairs of input elements {xi,xj}∈(X2)\{x_{i},x_{j}\}\in\binom{\mathcal{X}}{2}; we allow transformers unbounded computational power when processing the individual elements xi∈Xx_{i}\in\mathcal{X}. This is motivated by modern scaling regimes where the context length NN has rapidly increased, the self-attention embedding dimension mm remains much smaller than NN, and the parameterization of multi-layer perceptrons (MLPs) that operate on individual elements is much larger than mm. Indeed, the largest GPT-3 model (Brown et al., 2020) features a context length N=2048N=2048, an embedding dimension m=128m=128, and MLPs with a 12288-dimensional parameterization; the context length of GPT-4 is as large as N=32000N=32000. As such, we are interested in the capabilities of transformers with No(1)N^{o(1)} total “size”, as opposed to NΩ(1)N^{\Omega(1)}. The nature of the bottleneck in our model makes the tools of communication complexity indispensable for formalizing computational limits.

We frame standard transformer architectures as being able to efficiently represent functions that are decomposable into sparse pairwise interactions between inputs. To do so, we introduce two sequential tasks and prove a collection of constructions and hardness results that characterize the abilities of transformers to solve these tasks.

For both tasks, note that the output is an NN-dimensional vector whose iith element is 1 if and only if the sequence XX includes a pair or triple containing xix_{i}. In this sense, the problems differ from 2SUM and 3SUM, which are not sequence-to-sequence tasks.

In Appendices C.5 and C.6, we give a heuristic information-theoretic argument to support this conjecture, prove a matching upper-bound, and finally prove analogous results for graph-augmented transformers with respect to the problem of cycle detection in directed and undirected graphs.

2 Related work

Several computational and learning-theoretic aspects of transformers, distinct from but related to the specific aims of the present paper, have been mathematically studied in previous works.

To demonstrate the power of transformers, universal approximation results for transformers (Yun et al., 2020; Wei et al., 2022)—analogous to results for feedforward networks (Hornik et al., 1989)—establish the capability for sufficiently large networks to accurately approximate general classes of functions. Note, however, that the precise minimal dependence of the required size (e.g., number of attention units, depth of the network) as a function of the input size NN does not directly follow from such results, and it is complicated by the interleaving of other neural network elements between attention layers. (Approximate) Turing-completeness of transformers demonstrates their power in a different manner, and such results have been established, first assuming infinite precision weights (Pérez et al., 2019) and later also with finite-precision (Wei et al., 2022). Such results are more closely aligned with our aims, because Turing machines represent a uniform model of computation on inputs of arbitrary size. Wei et al. (2022) showed that Turing machines that run for TT steps can be approximated by “encoder-decoder” transformers of depth log⁡(T)\log(T) and size polynomial in log⁡(T)\log(T) and the number of states of the Turing machine (but the decoder runs for TT steps).

The ubiquity of transformers in natural language understanding has motivated the theoretical study of their ability to recognize formal languages. On the positive side, Bhattamishra et al. (2020) constructed transformers that recognize counter languages, and Yao et al. (2021) showed that transformers of bounded size and depth can recognize Dyck languages that have bounded stack depth. Liu et al. (2022) showed that the computations of finite-state automata on sequences of length NN can be performed by transformers of depth log⁡(N)\log(N) and size polynomial in the number of states. On the negative side, Hahn (2020) showed limitations of modeling distributions over formal languages (including Dyck) with fixed-size transformers (though this result does not imply quantitative lower bounds on the size of the transformer). Hahn (2020), as well as Hao et al. (2022), also establish the inability of “hard attention” Transformers to recognize various formal languages and circuit classes by leveraging depth reduction techniques from circuit complexity (Furst et al., 1984).

Graph neural networks (GNNs), like transformers, process very large inputs (graphs) using neural networks that act only on small collections of the input parts (vertex neighborhoods). Many classes of GNNs are universal approximators for classes of invariant and equivariant functions (Maron et al., 2019; Keriven and Peyré, 2019). At the same time, they are restricted by the distinguishing power of certain graph isomorphism tests (Xu et al., 2018; Morris et al., 2019; Chen et al., 2019), and lower bounds have been established on the network size to approximate such tests (Aamand et al., 2022). Loukas (2019) established a connection between GNNs and the Local (Angluin, 1980) and Congest (Peleg, 2000) models for distributed computation, and hence directly translates lower bounds for Congest—notably cycle detection problems—into size lower bounds for GNNs. Our lower bounds for cycle detection using transformers also leverage a connection to the Congest model. However, transformers do not have the same limitations as GNNs, since the computational substrate of a transformer does not depend on the input graph in the way it is with GNNs. Thus, we cannot directly import lower bounds for Congest to obtain lower bounds for transformers.

3 Conclusion and future work

Preliminaries

We first introduce the concept of self-attention, which is used as the building block of all transformer architectures included in this paper.

Let Ad,m,d′,p={fQ,K,V:Q,K,V}\mathcal{A}_{d,m,d^{\prime},p}=\{f_{Q,K,V}:Q,K,V\} denote all such self-attention units.

Self-attention units can be computed in parallel to create multi-headed attention.

Transformer models are composed of two components: multi-headed attention layers (as above) and element-wise multi-layer perceptrons. Due to universal approximation results, we model multi-layer perceptrons as arbitrary functions mapping fixed-precision vectors to themselves.

We concatenate the notation of each class of functions to denote function composition. For example, for output dimension d′d^{\prime}, we use Ad,m,d′,p′:=Am,m,d′,pΦd,m,p\mathcal{A}^{\prime}_{d,m,d^{\prime},p}:=\mathcal{A}_{m,m,d^{\prime},p}\Phi_{d,m,p} and Ad,m,d′,pH′:=Am,m,d′,pHΦd,m,p\mathcal{A}^{H\prime}_{d,m,d^{\prime},p}:=\mathcal{A}^{H}_{m,m,d^{\prime},p}\Phi_{d,m,p} to represent single-headed and multi-headed attention units with an input MLP respectively. (The capabilities and limitations of these models are studied in Section 3.) For depth DD, we let

represent a full transformer model comprising DD layers of HH-headed self-attention with interspersed MLPs.

While two key features of transformer architectures—the residual connection and the positional embedding—are conspicuously missing from this formalism, the two can be implemented easily under the framework. We can include a positional embedding by encoding the index as a coordinate of the input, i.e. xi,1=ix_{i,1}=i. Then, the subsequent MLP transformation ϕ(X)\phi(X) can incorporate ii suitably into the embedding. A residual connection can be included additively as input to a multi-layer perceptron layer (as is standard) by implementing an “approximate identity” attention head ff with Q,KQ,K and V=ImV=I_{m} set to ensure that f(X)≈Xf(X)\approx X.A simple construction involves letting XQ=XKXQ=XK with iid Gaussian columns fixed for every index ii. Then, the diagonals of XQKTXTXQK^{\scriptscriptstyle\mathsf{T}}X^{\scriptscriptstyle\mathsf{T}} are far larger than all other entries and its softmax is approximately INI_{N}.

We periodically consider transformers implemented with real-valued arithmetic with infinite bit complexity; in those cases, we omit the bit complexity pp from the notation.

Sparse averaging with attention units

We present the sparse averaging task to highlight the ability of transformer architectures to simulate a wide range of meaningful interactions between input elements. This task demonstrates how the embedding dimension of a self-attention unit modulates the expressive capabilities of the architecture, while showcasing the inabilities of fully-connected and recurrent neural networks to capture similar interactions (see Appendix A).

To do so, let each key KTϕ(xi)K^{\scriptscriptstyle\mathsf{T}}\phi(x_{i}) represent a fixed vertex on a convex polytope, which depends only on index ii and is constructed from random binary vectors. We select each query QTϕ(xi)Q^{\scriptscriptstyle\mathsf{T}}\phi(x_{i}) to ensure that ϕ(xi)TQKTϕ(xj)\phi(x_{i})^{\scriptscriptstyle\mathsf{T}}QK^{\scriptscriptstyle\mathsf{T}}\phi(x_{j}) is a fixed large value if j∈yij\in y_{i} and a slightly smaller value otherwise. We obtain the precise query, key, and value embeddings by employing tools from dual certificate analysis from the theory of compressed sensing.

The logarithmic dependence of the embedding dimension mm on the sequence length NN can be eliminated by considering self-attention units with real-valued arithmetic with infinite bit complexity.

We show that the construction used to prove Theorem 2 is nearly optimal.

(By choosing p=O(log⁡(qlog⁡N))p=O(\log(q\log N)), Theorem 2 is shown to be optimal up to logarithmic factors of qq and doubly-logarithmic factors of NN.)

The proof of Theorem 4 employs a standard communication complexity argument based on a reduction from the following set disjointness problem in the two-party communication model, in which each party possesses a subset of an nn element domain (encoded as nn-bit strings), and they wish to jointly determine whether their subsets are disjoint. We note that communication complexity is commonly-used technique for proving lower bounds on the representational power of circuits and feedforward neural networks (see, e.g., Karchmer and Wigderson, 1988; Ben-David et al., 2002; Martens et al., 2013; Vardi et al., 2021).

Alice encodes her input aa in a single subset by letting y2q+1={2i+ai−1:i∈[q]}y_{2q+1}=\{2i+a_{i}-1:i\in[q]\}.

Bob uses his input bb to assign z2i−1z_{2i-1} to 2bi−12b_{i}-1 and z2i=−1z_{2i}=-1 for all i∈[q]i\in[q].

All other input components are set to constant values known by both parties.

Standard transformer models can only efficiently represent intrinsically pairwise functions

The proof, given in Section C.1 uses both a “blank token” and a trigonometric positional embedding, which ensures that

The second focuses on localized sums, where are all components of a triple must be within a fixed range of constant width K≪NK\ll N: for each i∈[N]i\in[N],

We show that the two can be efficiently represented using compact standard transformer models.

We implement the localized construction by using Theorem 2 to construct a specific sparse simultaneous average of the inputs with q:=2K+1q:=2K+1 and d′:=2K+1d^{\prime}:=2K+1. To do so, we use the input MLP to convert xix_{i} to the embedding (zi;yi;i)(z_{i};y_{i};i), for zero-padded input

This construction ensures that the iith element of self-attention output computes (a rotation of) (xi−K,xi−K+1,…,xi+K)(x_{i-K},x_{i-K+1},\dots,x_{i+K}). An output MLP can then verify whether any matching triples involving xix_{i} exist among those vectors. ∎

We are grateful for many discussions with and feedback from Navid Ardeshir, Peter Bartlett, Alberto Bietti, Yuval Efron, Shivam Nadimpalli, Christos Papadimitriou, Rocco Servedio, Yusu Wang, and Cyril Zhang. This work was supported in part by NSF grants CCF-1740833 and IIS-1563785, a JP Morgan Faculty Award, and an NSF Graduate Research Fellowship.

References

z=(0⃗;… ;0⃗;ξq;… ;ξN)z=(\vec{0};\dots;\vec{0};\xi_{q};\dots;\xi_{N}), and z′=(0⃗;… ;0⃗;−ξq;… ;−ξN)z^{\prime}=(\vec{0};\dots;\vec{0};-\xi_{q};\dots;-\xi_{N}). Then,

Vjzj=Vjzj′=0V_{j}z_{j}=V_{j}z_{j}^{\prime}=0 for all j∈[N]j\in[N]; and

∥zj∗−zj∗′∥2=2\|z_{j^{*}}-z_{j^{*}}^{\prime}\|_{2}=2 for some j∗∈{q,…,N}j^{*}\in\{q,\dots,N\}.

Therefore, for any y1,…,yN∈([N]q)y_{1},\dots,y_{N}\in\binom{[N]}{q}, respective x=(1;… ;N;y1;… ;yN;z1;… ;zN)x=(1;\dots;N;y_{1};\dots;y_{N};z_{1};\dots;z_{N}) and x′=(1;… ;N;y1;… ;yN;z1′;… ;zN′)x^{\prime}=(1;\dots;N;y_{1};\dots;y_{N};z_{1}^{\prime};\dots;z_{N}^{\prime}) satisfy f(x)=f(x′)f(x)=f(x^{\prime}). Consider yy with yj=(1,…,q−1,j)y_{j}=(1,\dots,q-1,j) for each j∈{q,…,N}j\in\{q,\dots,N\}. Then,

Alice constructs inputs xi=(zi,∅,i)x_{i}=(z_{i},\emptyset,i) for i=1,…,n+1i=1,\dotsc,n+1, where for each i=1,…,ni=1,\dotsc,n,

Bob constructs inputs xn+1+i=(0,yn+1+i,n+1+i)x_{n+1+i}=(0,y_{n+1+i},n+1+i) for i=1,…,ni=1,\dotsc,n, where

Observe that, for this input X=(x1,…,x2n+1)X=(x_{1},\dotsc,x_{2n+1}), we have

Alice simulates the memory-bounded algorithm on the first n+1n+1 inputs x1,…,xn+1x_{1},\dotsc,x_{n+1}, and sends Bob the mm-bit memory state hn+1h_{n+1}. This requires mm bits of communication.

Starting with hn+1h_{n+1}, Bob continues the simulation of the memory-bounded algorithm on these nn additional inputs xn+2,…,x2n+1x_{n+2},\dotsc,x_{2n+1}.

If any output f(X)n+1+if(X)_{n+1+i} for i=1,…,ni=1,\dotsc,n satisfies

then Bob outputs 11 (not disjoint); otherwise Bob outputs (disjoint).

Appendix B Supplementary results for Section 3

We analyze the output of the softmax. If i′∈yii^{\prime}\in y_{i}, then

Likewise, if i′∉yii^{\prime}\not\in y_{i}, then

We thus conclude that that we meet the desired degree of approximation for such mm:

The first result shows the existence of a sign-valued matrix UU that satisfies the desired distance-preserving property.

Sparse subsets of the columns of such a UU can then be linearly separated from all other columns.

B.2 Proof of Theorem 3

The proof relies on the properties of neighborly polytopes, which we define.

A polytope PP is qq-neighborly if every subset of q′≤qq^{\prime}\leq q vertices forms a (q′−1)(q^{\prime}-1)-face.

The proof of Theorem 3 is immediate from the aforementioned fact and the following lemma.

The construction employs a similar look-up table MLP ϕ\phi to the one used in the proof of Theorem 2. We let the key and value embeddings be

For α>0\alpha>0, let ϕ(xi)TQ=αwy=α(wy′,by)\phi(x_{i})^{\scriptscriptstyle\mathsf{T}}Q=\alpha w_{y}=\alpha(w^{\prime}_{y},b_{y}).

B.3 Proof of Theorem 4

All other inputs are set arbitrarily. Then,

Bob determines z1,…,z2qz_{1},\dots,z_{2q} from bb. Using those and the information from Alice, he computes f(X)2q+1f(X)_{2q+1}. He returns 11 if and only if f(X)2q+1Te1≥−1+1qf(X)_{2q+1}^{\scriptscriptstyle\mathsf{T}}e_{1}\geq-1+\frac{1}{q}.

B.4 Optimality of Theorem 3 under restricted architectures

While the near-optimality of the bounded-precision self-attention construction in Theorem 2 is assured by the communication complexity argument of Theorem 4, it is not immediately apparent whether Theorem 3 is similarly optimal among infinite-precision self-attention models. Theorem 16 proves that this is indeed the case for a restricted family of architectures that resembles cross-attention rather than self-attention.

The architectural assumptions of this statement are strong. For each element xi=(zi;yi;i)x_{i}=(z_{i};y_{i};i), its value embedding must reproduce its target ziz_{i}; its key embedding depends exclusively on the index ii; and its query embedding only on the indices yiy_{i} and ii. Indeed this attention unit more closely resembles cross-attention rather than self-attention, in which the problem is formulated as two sequences ((z1,1),…,(zN,N))((z_{1},1),\dots,(z_{N},N)) and (y1;1),…,(yN;N)(y_{1};1),\dots,(y_{N};N) that are passed to the key and value inputs and the query inputs respectively. We leave open the problem of generalizing this result to include all infinite-precision cross-attention or self-attention architectures, but we note that the constructions in Theorems 2 and 3 can be implemented under such architectural assumptions.

The proof relies on a geometric argument about how the convex hull of fixed key embeddings U=(u(1),…,u(N))U=(u(1),\dots,u(N)) lacks neighborliness and hence cannot separate every size-qq subsets of values embeddings z1,…,zNz_{1},\dots,z_{N} from the other values.

It suffices to show that for any fixed key embedding UU, there exists some yiy_{i} and setting of z1,…,zNz_{1},\dots,z_{N} such that

By the Sauer-Shelah Lemma (Sauer, 1972; Shelah, 1972; Vapnik and Chervonenkis, 1968) and the fact that the VC dimension of m′m^{\prime}-dimensional linear thresholds is m′+1m^{\prime}+1, the maximum number of partitions of the columns of UU that can be linearly separated is at most

for a sufficiently large choice of CC given universal constant C′C^{\prime}. If the fact were to be false, then at least (Nq)≥(Nq)q{N\choose q}\geq(\frac{N}{q})^{q} such partitions must exist, which contradicts the above bound. ∎

Appendix C Supplementary results for Section 4

VTϕ(xi)=1⃗V^{\scriptscriptstyle\mathsf{T}}\phi(x_{i})=\vec{1}, QTϕ(x′)=0⃗Q^{\scriptscriptstyle\mathsf{T}}\phi(x^{\prime})=\vec{0}, KTϕ(x′)=e3K^{\scriptscriptstyle\mathsf{T}}\phi(x^{\prime})=e_{3}, and VTϕ(x′)=0⃗V^{\scriptscriptstyle\mathsf{T}}\phi(x^{\prime})=\vec{0}. By elementary trigonometric identities, the following is true about the corresponding inner products:

As a result, (QTϕ(xi))TKTϕ(xj)=cd(Q^{\scriptscriptstyle\mathsf{T}}\phi(x_{i}))^{\scriptscriptstyle\mathsf{T}}K^{\scriptscriptstyle\mathsf{T}}\phi(x_{j})=cd if and only if xi+xj=0(modM)x_{i}+x_{j}=0\kern-8.0pt\pmod{M}. Otherwise, (QTϕ(xi))TKTϕ(xj)≤c(1−1M2)(Q^{\scriptscriptstyle\mathsf{T}}\phi(x_{i}))^{\scriptscriptstyle\mathsf{T}}K^{\scriptscriptstyle\mathsf{T}}\phi(x_{j})\leq c(1-\frac{1}{M^{2}}). (Here, the O(log⁡M)O(\log M)-bit fixed-precision arithmetic is sufficient to numerically distinguish the two cases.) For each i∈[N]i\in[N] let

represent the total number of matches the input belongs to. If we take c=M2log⁡(6N)c=M^{2}\log(6N), then

where ≤\leq is a partial ordering with v≤v′v\leq v^{\prime} if vi≤vi′v_{i}\leq v^{\prime}_{i} for all ii. Since the latter case holds only when βi≥1\beta_{i}\geq 1, the final step of the proof is design an output MLP ψ\psi such that ψ(z)=1\psi(z)=1 if z≥13z\geq\frac{1}{3} and ψ(z)=0\psi(z)=0 if z≤16z\leq\frac{1}{6}, which can be crafted using two ReLU gates. ∎

C.2 Proof of Theorem 7

Alice and Bob compute (x2,…,xN+12)(x_{2},\dots,x_{\frac{N+1}{2}}) and (xN+32,…,xN)(x_{\frac{N+3}{2}},\dots,x_{N}) from aa and bb respectively.

Alice computes an O(plog⁡log⁡N)O(p\log\log N)-bit approximation of the logarithm of the first half of the softmax normalization term for each attention head and sends the result to Bob. That is, she sends Bob

for each h∈[H]h\in[H]. This requires transmitting O(pHlog⁡log⁡N)O(pH\log\log N) bits.

Bob finishes the computation of normalization terms

for each hh and sends the result back to Alice (up to O(plog⁡log⁡N)O(p\log\log N)-bits of precision). This again requires transmitting O(pHlog⁡log⁡N)O(pH\log\log N) bits.

Alice computes the partial convex combination of the first N+12\frac{N+1}{2} value vectors stipulated by the attention matrix

for each hh and sends the partial combinations to Bob. This requires transmitting O(mpHlog⁡log⁡N)O(mpH\log\log N) bits (using the same precision as above).

Bob finishes the computation of the convex combinations

Bob concludes the protocol by computing and outputting f(X)1f(X)_{1}, using his knowledge of each fh(X)f_{h}(X) and of ψ\psi.

C.3 Higher-order tensor attention

The following generalizes the definition of self-attention.

The input to the row-wise softmax is an N×Ns−1N\times N^{s-1} matrix. Let Ad,m,d′,p⊗s\mathcal{A}_{d,m,d^{\prime},p}^{\otimes s} denote the set containing all such attention units.

Note that Ad,m,d′,p⊗2=Ad,m,d′,p\mathcal{A}_{d,m,d^{\prime},p}^{\otimes 2}=\mathcal{A}_{d,m,d^{\prime},p}. Because ss-order self-attention units have the same domain and codomain as standard self-attention, multiple units can be analogous combined to construct multi-headed attention units and full transformer models. We define Ad,m,d′,pM,⊗s\mathcal{A}_{d,m,d^{\prime},p}^{M,\otimes s} and Td,m,d′,pD,H,⊗s\mathcal{T}_{d,m,d^{\prime},p}^{D,H,\otimes s} accordingly.

The purpose of the ss-order transformer model as a theoretical construct is to posit how strictly generalizing the architecture in order to permit higher order outer products transfers the expressive powers of standard transformer architectures to more sophisticated interactions among elements of the input sequence XX. The model is not defined to be immediately practical, due to its steep computational cost of evaluation.

However, the trade-offs involved in using such architectures resemble those already made by using transformer models instead of fully-connected networks. Transformers are already computationally wasteful relative to the number of the parameters, and these models likely succeed only because extremely efficient factorized parameterization exist. Likewise, third-order transformers could indeed be practical if even more factorization proves useful, since the computational costs may prove mild if the embedding dimension mm, number of heads HH, and depth DD necessary to succeed on a task exceed the sequence length NN for standard second-order transformers.

The proof is almost identical to that of Theorem 6, except that we instead use a different key and query transforms to express a different trigonometric function:

Together, these ensure that the resulting tensor products reduce to a trigonometric expression that is maximized when xi+xj1+xj2=0(modM)x_{i}+x_{j_{1}}+x_{j_{2}}=0\pmod{M}. That is,

We similarly let V1ϕ(xi)=V2ϕ(xi)=1⃗V^{1}\phi(x_{i})=V^{2}\phi(x_{i})=\vec{1} and V1ϕ(x′)=V2ϕ(x′)=0⃗V^{1}\phi(x^{\prime})=V^{2}\phi(x^{\prime})=\vec{0}. The remaining choice of cc and the output MLP, and the analysis of the softmax proceeds identically to the previous proof. ∎

C.5 Heuristic argument for Informal Conjecture 1

Note that under event E1E_{1}, a three matching elements exist with probability at most 1N\frac{1}{N}, and

Under D\mathcal{D}, any subset of {x1,…,xN}\{{\mathbf{x}}_{1},\dots,{\mathbf{x}}_{N}\} consists of iid integers drawn uniformly from [M][M], unless all of xj1,xj2,xj3{\mathbf{x}}_{j_{1}},{\mathbf{x}}_{j_{2}},{\mathbf{x}}_{j_{3}} appear in the subset. Consider a transformer architecture with pp-bit precision, mm-dimensional embeddings, HH heads per layer, and DD layers. We argue informally that a single-element output of a self-attention unit can take into account information about mpmp more inputs x1,…,xN{\mathbf{x}}_{1},\dots,{\mathbf{x}}_{N} than that it had in the previous layer. By induction, after DD layers of HH-headed self-attention with interleaved MLPs, each element is a function of at most mpHDmpHD inputs. Until an element exists that is a function of at least two of the three of xj1,xj2,xj3{\mathbf{x}}_{j_{1}},{\mathbf{x}}_{j_{2}},{\mathbf{x}}_{j_{3}}, we assume that the elements “known” by each output are chosen independently of the indices j1,j2,j3j_{1},j_{2},j_{3}. (Given two elements of the triple, the third element can be identified with a single self-attention unit.) Hence, we argue that it suffices to show that the probability any two elements of the triple j1,j2,j3j_{1},j_{2},j_{3} occurring within any of the NN sets of mpHDmpHD inputs is vanishingly small for sufficiently large transformer parameters. The probability of single collection having any of two of the three inputs is at most

Thus, the probability that any collection has all three inputs is no more than 3(empHD)2/N3(empHD)^{2}/N. If mpHD=O(N)mpHD=O(\sqrt{N}), then the randomly chosen triple will not jointly appear as the outcome of a single element of a self-attention unit with probability at least 0.90.9, and the transformer will be unexpected to successfully distinguish between the two cases.

We construct an architecture that collects a group of candidate pairs in each layer of single-headed self-attention and verifies whether there exists a triple incorporating each pair that satisfies the summation property. Then, all candidate triples are disposed of, and the subsequent layer collects a new family of candidates.

for sufficiently large cc. We additionally let

We repeat this construction DD times, with the only modifications being the replacement of P1P_{1} and the fact that the second dimension of the embedding remains 1 after being set to that value. After DD layers, the final MLP outputs the value of the second dimension, which will be 1 if and only if the respective xix_{i} belongs to a three-way match. ∎

C.6 Sharper separations for embedded subgraph detection problems

In pursuit of proving separations analogous to the one between Theorem 18 and Conjecture 19, we draw techniques for proving lower bounds for graph problems in the Congest model of distributed computation with restricted bandwidth (Peleg, 2000).At a high level, the Congest model features NN players that communicate in synchronous rounds over a network (an undirected graph with [N][N] as its vertices) to solve a computational problem Peleg (2000). In each round, each player can send a message to each of its neighbors. The computation that each player does with the messages received from its neighbors is unrestricted; the primary resources considered in Congest is the number of rounds of communication and the message sizes. Although Congest is often studied for solving computational problems on input graphs with vertices [N][N], the input graph need not be the same as the communication network.

The former treats XX as a directed graph (where XX need not be symmetric) and asks whether each input belongs to a directed 3-cycle. The latter insists that XX be an undirected graph by enforcing symmetry and determines membership in (undirected) 5-cycles.

However, solving these problems with any transformer model of constant order trivially requires having the product of the precision pp, embedding dimension mm, heads per layer HH, and depth DD grow polynomially with NN, since each attention unit is limited to considering at most pmpm bits of information from each input. Such a lower bound is not interesting for dense graphs, where every vertex may have Ω(N)\Omega(N) incident edges; the bottleneck is not due to any feature of standard attention units (and would persist with higher-order attention).

To circumvent this issue, we consider an augmented self-attention unit, which permits each element of the self-attention tensor to depend on both its respective inner product and on the presence of edges among corresponding inputs.

Let AGd,m,d′,p⊗s\mathcal{AG}_{d,m,d^{\prime},p}^{\otimes s} and TGd,m,d′,pD,H,⊗s\mathcal{TG}_{d,m,d^{\prime},p}^{D,H,\otimes s} denote all such attention units and all such transformers respectively.

Now, we provide four results that exhibit separations between orders of graph self-attention.

The proofs of Theorems 22 and 24 are immediate from the construction. Because each cell of the self-attention tensor has explicit access the the existence of all relevant edges, κ\kappa can be configured to ensure that cell’s value is large if and only if the requisite edges for the desired structure all exist. Taking a softmax with a blank element (like in Theorem 6) ensures that the outcome of the self-attention unit for a given element distinguishes between whether or not it belongs to a 5-cycle or a directed 3-cycle. The output MLP ensure that the proper output is returned.

The key principle of our analysis is that the predominant limitation of a transformer model is in its communication bandwidth and not its computational abilities. We model transformers as having element-wise multi-layer perceptron units with unbounded computational ability (but bounded precision inputs and outputs) and self-attention units, which compute linear combinations of inputs in a carefully regimented way that limits the ability of individual elements to share information with one another. Here, we introduce a specific Congest graph for each sequence length NN and show that every transformer has a communication protocol that simulates its computation in this graph.

For fixed NN, we design an undirected Congest graph GN=(VN,EN)G^{N}=(V^{N},E^{N}) with O(N2)O(N^{2}) nodes, each having degree at most 3. (Note that this graph is not the same as the graph provided as input XX to a transformer; this graph is consistent across all transformers taking input of sequence size NN.) Let u1,…,uNu_{1},\dots,u_{N} be nodes in VNV^{N} corresponding to each input. For every pair i,j∈[N]i,j\in[N], let vi,jv_{i,j} be a node as well. For each i∈[N]i\in[N], let Bi=(Vi,Ei)B_{i}=(V_{i},E_{i}) be a balanced binary trees having root uiu_{i} and leaves vi,1,…,vi,N,v1,i,…,vN,iv_{i,1},\dots,v_{i,N},v_{1,i},\dots,v_{N,i}. Hence, each BiB_{i} has O(N)O(N) vertices of degree 3 and is of depth O(log⁡N)O(\log N). Let VN=V1∪⋯∪VNV^{N}=V_{1}\cup\dots\cup V_{N} and EN=E1∪⋯∪ENE^{N}=E_{1}\cup\dots\cup E_{N}. Noting that E1,…,ENE_{1},\dots,E_{N} are disjoint and that V1,…,VNV_{1},\dots,V_{N} are disjoint, except for leaves vi,jv_{i,j}, we ascertain that GNG^{N} contains O(N2)O(N^{2}) vertices of degree at most 3 and has diameter O(log⁡N)O(\log N). We visualize the graph GNG^{N} with a highlighted tree B1B_{1} in Figure 4.

Before any communication begins, each node uiu_{i} is provided with xix_{i} and each node vi,jv_{i,j} is provided with xi,jx_{i,j} and xj,ix_{j,i}.

After T=O(HD(m+log⁡N))T=O(HD(m+\log N)) rounds of communication, each node uiu_{i} outputs f(X)if(X)_{i}.

Each vi,jv_{i,j}, using their knowledge of xi,jx_{i,j} and xj,ix_{j,i}, computes αi,j:=exp⁡(κ(xi,j,xj,i,yiTQKTyj))\alpha_{i,j}:=\exp(\kappa(x_{i,j},x_{j,i},y_{i}^{\scriptscriptstyle\mathsf{T}}QK^{\scriptscriptstyle\mathsf{T}}y_{j})). This takes zero rounds.

Each uiu_{i} computes ∑j=1Nαi,j\sum_{j=1}^{N}\alpha_{i,j} by propagating each αi,j\alpha_{i,j} in vi,jv_{i,j} up BiB_{i} to uiu_{i}, iteratively summing terms passed up. This takes O(log⁡N)O(\log N) rounds.

Similarly, uiu_{i} computes ∑j=1Nαi,jVTyj\sum_{j=1}^{N}\alpha_{i,j}V^{\scriptscriptstyle\mathsf{T}}y_{j} in O(mlog⁡N)O(m\log N) rounds. Then, it computes

which is the target output of the self-attention unit.

Because all steps are achievable in parallel with O(m+log⁡N)O(m+\log N) rounds, the claim follows. ∎

C.6.2 Reduction from set disjointness

Before proving Theorems 21 and 23 by embedding an instance of a transformer model into an instance of each subgraph identification problem, we first introduce a partition of the vertices VNV^{N} of the Congest graph into those possessed by Alice and Bob for use in a two-party communication protocol. We call those two sets VaNV_{a}^{N} and VbNV_{b}^{N}.

Otherwise, let wp∈VaNw_{p}\in V_{a}^{N} if and only if root ui∈VaNu_{i}\in V_{a}^{N}.

This partition, which we visualize in Figure 5, bounds the number of bits Alice and Bob can exchange by simulating a protocol on Congest graph GNG^{N}.

Suppose Alice and Bob simulate an RR-round pp-bit protocol on Congest communication graph GNG^{N} where Alice has access to all vertices VaNV_{a}^{N} and Bob VbNV_{b}^{N}. No other communication is permitted besides sharing bits as permitted by the Congest protocol between neighboring vertices. Then, Alice and Bob exchange at most O(pRNlog⁡N)O(pRN\log N) bits.

It suffices to show that the partition VaN,VbNV_{a}^{N},V_{b}^{N} induces a cut of size at most O(Nlog⁡N)O(N\log N); this ensures that each can send no more than O(pNlog⁡N)O(pN\log N) bits per round.

If i∈(0,N5]i\in(0,\frac{N}{5}] and j∈(N5,2N5]j\in(\frac{N}{5},\frac{2N}{5}], then xi,j=xj,i=ai,j−N/5x_{i,j}=x_{j,i}=a_{i,j-N/5}.

If i∈(N5,3N5]i\in(\frac{N}{5},\frac{3N}{5}] and j∈(2N5,4N5]j\in(\frac{2N}{5},\frac{4N}{5}], then xi,j=xj,i=δi,j−N/5x_{i,j}=x_{j,i}=\delta_{i,j-N/5}.

If i∈(3N5,4N5]i\in(\frac{3N}{5},\frac{4N}{5}] and j∈(4N5,N]j\in(\frac{4N}{5},N], then xi,j=xj,i=bj−4N/5,i−3N/5x_{i,j}=x_{j,i}=b_{j-4N/5,i-3N/5}.

If i∈(4N5,N]i\in(\frac{4N}{5},N] and j∈(0,N5]j\in(0,\frac{N}{5}], then xi,j=xj,i=δi,j+4N/5x_{i,j}=x_{j,i}=\delta_{i,j+4N/5}.

This ensures that XX has a 5-cycle if and only there exist i,j∈(0,N5]i,j\in(0,\frac{N}{5}] such that ai,jbi,j=1a_{i,j}b_{i,j}=1We consider 5-cycles rather than 4-cycles because a spurious 4-cycle could exist among edges {xi,j:i∈(0,N5], j∈(N5,2N5]}\{x_{i,j}:i\in(0,\frac{N}{5}],\ j\in(\frac{N}{5},\frac{2N}{5}]\}. . In addition, note that under the protocol in Lemma 25, Alice’s and Bob’s inputs aa and bb are known exclusively by nodes belonging to VaNV_{a}^{N} and VbNV_{b}^{N} respectively.

If i∈(0,N4]i\in(0,\frac{N}{4}] and j∈(N2,3N4]j\in(\frac{N}{2},\frac{3N}{4}], then xi,j=ai,j−N/2x_{i,j}=a_{i,j-N/2}.

If i∈(N2,3N4]i\in(\frac{N}{2},\frac{3N}{4}] and j∈(3N4,N]j\in(\frac{3N}{4},N], then xi,j=bj−3N/4,i−N/2x_{i,j}=b_{j-3N/4,i-N/2}.

If i∈(3N4,N]i\in(\frac{3N}{4},N] and j∈(0,N4]j\in(0,\frac{N}{4}], then xi,j=δi,j+3N/4x_{i,j}=\delta_{i,j+3N/4}.

This construction ensures that a directed 3-cycle exists if and only if a corresponding pair of elements in aa and bb are both 1. ∎

Appendix D Experiment details

As such, the total dimension of a sequence element is d1+(q+1)d0=32d_{1}+(q+1)d_{0}=32. The architectures are detailed as follows.

The attention is identical to the description in the paper body, with the additional detail of the width and embedding dimension mm being fixed to 100100.

Figure 6 also contains an MLP, which first flattens the input, then has a single hidden ReLU layer of width 256, before a final linear layer and an output reshaping to match the desired output sequence shapes.

Figure 6 also contains an LSTM, which is a standard pytorch LSTM with 22 layers and a hidden state size 800800, which is 200200 times larger than the target output dimension 44.

Experiments fit the regression loss using Adam and a minibatch size of 32, with default precision, and take a few minutes to run on an NVIDIA TITAN XP, and would be much faster on standard modern hardware.

In Figure 2 and Figure 7, we plot (post-softmax) alignment matrices after T∈{0,1000,40000}T\in\{0,1000,40000\} iterations of Adam. The alignment matrices in Figure 2 are taken from the training example whose loss is the median loss across all examples. Figure 7 is similar, but additionally shows the examples of minimal and maximal loss.

Figure 6 plots training and testing error curves for the same attention architecture as in Figure 2, but with further MLP and LSTM architectures as described above. but also an MLP trained on flattened (vectorized) error bars reflect 55 separate training runs from random initialization. A few variations of these architectures were attempted, however curves did not qualitatively change, and in particular, only the attention layer achieves good generalization across all attempts.