Linear Transformers Are Secretly Fast Weight Programmers

Imanol Schlag, Kazuki Irie, Jürgen Schmidhuber

Introduction

Transformers (Vaswani et al., 2017) have achieved impressive results in a myriad of sequence processing tasks, including machine translation, language modelling (Al-Rfou et al., 2019; Dai et al., 2019; Baevski & Auli, 2019; Radford et al., 2019), and question answering (Devlin et al., 2019), domains previously dominated by recurrent neural networks (Graves, 2013; Bahdanau et al., 2015).

The core component of a Transformer is the self-attention mechanism (Cheng et al., 2016; Parikh et al., 2016; Lin et al., 2017) which was recently connected to the modern Hopfield network (Ramsauer et al., 2021; Krotov & Hopfield, 2016; Demircigil et al., 2017). It extends a form of attention (Bahdanau et al., 2015) originally introduced to complement recurrent neural networks, e.g., (Hochreiter & Schmidhuber, 1997). While relinquishing the recurrence property, all computations across the time axis can be parallelised. However, this comes with drawbacks: self-attention computations scale quadratically with sequence length while the memory of the model grows linearly. Therefore, practitioners are forced to limit the context window to a reasonable size, which in turn makes it impossible to capture longer-term dependencies.

Recent work proposed “linear Transformers” with constant size memory and time complexity linear in sequence length (Katharopoulos et al., 2020; Choromanski et al., 2021; Peng et al., 2021; Shen et al., 2018). This complexity reduction is mainly due to a linearisation of the softmax (reviewed in Sec. 3.2).

Here we emphasize the formal equivalence of this family of linear Transformers and the Fast Weight Controllers or Fast Weight Programmers (FWPs) from the ’90s (Schmidhuber, 1991, 1992, 1993, AI Blog, 2021) (apart from normalisation). The memories of such FWPs contain key-value associations, and an FWP can learn to reprogram them through sequences of differentiable elementary instructions (also called update rules), which are additive outer products between keys and values invented by the FWP.

This view allows us to derive a limitation of the memory capacity of linear Transformers and similar models. When the sequence length exceeds storage capacity, the model may end up in an overcapacity regime (discussed in depth in Sec. 4.1). To properly operate under such a regime, the model should learn to dynamically interact with the memory contents and selectively decide which key-value associations to keep and which ones to delete. The purely additive instruction may be inappropriate for this purpose. Therefore, inspired by recent work on FWPs (Schlag et al., 2021), we introduce an improved programming instruction akin to the famous error-correcting delta-rule (Widrow & Hoff, 1960).

Furthermore, softmax linearisation techniques for Transformers are still underexplored. The existing techniques are either very simplistic (Katharopoulos et al., 2020) or mathematically well explained but complex (Choromanski et al., 2021; Peng et al., 2021). We provide a comprehensive comparison and propose a new method which is both simple and effective.

We demonstrate the benefits of the proposed methods on our own synthetic retrieval dataset (Sec. 6.1), the standard WMT14 English to German machine translation task (Sec. 6.2), and the Wikitext-103 (Merity et al., 2017) language modelling task (Sec. 6.3)Source code used in this paper is available at github.com/ischlag/fast-weight-transformers..

Background on Fast Weight Programmers

Here we review the concepts of Fast Weight Programmers (FWPs) before relating them to linear Transformer variants in Sec. 3.

In standard neural networks, the weights remain fixed after training, unlike the activations, which change depending on the inputs at test time. The general idea of fast weights is to make the weights also variable and input-dependent. This concept was called synaptic modulation (von der Malsburg, 1981), a method for variable binding in neural networks (see e.g. the recent survey by Greff et al. (2020)), or dynamic connections (Feldman, 1982). Von der Malsburg defines the effective weights as a (multiplicative) superposition of conventional, context-independent slow weights, and fast changing, context-dependent fast weights. Hinton & Plaut (1987) studied a net with (additive) superposition of two sets of weights with two different learning rates in a scenario of model retraining. Before 1991, however, no network learned by gradient descent to quickly compute the changes of the fast weight storage of another network or of itself.

