概述

从零开始写Qwen3(一)模型结构分析中,已经识别了Qwen3的基本结构,了解了使用哪些关键组件,现在开始实现,将模型分为这几个模块:

  • attn,注意力计算,不包含短接和QKV嵌入,无参数
  • self_attn,自注意力,调用attn
  • norm,主要实现RMSNorm
  • rope
  • feedback,实现SwiGLU
  • transformer_block,一个完整的Transformer Decoder层
  • qwen3,完整的模型

本章的目标就是快速搭建一个模型,能够推理即可,之后再自己手写组件

从零开始写Qwen3目录

模型组件

Attn

直接使用现成的

from torch.nn.functional import scaled_dot_product_attention
        return scaled_dot_product_attention(
            q,
            k,
            v,
            is_causal=is_causal,
            enable_gqa=True,
            scale=self.n_head_embed**-0.5,
        )

需要注意:

  1. qkv以及输出o的维度是 BxHxSxD,H是多头,在enable_gqa时,Q的头可以是KV的整倍数,O的头数和Q相同
  2. is_causal,表示是否开启遮罩。
    a. 这里非常奇怪的一点,理论上在decode阶段,开启KVCache,Q长度只有1,这个参数有没有无所谓,但如果传了True会导致结果错误,所以我们在decode阶段将其手动设置为False
  3. scale就是自注意力中的那个缩放, D \sqrt{D} D

norm

有现成的

def __init__(...):
        self.weight = nn.Parameter(torch.ones(norm_dim))
