A Practical Sparse Approximation for Real Time Recurrent Learning
Jacob Menick, Erich Elsen, Utku Evci, Simon Osindero, Karen Simonyan, Alex Graves
Introduction
Recurrent neural networks (RNNs) have been successfully applied to a wide range of sequence learning tasks, including text-to-speech , language modeling , automatic speech recognition , translation and reinforcement learning . RNNs have greatly benefited from advances in computational hardware, dataset sizes, and model architectures. However, the algorithm used to compute their gradients remains unchanged: Back-Propagation Through Time (BPTT). The key limitation of BPTT is that the entire state history must be stored, meaning that the memory cost grows linearly with the sequence length. For sequences too long to fit in memory, as often occurs in domains such as language modelling or long reinforcement learning episodes, truncated BPTT (TBPTT) can be used. Unfortunately the truncation length used by TBPTT also limits the duration over which temporal structure can be realiably learned.
Forward-mode differentiation, or Real-Time Recurrent Learning (RTRL) as it is called when applied to RNNs , solves some of these problems. It doesn’t require storage of any past network states, can theoretically learn dependencies of any length and can be used to update parameters at any desired frequency, including every step (i.e. fully online). However, its fixed storage requirements are , where is the state size and is the number of parameters in the core. Perhaps even more daunting, the computation it requires is . This makes it impractical for even modestly sized networks. The advantages of RTRL have led to a search for more efficient approximations that retain its desirable properties, whilst reducing its computational and memory costs. One recent line of work introduces unbiased, but noisy approximations to the influence update. Unbiased Online Recurrent Optimization (UORO) is an approximation with the same cost as TBPTT – – however its gradient estimate is severely noisy and its performance has in practice proved worse than TBPTT . Less noisy approximations with better accuracy on a variety of problems include both Kronecker Factored RTRL (KF-RTRL) and Optimal Kronecker-Sum Approximation (OK) . However, both increase the computational costs to .
The last few years have also seen a resurgence of interest in sparse neural networks – both their properties and new methods for training them . A number of works have noted their theoretical and practical efficiency gains over dense networks . Of particular interest is the finding that scaling the state size of an RNN while keeping the number of parameters constant leads to increased performance .
In this work we introduce a new sparse approximation to the RTRL influence matrix. The approximation is biased but not stochastic. Rather than tracking the full influence matrix, we propose to track only the influence of a parameter on neurons that are affected by it within steps of the RNN. The algorithm is strictly less biased but more expensive as increases. The cost of the algorithm is controlled by and the amount of sparsity in the Jacobian of the recurrent cell. Larger can be coupled with concomitantly higher sparsity to keep the cost fixed. The approximation approaches full RTRL as increases. Our contributions are as follows:
We propose SnAp – a practical approximation to RTRL, which is is applicable to both dense and sparse RNNs, and is based on the sparsification of the influence matrix.
We show that parameter sparsity in RNNs reduces the costs of RTRL in general and SnAp in particular.
We carry out experiments on both real-world and synthetic tasks, and demonstrate that the SnAp approximation: (1) works well for language modeling compared to the exact unapproximated gradient; (2) admits learning long-term dependencies on a synthetic copy task and (3) can learn faster than BPTT when run fully online.
Background
The recursive expansion is the backpropagation rule. The slightly nonstandard notation refers to the copy of the parameters used at time , but the weights are shared for all timesteps and the gradient adds over all copies.
Real Time Recurrent Learning computes the gradient as:
This can be viewed as an iterative algorithm, updating from the intermediate quantity . To simplify equation 2 we introduce the following notation: have , and . stands for “Jacobian”, for “immediate Jacobian”, and for “dynamics”. We sometimes refer to as the “influence matrix”. The recursion can be rewritten .
2 Truncated RTRL and stale Jacobians
In analogy to Truncated BPTT, one can consider performing a gradient update partway through a training sequence (at time ) but still passing forward a stale state and a stale influence Jacobian rather than resetting both to zero after the update. This enables more frequent weight updating at the cost of a staleness bias. The Jacobian becomes “stale” because it tracks the sensitivity of the state to old parameters. Experiments (section 5.2) show that this tradeoff can be favourable toward more frequent updates in terms of data efficiency. In fact, much of the RTRL literature assumes that the parameters are updated at every step (“fully online”) and that the influence Jacobian is never reset, at least until the start of a new sequence.
3 Sparsity in RNNs
One of the early explorations of sparsity in the parameters of RNNs (i.e. many entries of are exactly zero) was Ström , where one-shot pruning based on weight magnitude with subsequent retraining was employed in a speech recognition task. The current standard approach to inducing sparsity in RNNs remains similar, except that magnitude based pruning happens slowly over the course of training so that no retraining is required.
Kalchbrenner et al. discovered a powerful property of sparse RNNs in the course of investigating them for text-to-speech – for a constant parameter and flop budget sparser RNNs have more capacity per parameter than dense ones. This property has so far only been shown to hold when the sparsity pattern is adapted during training (in this case, with pruning). Note that parameter parity is achieved by simultaneously increasing the RNN state size and the degree of sparsity. This suggests that training large sparse RNNs could yield powerful sequence models, but the memory required to store the history of (now much larger) states required for BPTT becomes prohibitive for long sequences.
The Sparse n-Step Approximation (SnAp)
Our main contribution in this work is the development of an approximation to RTRL called the Sparse n-Step Approximation (SnAp) which reduces RTRL’s computational requirements substantially.
SnAp imposes sparsity on even though it is in general dense. We choose the sparsity pattern to be the locations that are non-zero after steps of the RNN (Figure 1). We also choose to use the same pattern for all steps, though this is not a requirement. This means that the sparsity pattern of is known and can be used to reduce the amount of computation in the product . See Figure 3 for a visualization of the process. The costs of the resulting methods are compared in Table 3. We note an alternative strategy would be to perform the full multiplication of and then only keep the top-k values. This would reduce the bias of the approximation but increase its cost.
More formally, we adopt the following approximation for all :
Even for a fully dense RNN, each parameter will in the usual case only immediately influence the single hidden unit it is directly connected to. This means that the immediate Jacobian tends to be extremely sparse. For example, a Vanilla RNN will have only one nonzero element per column, which is a sparsity level of . Storing only the nonzero elements of already saves a significant amount of memory without making any approximations; is the same shape as the daunting matrix whereas the nonzero values are the same size as .
can become more dense in architectures (such as GRU and LSTM) which involve the composition of parameterised layers within a single core step (see section 3 for an in-depth discussion of the effect of gating architectures on Jacobian sparsity). In the Sparse One-Step Approximation, we only keep entries in if they are nonzero in . After just two RNN steps, a given parameter has influenced every unit of the state through its intermediate influence on other units. Thus only SnAp with is efficient for dense RNNs because does not result in any sparsity in , i.e.: for dense networks SnAp-2 already reduces to full RTRL. (N.b.: SnAp-1 is also applicable to sparse networks.) Figure 1 depicts the sparse structure of the influence of a parameter for both sparse and fully dense cases.
SnAp-1 is effectively diagonal, in the sense that the effect of parameter on hidden unit is maintained throughout time, but ignoring the indirect effect of parameter on unit via paths through other units . More formally, it is useful to define as the one component in the state connected directly to the parameter (which has at the other end of the connection some other entry within or ). Let . The imposition of the one-step sparsity pattern means only the entry in row will be kept for column in . Inspecting the update for this particular entry, we have
The equality follows from the assumption that if . Diagonal entries in are thus crucial for this approximation to be expressive, such as those arising from skip connections.
2 Optimizations for full RTRL with sparse networks
When the RNN is sparse, the costs of even full (unapproximated) RTRL can be alleviated to a surprising extent; we save computation proportional to a factor of the sparsity squared. Assume a proportion of the entries in both and are equal to zero and refer to this number as “the level of sparsity in the RNN”. For convenience, . With a Vanilla RNN, this correspondence between parameter sparsity and dynamics sparsity holds exactly. For popular gating architectures such as GRU and LSTM the relationship is more complicated so we include empirical measurements of the computational cost in FLOPS (Table 1) in addition to the theoritical calculations here. More complex recurrent architectures involving attention would require an independent mechanism for inducing sparsity in ; we leave this direction to future work and assume in the remainder of this derivation that sparsity in corresponds to sparsity in .
3 Sparse N𝑁N Step Approximation (SnAp-N𝑁N)
Jacobian Sparsity of GRUs and LSTMs
Unlike vanilla RNNs whose dynamics Jacobian has sparsity exactly equal to the sparsity of the weight matrix, GRUs and LSTMs have inter-cell interactions which increase the Jacobians’ density. In particular, the choice of GRU variant can have a very large impact on the increase in density. This is relevant to the “dynamics” jacobian and the parameter jacobians and .
Looking at LSTM’s update equations, we can see that an individual parameter will only directly affect one entry in each gate (, , ) and the candidate cell . These in turn produce the next cell and next hidden state with element-wise operations ( is the sigmoid function applied element-wise and is usually hyperbolic tangent). In this case Figure 1 is an accurate depiction of the propagation of influence of a parameter as the RNN is stepped.
However, for a GRU there are multiple variants in which a parameter or hidden unit can influence many more units of the next state. The original variant is as follows:
For our purposes the main thing to note is that the parameters influencing further influence every unit of because of the matrix multiplication by . They therefore influence every unit of within one recurrent step, which means that the dynamics jacobian is fully dense and the immediate parameter jacobian for , , and are all fully dense as well.
An alternative formulation which was popularized by Engel , and also used in the CuDNN library from NVIDIA is given by:
The second variant has moved the reset gate after the matrix multiplication, thus avoiding the composition of parameterized linear maps within a single RNN step. As the modeling performance of the two variants has been shown to be largely the same, but the second variant is faster and results in sparser and , we adopt the second variant throughout this paper.
Related Work
It is possible to reduce the storage requirements of TBPTT using a technique known as “gradient checkpointing” or “rematerialization”. This reduces the memory requirements of backpropagation by recomputing states rather than storing them. First introduced in Griewank and Walther and later applied specifically to RNNs in Gruslys et al. , these methods are not compatible with the fully online setting where may be arbitrarily large as even the optimally small amount of re-computation can be prohibitive. For reasonably sized , however, rematerialization is a straightforward and effective way to reduce the memory requirements of TBPTT, especially if the forward pass can be computed quickly.
Experiments
We include experimental results on the real world language-modelling task WikiText103 and the synthetic ‘Copy’ task of simply repeating an observed binary string. Whilst the first is important for demonstrating that our methods can be used for real, practical problems, language modelling doesn’t directly measure a model’s ability to learn structure that spans long time horizons. The Copy task, however, allows us to parameterize exactly the temporal distance over which structure is present in the data. In terms of, respectively, task complexity and RNN state size (up to 1024) these investigations are considerably more “large-scale” than much of the RTRL literature.
All of our WikiText103 experiments tokenize at the character (byte) level and use SGD to optimize the log-likelihood of the data. We use the Adam optimizer with , , and . We train on randomly cropped sequences of length 128 sampled uniformly with replacement and do not propagate state across the end-of-sequence boundary (i.e. no truncation). Results are reported on the standard validation set.
In this section, we refrain from performing a weight update until the end of a training sequence (see section 2.2) so that BPTT is the gold standard benchmark for performance, assuming the gradient is the optimal descent direction. The architecture is a Gated Recurrent Unit (GRU) network with 128 recurrent units and a one-layer readout MLP mapping to 1024 hidden units before the final 256-unit softmax layer. Learning curves in Figure 4 (Left) show that SnAp-1 outperforms RFLO and UORO, and that in this setting UORO fails to match the surprisingly strong baseline of not training the recurrent parameters at all and instead leaving them at their randomly initialized value.
1.2 Language Modeling with Sparse RNNs: SnAp-1 and SnAp-2
Here we use the same architecture as in section 5.1.1, except that we introduce 75% sparsity into the weights of the GRU, in particular the weight matrices (more sparsity levels are considered in later experiments). Biases are always kept fully dense. In order to induce sparsity, we generate a sparsity pattern uniformly at random and fix it throughout training. As would be expected because it is strictly less biased, Figure 4 (Right) shows that SnAp-2 outperforms SnAp-1 but only slightly. Furthermore, both closely match the (gold-standard) accuracy of a model trained with BPTT. Table 1 shows that SnAp-2 actually costs about 600x more FLOPs than BPTT/SnAp-1 at 75% sparsity, but higher sparsity substantially reduces FLOPs. It’s unclear exactly how the cost compares to UORO, which though does have constant factors required for e.g. random number generation, and additional overheads when approximations use rank higher than one.
Our experiments do not use state-of-the-art strategies for inducing sparsity because there is no such strategy compatible with SnAp at the time of writing. The requirement of a dense gradient in Evci et al. and Zhu and Gupta prevents the use of the optimization in Equation 4, which is strictly necessary to fit the RTRL training computations on accelerators without running out of memory.
To further motivate the development of sparse training strategies that do not require dense gradients, we show that larger sparser networks trained with BPTT and magnitude pruning monotonically outperform their denser counterparts in language modelling, when holding the number of parameters constant. This provides more evidence for the scaling law observed in Kalchbrenner et al. .
The experimental setup is identical to the previous section except that all networks are trained with full BPTT. To hold the number of parameters constant, we start with a fully dense 128-unit GRU. We make the weight matrices 75% sparse when the network has 256 units, 93.8% sparse when the network has 512 units, 98.4% when the network has 1024 units, and so on. The sparsest network considered has 4096 units and over 99.9% sparsity, and performed the best. Indeed it performed better than a dense network with 6.25x as many parameters (Figure 6). Pruning decisions are made on the basis of absolute value every 1000 steps, and the final sparsity is reached after 350,000 training steps.
2 Copy Task
Our experiments on the Copy task aim to investigate the ability of the proposed sparse RTRL approximations to learn about long-term temporal structure. We follow and adopt a curriculum-learning approach over the length of sequences to be copied, starting with . When the average bits per character of a training minibatch drops below 0.15, we increment by one. We sample the length of target sequences uniformly between as in previous work. After some number of training steps, we compare algorithms on the basis of the level of they have reached because a model which has reached a greater has more rapidly learned how to capture structure over a time horizon of stepsacutally because of start/end flags and because the observation/target are presented in sequence. We measure performance versus ‘data-time’, i.e. we give each algorithm a time budget in units of the cumulative number of tokens seen throughout training. A consequence of this scheme is that full BPTT is no longer an upper bound on performance because, for example, updating once on a sequence of length 10 with the true gradient may yield slower learning than updating twice on two consecutive sequences of length 5, with truncation.
In these experiments we examine SnAp performance for multiple sparsity levels and recurrent architectures including Vanilla RNNs, GRU, and LSTM. Table 1 includes the architectural details. The sparsity pattern is again chosen uniformly at random. As a result, comparison between sparsity levels is discouraged. For each configuration we sweep over learning rates in and compare average performance over three seeds with the best chosen learning rate (all methods performed best with learning rate ). The minibatch size was 16. We train with either full unrolls or truncation with . This means that the RTRL approximations update the network weights at every timestep and persist a stale Jacobian (see section 2.2).
One striking observation is that Truncated BPTT completely fails to learn long-term structure in the fully online () regime. Interestingly, the SnAp methods perform better with more frequent updates. Compare solid versus dotted lines of the same color in Figure 7. Fully online SnAp-2 and SnAp-3 mostly outperform or match BPTT for training LSTM and GRU architectures despite the “staleness” noted in Section 2.2. We attribute this to the hypothesis advanced in the RTRL literature that Jacobian staleness can be mitigated with small learning rates but leave a more thorough investigation of this phenomenon to future work.
For SnAp there is a tradeoff between the biasedness of the approximation and the computational costs of the algorithm. We see that correspondingly, SnAp-1 is outperformed by SnAp-2, which is in turn outperformed by SnAp-3 in the Copy experiments. The RFLO baseline is even more biased than SnAp-1, but both methods have comparable costs. SnAp-1 significantly outperforms RFLO in all of our experiments.
Here we augment the asymptotic cost calculations from Table 3 with empirical measurements of the FLOPs, broken out by architecture and sparsity level in Table 1. Gating architectures require a high degree of parameter sparsity in order to keep a commensurate amount of of Jacobian sparsity due to the increase in density brought about by composing linear maps with different sparsity patterns (see section 3).
For instance, the 75% sparse GRU considered in the experiments from Section 5.1.2 lead to SnAp-2 parameter Jacobian that is only 70.88% sparse. With SnAp-3 it becomes much less sparse – only 50%. This may partly explain why SnAp performs best compared to BPTT in the LSTM case (Figure 7), though it still significantly outperforms BPTT in the high sparsity regime when SnAp-2 becomes practical. Also, LSTM is twice as costly to train with RTRL-like algorithms because it has two components to its state, requiring the maintenance of twice as many jacobians and the performance of twice as many jacobian multiplications (Equations 3/5). For a 75% sparse LSTM, the SnAp-2 Jacobian is much denser at 38.5% sparsity and SnAp-3 has essentially reached full density (so it is as costly as RTRL).
Figure 7 also shows that for Vanilla RNNs, increasing improves performance, but SnAp does not outperform BPTT with this architecture. In summary, Increasing improves performance but costs more FLOPs.
3 Analysis of the bias introduced by SnAp
Finally, we examine the empirical magnitudes of entries which are nonzero in the true, unapproximated influence matrix but set to zero by SnAp. For the benefit of visualization we train a small network on a non-curriculum variant of the Copy-task with target sequences fixed in length to 16 timesteps. This enables us to measure and display the bias of SnAp.
This particular run is a 8-unit GRU with 75% sparsity. The influence matrix considered is the final value after processing an entire sequence, which has 35 elements including the observation, target, and start/end flags. The network is optimized with full (untruncated) BPTT. We find (Table 9) that at the beginning of training the influence entries ignored by SnAp are small in magnitude compared to those kept, even after the influence has had many RNN iterations to fill in.
This analysis complements the experimental results concerning how useful the approximate gradients are for learning; instead it shows where — and by how much — the sparse approximation to the influence differs from the true accumulated influence. Interestingly, despite the strong task performance of SnAp, the magnitude of ignored entries in the influence matrix is not always small (see Figure 9). The accuracy, as measured by such magnitudes, trends downward over the course of training. We speculate that designing methods to temper the increased bias arising later in training may be beneficial but leave this to future work.
Conclusion
We have shown how sparse operations can make a form of RTRL efficient, especially when replacing dense parameter Jacobians with approximate sparse ones. We introduced SnAp-1, an efficient RTRL approximation which outperforms comparably-expensive alternatives on a popular language-modeling benchmark. We also developed higher orders of SnAp including SnAp-2 and SnAp-3, approximations tailor-made for sparse RNNs which can be efficient in the regime of high parameter sparsity, and showed that they can learn temporal structure considerably faster than even full BPTT.
Our results suggests that training very large, sparse RNNs could be a promising path toward more powerful sequence models trained on arbitrarily long sequences. This may prove useful for modelling whole documents such as articles or even books, or reinforcement learning agents which learn over an entire lifetime rather than the brief episodes which are common today. A few obstacles stand in the way of scaling up our methods further:
The need for a high-performing sparse training strategy that does not require dense gradient information.
Sparsity support in both software and hardware that enables better realization of the theoretical efficiency gains of sparse operations.
It may also be fruitful to further develop our methods for hybrid models combining recurrence and attention or even feedforward architectures with tied weights .
Acknowledgements
We wish to thank Max Jaderberg, Siddhant Jayakumar, Jack Rae, Greg Wayne, Daan Wierstra, Tim Harley, James Martens, Corentin Tallec, Danilo Rezende, Charles Blundell, Tim Lillicrap, Peter Humphreys, Vlad Firoiu, Matthew Johnson, Jeffrey De Fauw, and Maneesh Sahani for stimulating discussions. We also thank the Jax and Haiku teams for creating infrastructure that made this very fun to implement.