where ⊗\otimes denotes the outer product, σ\sigma is an activation function, Wa{\bm{W}}_{a} and Wb{\bm{W}}_{b} are trainable slow weights, while the fast weights W(i){\bm{W}}^{(i)} are generated at each time step ii and serve as a short-term memory. This is a key-value associative memory model in which the write operation is based on a summation (Eq. 2) and the retrieval is a matrix-vector multiplication (Eq. 3). Schmidhuber (1993) describes a recurrent version and discusses “internal spotlights of attention” (such attention terminology is now widely used in the context of transformers). The use of outer products results in a model of associations similar to tensor product presentations (Smolensky, 1990). In fact, outer-product based associative memory can be found in numerous works since Hebb’s informal rule (Hebb, 1949) and its more concrete formal variants (Steinbuch, 1961; Steinbuch & Piske, 1963; Kohonen, 1972; Palm, 1980) including Hopfield networks (Hopfield, 1982; Little, 1974) and bi-directional associative nets (Kosko, 1988). However, these authors described pre-wired rules to associate given patterns with each other. Their systems did not learn to use such rules for associating self-invented patterns like the FWPs since 1991.

The concept of FWPs has been revisited recently (Ba et al., 2016; Schlag & Schmidhuber, 2017), also under different names, e.g., hypernetworks (Ha et al., 2017; Perez et al., 2018; Galanti & Wolf, 2020), dynamic plasticity (Miconi et al., 2018, 2019), dynamic convolution (Klein et al., 2015; Noh et al., 2016; Jia et al., 2016), or lambda networks (Bello, 2021) used for applications including meta-learning (Munkhdalai & Yu, 2017; Munkhdalai & Trischler, 2018; Munkhdalai et al., 2019; Kirsch & Schmidhuber, 2020). FWPs recently also improved memory models through explicit mechanisms for facilitating the replacement of deprecated information and updating associations (Schlag & Schmidhuber, 2018; Schlag et al., 2021).

Relation to Transformers

Ba et al. (2016) have already pointed out a relation between a variant of outer product-based FWPs (Schmidhuber, 1993) and attention (Bahdanau et al., 2015). Katharopoulos et al. (2020) have analysed linearised transformers. We review these derivations, emphasising the relation between Transformers and the FWPs of the previous section.

Now if we remove the softmax in Eq. 7 we obtain:

Denoting by W(i){\bm{W}}^{(i)} the corresponding weight matrix generated from key and value vectors:

we can rewrite Eqs. 4-7 such that they directly relate to Eqs. 1-3 where the activation function σ\sigma is the identity function and without query projection Wq{\bm{W}}_{q}:

2 Linearising Self-Attention

Instead of removing the softmax as in Sec. 3.1, prior works have introduced techniques for linearising the softmax (Tsai et al., 2019), which has been shown to improve computational efficiency of self-attention for long sequences (Katharopoulos et al., 2020; Choromanski et al., 2021; Peng et al., 2021).

By writing the softmax explicitly, Eq. 7 can be written as:

Using the outer-product notation, the numerator is analogous to the case without softmax (Sec. 3.1):

By introducing the fast weight matrix W(i){\bm{W}}^{(i)} and an additional vector z(i){\bm{z}}^{(i)} for the denominator,

forward computations of linear Transformers can be written as (Katharopoulos et al., 2020):

which is a Fast Weight Programmer (Sec. 2) with normalisation. Hence, the core of linear Transformer variants are outer product-based Fast Weight Programmers.

Analysing and Improving Linear Transformers as Fast Weight Programmers

Viewing linear Transformer variants as Fast Weight Programmers provides us with two insights which we investigate in this work: their capacity limits as associative memories (Sec. 4.1), and their ineptness to edit previously stored associations (Sec. 4.2).

Endlessly adding new associations to a memory of finite size, as in Eq. 17, inevitably will reach a limit. In linear attention, information is stored in a matrix and is retrieved using matrix multiplication (see Eq. 19). As a consequence, to prevent associations from interfering with each other upon retrieval, the respective keys need to be orthogonal. Otherwise, the dot product will attend to more than one key and return a linear combination of values. With keys embedded in a ddotd_{\text{dot}} space, there cannot be more than ddotd_{\text{dot}} orthogonal vectors. That is, storing more than ddotd_{\text{dot}} associations will result in a retrieval error. In linear Transformers, when the length of the sequence is longer than ddotd_{\text{dot}}, the model might be in such an overcapacity regime. While we experimentally demonstrate this effect on toy tasks (Sec. 6.1), prior work on tensor product representations allows for a more formal discussion.

Early work in connectionist research investigated the usage of distributed representations as a means for storing symbolic structures. One highly-influential work is the tensor-product-based variable binding mechanism (Smolensky, 1990). A tensor product representation (TPR) of a structured symbolic system consisting of a set of variables and values constructed from outer products of the so called role and filler vectors. These terms directly translate into keys and values in our context. The fast weight memories of Eq. 17 are the most basic form of such representations (second order tensors). Therefore, many results discussed in Smolensky’s work transfer to our model. In particular, Theorem 3.3 and 3.1 of Smolensky (1990) discuss more formally the crosstalk and retrieval error intuitively described in the previous paragraph.