def forward(...):
        return nn.functional.rms_norm(
            x, self.weight.shape, self.weight, self.eps

eps是为了防止分母为0添加的一个数,这个参数一般模型config.json会提供,Qwen3-0.6B是1e-6

rope

这个没有现成的,需要手写

原理

RoPE是模仿向量内积的那个夹角 cos ⁡ θ \cos\theta cosθ弄的,让两个向量旋转它们位置对应的角度,然后内积就可以体现它们之间相对位置了

对于二维向量可以在平面上执行旋转,但高维向量无法保证都在同一个平面上,其实并不需要旋转高维向量
x ⋅ y = ∑ i = 1 d x i y i = ∑ k = 1 d / 2 ( x 2 k − 1 y 2 k − 1 + x 2 k y 2 k ) \mathbf{x} \cdot \mathbf{y} = \sum_{i=1}^{d} x_i y_i = \sum_{k=1}^{d/2} \left( x_{2k-1} y_{2k-1} + x_{2k} y_{2k} \right) xy=i=1dxiyi=k=1d/2(x2k1y2k1+x2ky2k)
两个高维向量的内积,其实可以看做多个二维向量内积之和(大模型隐藏层维度一般是偶数,如果真是奇数,可以补充一个0上去)

rope的做法就是把特征向量当成多个二维向量,对每个二维向量进行旋转
对一个二维向量的旋转就是乘一个矩阵
R ( θ ) = [ cos ⁡ θ − sin ⁡ θ sin ⁡ θ cos ⁡ θ ] R(\theta) = \begin{bmatrix} \cos\theta & -\sin\theta\\ \sin\theta & \cos\theta \end{bmatrix}\\ R(θ)=[cosθsinθsinθcosθ]
实际上,每个小向量可以旋转不同的角度,比如低维度的可以旋转大一点的角度,高维度旋转小一点的角度,这样在相同 cos ⁡ θ \cos \theta cosθ 时,高维度必须距离更远才能达到和低维度相同的系数,低维度可以捕获短距离、局部的关系,高维度则捕获长距离、全局的关系,一般是这样生成基础旋转角的
θ = β − 2 i / d \theta=\beta^{-2i/d} θ=β2i/d
i是维度组,向下除以2取整,d是维度,beta是10000乃至更大,这个参数也在config.json中有记录

对每个在位置n的词元嵌入,旋转 n θ n\theta nθ

实现

注意到旋转之后有这个
x ′ = cos ⁡ θ x − sin ⁡ θ y y ′ = sin ⁡ θ x + cos ⁡ θ y x ′ ⃗ = [ x ′ ; y ′ ] = cos ⁡ θ x ⃗ + sin ⁡ θ σ ( x ⃗ ) x'=\cos \theta x -\sin\theta y\\ y'=\sin \theta x + \cos \theta y\\ \vec{x'}=[x';y']=\cos \theta \vec{x} + \sin \theta \sigma(\vec{x}) x=cosθxsinθyy=sinθx+cosθyx =[x;y]=cosθx +sinθσ(x )
这里 σ \sigma σ是奇偶交换,并把交换后的偶数部分取负

实际上, σ \sigma σ除了奇偶交换(原始版本),还有其他方式,比如GPT-Neox的交换前一半和后一半,Qwen3也用这个
在llama.cpp中,将交换奇偶的版本叫做normal,交换前后一半的叫做neox

将上面的公式扩展后就有了rope的实现
x ′ ⃗ = cos ⁡ θ ⃗ ∘ x ⃗ + sin ⁡ θ ⃗ ∘ σ ( x ⃗ ) \vec{x'}=\cos \vec{\theta} \circ \vec{x} + \sin \vec{\theta}\circ \sigma(\vec{x}) x =cosθ x +sinθ σ(x )
这个rope在单次前向计算中对所有层,不管是K还是V,这个 cos ⁡ θ ⃗ , sin ⁡ θ ⃗ \cos\vec{\theta},\sin\vec{\theta} cosθ ,sinθ 都是相同的,所以全局计算一次即可,如果可以保证生成长度不会超过最大长度,那甚至可以预先计算最大长度的这两个向量,然后后面直接截取其中一部分使用即可

简单计算

    def build_cos_sin_embed(
        self, dtype, position_ids: torch.Tensor
    ) -> tuple[torch.Tensor, torch.Tensor]:
        inv_freq = 1.0 / (
            self.base_freq
            ** (
                torch.arange(
                    0, self.head_dim, 2, device=position_ids.device
                ).float()
                / self.head_dim
            )
        ).unsqueeze(0)
        freqs = torch.einsum("bj,bk->bjk", position_ids, inv_freq)
        emb = torch.cat([freqs, freqs], dim=-1)
        return (emb.cos().to(dtype), emb.sin().to(dtype))

这里做了一个乘法,其实就是 n θ n\theta nθ,为每个位置生成一个旋转向量

einsum是爱因斯坦求和,用于张量计算的,bj,bk->bjk 的意思是输入两个矩阵,下标分别是bj和bk,然后乘出来一个三维数组,下标是 bjk,也就是相同 b 的数乘到一起 C b j k = A b j B b k C_{bjk}=A_{bj}B_{bk} Cbjk=AbjBbk,如果是 jb,bk->jk,那就是常规的矩阵乘法了 C j k = ∑ j A b j B j k C_{jk}=\sum_j A_{bj}B_{jk} Cjk=jAbjBjk,通常后面这个求和符号可以省略

selfAttn

这个需要自己手写,也简单:

  • QKV进行投影,QK需要进行后norm
  • 然后把QKV通过view+tranpose转为 BHSD 的格式
  • 对QK执行rope
  • 然后计算attn得到o
  • o转置到BSHD,然后reshape到BSD后投影输出

所有的投影直接使用Linear即可

其他

feedback其实没什么特别的组件,Linear和Swish激活函数就没了

transformer_block 就是上面组件拼接即可

qwen3在从零开始写Qwen3(一)模型结构分析中展示了forward,就是一个Embedding嵌入,执行多层TransformerBlock,最后Norm+Linear就行

运行推理

加载参数

HuggingFace下载的模型文件夹下有safetensors文件(Qwen3-0.6B较小,只有一个文件),这就是模型参数,里面包含所有张量和张量的名字,通过张量的名字来加载到参数中:

实际上这些张量的名字就是transformers中成员变量的名字,这些成员要么嵌套包含其他组件,要么包含一个叫做.weight的成员,类型是nn.Parameter,后者就是要加载的参数,实际上所有的张量名字都是以.weight结尾的

比如一个张量的名字是model.layers.8.self_attn.k_norm.weight,它就是编号为8(从0开始)的解码层的自注意力组件的k_norm的参数

执行推理

有个完整的模型,就可以推理了,整个模型的输入是嵌入id列表,可以直接使用from tokenizers import Tokenizer这个库,这个库和 transformersAutoTokenizer效果是一样的,因为都是加载了模型目录下的tokenizer.json

        tokenizer = Tokenizer.from_file(model_path + "/tokenizer.json")
        print(prompt)
        inputs = tensor.tensor([tokenizer.encode(prompt)])

输出参数是BxSxD的logits,只需要最后一个S的特征向量,它的长度就是字典长度,每个值代表一个词元,可以简单argmax得到最可能的下一个词元id,然后把这个词元id拼接到 inputs,预测下一个词元,不断如此,直到达到最大长度或者遇到 eos

提示词可以直接写字符串,这是大模型的补全模式,但模型目录下还有一个tokenizer_config.json文件,其中有一个chat_template成员,这是一个Jinja语法的模板字符串,它可以填充模板实现多轮对话的格式,也是对话模式,可以直接加载上

    with open(model_path + "/tokenizer_config.json") as f:
        data = json.load(f)
        template = jinja2.Template(data["chat_template"])
        prompt = template.render(
            messages=[{"role": "user", "content": "介绍一下你自己"}]
        )
        tokenizer = Tokenizer.from_file(model_path + "/tokenizer.json")

这个模板实际上就是我们常用的ChatMessage格式。模板中还支持工具调用,写tool_calls就行,不知道Qwen3-0.6B的模型的工具调用能力如何,有机会写个Agent看看

这个模板基本是这样的(截取前半部分)

{%- if tools %}
    {{- '<|im_start|>system\n' }}
    {%- if messages[0].role == 'system' %}
        {{- messages[0].content + '\n\n' }}
    {%- endif %}
    {{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
    {%- for tool in tools %}
        {{- "\n" }}
        {{- tool | tojson }}
    {%- endfor %}
    {{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
{%- else %}
    {%- if messages[0].role == 'system' %}
        {{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }}
    {%- endif %}
{%- endif %}

这是Jinja语法,上面的内容大概是这样

  • 如果有tools这个变量传入才处理
    • 如果第一个messages(也就是可以传数组)的role是system
      • 把系统提示词放到总体提示词开头
    • 接着写入工具提示词,把工具描述放进去,直接json(但这个json通常都是有严格格式要求的)
  • 否则,也就是没有传入工具
    • 简单把系统提示词传入

注意到每轮对话都是以 <|im_start|>role开头,<|im_end|>结尾,eos就是这个<|im_end|>

现在没有做KVCache,运行速度很慢,而且越来越慢,这是我们下一节要解决的问题

Logo

欢迎加入 MCP 技术社区!与志同道合者携手前行,一同解锁 MCP 技术的无限可能!

更多推荐