Random Feature Attention
Hao Peng, Nikolaos Pappas, Dani Yogatama, Roy Schwartz, Noah A. Smith, Lingpeng Kong
Introduction
Transformer architectures (Vaswani et al., 2017) have achieved tremendous success on a variety of sequence modeling tasks (Ott et al., 2018; Radford et al., 2018; Parmar et al., 2018; Devlin et al., 2019; Parisotto et al., 2020, inter alia). Under the hood, the key component is attention (Bahdanau et al., 2015), which models pairwise interactions of the inputs, regardless of their distances from each other. This comes with quadratic time and memory costs, making the transformers computationally expensive, especially for long sequences. A large body of research has been devoted to improving their time and memory efficiency (Tay et al., 2020c). Although better asymptotic complexity and prominent gains for long sequences have been achieved (Lee et al., 2019; Child et al., 2019; Beltagy et al., 2020, inter alia), in practice, many existing approaches are less well-suited for moderate-length ones: the additional computation steps required by some approaches can overshadow the time and memory they save (Kitaev et al., 2020; Wang et al., 2020; Roy et al., 2020, inter alia).
This work proposes random feature attention (Rfa), an efficient attention variant that scales linearly in sequence length in terms of time and space, and achieves practical gains for both long and moderate length sequences. Rfa builds on a kernel perspective of softmax (Rawat et al., 2019). Using the well-established random feature maps (Rahimi & Recht, 2007; Avron et al., 2016; §2), Rfa approximates the dot-then-exponentiate function with a kernel trick (Hofmann et al., 2008): . Inspired by its connections to gated recurrent neural networks (Hochreiter & Schmidhuber, 1997; Cho et al., 2014) and fast weights (Schmidhuber, 1992), we further augment Rfa with an optional gating mechanism, offering a straightforward way of learning with recency bias when locality is desired.
Rfa and its gated variant (§3) can be used as a drop-in substitute for the canonical softmax attention, and increase the number of parameters by less than 0.1%. We explore its applications in transformers on language modeling, machine translation, and long text classification (§4). Our experiments show that Rfa achieves comparable performance to vanilla transformer baselines in all tasks, while outperforming a recent related approach (Katharopoulos et al., 2020). The gating mechanism proves particularly useful in language modeling: the gated variant of Rfa outperforms the transformer baseline on WikiText-103. Rfa shines in decoding, even for shorter sequences. In our head-to-head comparison on machine translation benchmarks, Rfa decodes around faster than a transformer baseline, without accuracy loss. Comparisons to several recent efficient transformer variants on three long text classification datasets show that Rfa is competitive in terms of both accuracy and efficiency. Our analysis (§5) shows that more significant time and memory efficiency improvements can be achieved for longer sequences: 12 decoding speedup with less than 10% of the memory for 2,048-length outputs.
Background
The attention mechanism (Bahdanau et al., 2015) has been widely used in many sequence modeling tasks. Its dot-product variant is the key building block for the state-of-the-art transformer architectures (Vaswani et al., 2017). Let denote a sequence of query vectors, that attend to sequences of key and value vectors. At each timestep, the attention linearly combines the values weighted by the outputs of a softmax:
is the temperature hyperparameter determining how “flat” the softmax is (Hinton et al., 2015). in self-attention; they may differ, e.g., in the cross attention of a sequence-to-sequence model.
Calculating attention for a single query takes time and space. For the full sequence of queries the space amounts to . When the computation cannot be parallelized across the queries, e.g., in autoregressive decoding, the time complexity is quadratic in the sequence length.
2 Random Feature Methods
The theoretical backbone of this work is the unbiased estimation of the Gaussian kernel by Rahimi & Recht (2007). Based on Bochner’s theorem (Bochner, 1955), Rahimi & Recht (2007) proposed random Fourier features to approximate a desired shift-invariant kernel. The method nonlinearly transforms a pair of vectors and using a random feature map ; the inner product between and approximates the kernel evaluation on and . More precisely:
When -dimensional random vectors are independently sampled from ,
Variance of the estimation is inversely proportional to (Appendix A.2; Yu et al., 2016).
Random feature methods proved successful in speeding up kernel methods (Oliva et al., 2015; Avron et al., 2017; Sun, 2019, inter alia), and more recently are used to efficiently approximate softmax (Rawat et al., 2019). In §3.1, we use it to derive an unbiased estimate to and then an efficient approximation to softmax attention.
Model
This section presents Rfa (§3.1) and its gated variant (§3.2). In §3.3 we lay out several design choices and relate Rfa to prior works. We close by practically analyzing Rfa’s complexity (§3.4).
Rfa builds on an unbiased estimate to from Theorem 1, which we begin with:
denotes the outer product between vectors, and corresponds to the temperature term in Eq. 1.
Rfa can be used as a drop-in-replacement for softmax-attention.
The input is revealed in full to cross attention and encoder self-attention. Here Rfa calculates attention using Eq. 5.
denotes the size of Appendix A.1 summarizes the computation procedure of Rfa, and Figure 1 compares it against the softmax attention. Appendix A.3 derives causal Rfa in detail.
Analogously to the softmax attention, Rfa has its multiheaded variant (Vaswani et al., 2017). In our experiments we use causal Rfa in a transformer language model (§4.1), and both cross and causal Rfa in the decoder of a sequence-to-sequence machine translation model.
2 Rfa-Gate: Learning with Recency Bias
The canonical softmax attention does not have any explicit modeling of distance or locality. In learning problems where such inductive bias is crucial (Ba et al., 2016; Parmar et al., 2018; Miconi et al., 2018; Li et al., 2019, inter alia), transformers heavily rely on positional encodings. Answering to this, many approaches have been proposed, e.g., learning the attention spans (Sukhbaatar et al., 2019; Wu et al., 2020), and enhancing the attention computation with recurrent (Hao et al., 2019; Chen et al., 2019) or convolutional (Wu et al., 2019; Mohamed et al., 2019) components.
Rfa faces the same issue, but its causal attention variant (Eq. 6) offers a straightforward way of learning with recency bias. We draw inspiration from its connections to RNNs, and augment Rfa with a learned gating mechanism (Hochreiter & Schmidhuber, 1997; Cho et al., 2014; Peng et al., 2018, inter alia):
and are learned parameters, and is the input representation at timestep . In multihead attention (Vaswani et al., 2017), and are calculated from using learned affine transformations. By multiplying the learned scalar gates against the hidden state , history is exponentially decayed, favoring more recent context.
The gating mechanism shows another benefit of Rfa: it would be otherwise more difficult to build similar techniques into the softmax attention, where there is no clear sense of “recurrence” (Appendix A.5). It proves useful in our language modeling experiments (§4.1).
3 Discussion
On query and key norms, and learned random feature variance. Eq. 5 assumes both the query and keys are of norm-1. It therefore approximates a softmax attention that normalizes the queries and keys before multiplying them, and then scales the logits by dividing them by . Empirically, this normalization step scales down the logits (Vaswani et al., 2017) and enforces that . In consequence, the softmax outputs would be “flattened” if not for , which can be set a priori as a hyperparameter (Yu et al., 2016; Avron et al., 2017; Sun, 2019, inter alia). Here we instead learn it from data with the reparameterization trick (Kingma & Welling, 2014):
is the identity matrix, and denotes elementwise product between vectors. -dimensional vector is learned, but random vectors are not. This departs from Eq. 2 by lifting the isotropic assumption imposed on the Gaussian distribution: note the difference between the vector in Eq. 8 and the scalar in Eq. 3. We find this improves the performance in practice (§4), even though the same result in Theorem 1 may not directly apply.
This norm-1 constraint is never mandatory. Rather, we employ it for notation clarity and easier implementation. In preliminary experiments we find it has little impact on the performance when is set properly or learned from data. Eq. 12 in Appendix A presents Rfa without imposing it.
Going beyond the Gaussian kernel. More broadly, random feature methods can be applied to a family of shift-invariant kernels, with the Gaussian kernel being one of them. In the same family, the order-1 arc-cosine kernel (Cho & Saul, 2009) can be approximated with feature map: (Alber et al., 2017). Apart from replacing the sinusoid functions with , it constructs in the same way as Eq. 8. In our experiments, the Gaussian and arc-cosine variants achieve similar performance. This supplements the exploration of alternatives to softmax in attention (Tsai et al., 2019; Gao et al., 2019).
Relations to prior work. Katharopoulos et al. (2020) inspire the causal attention variant of Rfa. They use a feature map based on the exponential linear unit activation (Clevert et al., 2016): . It significantly underperforms both the baseline and Rfa in our controlled experiments, showing the importance of a properly-chosen feature map. Random feature approximation of attention is also explored by a concurrent work (Choromanski et al., 2020), with applications in masked language modeling for proteins. They propose positive random features to approximate softmax, aiming for a lower variance in critical regions. Rfa instead normalizes the queries and keys before random projection to reduce variance. Going beyond both, Rfa establishes the benefits of random feature methods as a more universal substitute for softmax across all attention variants, facilitating its applications in, e.g., sequence-to-sequence learning.
There are interesting connections between gated Rfa and fast weights (Schmidhuber, 1992; 1993; Ba et al., 2016; Miconi et al., 2018, inter alia). Emphasizing recent patterns, they learn a temporal memory to store history similarly to Eqs. 7. The main difference is that Rfa additionally normalizes the output using as in Eq. 6, a by-product of approximating softmax’s partition function. It is intriguing to study the role of this normalization term, which we leave to future work.
4 Complexity Analysis
Time. Scaling linearly in the sequence lengths, Rfa needs less computation (in terms of number of operations) for long sequences. This implies speedup wherever the quadratic-time softmax attention cannot be fully-parallelized across time steps. More specifically:
Significant speedup can be expected in autoregressive decoding, both conditional (e.g., machine translation) and unconditional (e.g., sampling from a language model). For example, 1.9 speedup is achieved in our machine translation experiments (§4.2); and more for longer sequences (e.g., 12 for 2,048-length ones; §5).
Some applications (e.g., language modeling, text classification) reveal inputs to the model in full.A causal masking is usually used to prevent the model from accessing future tokens in language models. When there are enough threads to parallelize softmax attention across time steps, hardly any speedup from Rfa can be achieved; when there are not, typically for very long sequences (1,000), substantial speed gain is possible. For example, Rfa does not achieve any speedup when working with 512-length context (§4.1), but achieves a speedup with 4,000-length context (§4.3).
Memory. Asymptotically, Rfa has a better memory efficiency than its softmax counterpart (linear vs. quadratic). To reach a more practical conclusion, we include in our analysis the cost of the feature maps. ’s memory overhead largely depends on its size . For example, let’s consider the cross attention of a decoder. Rfa uses space to store , , and (Eq. 5; line 12 of Algo. 2). Rfa never constructs the tensor , but sequentially processes the sequence. In contrast, softmax cross attention stores the encoder outputs with memory, with being the source length. In this case Rfa has a lower memory overhead when . Typically should be no less than in order for reasonable approximation (Yu et al., 2016); In a transformer model, is the size of an attention head, which is usually around 64 or 128 (Vaswani et al., 2017; Ott et al., 2018). This suggests that Rfa can achieve significant memory saving with longer sequences, which is supported by our empirical analysis in §5. Further, using moderate sized feature maps is also desirable, so that its overhead does not overshadow the time and memory Rfa saves. We experiment with at and ; the benefit of using is marginal.
Appendix A.6 discusses the time and space complexity in more detail, and Appendix C.2 studies the effect of random feature size on performance.
Experiments
We evaluate Rfa on language modeling, machine translation, and long text classification.
Setting. We experiment with WikiText-103 (Merity et al., 2017). It is based on English Wikipedia. Table 5 in Appendix B summarizes some of its statistics. We compare the following models:
Base is our implementation of the strong transformer-based language model by Baevski & Auli (2019).
Rfa builds on Base, but replaces the softmax attention with random feature attention. We experiment with both Gaussian and arc-cosine kernel variants.
Rfa-Gate additionally learns a sigmoid gate on top of Rfa (§3.2). It also has a Gaussian kernel variant and a arc-cosine kernel one. This gating technique is specific to Rfa variants, in the sense that it is less intuitive to apply it in Base.
is a baseline to Rfa. Instead of the random feature methods it uses the feature map, as in Katharopoulos et al. (2020).
To ensure fair comparisons, we use comparable implementations, tuning, and training procedure. All models use a 512 block size during both training and evaluation, i.e., they read as input a segment of 512 consecutive tokens, without access to the context from previous mini-batches. Rfa variants use 64-dimensional random feature maps. We experiment with two model size settings, small (around 38M parameters) and big (around 242M parameters); they are described in Appendix B.1 along with other implementation details.
Closing this section, we explore a “stateful” variant of Rfa-Gate-Gaussian. It passes the last hidden state to the next mini-batch during both training and evaluation, a technique commonly used in RNN language models (Merity et al., 2018). This is a consequence of Rfa’s RNN-style computation, and is less straightforward to be applicable in the vanilla transformer models.Some transformer models use a text segment from the previous mini-batch as a prefix (Baevski & Auli, 2019; Dai et al., 2019). Unlike Rfa, this gives the model access to only a limited amount of context, and significantly increases the memory overhead. From the last row of Table 1 we see that this brings a more than 1.5 test perplexity improvement.
2 Machine Translation
Datasets. We experiment with three standard machine translation datasets.
WMT14 EN-DE and EN-FR (Bojar et al., 2014). Our data split and preprocessing follow those of Vaswani et al. (2017). We share the source and target vocabularies within each language pair, with 32,768 byte pair encoding types (BPE; Sennrich et al., 2016).
IWSLT14 DE-EN (Cettolo et al., 2014) is based on TED talks. The preprocessing follows Edunov et al. (2018). Separate vocabularies of 9K/7K BPE types are used for the source and target.
Table 5 in Appendix B summarizes some statistics of the datasets.
Setting. We compare the Rfa variants described in §4.1. They build on a Base model that is our implementation of the base-sized transformer (Vaswani et al., 2017). All Rfa models apply random feature attention in decoder cross and causal attention, but use softmax attention in encoders. This setting yields the greatest decoding time and memory savings (§3.4). We use 128/64 for in cross/causal attention. Rfa-Gate learns sigmoid gates in the decoder causal attention. The baseline uses the same setting and applies feature map in both decoder cross and causal attention, but not in the encoders. Further details are described in Appendix B.2.
Results. Table 2 compares the models’ test set BLEU on three machine translation datasets. Overall both Gaussian and arc-cosine variants of Rfa achieve similar performance to Base on all three datasets, significantly outperforming Katharopoulos et al. (2020). Differently from the trends in the language modeling experiments, here the gating mechanism does not lead to substantial gains. Notably, all Rfa variants decode more than faster than Base.
3 Long Text Classification
We further evaluate Rfa’s accuracy and efficiency when used as text encoders on three NLP tasks from the recently proposed Long Range Arena benchmark (Tay et al., 2021), designed to evaluate efficient Transformer variants on tasks that require processing long sequences. https://github.com/google-research/long-range-arena
Experimental setting and datasets. We compare Rfa against baselines on the following datasets:
ListOps (LO; Nangia & Bowman, 2018) aims to diagnose the capability of modelling hierarchically structured data. Given a sequence of operations on single-digit integers, the model predicts the solution, also a single-digit integer. It is formulated as a 10-way classification. We follow Tay et al. (2021) and consider sequences with 500–2,000 symbols.
Character-level text classification with the IMDb movie review dataset (Maas et al., 2011). This is a binary sentiment classification task.
Character-level document retrieval with the ACL Anthology Network (AAN; Radev et al., 2009) dataset. The model classifies whether there is a citation between a pair of papers.
To ensure fair comparisons, we implement Rfa on top of the transformer baseline by Tay et al. (2021), and closely follow their preprocessing, data split, model size, and training procedure. Speed and memory are evaluated on the IMDb dataset. For our Rfa model, we use for the IMDb dataset, and for others. We refer the readers to Tay et al. (2021) for further details.
Results. From Table 3 we can see that Rfa outperforms the transformer baseline on two out of the three datasets, achieving the best performance on IMDb with 66% accuracy. Averaging across three datasets, Rfa outperforms the transformer by 0.3% accuracy, second only to Zaheer et al. (2020) with a 0.1% accuracy gap. In terms of time and memory efficiency, Rfa is among the strongest. Rfa speeds up over the transformer by –, varying by sequence length. Importantly, compared to the only two baselines that perform comparably to the baseline transformer model (Tay et al., 2020a; Zaheer et al., 2020), Rfa has a clear advantage in both speed and memory efficiency, and is the only model that is competitive in both accuracy and efficiency.
Analysis
Decoding time and memory varying by sequence length. §3.4 shows that Rfa can potentially achieve more significant speedup and memory saving for longer sequences, which we now explore.
We use a simulation conditional generation experiment on to compare Rfa’s sequence-to-sequence decoding speed and memory overhead against the baseline’s. Here we assume the input and output sequences are of the same length. The compared models are of the same size as those described in §4.2, with 6-layer encoders and decoders. Other hyperparameters are summarized in Appendix B.2. All models are tested using greedy decoding with the same batch size of 16, on a TPU v2 accelerator.
From Figures 2 (a) and (b) we observe clear trends. Varying the lengths, both Rfa variants achieve consistent decoding speed with nearly-constant memory overhead. In contrast, the baseline decodes slower for longer sequences, taking an increasing amount of memory. Notably, for 2,048-length sequences, Rfa decodes around 12 faster than the baseline while using less than 10% of the memory. Rfa- slightly outperforms Rfa-Gaussian in terms of speed and memory efficiency. This is because when using the same (as we do here), the is half the size of . These results suggest that Rfa can be particularly useful in sequence-to-sequence tasks with longer sequences, e.g., document-level machine translation (Miculicich et al., 2018).
Figure 3 in Appendix C.1 compares the speed and memory consumption in unconditional decoding (e.g., sampling from a language model). The overall trends are similar to those in Figure 2.
Notes on decoding speed. With a lower memory overhead, Rfa can use a larger batch size than the baseline. As noted by Katharopoulos et al. (2020) and Kasai et al. (2021), if we had used mini-batches as large as the hardware allows, Rfa could have achieved a more significant speed gain. Nonetheless, we control for batch size even though it is not the most favorable setting for Rfa, since the conclusion translates better to common applications where one generates a single sequence at a time (e.g., instantaneous machine translation). For the softmax attention baseline, we follow Ott et al. (2018) and cache previously computed query/key/value representations, which significantly improves its decoding speed (over not caching).
Further analysis results. Rfa achieves comparable performance to softmax attention. Appendix C.3 empirically shows that this cannot be attributed to Rfa learning a good approximation to softmax: when we train with one attention but evaluate with the other, the performance is hardly better than randomly-initialized untrained models. Yet, an Rfa model initialized from a pretrained softmax transformer achieves decent training loss after a moderate amount of finetuning steps (Appendix C.4). This suggests some potential applications, e.g., transferring knowledge from a pretrained transformer (e.g., GPT-3; Brown et al., 2020) to an Rfa model that is more efficient to sample from.
Related Work
One common motivation across the following studies, that is shared by this work and the research we have already discussed, is to scale transformers to long sequences. Note that there are plenty orthogonal choices for improving efficiency such as weight sharing (Dehghani et al., 2019), quantization (Shen et al., 2020), knowledge distillation (Sanh et al., 2020), and adapters (Houlsby et al., 2019). For a detailed overview we refer the reader to Tay et al. (2020c).
Sparse attention patterns. The idea behind these methods is to limit the reception field of attention computation. It motivates earlier attempts in improving attention’s efficiency, and still receives lots of interest. The sparse patterns can be set a priori (Liu et al., 2018; Qiu et al., 2020; Ho et al., 2020; You et al., 2020, inter alia) or learned from data (Sukhbaatar et al., 2019; Roy et al., 2020, inter alia). For most of these approaches, it is yet to be empirically verified that they are suitable for large-scale sequence-to-sequence learning; few of them have recorded decoding speed benefits.
Compressed context. Wang et al. (2020) compress the context along the timesteps so that the effective sequence length for attention computation is reduced. Another line of work aims to store past context into a memory module with limited size (Lee et al., 2019; Ainslie et al., 2020; Rae et al., 2020, inter alia), so that accessing longer history only moderately increases the overhead. Reminiscent of RNN language models, Rfa attends beyond a fixed context window through a stateful computation, without increasing time or memory overhead.
Conclusion
We presented random feature attention (Rfa). It views the softmax attention through the lens of kernel methods, and approximates it with random feature methods. With an optional gating mechanism, Rfa provides a straightforward way of learning with recency bias. Rfa’s time and space complexity is linear in the sequence length. We use Rfa as a drop-in substitute for softmax attention in transformer models. On language modeling, machine translation, and long text classification benchmarks, Rfa achieves comparable or better performance than strong baselines. In the machine translation experiment, Rfa decodes twice as fast. Further time and memory efficiency improvements can be achieved for longer sequences.
Acknowledgments
We would like to thank Phil Blunsom, Chris Dyer, Nando de Freitas, Jungo Kasai, Adhiguna Kuncoro, Dianqi Li, Ofir Press, Lianhui Qin, Swabha Swayamdipta, Sam Thomson, the language team at DeepMind and the ARK group at the University of Washington for their helpful feedback. We also thank Tay Yi for helping run the Long Range Arena experiments, Richard Tanburn for the advice on implementations, and the anonymous reviewers for their thoughtful comments. This work was supported in part by NSF grant 1562364 and a Google Fellowship. Nikolaos Pappas was supported by the Swiss National Science Foundation under grant number P400P2_183911 “UNISON.”
References
Appendix A Random Feature Attention in More Detail
Algorithms 1 and 2 describe causal and cross random feature attention’s computation procedures.
A.2 Variance of Random Fourier Features
The following result is due to Yu et al. (2016). Using the same notation as in §2.2:
where .
A.3 Derivation of Causal Rfa
This section presents a detailed derivation of causal Rfa as in §3.1. Following Eq. 5 but changing the attended keys and values to the prefix:
Let , and ; both can be calculated recurrently. Assuming and :
This completes the derivation of causal Rfa as in §3.1.
A.4 Rfa without Norm-1 Constraints
§3.1 assumes that the queries and keys are unit vectors. This norm-1 constraint is not a must. Here we present a Rfa without imposing this constraint. Let . From Eq. 4 we have
The specific attention computation is similar to those in §3.1. In sum, lifting the norm-1 constraint brings an additional scalar term .
A.5 Relating Rfa-Gate to Softmax Attention
Drawing inspiration from gated RNNs, §3.2 introduces a gated variant of Rfa. Now we study its “softmax counterpart.”
is the output at timestep and is used for onward computation.
At each step, all prefix keys and values are decayed by a gate value before calculating the attention. This implies that the attention computation for cannot start until that of is finished. Combined with the linear complexity of softmax normalization, this amounts to quadratic time in sequence length, even for language modeling training.
The above model is less intuitive and more expensive in practice, without the Rfa perspective. This shows that Rfa brings some benefits in developing new attention models.
A.6 Detailed Complexity Analysis
Table 4 considers a sequence-to-sequence model, and breaks down the comparisons to training (with teacher forcing; Williams & Zipser, 1989) and autoregressive decoding. Here we assume enough threads to fully parallelize softmax attention across timesteps when the inputs are revealed to the model in full. Rfa has a lower space complexity, since it never explicitly populates the attention matrices. As for time, Rfa trains in linear time, and so does the softmax attention: in teacher-forcing training a standard transformer decoder parallelizes the attention computation across time steps. The trend of the time comparison differs during decoding: when only one output token is produced at a time, Rfa decodes linearly in the output length, while softmax attention decodes quadratically.
Appendix B Experimental Details
Table 5 summarizes some statistics of the datasets used in our experiments. Our implementation is based on JAX. https://github.com/google/jax.
During training, we sample a different random projection matrix for each attention head. Preliminary experiments suggest this performs better than using the same random projection throughout training (Table 6). Our conjecture is that this helps keep the attention heads from “over committing” to any particular random projection (Peng et al., 2020). To avoid the overhead of sampling from Gaussian during training, we do this in an offline manner. I.e., before training we construct a pool of random matrices (typically 200), at each training step we draw from the pool. At test time each attention head uses the same random projection, since no accuracy benefit is observed by using different ones for different test instances.
B.2 Machine Translation
Appendix C More Analysis Results
Figure 3 compares the Rfa’s unconditional decoding speed and memory against the softmax attention. The setting is the same as that in §5 except that here the models do not have an encoder. This experiment aims to simulate the applications such as sampling from a language model.
C.2 Effect of Random Feature Size
This section studies how the size of affects the performance. Table 9 summarize Rfa-Gaussian’s performance on WMT14 EN-DE development set. The model and training are the same as that used in §4.2 except random feature size. Recall from §2.2 that the size of is for Rfa-Gaussian. When the size of is too small (32 or 64 for cross attention, 32 for causal attention), training does not converge. We observe accuracy improvements by using random features sufficiently large (256 for cross attention and 128 for causal attention); going beyond that, the benefit is marginal.
C.3 Train and Evaluate with Different Attention Functions
Rfa achieves comparable performance to its softmax counterpart. Does this imply that it learns a good approximation to the softmax attention? To answer this question, we consider:
an Rfa-Gaussian model initialized from a pretrained softmax-transformer;
a softmax-transformer initialized from a pretrained an Rfa-Gaussian model.
If Rfa’s good performance can be attributed to learning a good approximation to softmax, both, without finetunining, should perform similarly to the pretrained models. However, this is not the case on IWSLT14 DE-EN. Both pretrained models achieve more than 35.2 development set BLEU. In contrast, (i) and (ii) respectively get 2.3 and 1.1 BLEU without finetuning, hardly beating a randomly-initialized untrained model. This result aligns with the observation by Choromanski et al. (2020), and suggests that it is not the case that Rfa performs well because it learns to imitate softmax attention’s outputs.
C.4 Knowledge Transfer from Softmax Attention to RFA
We first supplement the observation in Appendix C.3 by finetuning (i) on the same pretraining data. Figure 4 plots the learning curves. It takes Rfa roughly 1,500 steps to reach similar training loss to the pretrained model. As a baseline, “Rfa Reset” resets the multihead attention parameters (i.e., those for query, key, value, and output projections) to randomly initialized ones. Its learning curve is similar to that of (i), suggesting that the pretrained multihead attention parameters are no more useful to Rfa than randomly initialized ones. To further confirm this observation, “softmax Reset” resets the multihead attention parameters without changing the attention functions. It converges to the pretraining loss in less than 200 steps.
Takeaway. From the above results on IWSLT14, pretrained knowledge in a softmax transformer cannot be directly transferred to an Rfa model. However, from Figure 4 and a much larger-scale experiment by Choromanski et al. (2020), we do observe that Rfa can recover the pretraining loss, and the computation cost of finetuning is much less than training a model from scratch. This suggests some potential applications. For example, one might be able to initialize an Rfa language model from a softmax transformer pretrained on large-scale data (e.g., GPT-3; Brown et al., 2020), and finetune it at a low cost. The outcome would be an Rfa model retaining most of the pretraining knowledge, but is much faster and more memory-friendly to sample from. We leave such exploration to future work.