However, we also note an important difference: the classic TPRs of Smolensky (1990) are constructed with a priori knowledge of the symbolic structure. In contrast, our FWPs since 1991, including recent FWPs (Schlag & Schmidhuber, 2018), learn all the vectors involved in constructing such a representation.

2 Improving the FWP’s Programming Instruction

Sec. 4.1 argues that the linear Transformers can end up in an overcapacity regime, if the sequence length LL exceeds the dimension ddotd_{\text{dot}} of the keys. Once in overcapacity, an ideal memory model should dynamically interact with the memory contents and selectively determine which associations to remember or to forget. This is in stark contrast to the standard Transformer which stores immutable pairs of key and value vectors by concatenation, thus increasing the storage size. While such models work well in practice, we consider a model’s capability to update previously acquired knowledge to be critical for many problems. Hence, from the perspective of dynamic interaction with the memory, the purely additive update rule of Eqs. 17 may be sub-optimal. This motivates us to improve the elementary differentiable programming instruction (i.e. the update rule) of FWPs.

As shown in Eq. 24, our programming instruction or update rule is effectively a delta rule with a dynamic learning rate β(i)\beta^{(i)}. The model thus learns to correct the current key to value association. In Appendix B, we formally show the advantage of this approach over the gated update rule concurrently proposed by Peng et al. (2021).

In the equations above, no normalisation is applied to the value we retrieve. A straightforward normalisation can be obtained by following the derivation in Sec. 3.2, i.e. by introducing an accumulator:

and replacing Eqs. 20 and 25 respectively by:

where we define vˉ(1)=0\bar{{\bm{v}}}^{(1)}=0. In this approach, the output y(i){\bm{y}}^{(i)} is a weighted average of β(j)(v(j)−vˉ(j))\beta^{(j)}({\bm{v}}^{(j)}-\bar{{\bm{v}}}^{(j)}) for 1≤j≤i1\leq j\leq i. We refer to this approach as attention normalisation.

This approach, however, has drawbacks. First, the accumulation of positive values in Eq. 26 always grows with the number of steps, and may result in instability. Second, specifically for our update rule, this normalisation is not sufficient to balance the weights between write and remove operations in Eq. 23 (see derivations in Appendix A.2). Here we propose a better approach based on simple normalisation. We divide the effective key and query vectors ϕ(k(i))\phi({\bm{k}}^{(i)}) and ϕ(q(i))\phi({\bm{q}}^{(i)}) by the sum of its components, e.g., for the query:

before applying Eqs. 20-25. A general consequence of this normalisation is intuitively understood by noticing that the output of any matrix-vector operations (like Eq. 25) is a weighted sum of columns of the matrix where weights are the components of the vector; thus, if the vector components sum up to one, the operation can be viewed as an attention over the columns of the matrix. We provide further explanations and precise implications for our FWP in Appendix A.2. We refer to this approach as sum normalisation.

Since this is a simple substitution of ϕ(k(i))\phi({\bm{k}}^{(i)}) and ϕ(q(i))\phi({\bm{q}}^{(i)}) in Eqs. 20-25, one might still ask whether additional attention normalisation is needed. In language modelling experiments (Sec. 6.3), we show that this is not the case.

Linear Attention Functions

For Eq. 13 to define proper attention weights between 0 and 1, the codomain of ϕ\phi should be positive. Another property of ϕ\phi derives from the discussion of memory capacity in Sec. 4.1. The dimensionality of its codomain ddotd_{\text{dot}} defines the model’s capacity. Therefore, by including a transformation which projects the input dimension dkeyd_{\text{key}} to a larger dimension ddotd_{\text{dot}}, the ϕ\phi function can potentially increase the upper bound of the capacity.

2 Katharopoulos’ Linear Attention

3 FAVOR+

In contrast to Katharopoulos et al. (2020)’s ϕ\phi function which merely satisfies positivity (and a good gradient) property, Choromanski et al. (2021) propose a mathematically rigorous method to approximate the softmax with random features. They propose the following ϕ\phi function:

With FAVOR+, the dimension of the codomain ddotd_{\text{dot}} is 2m2m which increases the theoretical capacity of the memory if 2m>dkey2m>d_{\text{key}}. At the same time, the model’s capacity is still limited, and equals the infinite capacity of the softmax memory only when mm goes to infinity, which is never achieved in practice. During training, we redraw these mm random vectors for each mini-batch. During evaluation, we draw a set of mm random vectors once, and keep them fixed. mm is the only hyperparameter of FAVOR+ and influences the quality of the softmax approximation. Choromanski et al. (2021) suggest to choose mm in the order of dkeylog⁡(dkey)d_{\textit{key}}\log(d_{\textit{key}}). This sampling process is the main drawback of FAVOR+ as it introduces variance into the model’s output.

