发布于 2026-08-21 13:44:18

古法手撕简易LLM推理器

深度学习学完Transformer之后,决定古法手撕一个推理器出来,亲自做一遍可以验证自己的理解是否存在偏差,并且对各类细节有更深入和更系统的了解。

我选择Qwen3-0.6B模型实现推理器,并在我的3070Ti(8GB)上验证。

整体仅使用pytorch的张量计算,不使用任何其他高级封装函数。但在参数读取和tokenize方面,使用了Hugging Face的safetensors库读取模型权重,以及tokenizers库实现文本与token id之间的转换,省去了基于tokenizer.json处理BPE词表的merges的步骤,更关注于推理器的实现。总代码约200多行。

完整代码>>

完整输出结果(Jupyter Notebook)>>

模型结构

首先要看模型结构,图片来自LLM Architecture Gallery

image.png

相比Transformer原始论文中的实现,主要有以下几点变化:

  • 采用RMSNorm而非LayerNorm,不再减去均值
  • 残差连接与RMSNorm的位置与原文有略微差异
  • 采用旋转位置编码RoPE而非三角函数位置编码
  • 额外引入对Q/K的RMSNorm
  • Q与K/V的注意力头数不一样,需要对K/V做交错式repeat
  • 前馈神经网络FFN采用SwiGLU门控前馈网络,引入了门控机制并采用SiLU作为激活函数

实现

整体拆解成下面这些步骤和模块:

tokenize逻辑

借助Hugging Face的tokenizers库快速实现文本与token id之间的转换。

DIR = "../models/HuggingFace/Qwen3-0.6B"

from tokenizers import Tokenizer

tokenizer = Tokenizer.from_file(DIR + "/tokenizer.json")
output = tokenizer.encode("测试一下tokenizer的效果")
print(output.ids)
tokenizer.decode(output.ids)

处理参数和超参数

model.safetensorsconfig.json中读取参数和超参数。

from safetensors import safe_open
import json
import torch

# 常数和超参数
eos_token_id: int  # 用于判断模型输出完的token id
rms_norm_eps: float  # RMSNorm防止归一化分母为0的eps
rope_theta: int  # RoPE的theta的底数部分
num_hidden_layers: int  # transformer层数
num_attention_heads: int  # Q的注意力头数
num_key_value_heads: int  # KV的注意力头数
head_dim: int  # 单头注意力维数
max_position_embeddings: int  # 允许位置编码的最大序列长度

# 参数
E: torch.Tensor  # 嵌入层
Gamma_finalRMSNorm: torch.Tensor  # 最终RMSNorm的缩放系数
W_output: torch.Tensor  # 最终投影到词表大小的线性层
Gamma_inputRMSNorm: list[torch.Tensor] = []  # transformer层中对输入的RMSNorm的缩放系数
Gamma_postATTNRMSNorm: list[torch.Tensor] = []  # transformer层中对自注意力结果的RMSNorm的缩放系数
W_FFNGate: list[torch.Tensor] = []  # transformer层中FFN的门控层
W_FFNUp: list[torch.Tensor] = []  # transformer层中FFN的升维层
W_FFNDown: list[torch.Tensor] = []  # transformer层中FFN的降维层
Gamma_QRMSNorm: list[torch.Tensor] = []  # transformer层中Q投影后、RoPE前的每个头的RMSNorm的缩放系数
Gamma_KRMSNorm: list[torch.Tensor] = []  # transformer层中K投影后、RoPE前的每个头的RMSNorm的缩放系数
W_Q: list[torch.Tensor] = []  # transformer层中Q投影矩阵
W_K: list[torch.Tensor] = []  # transformer层中K投影矩阵
W_V: list[torch.Tensor] = []  # transformer层中V投影矩阵
W_O: list[torch.Tensor] = []  # transformer层中O投影矩阵

# 读取常数和超参数
with open(DIR + "/config.json", "r", encoding="utf-8") as f:
    data = json.load(f)

    def readParam(name, t):
        d = data.get(name)
        if not isinstance(d, t):
            raise ValueError(f"get {name} error")
        return d

    eos_token_id = readParam("eos_token_id", int)  # 151645
    rms_norm_eps = readParam("rms_norm_eps", float)  # 1e-06
    rope_theta = readParam("rope_theta", int)  # 1000000
    num_hidden_layers = readParam("num_hidden_layers", int)  # 28
    num_attention_heads = readParam("num_attention_heads", int)  # 16
    num_key_value_heads = readParam("num_key_value_heads", int)  # 8
    head_dim = readParam("head_dim", int)  # 128
    max_position_embeddings = readParam("max_position_embeddings", int)  # 40960

