top of page
Search

Multi-Token Prediction from Scratch: Why it is needed for Edge devices?

Writer: Dhamo Dharan
Dhamo Dharan
Aug 16
6 min read
Autoregressive models like Llama, Gemma, and Qwen operate on edge, compute, and server. The main issue with autoregression decoding is that it is memory bound, not compute bound.
This impacts latency and throughput and incurs high costs for inference as edge devices have low-memory budget. The new technique of multiple token prediction, introduced in the paper “Better & Faster Large Language Models via Multi-token Prediction” at ICML 24, has advanced the ability to generate multiple tokens in a single decoding step.
  • This blog provides a foundational understanding of the Autoregressive decode loop and its costs. However, we will concentrate primarily on attention projection and linear projection layers rather than ROPE/Positional embedding, tokenizer, or FFN. We will discuss step by step how we can achieve token and KV effects, and how MTP can assist.


  • This will be a more extensive blog. I have provided each explanation in segments. But still, I suggest reading "LLM from Scratch" by Sebastin for a more detailed implementation of LLM architecture fully.


  • We will utilize notebook-based code to explain the concept and code. We will start with a small model, referred to as TinyLLM, which has no prebuilt weights or trained parameters, solely to build intuition.


  • Later, we will enhance TinyLLM with MTP heads to demonstrate the benefits. Our approach will use Python and Torch, focusing on illustrating functionality and advantages over traditional decoders, rather than high-performance optimization. Each hardware may require its own HPC optimizations.




Let's start with a few additional assumptions: the aim of this tutorial blog is to explore how MTP is advantageous compared to a single-step decoder. We will not delve into the reasoning behind the design of the LLM architecture; instead, we will accept the design as it stands and focus on examining the mathematics and computations it involves, as well as its impact on memory usage!


More importantly how to make fully use of this tutorial blog:


Cell 1: Setup & configuration:


We deliberately use a tiny model. It is intended to make the execution mechanics visible. Our model will essentially be Token > Embeddeing > Transformer > Hidden state > LM Head(Loop)-Next token. And Later Hidden state > [MTP 1, MTP 2, MTP 3, MTP 4].


We will break down all these mechanics in the upcoming cells. let's start importing packages and define the config class for TinyLLM


import torch
import torch.nn as nn
import torch.nn.functional as F
import time
import math
from dataclasses import dataclass

print("PyTorch:", torch.__version__)
print("Device:", "cuda" if torch.cuda.is_available() else "cpu")

@dataclass

class Config:
    vocab_size: int = 1000
    hidden_dim: int = 128
    num_heads: int = 4
    head_dim: int = 32
    num_layers: int = 2
    # Number of tokens MTP will try to predict
    mtp_k: int = 4
    # Maximum sequence length
    max_seq_len: int = 256

cfg = Config()
device = torch.device(
    "cuda" if torch.cuda.is_available() else "cpu"
)
print("\nExperiment configuration")
print("------------------------")
print("vocab size :", cfg.vocab_size)
print("hidden dim :", cfg.hidden_dim)
print("heads      :", cfg.num_heads)
print("layers     :", cfg.num_layers)
print("MTP tokens :", cfg.mtp_k)
print("device     :", device)

Cell 1: Build a tiny causal self-attention layer


Before MTP, we need to understand exactly what the normal autoregressive model is doing.


The most important object for this entire tutorial will be:

Kcache , Vcache


because MTP does not magically eliminate KV-cache work. It changes how we use the model around the cache