4 Deterministic Parameter-Free Projection (DPFP)

The two previous sub-sections highlight the sub-optimality of the existing ϕ\phi functions. Sampling introduces extra complexity to FAVOR+ (Sec. 5.3), while the Linear Transformer (Sec. 5.2) lacks the ability to project up the dot product dimension. Here we propose an alternative approach called deterministic parameter-free projection (DPFP). It is deterministic and easy to compute like Linear Transformers while increasing the dot product dimension without requiring FAVOR+’s random features.

Figure 1 illustrates this function. The elements of the 4-dimensional space are displayed as the zz component of the four coloured surfaces. The figure shows how each vector in the 2d plane will have a single non-zero component in the 4d space and equally splits the input space into four areas which will be orthogonal in the projected space.

where ν∈{1,2,..,dkey2−1}\nu\in\{1,2,..,d_{\text{key}}2-1\} is a capacity controlling hyperparameter. The codomain dimensionality of ϕ(k)\phi({\bm{k}}) is thus ddot=2dkeyνd_{\text{dot}}=2d_{\text{key}}\nu. Eq. 37 is highly parallelisable because each partial function can be computed independently. This can be implemented in few lines of code as we show in Appendix C.

Experimental Results

Now we present our experimental results on synthetic retrieval problems (Sec. 6.1.1 and 6.1.2), machine translation (Sec. 6.2), and language modelling (Sec. 6.3).

We illustrate the capacity issue (Sec. 4.1) of linear attention and the effectiveness of our new update rule (Sec. 4.2) on two synthetic problems.

In both settings, our toy problem consists of retrieving the correct value from a sequence of randomly sampled key-value associations when queried with one of the used keys. Crucially, the query is given at the end of the sequence, such that the model is not aware of it while processing the inputs. To succeed, the model has to learn to store the observed associations in its memory without interference.

Let K\mathcal{K} and V\mathcal{V} be the finite and fixed sets of keys and values and S=∣K∣=∣V∣S=|\mathcal{K}|=|\mathcal{V}|. Then, the input to the model is the sequence [(k,v)1,...,(k,v)L][(\mathsf{k},\mathsf{v})_{1},...,(\mathsf{k},\mathsf{v})_{L}] followed by q\mathsf{q} where every pair (k,v)∈K×V(\mathsf{k},\mathsf{v})\in\mathcal{K}\times\mathcal{V} is sampled randomly, and q\mathsf{q} is randomly chosen to be one of the LL keys.

In this setting, we experimentally demonstrate the capacity limit of linear attention (Sec. 4.1). We conduct experiments for the various ϕ\phi functions described in Sec. 5. We fix dkeyd_{\text{key}} to be 6464, while different ϕ\phi functions produce different ddotd_{\text{dot}}. We set the sequence length to be equal to the number of unique keys (L=SL=S), and sample the keys and values without replacement to generate the sequences. By varying the sequence length SS, our goal is to show that all linear attention models (using the simple sum update rule of Sec. 3.2) fail at retrieving when SS exceeds ddotd_{\text{dot}}.

All models are trained with a mini-batch size of 3232 until the evaluation loss falls below 0.0010.001 or until lack of progress for 10001000 steps. In Figure 2, the best validation set performance for each model and each SS is displayed (for the learning curves see Appendix D.1). The number of unique keys is initially S=20S=20 and is incremented by 2020 until S=600S=600. The following models are compared: Softmax, Linear-Attention, FAVOR+ with 64, 128, and 512 random features, DPFP-ν\nu with ν∈{1,2,3}\nu\in\{1,2,3\}.

The results support our theoretical analysis. Linear-Attention has a capacity of 6464 due to the choice of dkey=ddot=64d_{\text{key}}=d_{\text{dot}}=64. Experimentally, Linear-Attention begins to accumulate errors with 6060 or more associations. Similarly, DPFP projections 1, 2 and 3 start to accumulate errors as they approach their respective limits at 128128, 256256, and 384384. FAVOR+, on the other hand, fails to achieve a loss of 0 in any experiment. Finally, as expected, softmax attention is outperforming all ϕ\phi functions, although it struggles to fully converge with more than 500 keys.

1.2 Setting 2: Comparing Update Rules

