Visual explainer of empirical Neural Tangent Kernels
There is a new mech interp method in town. I found it difficult to get my head around at first, but once I understood it, I found it beautiful. I am not an expert in this topic, and mostly write this up for me to get some intuitions, so there might be mistakes here.
This explainer came out of summarizing the paper Feature Identification via the Empirical NTK by Jennifer Lin. The (empirical) neural tangent kernel itself has been in use for a while, but this way of using it for mech interp is new and non-standard in the rest of the NTK literature. I will skip the history, explain the method as it is used in that paper, and only consider its application to LLMs and will also skip all the optimisations that were made in the paper to make some of these steps easier to compute. Some of the intuitions are my own.
What does the eNTK try to do?
Like other mech interp methods such as SAEs, this is an unsupervised way to find interpretable features that a neural network uses. Unlike those methods, though, we do not take activations and cluster them in some clever way. Instead, we look at what kind of updates a model would make if it were trained on some data: which behaviours move together, and how an update on one data point affects the model's behaviour on some other data point.
What we get out at the end is a weighted list of input–output pairs that make the model update in a particular direction. We can translate this into a vector in weight space, and potentially in activation space.
How do we compute the eNTK?
So how do we get these features? We take our model (an LLM), a text corpus of inputs , and a part of the network we want to analyze. Let's say a particular weight matrix . The paper uses the MLP-out matrix of one layer at a time. In theory, you could also take several parts of the network at once, or even all of its weights. One input is a whole context window, and the outputs are the next-token logits at the end of that window.
Now we feed an input into the network, pick one token from the vocabulary, and calculate the gradient of the logit of token with respect to the weights . This gradient has the same shape as , and its meaning is pretty straightforward: entry says in which direction we would have to nudge weight so that, for this input, the logit of token goes up. Raising a logit roughly makes the token more likely, but strictly speaking we are taking gradients of logits, not probabilities. We do not care about the shape of here, so we flatten the gradient into one long vector.
We do this for every token in the vocabulary, so each input gives us gradient vectors of length .
Then we do the same for every input in our corpus and stack all these vectors into one matrix of shape . Each row of is one input–output behaviour: it pairs one input with one output token, and gives the direction in which we would have to update our parameters to raise this output logit given this input.
The next thing we are interested in is how these updates support or hinder each other. We take the inner product , which gives us a matrix of shape . This is the (flattened) empirical NTK. Each entry is the dot product between the update that promotes input–output pair and the update that promotes pair . If it is positive, the two behaviours cohere: training on one of them also increases performance on the other. Similarly, if the dot product is negative, updating on one pushes the network against the other.
Let us consider what it means to multiply a vector with this matrix. We take a vector with one value for each input–output pair. We can imagine it as a push on each pair: a weighted training distribution built from our corpus, where negative entries mean pushing that logit down (like rejected completions in DPO). What we get out, , is, how much every logit changes if we take a gradient step with this training data, again with one value per pair.
Let's make up an example to get a feeling for this. Say we train on a set of happy text snippets that all reinforce each other, plus one odd one out: a sad text snippet. The resulting change pushes the happy pairs up, but the odd one out actually ends up pushed down, because the model generalizes to output happy text in general since the majority of its training data was happy. has a different shape than .
An eigenvector is a pattern of pushes where this does not happen: the resulting change is proportional to the push itself. Every pair comes back scaled by the same factor . so training on this pattern purely reinforces itself. Eigenvector entries can be negative, so a pattern can push some pairs up and others down at the same time. And since , all eigenvalues are , so training on a pattern never moves the model against that pattern.
The top eigenvalues belong to the patterns that the model amplifies the most when trained on them. Intuitively these eigenvectors correspond to the features the model uses for generalization: along which lines the model would update if trained on which data.
So far, these features live in the space of inputs and outputs, but we can easily find them in weight space too, by calculating the gradient step the model would take if trained on this pattern, which is . We might also find directions in activation space by taking an additive bias term as the subject of the investigation, since the gradient with respect to an added bias equals the gradient with respect to that activation.
Does this work?
The paper tests this on three models. In a 1-layer MLP and a 1-layer Transformer, both trained on modular addition, where the ground truth features the model uses for computation is known. Here the subspaces spanned by the ground-truth features could be recovered. In Gemma-3-270M on TinyStories, eNTK eigenvectors pick out grammatical features of the next word the model predicts, like nouns, verbs, adverbs, past-tense verbs or infinitives, better than PCA on activations with the same budget. An earlier version of the work, the LessWrong post Finding Features in Neural Networks with the Empirical NTK, also shows that the method recovers the ground-truth features in the toy models of superposition.
- In equations: write for the vector of all logits, one entry per input–output pair, so . Training on means taking a gradient step that increases , i.e. on the loss . The step is , and to first order the logits change by .
- This is the hypothesis the paper investigates. It presents some evidence for it, but this is not settled science.
- The method finds the correct subspace in which these features live. The eigenvalues within that subspace are nearly equal, so the eigenvectors can be rotated within it, and the split into individual feature directions isn't pinned down.