Organizations: Computing and Mathematical Sciences California Institute of Technology · Lawrence Berkeley National Laboratory University of California, Berkeley ICSI
Operator learning on probability measures can be accomplished with transformers. For measures with polynomial tails, the exponential weighting in softmax can make the corresponding measure-level attention integrals diverge. This motivates replacing the exponential with slower-growing functions. We construct two benchmarks for operator learning on measures with closed-form targets. We use these benchmarks to study attention kernel growth and data transformation in post-norm transformers. Without data transformation, the softmax models exhibit ensemble collapse on both heavy-tailed benchmarks, while the three slower-growing kernels avoid collapse. Symlog preprocessing allows softmax to avoid collapse on the matrix inverse task but not on the sheared swap task. On the Gaussian control, all four kernels perform similarly. We also examine how sample size affects the sensitivity of empirical energy and Wasserstein distances to tail differences. These results support slower-growing attention kernels as an effective design choice for post-norm transformers learning from heavy-tailed ensembles.
Figures & tables
kernel
g(s) , s>0
g(s) , s≤0
growth
softmax
es
es
exponential
squared-ELU
(s+1)2
e2s
quadratic
linear-ELU
s+1
es
linear
log-ELU
log(1+s)+1
(1−s)−1
logarithmic
Table 1: The attention kernel suite, ordered by the growth rate of the positive branch, given in the last column as s→∞ . Every member is strictly positive and C1 at the origin with g(0)=1 .
Figure 1: Sensitivity of the energy distance, 2-Wasserstein distance, and Hill gap to tail-index differences at varying sample budgets N . Each plotted point reports a value computed from two sample sets drawn from a measure pair as described in Section 4.1 . The rows report the energy distance, 2-Wasserstein distance, and absolute difference between the two sample Hill estimates (Hill gap), respectively. Each plot displays the value against a∗ . Every value is averaged over multiple draws, then shifted and scaled as described in Appendix F.2 , so that a∗=aref corresponds to a value of 0. The rows use separate vertical scales. The colors denote the quantile value q0 of the seam.
Figure 2: Approximation results for the sheared swap operator using post-norm at N=16384 . The grid organizes models trained with the symlog transform (right), without the symlog transform (left), on mirrored Pareto marginals (top), and on shifted Gaussian marginals (bottom). Each plot reports a median and mean energy distance for every attention kernel. The vertical lines mark the matching i.i.d. sampling reference under the same summary statistic, in the color of that statistic. Bars are 95% bootstrap intervals over evaluation trials (Appendix H.2 ). Scores below the sampling reference are possible for model-generated ensembles, as explained in Appendix H.1 . The label ensemble collapse marks runs with degenerate output ensembles; Appendix I shows an example.
Figure 3: Matrix inverse results on mirrored Pareto marginals, with the same setup as Figure 2 .
Appendix figures & tables21 assets
Supplementary material from the paper’s appendix.
Appendix
Figure 4: The Pareto, mirrored Pareto, and Gaussian densities centered at 1. Left: the Pareto density ( 5 ) on log–log axes for two tail indices. Center: the mirrored Pareto density ( 6 ) at the same two tail indices, with the seam z=λ dotted, only the vertical axis is plotted on a log scale. Right: the mirrored Pareto density against a shifted Gaussian density, only the vertical axis is plotted on a log scale.
Figure 5: The transport a trained model realizes. Arrows run from each input point to the model’s output, for the linear-ELU post-norm model evaluated at j=k=0 , where the shear is the identity and the sheared swap reduces to a coordinate exchange. 140 arrows are drawn from a set of 16384 points, sampled uniformly from the 94% of the set that falls inside the window; for this parameter draw the two marginals carry tail indices a1=2.50 and a2=1.76 .
marginal family
sheared swap M
matrix inverse B
decided by
mirrored Pareto
yes
yes
Propositions 4 and 5
shifted Gaussian
yes
no
Propositions 4 and 6
Appendix
Table 2: Identifiability of each operator paired with each marginal family, in the every-pair sense of Definition 2 .
Figure 6: The attention kernel suite of Table 1 . Left: the four maps on a linear vertical axis over the scores where they differ most visibly, with the common value g(0)=1 marked. Center: the same maps on a logarithmic vertical axis over a wider positive range, where the ordering of the last column of Table 1 — exponential, quadratic, linear, logarithmic — appears as a difference in slope. Right: The same maps over a wider negative range. Every map is strictly positive, so the row normalizer in ( 7 ) cannot vanish.
hyperparameter
value
transformer blocks
6
model dimension dmodel
256
feed-forward width
512
attention heads
1
feed-forward activation
gelu
normalization type
layer
Appendix
Table 3: Shared model and optimizer settings. The main experiments vary the attention kernel and data transform; the supplementary experiments also vary normalization placement. Attention arithmetic is described in Appendix E.4 .
run family
s / step
peak memory
hours for 50,000 steps
softmax (FlashAttention, bf16)
1.87
69.4 GiB
26
sub-exponential (chunked + recompute)
2.51
78.8 GiB
35
Appendix
Table 4: Measured on an H200 at the trained shapes. Per-step times are end-to-end training steps and peak memory is the whole step. The sub-exponential kernels are more expensive to train since they lack a fused kernel.
Figure 7: The splice ( 31 ) at q0=0.9 , reference index aref=1.5 and scale λ=1 , on log–log axes. Left: densities. Right: survival functions. The shaded region lies below the seam u0 (dotted, denoted "knot"), where every candidate coincides with the reference; the candidates differ from it only to the right of the seam, over a fixed mass 1−q0 .
mirrored Pareto
shifted Gaussian
norm
kernel
none
signed-log
none
signed-log
pre-norm
Softmax ( es )
0.857
0.817
1.123
1.101
Squared ELU ( (s+1)2 )
0.844
0.827
1.127
1.143
Linear ELU ( s+1 )
0.815
0.845
1.123
1.074
Log ELU ( log(1+s)+1 )
0.795
0.806
1.057
1.097
post-norm, res. scale
Softmax ( es )
collapse
collapse
1.119
1.097
Appendix
Table 5: Sheared swap energy distance as a multiple of the noise floor at N=16384 , under the median (top) and the mean (bottom). Bold marks ratios whose bootstrap interval contains or lies below one. collapse marks runs with degenerate output ensembles, as in Figure 2 .
Figure 8: The equivalent of Figure 2 for pre-norm.
mirrored Pareto
norm
kernel
none
signed-log
pre-norm
Softmax ( es )
1.085
1.053
Squared ELU ( (s+1)2 )
1.145
1.133
Linear ELU ( s+1 )
1.143
1.092
Log ELU ( log(1+s)+1 )
1.181
1.164
post-norm, res. scale
Softmax ( es )
collapse
1.120
Appendix
Table 6: The counterpart of Table 5 for the matrix inverse operator on mirrored Pareto marginals.
Figure 9: The equivalent of Figure 3 for pre-norm.
Figure 10: The empirical CDF of the per-trial energy distance for the sheared swap operator on mirrored Pareto marginals at N=16384 , organized by normalization placement (rows) and symlog transform (columns). The horizontal axis is the energy distance on a log scale and the vertical axis is the cumulative count of each curve’s 128 trials, so a curve further left has lower energy distances. The black curve is the noise floor. Softmax is absent from both post-norm panels because it collapses there.
Figure 11: The equivalent of Figure 10 for the matrix inverse operator. Here only the post-norm softmax curve without the symlog transform is omitted for collapse.
Figure 12: Model outputs for the sheared swap operator on mirrored Pareto marginals, using post-norm and no symlog transform, for a single evaluation instance of 16384 points. Left: the target measure. Center: the softmax output, concentrated near the single location marked by the arrow. Right: the log-ELU output. All three panels share symmetric-log axes on both coordinates.
Figure 13: Training curves for the four attention kernels of Table 1 on the sheared swap operator with mirrored Pareto marginals, using post-norm, with and without the symlog transform. Faint traces are the per-step loss on a fresh batch and heavy traces are a rolling median. Softmax (red) descends alongside the other three kernels, reaching 1.8×10−2 at step 733 without the transform and 6.3×10−4 at step 3931 with it, then diverges abruptly at step ≈3500 and ≈4500 to the collapsed value it holds for the remaining 45,000 steps. The other three kernels continue to descend. Collapse is therefore not a consequence of stopping training early.
mirrored Pareto
norm
kernel
none
signed-log
pre-norm
Softmax ( es )
7.6×10−3
2.5×10−3
Squared ELU ( (s+1)2 )
5.6×10−3
1.5×10−3
Linear ELU ( s+1 )
2.1×10−3
6.7×10−4
Log ELU ( log(1+s)+1 )
8.1×10−3
9.3×10−4
post-norm, res. scale
Softmax ( es )
4.3×108
1.8×105
Appendix
Table 7: E(e,e^) from ( 41 ) at N=16384 for the sheared swap operator. The statistic has expectation 0 when the distributions of candidate and reference Hill errors agree. Axis 0 carries the lighter tail ( a∈[2,3] ) and axis 1 the heavier ( a∈[1+ϵ,2] ). The two softmax post-norm entries are the collapsed runs of Table 5 .
mirrored Pareto
norm
kernel
none
signed-log
pre-norm
Softmax ( es )
1.8×10−2
5.1×10−4
Squared ELU ( (s+1)2 )
8.8×10−3
4.7×10−5
Linear ELU ( s+1 )
4.5×10−3
3.5×10−3
Log ELU ( log(1+s)+1 )
8.1×10−3
−3.4×10−4
post-norm, res. scale
Softmax ( es )
2.1×108
−2.5×10−4
Appendix
Table 8: The equivalent of Table 7 for the matrix inverse operator.
Figure 14: E(e,e^) for the sheared swap operator without the symlog transform, using pre-norm (top) and post-norm (bottom). The vertical line marks 0 . Larger population values indicate a greater difference between the two distributions of Hill errors; the finite-trial estimate can be negative. The two axes are scaled separately. Appendix J.1 defines the sample sizes and tail fractions. Runs that exhibited ensemble collapse in Section 4.2 are again labeled "ensemble collapse".
Figure 15: The equivalent of Figure 14 with the symlog transform.
Figure 16: The equivalent of Figure 14 for the matrix inverse operator.
Figure 17: The equivalent of Figure 14 for the matrix inverse operator and symlog transform.
Learning mappings between infinite-dimensional function spaces, or operator learning, is essential for many machine learning applications. Although transformer-based operators are popular, they often rely on token-wise attention. These methods treat continuous fields as discrete tokens and usually ignore the global functional structure. We introduce \emph{Functional Attention}, which reinterprets attention as a functional correspondence between adaptive bases. Inspired by geometric functional maps, our method replaces softmax affinities with structured linear operators. This yields a compact, generalizable, resolution-invariant representation that explicitly captures global dependencies. Experiments demonstrate that \emph{Functional Attention} can match state-of-the-art performance in many operator learning tasks, including solving PDEs, 3D segmentation, and regression, while remaining robust to varying discretizations. Project page is available at https://github.com/xjffff/FUNCATTN.
Jiefang Xiao, Maolin Gao, Simon Weber +2
Technical University of Munich, Germany · Munich Center for Machine Learning (MCML), Germany · PIXL, Department of Computer Science, University of Oxford, United Kingdom +1
Low-precision Transformer systems increasingly quantize attention matrix multiplications, while softmax often remains at higher precision. During pretraining, an approximate softmax changes the gradients that train the model as well as its forward computation. We study this interaction with K-interval attention, which approximates the exponential using K+1 grid values. We vary per-row grid calibration, interpolation versus hard rounding, and the placement of a straight-through surrogate relative to normalization. We derive the corresponding backward rules, including calibration derivatives, and compare these choices in pretraining experiments matched on model, data, and optimizer. Detaching the row extrema leaves the forward computation unchanged but produces a delayed increase in validation loss. With hard rounding at K=4, min-max calibration and a pre-normalization surrogate incur a large loss gap; changing either choice substantially reduces it. At 124M parameters and 2.5B training tokens, fixed-window calibration with a post-normalization surrogate yields a validation loss gap of +0.019 nats relative to softmax at K=4, and with a pre-normalization surrogate yields +0.004 nats at K=16.
Deep decoder-only Transformers often replace the original Post-Norm architecture with Pre-Norm variants because Post-Norm training is highly sensitive to warmup and learning rate under conventional initialization schemes. Although prior work has identified rank collapse and gradient vanishing as related symptoms, it remains poorly understood how causal attention creates high-similarity representations and why training dynamics fail to repair them. We give a two-stage analysis of Post-Norm rank collapse using token similarity as a scalar state variable. First, at initialization, causal attention acts approximately as a prefix-averaging operator that increases token similarity across depth, while the SwiGLU branch contributes only a smaller damping effect. Second, once training enters a high-similarity regime, growth of pre-normalization residual norms makes the RMSNorm backward factor contractive; under mild conditions, gradients to earlier layers decay geometrically. As a complementary result, we characterize the properties of a collapsed network: its best predictor is frequency distribution with relatively high loss floor, and gradients in collapsed layers vanish at frequency distribution. Experiments on 48-layer decoder-only Transformers trained on C4 dataset match the predicted initialization-time similarity growth and collapse-time gradient contraction, and show that collapsed runs stay near the predicted frequency loss. Together, these results distinguish the forward similarity amplification and backward repair incapacity in Post-Norm collapse, while also characterizing the behavior of collapsed networks.
Xingjian Wang, Qingyu Han, Xiaodong Luo +1
The Chinese University of Hong Kong, Shenzhen · Shenzhen Research Institute of Big Data