Retrieval-Centric Deep Learning in Growing Nonparametric Neural Networks
Abstract
We investigate a general-purpose layer for deep learning that, instead of compressing arbitrary-size training data into fixed-size weight matrices, stores a new pair of key-value representations for every data point during training, and retrieves and recombines these representations through an attention mechanism at inference time---resulting in a growing neural net (NN). We derive such a "retrieval-centric deep learning" (RCDL) paradigm from the classic duality of a linear layer trained by gradient descent as a type of linear attention (LA) over the training data points---which motivates us to replace the corresponding LA by a more powerful form of attention from recent work on sequence models, namely, kernelized attention using radial basis function (RBF) and softmax-like kernels, as well as more advanced linear attention variants such as MesaNet and DeltaNet. We theoretically derive learning algorithms for the growing layers with kernelized attention, and empirically demonstrate their promising performance and sample-efficiency on classic image classification tasks and synthetic teacher-student learning datasets. In the MesaNet/DeltaNet-inspired extensions, we show a formal connection to a recently proposed optimizer and derive highly sample-efficient optimizers for conventional fixed-size NNs. Despite open challenges, RCDL represents a promising paradigm for nonparametric machine learning.