Maximum-Margin Classification with Jointly Trained Softmax Attention
Abstract
We study the implicit bias of joint training in a single-layer, single-head softmax attention model for binary classification, with trainable query, key, value, output, and classifier factors. For exponential and logistic losses, we show that, once the training loss falls below a specified threshold, gradient flow and sufficiently small fixed-step gradient descent drive the loss to zero and monotonically increase a smoothed normalized margin. We further prove that the effective classifier converges in direction. More precisely, after normalization to unit margin, it converges to the unique Euclidean minimum-norm separator of all limiting signed features generated by the training trajectory. Under additional geometric conditions on the data and initialization, we prove for both losses that attention selects prescribed tokens, identifying the limiting classifier as the hard-margin SVM on those tokens. These results establish a maximum-margin characterization for joint attention and classifier training even though the model is not globally homogeneous. Experiments on MNIST illustrate normalized-margin growth and attention concentration, while a controlled example demonstrates token selection and classifier alignment with the predicted SVM.
Then back it, or bet against it.
Related papers
Open the market on this paper to see 7 more related papers.