# 读取参数
with safe_open(DIR + "/model.safetensors", framework="pt", device=0) as f:
    E = f.get_tensor("model.embed_tokens.weight")  # torch.Size([151936, 1024])
    Gamma_finalRMSNorm = f.get_tensor("model.norm.weight")  # torch.Size([1024])
    W_output = f.get_tensor("lm_head.weight")  # torch.Size([151936, 1024])
    for i in range(num_hidden_layers):
        Gamma_inputRMSNorm.append(f.get_tensor(f"model.layers.{i}.input_layernorm.weight"))  # torch.Size([1024])
        Gamma_postATTNRMSNorm.append(f.get_tensor(f"model.layers.{i}.post_attention_layernorm.weight"))  # torch.Size([1024])
        W_FFNGate.append(f.get_tensor(f"model.layers.{i}.mlp.gate_proj.weight"))  # torch.Size([3072, 1024])
        W_FFNUp.append(f.get_tensor(f"model.layers.{i}.mlp.up_proj.weight"))  # torch.Size([3072, 1024])
        W_FFNDown.append(f.get_tensor(f"model.layers.{i}.mlp.down_proj.weight"))  # torch.Size([1024, 3072])
        Gamma_QRMSNorm.append(f.get_tensor(f"model.layers.{i}.self_attn.q_norm.weight"))  # torch.Size([128])
        Gamma_KRMSNorm.append(f.get_tensor(f"model.layers.{i}.self_attn.k_norm.weight"))  # torch.Size([128])
        W_Q.append(f.get_tensor(f"model.layers.{i}.self_attn.q_proj.weight"))  # torch.Size([2048, 1024])
        W_K.append(f.get_tensor(f"model.layers.{i}.self_attn.k_proj.weight"))  # torch.Size([1024, 1024])
        W_V.append(f.get_tensor(f"model.layers.{i}.self_attn.v_proj.weight"))  # torch.Size([1024, 1024])
        W_O.append(f.get_tensor(f"model.layers.{i}.self_attn.o_proj.weight"))  # torch.Size([1024, 2048])

model_dtype = E.dtype

工具函数

主要是RMSNorm、RoPE、SiLU、softmax、masked_softmax等工具函数的实现。

def RMSNorm(X, gamma):
    return X / torch.sqrt(torch.mean(X[..., :] ** 2, -1, keepdim=True) + rms_norm_eps) * gamma

