深度学习学完Transformer之后,决定古法手撕一个推理器出来,亲自做一遍可以验证自己的理解是否存在偏差,并且对各类细节有更深入和更系统的了解。
我选择Qwen3-0.6B模型实现推理器,并在我的3070Ti(8GB)上验证。
整体仅使用pytorch的张量计算,不使用任何其他高级封装函数。但在参数读取和tokenize方面,使用了Hugging Face的safetensors库读取模型权重,以及tokenizers库实现文本与token id之间的转换,省去了基于tokenizer.json处理BPE词表的merges的步骤,更关注于推理器的实现。总代码约200多行。
模型结构
首先要看模型结构,图片来自LLM Architecture Gallery:

相比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.safetensors和config.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使用的就是该方式。

欢迎留言