Simplified Reversible Residual Networks
Erland B Olsson ⋅ Zhirong Yang
Abstract
The de facto standard way of training neural networks requires storing activations at each layer to memory, in order to compute gradients in backpropagation. As progress has often been made by stacking larger and more layers, memory consumption has now become a major bottleneck for making frontier models available to everyone. In this work, we take advantage of the Residual Network (ResNet) layer formulation that is now used in many state-of-the-art neural network architectures to recompute the activations at each layer on-the-fly during backpropagation instead of storing them in memory. We propose a modification to Reversible Residual Network (RevNet) that overcomes its limitation of requiring to replace the ResNet layer network function with two new network functions, by setting either one to an identity function. This modification is simple yet has a significant advantage because now there is no need to change the original ResNet architecture. Our method can work as a drop-in replacement for layers with residual connections such as in ordinary ResNets or Transformers with practically no loss in performance while offering significant reductions in memory consumption. With this simple modification, we are able to maintain the same computational cost and ease of use as activation checkpointing, which requires no architectural changes, while leveraging a reversible procedure such that we only need to store the activations from the last layer. The method is called Simplified RevNet, and compared to previous work in reversible architectures, we here propose a simpler and more streamlined approach that comes in two variants based on which of $F$ or $G$ in RevNet becomes an identity function. Empirically, we demonstrate performance at practically the same level as the non-reversible counterparts on ImageNet image classification with ResNets and on OpenWebText language modeling with Transformers.
Chat is not available.
Successful Page Load