In the second setting, we compare variations of the update rule. Unlike in setting 1, keys and values will be sampled with replacement and sequence length L=2SL=2S. As a result, in the same sequence, multiple keys can be re-assigned to a new value more than once. The expected value to retrieve is the most recent one associated with the query. With every new key, the previous value associated with this key deprecates and the model is required to update its finite size memory. The ability to update values associated with keys is essential to bind context-specific values to a key.

We use DPFP-1 as the ϕ\phi function. The sequence length is fixed at 40 with 20 unique keys and values. While this setting does not exceed the capacity of DPFP-1, our result is independent of the capacity regime (see results for different SS and ϕ\phi in Appendix D.2).

We compare the proposed fast weight memory programming instruction with normalisation of Sec. 4.2 (denoted here by ours) to three baselines: the sum update rule of Sec. 3 (sum rule), and two variants of previous update rules (Schlag et al., 2021): Schlag (2021) and Schlag (2021) with DPFP. Schlag (2021) is simply the model from Schlag et al. (2021) ported to this setting (i.e. without the LSTM layer). Schlag (2021) has neither a ϕ\phi function, nor the sum normalisation term of Sec. 4.2. Instead it uses a tanh⁡\tanh nonlinearity for its key representations. As an ablation we replace it with our DPFP-1 but we don’t use the normalisation term of Sec. 4.2, which we refer to as Schlag (2021) with DPFP.

Figure 3 presents the learning curves. They demonstrate that our new update rule outperforms all other variants. As expected, the baseline sum update rule fails.

2 Machine Translation Experiments

Here we compare ϕ\phi functions on the standard machine translation task. We compare Linear Transformer (Katharopoulos et al., 2020), Performer (Choromanski et al., 2021) and our ϕ\phi function DPFP (Sec. 5.4) to the regular Transformer, complementing prior comparisons, e.g., Tay et al. (2021).

We use the standard WMT14 English to German Translation dataset and standard data setups (Ott et al., 2018; Vaswani et al., 2017). We adapt the recipe of Ott et al. (2019) (see Appendix E) and train Vaswani et al. (2017)’s “big” models for about 4 days on three V100 GPUs. We use the exact same training configurations for all models without model-specific hyper-parameter tuning. We only vary the model hyper-parameters mm in Performers and ν\nu in DPFP models.

Table 1 shows the Bleu score (Papineni et al., 2002; Post, 2018) results. The Performer is as good as the basic Transformer when the number of samples mm is large enough (for ddot=512d_{\text{dot}}=512, we have m=256m=256). In fact, with dkey=64d_{\text{key}}=64, the recommended value for mm is ddotlog⁡(ddot)=266d_{\text{dot}}\log(d_{\text{dot}})=266. Our DPFP model outperforms the Linear Transformer as well as the Performer when ddotd_{\text{dot}} is relatively small; providing a good trade-off between simplicity and performance.

3 Language Modelling Experiments

Toy experimental Setting 2 (Sec. 6.1.2) illustrated the effect of our update rule. Now our goal is to confirm its effectiveness on a large-vocabulary word-level language modelling task, and investigate its further potential.

Our update rule should be evaluated on a dataset with sufficiently long contextual dependencies. We use the standard WikiText-103 (Merity et al., 2017) dataset. WikiText-103 consists of long articles from Wikipedia; the training set contains about 28 K articles with a total of 103 M running words. This results in contextual text blocks of about 3600 words. The validation and test sets also contain similarly long dependencies, respectively with 218 K and 246 K running words for 60 articles each. The vocabulary size is about 268 K words.

We split the training data into LL-word long segments (which is the backpropagation span). Unless stated otherwise, we treat these segments independently during training. For evaluation, we use a batch size of one, and go through the text with a sliding window of size LL, taking into account only the last position for computing perplexity (except in the first segment where all positions are evaluated). This is usually done for Transformers with a limited context (Al-Rfou et al., 2019). Appendix F provides further experimental details.

We first evaluate our update rule in two configurations. In the small configuration, we set the model dimension (same for key, value, and query) DD to 128, and the training and evaluation context length LL to 256. We note that D=H∗ddotD=H*d_{\text{dot}} where HH is the number of heads. HH is set to 8. The feed-forward layer dimension is 2048. The number of layers is 16 in all configurations. In the medium configuration, we set D=256D=256 and L=384L=384. Both configurations represent an overcapacity regime. We evaluate both Linear Transformers (Katharopoulos et al., 2020) and Performers (Choromanski et al., 2021). However, to keep the comparison simple, we set the capacity of Performers (Sec. 5.3) equal to the one of linear Transformers, by the right choice of projection dimension (m=8m=8 and m=16m=16, respectively, in small and medium configurations), even though this limits performance. We do not include DPFP here, since in both configurations even the smallest value for ν\nu provides enough capacity. Here we investigate the effect of the update rule in an overcapacity scenario (see Appendix D.3 for experimental results in a non-overcapacity regime including DPFP). All models can be trained using two V100 GPUs in less than four days. We refer to the Linear Transformer with our delta update rule as a Delta Network. Table 2 shows the perplexity results. In both configurations, our update rule provides convincing improvements over the models with the sum update rule.

