Skip to content
RWRui Wang / Ideas
← All notes

Learning - LLM

Transformers - Path to Deep Learning 3

Transformers - Path to Deep Learning 3 cover
Hello everyone, in this episode we will cover the Decoders that we didn't finish discussing last time.


In the previous episode, we mentioned that modern LLMs often discard Encoders and instead use Decoders. This time, we will explain how Decoders work from the perspective of LLMs.


First, let me show you the internal structure of Decoders. I believe that by this point, you can understand some of it. We have covered almost all the necessary theoretical knowledge, so I will just show the diagram directly.


transformer-arch

You may not fully understand it, but that's okay; this is the purpose of this episode.


First, I need to correct a mistake: in the episodes about Attention and Encoders, I incorrectly categorized positional encoding as a type of position encoding that can understand causal relationships. This is incorrect; our position encoding can only inform the model that this is the first xxx and this is the second xxx, and it does not carry the meaning of xxx. I apologize for the confusion; it was my miscommunication. My bad, my bad.


Alright, to understand how Decoders work, we first need to grasp some basic concepts:


1. First, the first one is Temperature, which literally translates to temperature. It is a hyperparameter. To make the description more vivid, I will show a diagram comparing the effects of high temperature and low temperature on the model's probabilities.

images (2)

As a viewer who has read this far in this channel, you should be able to understand it. After all, by this point, our understanding of LLMs is already higher than 99% of people.


Do you see that peak? It represents a very high probability for a certain Token, belonging to the Outlier level.


In contrast, the right side is much flatter, which is actually related to the model's variety. When T is a large value, the model tends to be relatively more diverse.


2. Now let's explain Unembedding, in this context, it meanshidden state mapping.


What does this mean? Let's take it literally; we have previously discussed Embedding, but Embd is actually quite different from Unemb, not a complete opposite relationship. But going back to the essence, the two things essentially do the opposite: one transforms a Token into a high-dimensional feature vector, while the other transforms a high-dimensional feature vector into Logits.


Let's talk about how Unemb transforms high-dimensional features into Logits.



After several layers (perhaps 10 layers) of Decoder Blocks (we will discuss this later, so let's not talk about it now), we obtain a comprehensive and good understanding of the input Token (high-dimensional features). Then, we introduce a weight library W_embd.


What is this? In fact, this is a weight library for our model's pre-training (we will discuss pre-training in the next episode). Our model inherently has a Vocabulary Bank with N words, and our W_embd is essentially a vector library with N rows, where each row is a pre-trained feature vector for a certain word.


To obtain the relationship of the current Token with other words, we perform a dot product between the current Token and all the values in this weight library, with the specific formula as follows:


\[ \underset{(1 \times V)}{\mathrm{Logits}} = \underset{(1 \times d_{\mathrm{model}})}{\mathbf{h}_{\mathrm{last}}} \cdot \underset{(d_{\mathrm{model}} \times V)}{\mathbf{W}_{\mathrm{unemb}}} \]


Among these several Logit values, we use a greedy algorithm.

Alright, we have now obtained the Logits values, which are the logical values.


Next, we convert them into probabilities P, introducing temperature T:


\(P_i = \frac{\exp\left(\frac{Z_i}{T}\right)}{\sum_{j=1}^{V} \exp\left(\frac{Z_j}{T}\right)}\)


Alright, this is roughly the concept of Unemb.


3. Causal Mask


We may have heard that one way to alleviate LLM hallucinations is to use a Causal Mask. This means that when calculating Attention, the machine is strictly forbidden from seeing the subsequent Tokens, for example, 'down', 'big', 'snow', 'already'; when the machine is looking at 'down', it must not see the next three tokens, which are 'big snow already'.


How does the causal mask achieve this? It's actually quite simple.


Still using the input token above as an example, our causal mask at this time is a 4*4 2D Array, like this


\(M = \begin{bmatrix} 0 & -\infty & -\infty & -\infty \\ 0 & 0 & -\infty & -\infty \\ 0 & 0 & 0 & -\infty \\ 0 & 0 & 0 & 0 \end{bmatrix}\)


Normal Attention looks like this:


\(\text{Attention}(Q, K, V) = \text{Softmax}\left( \frac{QK^T}{\sqrt{d_k}} \right) V\)


We add this causal mask M after applying Softmax


\(\text{Attention}(Q, K, V) = \text{Softmax}\left( \frac{QK^T}{\sqrt{d_k}} + M \right) V\)


This can be confusing; after adding the Mask, the upper right corner of the previous Attention becomes 0, meaning that the subsequent tokens are not considered in the Attention of this token at all, making it easy to complete the calculations for K, Q, and V.


4. KV - Cache.


What is this thing? Actually, you can guess a bit from its name.


Cache, as the name suggests, is a buffer; here we are actually sacrificing some GPU memory, which we commonly refer to as video memory, in exchange for significantly faster computation speed.


When our Unemb outputs, if there is no KV - cache, at the Nth token, it needs to compute K, Q, V for N-1 times, and at the N+1 token, it needs to compute K, Q, V for N times, which is a waste of time. As we mentioned earlier, the causal mask ensures that when processing the Ith token, the Attention of subsequent tokens is not considered, meaning we guarantee theconsistency of K, Q, and V.


Therefore, we can store the K, Q, V of the Nth token in video memory each time, allowing us to completely skip N-2 calculations of K, Q, V when it’s the turn of token sequence I.


In this way, while sacrificing some video memory usage (modern LLMs generally run on extremely high video memory, such as 256GB), we greatly save on computational power and time. This is also part of the reason why some improved LLMs with KV-Cache can achieve such high token output speeds, for example, Google Gemini 2 Flash has reached 2500+ token/s.


Having discussed these concepts, it’s time to fully present the workflow of Decoders.


Last time we talked about the part from Tokenization to Positional Encoding in Encoders, so I’ll skip that and provide a simple step diagram:


[用户文本 Prompt]
       │
       ▼ (Tokenizer + Embedding + RoPE)


Next, we move to the Decoder Exclusive phase:

There are mainly three stages:


1. Input Processing:


After embedding the tokens, we place them into several Decoder Blocks for preliminary Multi-head attention; since we discussed MTHA last time, we’ll skip it here.


Note that the Causal Mask is already in effect here; when processing Token T, it can only see tokens 1 to T.


Next, we put the results after MYHA into the FFN (Feed-Forward Network) for dimensional scaling and residual connection, and finally perform normalization.


2. Pre filling: Pre-warming


After completing the FFN, we activate the KV - Cache; at this point, every K, Q, V after the FFN is recorded and will not change due to subsequent operations.


This is where Unembedding comes into play, producing Token No.1.


3. Autoregressive self-loop token generation phase:


Through Unembedding, we obtained Token No.1, and now we enter the loop:


1. 仅输入上一步的 1 个 Token            
2. 计算新 k, v 追加至 [KV Cache]       
3. q_new 与历史全部 [KV Cache] 算 Attention     
4. FFN 抽取知识


Finally, when the Decoder calculates the special character, the stop symbol "<eos>", the loop ends. A complete sentence is produced.


In this way, a Decoder pipeline is completed, and we have obtained the results of the input tokens.


This episode ends here. In this article, I have planted a foreshadowing; in the next issue, we will explain Training - Training & Pre-training.


Thank you for reading this; your every visit is a support to me, and I sincerely appreciate it!


See you all next time~