AI infra 学习笔记(三):预处理·第一步——QKV 投影到底算了多少

文章目录

  • [AI infra 学习笔记(三):预处理·第一步------QKV 投影到底算了多少](#AI infra 学习笔记(三):预处理·第一步——QKV 投影到底算了多少)
    • [0. 从注意力公式倒推回去](#0. 从注意力公式倒推回去)
    • [1. 题目](#1. 题目)
    • [2. 这些数字是哪来的?翻 `config.json`](#2. 这些数字是哪来的?翻 config.json)
    • [3. 投影到底怎么算](#3. 投影到底怎么算)
    • [4. 交叉验证:去 `safetensors` 里看真实形状](#4. 交叉验证:去 safetensors 里看真实形状)
    • [5. 动手算:参数量](#5. 动手算:参数量)
      • [5.1 单层明细](#5.1 单层明细)
      • [5.2 全模型](#5.2 全模型)
    • [6. 顺手把计算量也算出来](#6. 顺手把计算量也算出来)
    • [7. 几个值得记住的结论](#7. 几个值得记住的结论)
    • 参考

AI infra 学习笔记(三):预处理·第一步------QKV 投影到底算了多少

0. 从注意力公式倒推回去

注意力的公式大家都背得出来:

Attention ( Q , K , V ) = softmax ( Q K T d k ) V \text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V Attention(Q,K,V)=softmax(dk QKT)V

但这里有个前置问题常被跳过:Q、K、V 本身是从哪来的?

答案很简单------它们都是同一个输入 Token 乘上三个不同的权重矩阵得到的 。这一步叫投影(Projection),是每个 Transformer 层里最先吃掉算力的地方。

先上一道题,一起实际看看!


1. 题目

在 Qwen3-8B 中,一个被输入进来的 Token 假设是 [1, 4096];其中查询(Q)和输出(O)投影的宽度是 4096 ,键(K)和值(V)的宽度是 1024 ,FFN 中间维度是 12288。

求四个注意力投影 + 三个 FFN 矩阵的参数量。

参考模型文件:https://www.modelscope.cn/models/Qwen/Qwen3-8B/files


2. 这些数字是哪来的?翻 config.json

上面的 4096 / 1024 / 12288 是在 config.json 里实际写到的。

json 复制代码
{
  "hidden_size": 4096,
  "head_dim": 128,
  "num_attention_heads": 32,
  "num_key_value_heads": 8,
  "intermediate_size": 12288,
  "num_hidden_layers": 36,
  "vocab_size": 151936,
  "torch_dtype": "bfloat16",
  "tie_word_embeddings": false
}

其中一些字段的意义是:

config.json 字段 值 对应概念
hidden_size 4096 隐状态维度:所有投影的输入 宽度,也是 Q、O 投影的输出宽度
num_attention_heads 32 Q 的头数
num_key_value_heads 8 K、V 的头数(GQA 分组查询,所以 K/V 比 Q 窄)
head_dim 128 每个注意力头的维度
intermediate_size 12288 FFN 中间维度(gate / up 的输出宽度)
num_hidden_layers 36 层数

于是可以直接推出三个关键宽度:

宽度 计算 结果
Q 的宽度 num_attention_heads × head_dim = 32 × 128 4096
K、V 的宽度 num_key_value_heads × head_dim = 8 × 128 1024
Token 的维度 hidden_size 4096

⚠️ 关于 GQA :因为 num_key_value_heads (8) < num_attention_heads (32),每 4 个 Q 头共享 1 组 K/V 头。这就是 K、V 投影比 Q 窄 4 倍的原因,也是显存能省下来的关键。如果 K/V 头数等于 32(标准 MHA),K、V 宽度就会是 4096 而不是 1024。


3. 投影到底怎么算

输入 Token 记作 x x x([1, 4096],1 个 token、4096 维),它分别和三个权重矩阵相乘:

Q = x W q , K = x W k , V = x W v Q = xW_q,\quad K = xW_k,\quad V = xW_v Q=xWq,K=xWk,V=xWv

由矩阵乘法的维度规则([1, in] × [in, out] = [1, out])可以反推每个权重矩阵的形状:

矩阵 形状 [in, out] 输出维度 推导
W q W_q Wq [4096, 4096] Q = [1, 4096] [1,4096] × [4096,4096] = [1,4096]
W k W_k Wk [4096, 1024] K = [1, 1024] [1,4096] × [4096,1024] = [1,1024]
W v W_v Wv [4096, 1024] V = [1, 1024] [1,4096] × [4096,1024] = [1,1024]

注意力算完之后,多头输出拼回 4096 维,还要再过一个输出投影 W o W_o Wo 才能回到残差流:

矩阵 形状 [in, out] 说明
W o W_o Wo [4096, 4096] 把拼接后的多头输出投回 hidden_size

FFN 部分(Qwen3 用的是 SwiGLU 结构)有三个矩阵:

矩阵 形状 [in, out] 作用
gate_proj [4096, 12288] 门控分支,过 SiLU 激活
up_proj [4096, 12288] 升维分支
down_proj [12288, 4096] 降维回 hidden_size

SwiGLU 的计算是 down ( SiLU ( x W g a t e ) ⊙ x W u p ) \text{down}\big(\text{SiLU}(xW_{gate}) \odot xW_{up}\big) down(SiLU(xWgate)⊙xWup),所以是"两升一降"三个矩阵,而不是经典 FFN 的两个。


4. 交叉验证:去 safetensors 里看真实形状

光靠 config 推导还不够放心,可以直接去权重文件里核对。

模型仓库里有一份 model.safetensors.index.json,它记录了每一层的每一个权重存在哪个分片文件里:

https://www.modelscope.cn/models/Qwen/Qwen3-8B/file/view/master/model.safetensors.index.json

内容大致长这样:

json 复制代码
"model.layers.13.input_layernorm.weight": "model-00002-of-00005.safetensors",
"model.layers.13.mlp.down_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.mlp.gate_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.mlp.up_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.post_attention_layernorm.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.k_norm.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.k_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.o_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.q_norm.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.q_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.13.self_attn.v_proj.weight": "model-00002-of-00005.safetensors",
"model.layers.14.input_layernorm.weight": "model-00002-of-00005.safetensors"

每一层的 k_proj / v_proj / q_proj / o_proj 以及 FFN 的三个矩阵,都能查到归属。想看第 13 层的 Q 权重具体维度,直接点开 model-00002-of-00005.safetensors:

💡 一个容易踩的坑 :HuggingFace / safetensors 里存储的 nn.Linear 权重形状是 [out_features, in_features] ,也就是上面表格的转置。所以你在文件里看到的是:

权重 文件中的形状 与本文记法
q_proj.weight [4096, 4096] 对称,一致
k_proj.weight [1024, 4096] W k T W_k^T WkT
v_proj.weight [1024, 4096] W v T W_v^T WvT
o_proj.weight [4096, 4096] 对称,一致
gate_proj.weight [12288, 4096] W g a t e T W_{gate}^T WgateT
up_proj.weight [12288, 4096] W u p T W_{up}^T WupT
down_proj.weight [4096, 12288] W d o w n T W_{down}^T WdownT

参数量是完全一样的(转置不改变元素个数),只是行、列含义对调了。看形状时别被绕晕。


5. 动手算:参数量

关键前提:向量与矩阵的点积,每个元素是"一次乘 + 一次加",即 2 次浮点运算(FLOPs)。

先算参数量。一个 [in, out] 的权重矩阵,参数量就是 in × out。

5.1 单层明细

模块 矩阵 形状 参数量
Attention q_proj 4096 × 4096 16,777,216
Attention k_proj 4096 × 1024 4,194,304
Attention v_proj 4096 × 1024 4,194,304
Attention o_proj 4096 × 4096 16,777,216
注意力小计 41,943,040
FFN gate_proj 4096 × 12288 50,331,648
FFN up_proj 4096 × 12288 50,331,648
FFN down_proj 12288 × 4096 50,331,648
FFN 小计 150,994,944
单层合计 192,937,984

即 单层约 1.93 × 10⁸ ≈ 193 M 参数。

5.2 全模型

Qwen3-8B 共 36 层:

192,937,984 × 36 = 6,945,767,424 ≈ 6.95 × 10 9 192{,}937{,}984 \times 36 = 6{,}945{,}767{,}424 \approx 6.95 \times 10^9 192,937,984×36=6,945,767,424≈6.95×109

这 7 个矩阵合计约 6.95 B 参数,占了模型的绝大部分。再算上词嵌入矩阵:

151,936 × 4,096 = 622,329,856 ≈ 0.62 B 151{,}936 \times 4{,}096 = 622{,}329{,}856 \approx 0.62\text{ B} 151,936×4,096=622,329,856≈0.62 B

加上各层 RMSNorm 的小参数,总量落在 7.6 B 左右,与"8B"这个名义规模吻合。


6. 顺手把计算量也算出来

FLOPs = 2 × 参数量 \text{FLOPs} = 2 \times \text{参数量} FLOPs=2×参数量

因为每个权重元素都要参与一次乘加。

单 token、单层:

2 × 192,937,984 = 385,875,968 ≈ 3.86 × 10 8 FLOPs 2 \times 192{,}937{,}984 = 385{,}875{,}968 \approx 3.86 \times 10^8 \text{ FLOPs} 2×192,937,984=385,875,968≈3.86×108 FLOPs

单 token、全模型(36 层)投影部分:

2 × 6,945,767,424 = 13,891,534,848 ≈ 13.9 GFLOPs 2 \times 6{,}945{,}767{,}424 = 13{,}891{,}534{,}848 \approx 13.9 \text{ GFLOPs} 2×6,945,767,424=13,891,534,848≈13.9 GFLOPs

写成原始式子就是:

2 \\times \\underbrace{(4096{\\times}4096 + 4096{\\times}1024 + 4096{\\times}1024 + 4096{\\times}4096)}_{\\text{QKV + O 投影}} * 2 \\times \\underbrace{3 \\times (4096{\\times}12288)}_{\\text{FFN 三个矩阵}}

Prefill 场景 :如果一次喂进 S S S 个 token,投影部分的计算量是 13.9 × S 13.9 \times S 13.9×S GFLOPs(矩阵乘可以批量做,近似线性放大)。


7. 几个值得记住的结论

  1. FFN 比注意力贵得多 。单层里 FFN 占 150.99 M,注意力占 41.94 M ------ FFN 是注意力的 3.6 倍。很多人以为 Transformer 的算力都在注意力上,其实在中小序列长度下,投影 + FFN 才是大头。

  2. GQA 省的是 K/V 投影 。如果换成 MHA(num_key_value_heads = 32),K、V 宽度会从 1024 涨到 4096,注意力部分参数会从 41.94 M 涨到 100.66 M。这还没算 KV Cache 的显存节省------那才是 GQA 真正的大头。

  3. 2 × 参数量 这个经验法则很好用。推理时估算算力,直接把参数量乘 2 就是每 token 的 FLOPs(不含注意力分数部分)。

  4. 注意力的分数计算被省略了 。本文只算投影和 FFN。加上 Q K T QK^T QKT 和 A V AV AV 后,每个 token 每层还要额外约 4 × 4096 × L 4 \times 4096 \times L 4×4096×L 次浮点运算( L L L 是序列长度)。在 L = 4096 L=4096 L=4096 时约 67 M,仍小于投影部分的 386 M------但序列一长,这一项就会反超。这也是长上下文推理变慢的根源。


参考

相关推荐
龙亘川1 小时前
政法工作系统数字化建设实践解读:以全周期闭环驱动政法工作现代化
大数据·人工智能·智慧城市·开源软件·数据可视化
小宋10211 小时前
RAG 知识库也会被投毒:恶意文档、间接 Prompt Injection 与入库审核
人工智能·prompt
AI创界者1 小时前
Python 进阶:重构经典设计模式(六)—— 享元模式与单例模式在 Python 3.10+ 中的内存优化与线程安全演化
人工智能·aigc
honsor1 小时前
PoE温湿度传感器:一根网线供电+通信,即插即用,机房/配电室/仓库温湿度监测首选
运维·服务器·网络·数据库·人工智能·安全
u0111026751 小时前
工具页案例图如何编写 alt让图片说明与处理场景对应
java·前端·javascript·图像处理·人工智能·算法·ai作画
sevenez1 小时前
AiMaMi 介绍及与 CC Switch 对比
笔记
喜欢打篮球的普通人1 小时前
MiniMind 学习笔记(十):优化器、学习率和数据设置——训练稳定性的三块基石
笔记·python·学习
安睿杰翻译(上海)有限公司1 小时前
项目复盘|上海涉外项目翻译服务商选型评估清单,规避标书与注册文档风险
人工智能
林伽一1 小时前
智能体安全下沉芯片层,推理效率与资本重估同场角力|2026年09月30日
人工智能·科技·安全·ai