We also conduct an ablation study to test the effect of the absolute positional encoding and an extra attention normalisation (Sec. 4.2). Table 3 shows the results. The sum normalisation (Sec. 4.2) is used in all cases: the models diverged otherwise. In contrast, better perplexities are obtained when no additional attention normalisation is applied. We also observe that the absolute positional encoding is not needed, confirming results of prior work (Irie et al., 2019a).

All methods we propose are within the framework of “linear Transformers”. Thus, there is no change to be discussed in terms of complexity which is constant in space and linear in time w.r.t. sequence length. However, our modified update rule introduces a few extra computations. The wall clock time and memory requirement (for the small LM setting) for the Linear Transformer with and without our delta update rule are: 63 K and 66 K words/sec, and 14 and 13 GB respectively in our implementation. The extra resource requirement is thus marginal. As we use custom CUDA kernels for these linear Transformers, they are faster than the regular Transformers implemented in PyTorch which process 33K words/sec and require 17 GB memory. The speed of the DPFP and Performer models (for Table 5 in Appendix with a larger ddotd_{\text{dot}}) are 63 K and 57 K words/sec. Performers are slower because of the sampling logic, which also motivates our DPFP.

Given the constant space requirements, we can feed inputs to linear Transformers for an arbitrary number of steps. To properly assess the model’s ability to process arbitrary long sequences, it is crucial to make the training consistent with the evaluation mode (Irie et al., 2019b). During training, we carry over the fast weight memory from one training segment to the following one, while still limiting the backpropagation span to be within the segment. We train a Delta Net, using neither positional encoding nor attention normalisation (the best setting from Table 3). It was crucial to remove the attention normalisation for the Delta Net since the accumulator blows up as indicated in Sec. 4.2, while for the Linear Transformer, removing it resulted in an even worse perplexity of over 1600. Table 4 shows the corresponding results. The Delta Net yields a slight improvement over the best model with a limited context window (Table 3), unlike the baseline Linear Transformer model with the naive sum update rule which breaks. We also train a Transformer-XL in our medium configuration as a baseline model specifically designed for this use case (Dai et al., 2019; Rae et al., 2020). We evaluate it using different state sizes by changing the Transformer XL’s memory and target segment lengths (see Appendix F for further details). Performance of the Delta Net does not yet match the performance of the Transformer XL when the latter is evaluated with a large state size (large attention window). However, when we take the state size into account (Table 4), we observe that the Delta Net performs very well with a small state size, which is a crucial property in some practical applications (Irie et al., 2020). These results are promising for future work on alternative Transformer models which can run for an unlimited number of steps.

Conclusion

We emphasise the connection between linearised self-attention and Fast Weight Programmers (FWPs, 1991) that program their fast weight memories through sequences of outer products between self-invented key and value patterns. The FWP perspective allows for discussing associative memory capacity limitations of linear attention, and for introducing an alternative differentiable elementary programming instruction that the FWP can use to dynamically edit the memory, akin to the famous delta rule, but such that the FWP can learn to use the rule wisely through gradient descent. We also propose and discuss a new method for linearising attention. Experiments on synthetic and real language tasks demonstrate the effectiveness of our proposals. The FWP perspective opens up new avenues for investigating even better programming instructions and designs for Transformers with finite memory.

Acknowledgements

We thank Sjoerd van Steenkiste, Hubert Ramsauer and Sepp Hochreiter for valuable comments and suggestions on the first version of the manuscript. This research was partially funded by ERC Advanced grant no: 742870, project AlgoRNN, and by Swiss National Science Foundation grant no: 200021_192356, project NEUSYM. We thank NVIDIA Corporation for donating several DGX machines, and IBM for donating a Minsky machine. We also thank Katharopoulos et al. (2020) for releasing their CUDA implementation of Linear Transformers, which was helpful to implement our models.

References

Appendix A Update Rule Derivation

Here we provide the intermediate steps from Eq. 23 to Eq. 24.

By grouping the last two terms, Eq. 23 becomes:

By using the definition of vnew(i){\bm{v}}^{(i)}_{\text{new}} from Eq. 22:

By substituting this expression to Eq. 38, we obtain Eq. 24 ∎.

A.2 Key Sum Normalisation

