6 ms·
TBH I don't find this exposition enlightening. If you don't already know the pytorch API, the code will be pretty obscure. And if you do understand the API, it'
by cshimmin 3y ago
TBH I don't find this exposition enlightening. If you don't already know the pytorch API, the code will be pretty obscure. And if you do understand the API, it's still not very clarifying.
Instead, it's a lot easier IMO to understand self attention in just plain English. If you understand the words, and also understand the torch API (or tensorflow or numpy, they're all more or less fungible), then it should also be straightforward to implement it.
Here's how I explain it to my grad students:
You have a collection of N things, with some fixed number of features each (doesn't matter how many, but call it d_f). Call that collection {x_i}. You have two learnable matrices Q and K, which learn to project those items into some collection of "questions" and "answers", respectively. Doesn't really matter what those questions or answers are, the NN will figure out what they should be during training.
These projection matrices map x_i -> k_i, and x_i -> q_i. k and q are each d_k-dimensional vectors (d_k is a number you choose) of features, so Q and K are d_k-by-d_f matrices. To measure the "compatibility" between questions q and answers k, we take the dot product of them. Ones which are similar (large, positive dot products) means the NN likes the answer k for question q.
So, here's what you do. For each token i, you compare it's "question" against every other token j's answer. I.e. you compute dot(q_i, k_j), which itself can be considered an NxN matrix of scalar numbers, qk_{ij}. Each element of this matrix contains the answer to the question "how well did item j answer the question of item i?"
Applying softmax to this matrix just converts the dot product values into a score from 0-1 over all the items. I.e., you convert the question to "on a scale from 0-1, how well did item j answer the question of item i?". The only subtlety here is that all the scores over all items j have to sum to 1 for each question i.
Finally, we have a third projection matrix, V, which maps x_j -> v_j. The v_j are d_v-dimensional vectors (doesn't have to equal d_k, and again something you arbitrarily choose). These represent the values that the NN would like to pass on from item j to any other item which decides it "likes" the answer of item j.
So now, for each item i we have a score for "how well does item j answer my question i?". And we also have the list of values v_j that each item j would like to pass to item i. So, to compute the new information to add to item i, we take the weighted sum of the values v_j that item i liked, using these compatibility scores as a weight. In math, this reads:
y_i = sum_over_j [ softmax(qk_ij) v_j ]
* Technically it's also empirically found that normalizing the argument of softmax by a factor 1/sqrt(d_k) helps. But it's not really germane to understanding and also not strictly necessary since a NN is of course free to simply learn to include an overall factor proportional to 1/sqrt(d_k) in the parameters of one of the projection matrices. Actually the use of softmax at all is also arbitrary, but encourages the attention to be sparse, which again is empirically known to be helpful.
- ntonozzi 3y agoThanks for the helpful explanation. What is the context length in this explanation?
- two_in_one 3y agomy guess it's N here
- cshimmin 3y agoCorrect. In the context of LLM's, the "items" I refer to would be tokens. As a particle physicist, in the transformers I work with the "items" often instead represent either particles, or detector measurements.
- two_in_one 3y agoThanks, I'm trying to understand all this mechanics. Matrix multiplication is an equivalent to passing through a single fully connected convolution layer. Which means we can probably beef up and make it a small network. Also it's easy to make it work not in fixed windows, but in scrolling, if the output is processed sequentially. Even if not it may make sense too. It can be implemented efficiently. The idea is that in switching window first and last elements are being influenced only from one side. Which is not a good thing in text processing. Only middle elements are influenced from both sides. With scrolling window all elements currently being processed are in the middle. Except for the ends of the dataset.
- two_in_one 3y agoCorrect me if I'm wrong. Here Q, K, and V are trainable. d_k and d_v are selectable metaparameters(?) d_f - input vectors' size, N is their number, or window's size. Next time we usually process another non-overlapping set of input vectors.
- fireworkcrayon 3y ago> You have a collection of N things, with some fixed number of features each (doesn't matter how many, but call it d_f). Call that collection {x_i}. You have two learnable matrices Q and K, which learn to project those items into some collection of "questions" and "answers", respectively. Doesn't really matter what those questions or answers are, the NN will figure out what they should be during training. These projection matrices map x_i -> k_i, and x_i -> q_i. k and q are each d_k-dimensional vectors (d_k is a number you choose) of features, so Q and K are d_k-by-d_f matrices. Wait: this is the plain English version? I’ve just been presented an alphabet soup of symbols, most of which haven’t been defined: N, d, f, x, i, Q, K, and q and k (lowercase). The only one I’m sure I understand is N. The different letters seem to be drawn from certain regions of the alphabet, so I guess that’s significant in some way? (I do understand vector, matrix, softmax, dimensionality, etc.) I’m embarrassed to admit this here, but I run into this issue almost any time concepts are explained in math, including in the textbooks that teach math. What am I missing?