class CausalSelfAttention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.num_heads = cfg.num_heads
        self.head_dim = cfg.head_dim
        self.hidden_dim = cfg.hidden_dim
        # Q, K, V projections
        self.q_proj = nn.Linear(
            cfg.hidden_dim,
            cfg.num_heads * cfg.head_dim,
            bias=False
        )
        self.k_proj = nn.Linear(
            cfg.hidden_dim,
            cfg.num_heads * cfg.head_dim,
            bias=False
        )
        self.v_proj = nn.Linear(
            cfg.hidden_dim,
            cfg.num_heads * cfg.head_dim,
            bias=False
        )
        # Final projection
        self.o_proj = nn.Linear(
            cfg.num_heads * cfg.head_dim,
            cfg.hidden_dim,
            bias=False
        )
    def forward(
        self,
        x,
        k_cache=None,
        v_ache=None,
    ):
        """
        x: [batch, seq_len, hidden_dim]
        k_cache: [batch, num_heads, cached_seq_len, head_dim]
        v_cache:
            [batch, num_heads, cached_seq_len, head_dim]
        """
        B, T, _ = x.shape
        # ---------------------------------------------------------
        # 1. Project hidden states -> Q, K, V
        # ---------------------------------------------------------
        q = self.q_proj(x)
        k = self.k_proj(x)
        v = self.v_proj(x)
        # [B, T, H*D]
        #        ↓
        # [B, H, T, D]
        q = q.view(
            B, T,
            self.num_heads,
            self.head_dim
        ).transpose(1, 2)
        k = k.view(
            B, T,
            self.num_heads,
            self.head_dim
        ).transpose(1, 2)
        v = v.view(
            B, T,
            self.num_heads,
            self.head_dim
        ).transpose(1, 2)
        # ---------------------------------------------------------
        # 2. Append new K/V to the KV cache
        # ---------------------------------------------------------
        if k_cache is not None:
            k = torch.cat([k_cache, k], dim=2)
        if v_cache is not None:
            v = torch.cat([v_cache, v], dim=2)
        # k/v now contain:
        #
        # [old tokens] + [new tokens]
        #
        # shape:
        # [B, H, total_seq_len, D]
        # ---------------------------------------------------------
        # 3. Attention
        # ---------------------------------------------------------
        # Q:
        # [B, H, T, D]
        #
        # K:
        # [B, H, total_seq_len, D]
        #
        # Q @ K^T:
        # [B, H, T, total_seq_len]
        scores = torch.matmul(
            q,
            k.transpose(-2, -1)
        )
        scores = scores / math.sqrt(self.head_dim)
        # ---------------------------------------------------------
        # 4. Causal mask
        # ---------------------------------------------------------
        total_seq_len = k.shape[2]
        # Position of the first new token
        past_len = total_seq_len - T
        q_positions = torch.arange(
            past_len,
            total_seq_len,
            device=x.device
        )
        k_positions = torch.arange(
            total_seq_len,
            device=x.device
        )
        causal_mask = (
            k_positions.unsqueeze(0)
            <= q_positions.unsqueeze(1)
        )
        scores = scores.masked_fill(
            ~causal_mask,
            float("-inf")
        )
        # ---------------------------------------------------------
        # 5. Softmax
        # ---------------------------------------------------------
        attention_weights = F.softmax(
            scores,
            dim=-1
        )
        # ---------------------------------------------------------
        # 6. Weighted sum of V
        # ---------------------------------------------------------
        out = torch.matmul(
            attention_weights,
            v
        )
        # [B, H, T, D]
        #      ↓
        # [B, T, H, D]
        #      ↓
        # [B, T, H*D]
        out = out.transpose(1, 2).contiguous()
        out = out.view(
            B,
            T,
            self.num_heads * self.head_dim
        )
        out = self.o_proj(out)
        return out, k, v

In the snippet above, we assume B=1, T=3, H=4, D=32, where B represents the batch size, T indicates the number of tokens, H denotes the number of heads, and D stands for the head dimension.


our input is "The cat sat" which is x = [1, 3, 128]. Then The Q/K/V projections transform this into:

Q = [1, 4, 3, 32]
K = [1, 4, 3, 32]
V = [1, 4, 3, 32]

The important transformation is:

[batch, sequence, hidden]

to:

[batch, heads, sequence, head_dim]

Why do we split the hidden dimension?

Our hidden dimension is 128, and we have 4 heads, with each head consisting of 32 dimensions. Therefore, we can express this relationship as 128 = 4 × 32. Instead of thinking of a token as a single 128-dimensional vector, attention mechanisms conceptualize it as a token that is divided into four heads, each with 32 dimensions. This distinction will become important when we discuss KV-cache bandwidth.

The critical part: KV cache: Imagine our prompt is:

"The cat sat"

We process all three tokens.

We obtain:

K_cache
[1, 4, 3, 32]

V_cache

[1, 4, 3, 32]

Now we generate:
"on"

We don't throw away the old K/V.

Instead:

old cache
K:
token1
token2
token3
       +

new K:
token4
       ↓
new cache
token1
token2
token3
token4
So:

K_cache:
[1, 4, 3, 32]
        ↓
[1, 4, 4, 32]
and similarly:
V_cache:
[1, 4, 3, 32]
        ↓
[1, 4, 4, 32]

Why is KV cache necessary?

Without KV cache, when generating token 4:

token1 token2 token3 token4

       ↓

     model

we would recompute K and V for:

token1
token2
token3

even though they haven't changed.

With KV cache:

token1 ──┐
token2 ──┤
token3 ──┤── stored K/V
token4 ──┘

Only token 4 needs new Q/K/V computation.

That's the fundamental reason autoregressive inference is feasible.

But notice something subtle

Even though we don't recompute old K/V, we still read the old K/V.

For one new token:

Q = [1, 4, 1, 32]
K_cache = [1, 4, T, 32]
V_cache = [1, 4, T, 32]

Attention does:

Q.KT

which means the new query has to interact with all previous keys.

And then:

Attention(Q,K,V)V

reads the values too.

So at long context:

          KV cache
             │
             │ read
             ▼
New Q ────► attention

             │
             ▼
           output

The computation becomes increasingly influenced by memory traffic.

This is exactly why your earlier LLM inference work around KV-cache bandwidth is important.

And now look at MTP


This is where things become interesting.

Normal autoregressive decoding:

             Q
             │
             ▼
       ┌─────────────┐
       │ KV cache │
       │ 1...T │
       └─────────────┘
             │
             ▼
          token T+1
Then:

             Q
             │
             ▼
       ┌─────────────┐
       │ KV cache │
       │ 1...T+1 │
       └─────────────┘
             │
             ▼
          token T+2


One expensive target step at a time.

With MTP/speculative decoding, we eventually want something conceptually closer to:

                 ┌── token T+1
                 │
hidden state ────┼── token T+2
                 │
                 ├── token T+3
                 │
                 └── token T+4

Then the target model can verify those positions together, but there is a huge conceptual trap: "If MTP generates 4 tokens, does the target model only need to read the KV cache once?" Not necessarily, and understanding exactly what happens to K/V for the speculative tokens is one of the most valuable parts of this notebook. We'll get there.
























 
 
 

Comments


bottom of page