where {w(1),...,w(i),...,w(dkey)}\{{\bm{w}}^{(1)},...,{\bm{w}}^{(i)},...,{\bm{w}}^{(d_{\text{key}})}\} are the column vectors of W{\bm{W}}. In the context of associative memory, we can interpret this expression as a set of associations with fixed keys e(i){\bm{e}}^{(i)} and the associated values w(i){\bm{w}}^{(i)}.

In this view, any update of W{\bm{W}} can be written as updates of each w(i){\bm{w}}^{(i)}. This perspective allows us to derive the sum normalisation of Sec. 4.2. For that, we start by deriving the update of w(i){\bm{w}}^{(i)}.

Given an arbitrary weight W{\bm{W}}, we consider updating it to W′{\bm{W}}^{\prime} by adding a new association (k,v)({\bm{k}},{\bm{v}}) using our update rule of Sec. 4.2 (where we omit β\beta):

Now by substituting W{\bm{W}} by its expression of Eq. 41:

We can explicitly write down vˉ\bar{{\bm{v}}} as:

which we can substitute in Eq. 48 to obtain:

In Eq. 51, the weight kik_{i} on the positive term v{\bm{v}} is in general not equal to the total weights on the negative terms ∑j=1dkeykikj\sum_{j=1}^{d_{\text{key}}}k_{i}k_{j}. We can force these weights to be balanced by introducing the normalisation: ∑j=1dkeykikj=ki\displaystyle\sum_{j=1}^{d_{\text{key}}}k_{i}k_{j}=k_{i}.

This corresponds to the sum normalisation we introduced in Sec. 4.2 ∎.

Appendix B Formal comparison to Peng et al. (2021)

Concurrently to our work, Peng et al. (2021) proposed the following gated update rule:

which is motivated by the gating mechanism in recurrent neural networks (Hochreiter & Schmidhuber, 1997). In contrast, our update rule of Eq. 24

is driven by an associative memory perspective, relates to the famous error-correcting delta rule, and offers a crucial property.

To illustrate a similarity and a crucial difference between the two update rules, we consider a fast weight matrix W{\bm{W}} which is constructed by two associations (k1,v1)({\bm{k}}_{1},{\bm{v}}_{1}) and (k2,v2)({\bm{k}}_{2},{\bm{v}}_{2}), i.e.

where we assume k1{\bm{k}}_{1} and k2{\bm{k}}_{2} to be orthonormal, and we omit ϕ\phi. Now we consider updating W{\bm{W}} to W′{\bm{W}}^{\prime} by adding a new association (k3,v3)({\bm{k}}_{3},{\bm{v}}_{3}) where k3=k2{\bm{k}}_{3}={\bm{k}}_{2}. Using Peng et al. (2021)’s update rule, we have:

This rule thus updates the value associated with the key k2=k3{\bm{k}}_{2}={\bm{k}}_{3} to be a convex combination of the old and the new values (1−β)v2+βv3(1-\beta){\bm{v}}_{2}+\beta{\bm{v}}_{3}:

However, it also modifies or in the worst case erases the value associated with the key k1{\bm{k}}_{1}:

In contrast, using our update rule, we have:

since vˉ=Wk3=Wk2=v2\bar{{\bm{v}}}={\bm{W}}{\bm{k}}_{3}={\bm{W}}{\bm{k}}_{2}={\bm{v}}_{2}. Our rule thus also updates the value associated with the key k2=k3{\bm{k}}_{2}={\bm{k}}_{3} to be a convex combination of the old and the new values (1−β)v2+βv3(1-\beta){\bm{v}}_{2}+\beta{\bm{v}}_{3}:

while crucially, it keeps the value associated with k1{\bm{k}}_{1} unmodified:

Our update rule thus differs from Peng et al. (2021)’s one on this property of updating associations while keeping other “unrelated” ones intact in an associative memory.

Appendix C DPFP-ν𝜈\nu Implementation

Listing 1 is a simple PyTorch implementation of DPFP-ν\nu (Eq. 37) which consist of two concatenations followed by one element-wise multiplication.

Appendix D Additional Experimental Results

In this section, we provide additional experimental results which we could not include in the main paper because of space limitations.

Figure 4 shows learning curves for the synthetic setting 1 (without replacement) with 600 unique keys and values. The scripts used to generate such figures can be found in our GitHub repository.

D.2 Synthetic Task Setting 2

Figure 5 is a capacity plot for setting 2 with an increasing number of unique keys and queries (analogous to Figure 2 of setting 1 apart from the log-scale of the y-axis). We did not include FAVOR+ in this plot, because its combination with our update rule resulted in not-a-number in this setting.