# interleaved: 交错配对theta[0,0,1,1,2,2,3,3]
# rope_theta_repeat2 = torch.pow(rope_theta, -torch.arange(head_dim // 2) * 2 / head_dim).repeat_interleave(2)
# rotate_half: 前后两半配对theta_0,1,2,3,0,1,2,3
rope_theta_repeat2 = torch.pow(rope_theta, -torch.arange(head_dim // 2) * 2 / head_dim).repeat(2)
rope_m_theta = torch.arange(max_position_embeddings).reshape(-1, 1) * rope_theta_repeat2.reshape(1, -1)
rope_sin_m_theta = torch.sin(rope_m_theta).cuda()  # torch.Size([40960, 128])
rope_cos_m_theta = torch.cos(rope_m_theta).cuda()  # torch.Size([40960, 128])
def RoPE(X, offset=0):
    Ls = X.shape[-2]
    d = X.shape[-1]
    cos, sin = rope_cos_m_theta[offset : offset + Ls], rope_sin_m_theta[offset : offset + Ls]
    # interleaved: 交错配对x[-1,0,-3,2]
    # X_rot = torch.stack((-X[..., 1::2], X[..., 0::2]), dim=-1).flatten(-2)
    # rotate_half: 前后两半配对x[-0,-1,-2,-3,0,1,2,3]
    X_rot = torch.cat((-X[..., d // 2 :], X[..., : d // 2]), dim=-1)
    return (X * cos + X_rot * sin).to(model_dtype)

def SiLU(X):
    return X / (1 + torch.exp(-X))

def softmax(X, temperature=1):
    X_shifted = X - torch.max(X / temperature, dim=-1, keepdims=True).values
    exp_X = torch.exp(X_shifted)
    return exp_X / torch.sum(exp_X, dim=-1, keepdims=True)

# 右上方三角被mask
def masked_softmax(X):
    # X: (batch_size, num_heads, num_steps, num_steps)
    maxlen = X.size(-1)
    mask = (torch.arange(maxlen).reshape(-1, 1) >= torch.arange(maxlen)).to(X.device)
    mask_not = ~mask
    return softmax(X * mask + (mask_not * -1e6).to(X.dtype))

核心计算函数与Transformer层

主要是FFN、缩放点积注意力、KVCache、单个Transformer层的计算逻辑。

由于投影矩阵是所有注意力头合在一起存放的,transpose_qkv/transpose_output主要用于在计算时的注意力头拆分与恢复。

KVCache主要存储经过投影、RMSNorm、RoPE的K,以及经过投影的V(V不需要RMSNorm、RoPE)。在prefill阶段一次性存储,在predict阶段每个token预测过程中一个个附加存储上去。

def FFN(X, gate, up, down):
    return (SiLU(X @ gate.T) * (X @ up.T)) @ down.T

# 通过转置实现Q/K/V张量按注意力头的拆分与恢复,便于一次性计算所有注意力头
def transpose_qkv(X, num_heads):
    # X_h: (batch_size, num_steps, num_heads, num_hiddens/num_heads)
    X_h = X.reshape(*X.shape[:-1], num_heads, -1)
    return X_h.transpose(-2, -3)
def transpose_output(X):
    # X_t: (batch_size, num_steps, num_heads, num_hiddens/num_heads)
    X_t = X.transpose(-2, -3)
    return X_t.reshape(*X_t.shape[:-2], -1)

def DotProductAttention(Q, K, V, offset=0):
    KVScale = num_attention_heads // num_key_value_heads
    # Q/K/V: (batch_size, num_heads, num_steps, num_hiddens/num_heads)
    # scores/A: (batch_size, num_heads, num_steps, num_steps)
    scores = Q @ torch.repeat_interleave(K.transpose(-1, -2), KVScale, -3) / torch.sqrt(torch.tensor(Q.shape[-1]))
    A = softmax(scores) if offset else masked_softmax(scores)
    return A @ torch.repeat_interleave(V, KVScale, -3)

KVCache: list[list[torch.Tensor]]

def ClearKVCache():
    global KVCache
    KVCache = [None] * num_hidden_layers

# KVCache中已经prefill或predict了多少个字符
def GetOffset():
    global KVCache
    return KVCache[0][0].size(-2)

# offset为0代表prefill阶段
def TransfomerLayer(layer, X, Gamma_inputRMSNorm, Gamma_postATTNRMSNorm, W_FFNGate, W_FFNUp, W_FFNDown, Gamma_QRMSNorm, Gamma_KRMSNorm, W_Q, W_K, W_V, W_O, offset=0):
    # X: (batch_size, num_steps, num_hiddens)
    X_inputRMSNorm = RMSNorm(X, Gamma_inputRMSNorm)
    # 投影 -> QK RMSNorm -> RoPE -> attention
    # Q/K/V: (batch_size, num_steps, num_heads, num_hiddens) -> (batch_size, num_heads, num_steps, num_hiddens/num_heads)
    Q = transpose_qkv(X_inputRMSNorm @ W_Q.T, num_attention_heads)
    K = transpose_qkv(X_inputRMSNorm @ W_K.T, num_key_value_heads)
    V = transpose_qkv(X_inputRMSNorm @ W_V.T, num_key_value_heads)
    Q_RMSNorm = RMSNorm(Q, Gamma_QRMSNorm)
    K_RMSNorm = RMSNorm(K, Gamma_KRMSNorm)
    Q_RoPE = RoPE(Q_RMSNorm, offset)
    K_RoPE = RoPE(K_RMSNorm, offset)
    # prefill阶段初始化KVCache
    global KVCache
    if offset == 0:
        KVCache[layer] = [K_RoPE, V]
    else:
        KVCache[layer][0] = torch.cat((KVCache[layer][0], K_RoPE), dim=-2)
        KVCache[layer][1] = torch.cat((KVCache[layer][1], V), dim=-2)
    # V_attn: (batch_size, num_heads, num_steps, num_hiddens/num_heads)
    V_attn = DotProductAttention(Q_RoPE, KVCache[layer][0], KVCache[layer][1], offset)
    X2 = X + transpose_output(V_attn) @ W_O.T
    # FFN前的RMSNorm -> FFN
    # O: (batch_size, num_steps, num_hiddens)
    O = X2 + FFN(RMSNorm(X2, Gamma_postATTNRMSNorm), W_FFNGate, W_FFNUp, W_FFNDown)
    return O

prefill与predict

由于prefill阶段不需要output,所以这里把Transformer层后面的输出逻辑单独摘出来为lastOutputToken函数,prefill不需要调用,predict需要调用。

def prefill(X):
    ClearKVCache()
    # X: (batch_size, num_steps)
    # O: (batch_size, num_steps, num_hiddens)
    O = E[X]
    for i in range(num_hidden_layers):
        O = TransfomerLayer(
            i, O, Gamma_inputRMSNorm[i], Gamma_postATTNRMSNorm[i], W_FFNGate[i], W_FFNUp[i], W_FFNDown[i], Gamma_QRMSNorm[i], Gamma_KRMSNorm[i], W_Q[i], W_K[i], W_V[i], W_O[i]
        )
    return O

# 多层Transformer得到的结果经过最终的RMSNorm、线性层和softmax得到概率
def lastOutputToken(O, temperature=1):
    O = RMSNorm(O, Gamma_finalRMSNorm)
    Y = O @ W_output.T
    P = softmax(Y, temperature)
    return torch.argmax(P, dim=-1)[..., -1:].to(torch.long)

def predict(X, temperature=1):
    offset = GetOffset()
    # O: (batch_size, 1)
    res = X.clone().cpu()
    while offset < max_position_embeddings:
        token_id = res[0][-1].numpy()
        print(tokenizer.decode([token_id], skip_special_tokens=False), end="")
        if token_id == eos_token_id:
            break
        O = E[X]
        for i in range(num_hidden_layers):
            O = TransfomerLayer(
                i, O, Gamma_inputRMSNorm[i], Gamma_postATTNRMSNorm[i], W_FFNGate[i], W_FFNUp[i], W_FFNDown[i], Gamma_QRMSNorm[i], Gamma_KRMSNorm[i], W_Q[i], W_K[i], W_V[i], W_O[i], offset
            )
        X = lastOutputToken(O, temperature)
        res = torch.cat((res, X.cpu()), dim=-1)
        offset += 1
    print()
    return res.numpy()

构造输入

解析tokenizer_config.json文件中的template(jinja2格式),将给定测试messages转为用于模型输入的文本,并通过tokenizer.encode得到token id序列。

import jinja2

with open(DIR + "/tokenizer_config.json", "r", encoding="utf-8") as f:
    data = json.load(f)
    chat_template = data["chat_template"]

tpl = jinja2.Environment().from_string(chat_template)

messages = [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "steam现在最火的游戏是什么"}]

rendered = tpl.render(
    messages=messages,
    add_generation_prompt=True,
    enable_thinking=True,
)
tokens = tokenizer.encode(rendered).ids

发起预测

先完成prefill,将最后一个时间步的Transformer结果,计算最后的token输出逻辑,将其作为predict的第一个输入,循环预测。

同时在整个过程中监测prefill阶段耗时以及predict阶段每秒输出的token数。

由于整个实现是支持多样本的,因此可以同时跑多条prefill和predict。

import time

batch_size = 50

start = time.perf_counter()
O = prefill(torch.tensor(tokens, dtype=torch.long).reshape(1, -1).repeat(batch_size, 1))
end = time.perf_counter()
print(f"prefill耗时: {end - start:.2f} 秒")

start = time.perf_counter()
X = lastOutputToken(O)
res = predict(X)
end = time.perf_counter()
print(f"predict: {len(res[0])*batch_size/(end - start):.2f} tokens/s")
print(res)

性能

由于整体的实现在优化上做得很少,在我3070Ti(8GB)上,能够实现约22 tokens/s的推理速度,在并行50个batch的情况下,能够实现925 tokens/s的推理速度。

这个性能与vLLM/SGLang应该没有数量级上的差距,但是在显存占用上差距比较大。在并行50个batch的情况下,8G显存已经快要打满。核心原因是各类中间张量过多,可以通过算子简化或者inplace运算来降低显存占用。

在过程中遇到的问题

在debug的过程中主要有这么几类问题:

  • device的处理:权重和计算过程中的数据一开始就是默认在GPU上的,但是一些arange的数据默认在CPU上,需要手动指定
  • dtype数据类型的问题:模型权重是bfloat16,而pytorch默认是float32,新创建的(不是与已有张量计算产生的)张量需要指定dtype
  • RoPE计算问题:这个下面单独讲

整体来讲最大的bug在RoPE这里,因为我之前看的Transformer原始实现用的是三角函数位置编码,RoPE不太熟悉,找的资料与模型真实实现其实有差异。反倒是Transformer的实现由于理论方面了解得特别清楚,写完就bug free了。

RoPE的细节

第一个要注意的是RoPE的theta的底( ),要从config.json拿,并不都是10000。

第二个要注意的点是:RoPE有interleaved和rotate_half两种效果等价的变换方式,但是推理时必须采用与训练时一样的变换方式。

interleaved的计算公式:

rotate_half的计算公式:

整体来讲,工业界更常用rotate_half的方式,因为它可以通过复数来简化计算。Qwen3-0.6B使用的就是该方式。

欢迎留言