The self-attention mechanism, at the heart of the Transformer model, is able to effectively model pairwise interactions between tokens. However, numerous recent works have shown that it is unable to perform basic tasks involving detecting triples of correlated tokens, or compositional tasks where multiple input tokens need to be referenced to generate a result. Some higher-dimensional alternatives to self-attention have been proposed to address this, including higher-order attention and Strassen attention, which can perform some of these polyadic tasks in exchange for slower, superquadratic running times. In this work, we define a vast class of generalizations of self-attention, which we call poly-attention mechanisms. Our mechanisms can incorporate arbitrary higher-order (tensor) computations as well as arbitrary relationship structures between the input tokens, and they include the aforementioned alternatives as special cases. We then systematically study their computational complexity and representational strength, including giving new algorithms and matching complexity-theoretic lower bounds on the time complexity of computing the attention matrix exactly as well as approximately, and tightly determining which polyadic tasks they can each perform. Our results give interesting trade-offs between different desiderata for these mechanisms, including a tight relationship between how expressive a mechanism is, and how large the coefficients in the model may be so that the mechanism can be approximated in almost-linear time. Notably, we give a new attention mechanism which can be computed exactly in quadratic time, and which can perform function composition for any fixed number of functions. Prior mechanisms, even for just composing two functions, could only be computed in superquadratic time, and our new lower bounds show that faster algorithms for them are not possible.
Figures & tables
Mechanism
Exact cc
Apx cc
Bound
Self-attention
n2+o(1)
n1+o(1)
o(logn)
t -Tensor
n3+o(1)
n1+o(1)
o((logn)1/t)
Strassen
nω+o(1)
n1+o(1)
o(logn)
Tree (new)
n2+o(1)
n1+o(1)
o(logn)
Poly (new)
nt+o(1)
n1+o(1)
o((logn)1/k)
Table 1: This summarizes the running times of both exact and approximate algorithms for these attention variants. For entry-wise approximation (Apx cc), the bound B is the maximum absolute value of the matrix entries such that we can entry-wise approximate the output matrix in near-linear time; the attention polynomial is in t variables and has degree k . Alman and Song (2023) ; Alman and Song (2024) proved bounds for self-attention and tensor-attention, while we prove the rest.
Mechanism
2-fold
3-fold
Self-attention
No
No
3-Tensor
Yes
No
Strassen
Yes
No
Tree (new)
Yes
Yes
Poly (new)
Yes
Yes
Table 2: Compositionality results showing support for function composition. Peng et al. (2024) prove impossibility bounds for self-attention, Kozachinskiy et al. (2025) simulate 2-fold with Strassen-attention, while we prove the rest.
Figure 1: Graphical representation for the tree polynomial h(x1,…,x7)=x1x2+x1x3+x1x4+x2x5+x2x6+x4x7
Figure 2: Accuracy per epoch for learning f1(f2(x)) for sequence length 51 , on a single layer of tree-attention, one layer self-attention and two layer self-attention.
Appendix figures & tables9 assets
Supplementary material from the paper’s appendix.
Appendix
Ri:=U(i,1:r)1A(W(i,1:r)3)T∈R.
Appendix
Algorithm 1 Algorithm to compute entry-wise approximation of Att(S)
Algorithm 2 Algorithm to compute tree attention Att(h)
Att=[d1K(1)(K(2)⊘…⊘K(t))T]eW(2)⊘…⊘W(t),
Appendix
Algorithm 3 Algorithm to compute an entry-wise approximation of Att(h)
Figure 3: Training loss per epoch, averaged over 10 seeds, for learning f1(f2(x)) for sequence length 51 , on a single layer of tree-attention, one layer self-attention and two layer self-attention. Tree-attention learns faster and has less fluctuations.
Figure 4: Accuracy per FLOP, averaged over 10 seeds, for tree-attention, 1-layer self-attention and 2-layer self-attention for learning function composition. Notice that tree-attention learns more efficiently and the learning is stable.
Seq len
1-layer SA (ms)
2-layer SA (ms)
1-layer tree (ms)
1-layer 3-tensor (ms)
1-layer Strassen (ms)
20
1.076±0.057
1.775±0.085
1.367±0.057
1.442±0.062
1.593±0.086
50
1.079±0.048
1.757±0.060
1.363±0.055
2.911±0.044
1.594±0.088
100
1.080±0.048
1.781±0.097
1.374±0.060
13.813±0.051
3.395±0.081
Appendix
Figure 5: Average running time of various attention schemes implemented on NVIDIA A100 GPU. Tree-attention performs as fast as self-attention, implying that hidden constants in the time complexity computations are small.
Tree-attention
Self-attention
Generalization token accuracy
0.727691±0.013486
0.723993±0.008649
Generalization exact match
0.264919±0.127609
0.239024±0.087350
Appendix
Figure 6: Table for mean accuracies and standard deviation over 10 random seeds.
Figure 7: Plots for mean ± one standard deviation over 10 random seeds for token accuracy and exact match accuracy on the generalization set. Tree-attention has higher accuracy than self-attention.
Figure 8: Exact match accuracies with each seed on the generalization set for tree-attention and self-attention. Tree-attention reaches ∼40% accuracy for 4 out of 10 random seeds.