D.3 Language Modelling

In Sec. 6.3, we evaluated our update rule when the model is under overcapacity regime. Here we present an extra language modelling experiment which evaluate the benefits of our update rule in non-overcapacity scenarios. This also allows us to include DPFP in the evaluation. We train both, Performer and DPFP, in the small setting (D=128D=128, L=256L=256) with m=16m=16 and ν=1\nu=1, resulting in ddot=256d_{\text{dot}}=256 for both cases. Table 5 shows the perplexity results. First we observe that the Performer and DPFP baseline models with the sum update rule do not outperform the Linear Transformer baseline from Table 2. In fact, language modelling might be less affected by the capacity issue than the synthetic retrieval task, as it might not require the exact retrieval. Second we observe that our update rule improves both variants of linear attention over the sum update-rule baselines even in this condition. This indicates the general benefits of our update rule in Fast Weight Programmers. We note that the improvement is larger for the DPFP model than for the Performer. This is similar to Table 2 where our update rule improves the deterministic Linear Transformers more than the Performers. Finally, we note that we also tried the DPFP and Performer models with an increased ddotd_{\text{dot}} by setting ν=2\nu=2 and m=32m=32 respectively. While this increases ddotd_{\text{dot}} by a factor of two, it was not beneficial for this language modelling setting.

Appendix E Details on Machine Translation Experiments

We implemented different ϕ\phi functions in the fairseq tookit (Ott et al., 2019). The Transformer architecture used in the experiment is the one referred to as big in the original Transformer paper (Vaswani et al., 2017): the model has 6 layers each in the encoder and the decoder, with a hidden layer size of 1024 with 16 attention heads, 4096-dimensional feed-forward layers, using 32 K byte-pair encoding sub-word units (Sennrich et al., 2016). fairseq provides a training configuration for the corresponding model (Ott et al., 2018), which we adapted for our infrastructure. We trained our models on three GPUs using a batch size of up to 3584 tokens per GPU and accumulating gradients over 16 batches for 45 epochs, and selected the best model based on the validation Bleu score. In Table 1, we directly report Bleu for different values of ddotd_{\text{dot}}; Table 6 provides the conversion from hyper-parameters mm of Performers or ν\nu in the DPFP to ddotd_{\text{dot}}.

Appendix F Details on Language Modelling Experiments

All our implementations are based on PyTorch (Paszke et al., 2019). Our base language modelling code has been developed by using the public code by Dai et al. (2019) for Transformer-XL as a starting point. For ϕ\phi functions, we ported the same implementation we used for our translation experiments. For the implementation of our update rule, we modified the CUDA kernel for the Linear Transformer made publicly available by Katharopoulos et al. (2020). We note that a custom implementation of the backward pass for fast weights is crucial for language modelling. A naive backward computation generated by automatic differentiation would store the fast weights for each time step, which can quickly hit the GPU memory limit. The custom implementation ensures that we need to store only one set of weights by recomputing the fast weights needed for computing the gradients for each time step in the backward pass (which still remains time-efficient as the operations involved in the computation of our fast weights are rather inexpensive).

Here we provide extra experimental details to complement the descriptions of Sec. 6.3. For the small and medium configurations, we use batch sizes of 96 and 56 sequences, respectively, and train for about 120 and 70 epochs. In both settings, we apply 10% dropout (Hanson, 1990; Srivastava et al., 2014), and train using the Adam optimiser (Kingma & Ba, 2014) with an initial learning rate of 0.00025 and 2000 learning rate warm-up steps. For further details, we refer the readers to our code. For experiments with Transformer-XL (Table 4), we train it with the same backpropagation span as our models (i.e. 384384 words in the medium configuration). The model is trained with memory and target segment lengths of 384. The models with different state sizes in Table 4 are obtained by using different Transformer-XL memory segment lengths at evaluation time. The models with state sizes of 1.05 M, 2.10 M, and 6.29 M are obtained by using memory and target lengths of 64, 128, and 384, respectively. The model with a state size of 0.13 M uses a memory length of 15 and a target length of 1. Like for other models, a batch size of 1 is used for evaluating the Transformer XL. The state sizes in Table 4 are computed as follows. The per-layer state size of the Linear Transformer and the Delta Net are: number of heads (here 8) ×\times fast weight matrix size which is per-head key dimension (here 32) ×\times per-head value dimension (here 32). This yields a total size of 8,192. The per-layer state size of the Transformer XL is: memory segment length ×\times target segment length ×\times (total key dimension, here 256 ++ total value dimension, here 256). We obtain the total state size we report in Table 4 by multiplying the per-layer state size by the number of layers which is 16 for all our models.