ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

从零手搓Qwen:大模型架构核心模块与PyTorch实现解析

从零手搓Qwen:大模型架构核心模块与PyTorch实现解析 大概三个月前我决定把大模型的学习路线彻底捋一遍不借助封装好的框架不看别人总结的API调用手册直接从最底层的模型架构入手一行一行把代码逻辑啃下来。当时选定的第一个研究对象就是Qwen系列。原因很简单大模型、Qwen、模型架构这三个关键词绑在一起几乎是中文社区里资料最全、模型最丰富、踩坑记录最多的组合。如果连Qwen都啃不下来其他开源模型更不用想。这个系列的第一篇我会把学习Qwen模型架构的全过程整理出来。内容包括为什么选Qwen而不是GPT或Llama作为学习起点、Qwen架构里有哪几个核心模块、每个模块在训练和推理时到底在做什么、怎么用PyTorch从零写出一个能跑的简化版Qwen推理代码、以及在复现过程中我踩过的坑和排查技巧。目标是让看完这篇文章的人哪怕之前只写过Python脚本也能对Transformer类大模型的结构有一个清晰的框架感。这篇文章不是那种导入transformers库然后跑通就结束的教程我会尽量把每个模块的来龙去脉讲清楚比如为什么Qwen要用RMSNorm替换LayerNorm、为什么注意力要分组、RoPE旋转变换到底做了什么。这些细节在官方文档里往往是一笔带过但恰恰是手搓模型时必须搞清楚的东西。1. 学习路线图Qwen凭什么成为入门首选1.1 从Qwen1到Qwen3一条清晰的演进线我接触Qwen的时候开源社区里最火的基本是LLaMA和Falcon但真正让我决定深入研究的是Qwen在模型尺寸上的完整度。从0.5B、1.5B、3B、7B、14B到72BQwen几乎覆盖了所有可跑的尺寸档位。这意味着什么呢意味着我可以在本地显卡上先用小尺寸把架构逻辑跑通再逐步切换到更大模型整个学习过程不会因为显存不够而中断。Qwen1时代模型原生上下文长度只有2048放到现在看确实有些局促。到了Qwen1.5上下文拉到了32768引入了分组查询注意力GQA。Qwen2在架构上基本定型支持131072的上下文分词器也换成了tiktoken格式。到了Qwen2.5官方把基础模型和指令微调模型的发布节奏、训练细节、稳定性都打磨得比较到位这时候社区里的教程、微调实战案例也多了起来。到了Qwen3已经不只是dense架构还出现了MoE版本混合推理模式、思考模式这些新特性也进来了。但如果一上来就学Qwen3很容易被那些新增的高级特性绕晕。我的建议是先吃透Qwen2.5这条线——它架构规整、资料齐全、技术栈稳定是理解大模型核心原理的最佳样本。1.2 和LLaMA系列相比Qwen的优势到底在哪很多人在选学习对象时会在LLaMA和Qwen之间犹豫。这两个模型家族的架构其实非常相似都用RMSNorm、SwiGLU激活、RoPE位置编码、GQA注意力。差异主要在细节上LLaMA的n_heads、n_layers这些超参数组合以及分词器的vocab大小和Qwen不一样。但对我这种中文学习者来说Qwen有一个无可替代的优势中文语料理解得更好。这不只是生成中文不别扭的问题更重要的是当我用中文去搜索相关源码解读、踩坑教程、训练细节时中文社区里关于Qwen的资料量远远超过LLaMA。学习过程中遇到一个报错随手一搜就能找到别人整理的解决方案这种效率是学LLaMA时很难享受到的。还有一个技术层面的原因Qwen官方发布的模型细节说明文档写得相对清楚包括注意力头数、层数、中间层大小、RoPE的theta参数等等都会在config.json里明确标出。对于想手搓模型的人来说这些参数就是标准答案可以随时比对自己实现的代码是否和官方一致。1.3 手搓之前先把这些前置概念补上在进入源码和公式之前有几个概念必须提前掌握否则后面会越看越懵。第一个是嵌入、注意力、前馈网络这三个Transformer基础模块各自的作用——简单说嵌入是把token变成向量注意力是在token之间建立依赖关系前馈网络是逐位置做非线性变换。这三个模块交替堆叠就是LLM的基本骨架。第二个是训练和推理的差异。训练阶段需要保存中间激活值用于反向传播所以显存占用极高推理阶段只需要单向逐token生成内存占用低得多。理解这个差异才能明白为什么手搓推理代码比手搓训练代码简单得多也才能理解为什么网上那些从零实现GPT的教程基本都是写推理逻辑。第三个是自回归生成的基本过程。Qwen生成文本时每次只预测下一个token的概率分布然后用某种采样策略挑一个token拼到输入序列末尾再继续预测下一个。这个过程在代码里就表现为一个for循环配合KV Cache来加速。这套逻辑搞明白了后面看代码就不会迷路。2. 把Qwen架构拆开看每个零件长什么样2.1 宏观视角token从输入到输出的完整旅行如果你把一个完整的Qwen-7B模型打开用PyTorch打印它的结构你会看到最外层是一个QWenLMHeadModel里面包含一个QWenModel和一个lm_head线性层。QWenModel才是真正的Transformer主干而lm_head负责把最后一层的隐藏状态映射成词表大小的logits。整个数据流是这样的输入文本先经过分词器变成token id列表比如你好世界可能变成[bich2678, 1003, 1003, 55...]这类整数序列然后进入embedding层变成shape为[seq_len, hidden_size]的浮点向量序列。接着这些向量经过L层TransformerBlock的加工每一层内部先做注意力、再做前馈网络、中间穿插归一化和残差连接最终输出同样shape的隐藏状态。把这组隐藏状态送给lm_head得到一个[seq_len, vocab_size]的logits矩阵再取最后一个位置的logits做softmax就能得到下一个token的概率分布。这里有一个值得注意的细节embedding层的权重矩阵和lm_head层经常共享参数。Qwen在实现时也用了tie_word_embeddings这个配置选项具体是否开启要看具体的config设置llama系列的默认做法是共享因为词表矩阵特别大共享能省很多显存。如果词表是15万hidden_size是4096那光是词表矩阵就有约六亿个参数不共享的话内存压力会明显上升。2.2 核心组件逐个看RMSNorm、SwiGLU、GQA、RoPEQwen架构里最核心的几个组件每个都要单独吃透。RMSNorm是Root Mean Square Layer Normalization的缩写它和传统的LayerNorm的区别是LayerNorm会先计算均值和方差然后做减去均值、除以方差的标准化操作RMSNorm干脆不做减均值这一步只除以均方根值然后再乘一个可学习的缩放向量。这样做的好处是什么计算量更小而且在很多实验里效果和LayerNorm持平甚至更好。Qwen和LLaMA基本上都是用它。SwiGLU激活函数是前馈网络里的关键。传统的FFN用的是ReLU或者GELUSwiGLU的做法是先把输入分别通过两个线性层然后把一个的Swish激活结果和另一个的原始输出做逐元素相乘最后再过第三个线性层。这种设计在论文里被证明能提升模型表现但代价是参数数量增加了大约三分之一。所以你会看到Qwen的中间层大小intermediate_size会比hidden_size大很多比如hidden_size是3584中间层可能到18944。GQA分组查询注意力是MHA多头注意力的改进版。标准MHA是每个头都有独立的Key、Query、Value投影矩阵MQA是所有的Query头共享一组Key和Value省显存但效果可能受损GQA则是把Query头分成几组每组内部共享一组Key和Value。Qwen2.5-7B的配比是28个Query头、4个KV头相当于每7个Query头共享一组KV。这样做的好处是既压缩了KV Cache的占用又保持了模型效果。RoPE旋转位置编码则是Qwen处理位置信息的方式。它不去给位置做加法而是把相邻位置的向量按一定角度做旋转让Query和Key在计算注意力分数时通过内积结果自然地携带相对位置信息。后面第3节我会用代码细讲这个旋转到底怎么实现。2.3 一个7B模型几十亿参数是怎么算出来的学习架构的时候动手算一遍参数数量是非常好的理解方式。拿Qwen2.5-7B-Instruct来说它的隐藏状态维度hidden_size是3584中间层大小是18944Transformer层数是28Query头数28、KV头数4词表大小约151936。每个TransformerBlock里的参数包括self_attn部分有q_proj、k_proj、v_proj、o_proj四个矩阵大小分别是[35843584]、[10243584]KV投影维度是3584/4等于896不对是hidden/4896但这里是1284还是算一下其实KV头的维度是head_dim乘以KV头数head_dim通常是hidden_size除以Query头数即3584/28128所以KV投影是1284512*不对k_proj是[hidden_size, num_kv_heads * head_dim] [3584, 4*128512]。MLP部分有三个矩阵gate_proj、up_proj是[3584, 18944]down_proj是[18944, 3584]。这样算下来单层Transformer的参数量大约是注意力部分两个3584*3584的矩阵再加上两小一大MLP部分三个矩阵叠加后约1亿多参数乘以28层再加上embedding矩阵15万×3584约5.4亿和lm_head共享后不用重复计算总的参数量就落在7亿级别附近不对7B应该是70亿左右。我重新估一下单个block大约2.2亿参数乘以28层约61亿加上embedding 5亿多确实接近70亿。说明这个量级是合理的。新手第一次算的时候很容易算错的是KV投影的维度因为你不能直接写成[hidden_size, hidden_size]要记得它只有K和V头的数量那么大。算一遍参数基本上就能把Qwen的维度关系理清了。3. 核心细节注意力机制与位置编码到底怎么运作3.1 从标准注意力到GQA少几个KV头真的没关系吗要理解GQA先得回到标准自注意力。假设输入序列长度是N每个token经过Q、K、V三个投影后得到Query、Key、Value向量。注意力分数就是Query和所有Key做点积然后除以根号下向量维度做缩放再经过softmax变成权重最后和对应的Value做加权求和。整个过程可以写成一行公式但在代码里会展开成好几个矩阵运算。标准MHA是每一个头都有自己的K、V投影矩阵这导致推理时的KV Cache特别大。KV Cache就是之前算出来的Key和Value矩阵缓存为了生成下一个token时不用重新计算前面所有token的Key和Value。序列越长KV Cache越大有时候甚至比模型权重本身占的显存还多。GQA的思路就是让多个Query头共享一组K、V。这样K、V头的总数就被压低KV Cache的大小也随之缩水。论文和Qwen的实际测试都表明在模型效果损失很小的情况下GQA能明显提升推理吞吐。这就是工程权衡的经典案例——越接近工程实战越能理解为什么大家不直接用理论上更好的标准MHA。3.2 RoPE旋转位置编码为什么旋转能让模型感知位置RoPE的实现思路很巧妙。它把每个token的Query或Key向量看成复数空间里的一个向量然后按照token的位置给它乘上不同角度的旋转矩阵。角度和位置的关系通常是线性增长的不同维度配不同的频率。这样做的好处是两个token的Query和Key在做点积的时候内积结果会自然地依赖于它们的相对位置差而不是绝对位置。你不需要把位置信息单独拼到输入里模型就能从内积里读出相对距离。这比那种直接加一个位置embedding的做法优雅得多也更容易泛化到训练时没见过的长序列上。我在手写RoPE的时候第一次试图用循环遍历每个位置去旋转结果发现速度极慢。后来才反应过来在PyTorch里可以用torch.polar或者直接构造cos、sin矩阵做逐元素旋转一次矩阵运算就能搞定所有位置。这种从能跑到跑得快的优化也是手搓过程中最有收获的部分。3.3 藏在config.json里的超参数每个都对应一个设计决策学习Qwen架构时有一个特别好的习惯打开任意一个Qwen2.5模型的config.json把每个字段都查一遍。这里面最关键的几个字段包括architectures、hidden_size、intermediate_size、num_hidden_layers、num_attention_heads、num_key_value_heads、max_position_embeddings、rms_norm_eps、rope_theta、vocab_size等。有一个字段要特别注意rope_theta。它控制RoPE里频率的基数Qwen2系列用的是1000000而LLaMA用的是10000。这个数值越大低维频率的周期越长模型能感知的相对位置范围就越大对长文本场景更友好。Qwen2把上下文拉到3万多和这个theta值有直接关系。理解这些超参数之间的关系是手搓模型的必修课。比如改了hidden_size那么注意力头的维度head_dim也会变KV投影的目标维度也跟着变。改了一个数字整个网络的参数量都会跳动。我建议在动手看代码之前先把config.json里的每个字段抄一遍在旁边标上它会影响哪个模块这样源码阅读会顺利很多。4. 实操用PyTorch从零复现Qwen核心模块4.1 环境准备与依赖选择动手写代码之前先把环境理清楚。我自己用的是PyTorch 2.1以上的版本配合transformers库作为参考实现但核心代码尽量不依赖transformers而是直接用nn.Module自己写。分词器暂时可以用transformers加载为的是把文本快速转成token id跟架构学习的主线不冲突。硬件方面如果只是跑0.5B的模型做验证一张8G显存的显卡足够了。如果没有GPU用CPU跑也不是不行只是生成速度会慢到让人怀疑人生。我建议至少准备一个能跑小模型的GPU环境卡着不动的话太打击学习积极性。主要依赖包列一下torch、transformers、tiktoken。tiktoken是Qwen2之后再用的分词库如果你用Qwen1分词器是另一种实现。装好之后先跑一个魔改版的最小测试确认能从HuggingFace拉取模型配置。4.2 核心模块手写实现与逐行解读下面给出核心模块的精简实现。这里特别声明这段代码是学习用途的简化版去掉了缓存、批量推理等工程优化目的在于把架构逻辑讲清楚。import torch import torch.nn as nn import math class RMSNorm(nn.Module): def __init__(self, hidden_size, eps1e-6): super().__init__() self.weight nn.Parameter(torch.ones(hidden_size)) self.eps eps def forward(self, x): # x: [batch, seq_len, hidden_size] variance x.pow(2).mean(-1, keepdimTrue) x_normed x / torch.sqrt(variance self.eps) return self.weight * x_normed def precompute_rope_cache(seq_len, dim, theta1000000.0, devicecpu): # 维度索引偶数维和奇数维配对不同频率 freqs 1.0 / (theta ** (torch.arange(0, dim, 2, devicedevice).float() / dim)) positions torch.arange(seq_len, devicedevice).float() angles torch.outer(positions, freqs) # [seq_len, dim/2] cos_cache angles.cos() sin_cache angles.sin() return cos_cache, sin_cache def rotate_half(x): x1 x[..., :x.shape[-1] // 2] x2 x[..., x.shape[-1] // 2:] return torch.cat([-x2, x1], dim-1) def apply_rope(x, cos_cache, sin_cache): # x: [batch, heads, seq_len, head_dim] seq_len x.shape[2] cos_c cos_cache[:seq_len].unsqueeze(0).unsqueeze(0) sin_c sin_cache[:seq_len].unsqueeze(0).unsqueeze(0) return x * cos_c rotate_half(x) * sin_c class GroupedQueryAttention(nn.Module): def __init__(self, hidden_size, num_heads, num_kv_heads): super().__init__() self.num_heads num_heads self.num_kv_heads num_kv_heads self.head_dim hidden_size // num_heads self.q_proj nn.Linear(hidden_size, num_heads * self.head_dim, biasFalse) self.k_proj nn.Linear(hidden_size, num_kv_heads * self.head_dim, biasFalse) self.v_proj nn.Linear(hidden_size, num_kv_heads * self.head_dim, biasFalse) self.o_proj nn.Linear(num_heads * self.head_dim, hidden_size, biasFalse) def forward(self, x): batch, seq_len, _ x.shape q self.q_proj(x).view(batch, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k self.k_proj(x).view(batch, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) v self.v_proj(x).view(batch, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) # 将KV头扩展到与Query头数量一致一组扩成多组 k k.repeat_interleave(self.num_heads // self.num_kv_heads, dim1) v v.repeat_interleave(self.num_heads // self.num_kv_heads, dim1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) scores torch.softmax(scores, dim-1) output torch.matmul(scores, v) output output.transpose(1, 2).contiguous().view(batch, seq_len, -1) return self.o_proj(output) class SwiGLUFFN(nn.Module): def __init__(self, hidden_size, intermediate_size): super().__init__() self.gate_proj nn.Linear(hidden_size, intermediate_size, biasFalse) self.up_proj nn.Linear(hidden_size, intermediate_size, biasFalse) self.down_proj nn.Linear(intermediate_size, hidden_size, biasFalse) def forward(self, x): return self.down_proj(nn.functional.silu(self.gate_proj(x)) * self.up_proj(x)) class QwenBlock(nn.Module): def __init__(self, hidden_size, num_heads, num_kv_heads, intermediate_size, eps1e-6): super().__init__() self.attn_norm RMSNorm(hidden_size, eps) self.attn GroupedQueryAttention(hidden_size, num_heads, num_kv_heads) self.ffn_norm RMSNorm(hidden_size, eps) self.ffn SwiGLUFFN(hidden_size, intermediate_size) def forward(self, x): # 残差连接 预归一化 x x self.attn(self.attn_norm(x)) x x self.ffn(self.ffn_norm(x)) return x class MiniQwen(nn.Module): def __init__(self, vocab_size, hidden_size, num_layers, num_heads, num_kv_heads, intermediate_size, max_seq_len2048): super().__init__() self.embed nn.Embedding(vocab_size, hidden_size) self.blocks nn.ModuleList([ QwenBlock(hidden_size, num_heads, num_kv_heads, intermediate_size) for _ in range(num_layers) ]) self.norm RMSNorm(hidden_size) self.lm_head nn.Linear(hidden_size, vocab_size, biasFalse) self.embed.weight self.lm_head.weight # 共享词表权重 precompute_rope_cache(max_seq_len, hidden_size // num_heads) def forward(self, input_ids): x self.embed(input_ids) for block in self.blocks: x block(x) x self.norm(x) logits self.lm_head(x) return logits这套代码里RMSNorm、GQA、SwiGLU、RoPE四个核心组件都实现了QwenBlock里加上残差和预归一化形成完整的基本单元MiniQwen把嵌入、层堆叠、输出层串起来。读代码的时候建议逐个模块对照它的数学定义来理解别只是把代码跑通就算完。4.3 用随机权重验证前向传播如果手里暂时没有正式的Qwen权重可以先用随机权重验证模块逻辑是否正确。随机初始化一个MiniQwen实例输入一段随机的token id观察输出logits的形状是不是[seq_len, vocab_size]loss能不能正常算出来梯度能不能正常反传。这些检查通过了说明整个网络的连接逻辑基本正确。接着再去HuggingFace下载Qwen2.5-0.5B-Instruct的权重把config里的维度提出来然后尝试用自己的MiniQwen加载真实权重。这一步需要特别小心因为PyTorch保存的state_dict里键名必须严格对齐否则加载时会报shape mismatch或missing key。我的做法是写一个简单的映射函数把官方权重里类似model.layers.0.self_attn.q_proj.weight这样的键名映射成我代码里blocks.0.attn.q_proj.weight的格式。等权重加载成功就可以真正和transformers库的官方推理做对比输入同样一段文本看看我手搓的模型输出logits和官方模型输出的差异是否在数值误差范围内。如果两个结果高度一致说明架构理解已经基本到位了。5. 实操过程中常见的坑与排查清单5.1 加载模型时内存溢出手搓模型最容易遇到的问题就是OOM尤其是加载较大模型或者推理长文本时。有一个被很多人忽视的细节RoPE的缓存矩阵会随着序列长度增加而线性增长如果在加载模型时直接把max_position_embeddings设得过大光cos、sin缓存就能吃掉几百兆显存。另一个常见问题是Transformer的中间激活值在推理时虽然不用保存但如果你用PyTorch的默认模式中间变量会不会被保留到计算图里要取决于你是否开了torch.no_grad()。推理阶段务必加上no_grad否则尝试生成100个token时计算图会越积越大最后直接把显存撑爆。排查方法其实很简单用nvidia-smi监控显存曲线发现内存持续增长而不是达到某个稳定峰值那就是计算图没有释放而不是模型太大。5.2 输出一堆乱码或重复字符很多第一次手搓模型的朋友都会遇到这个问题输入今天天气输出啊啊啊啊啊。出现这种情况大概率是分词器出了问题。Qwen的词表很大如果你在加载权重时用了不匹配的tokenizertoken id和文本的对应关系完全错位模型当然只会吐乱码。还有一种情况是采样策略不对。生成时直接用argmax逐token选概率最高的那个容易陷入重复循环用temperature过高输出又会变得语无伦次。新手可以先固定temperature0.7、top_p0.9跑几个例子找感觉再逐步调节。5.3 推理速度慢到怀疑人生纯Python手写的推理循环即便用PyTorch速度也比官方模型慢不少。原因很多没有用torch.compile、没有实现KV Cache、每次生成重复计算了前面所有token的注意力。我在复现时试过给GQA加上KV Cache生成速度提升了三倍以上。具体做法是在注意力层里维护一个缓存字典把每个token计算出的K、V存起来生成下一个token时只计算新token的K、V再和缓存里的旧K、V拼接。代码改动量不大但效果非常明显。这也是为什么官方推理框架速度那么快除了底层算子优化KV Cache是核心中的核心。现象可能原因排查建议OOM暴增的计算图、过长序列加no_grad限制序列长度观察显存曲线输出乱码tokenizer与模型不匹配加载对应版本的tokenizer检查词表ID重复生成sampling参数不合适降低temperature调节top_p检查是否用了no_repeat_ngram_size跑得太慢无KV Cache、无编译优化手写KV Cache考虑torch.compile使用half精度权重加载报错键名不匹配打印官方state_dict键名写映射函数逐一对应5.4 一个高效的学习进阶技巧光读懂Qwen本身的架构还不够我建议在写完MiniQwen之后再去对比一下LLaMA的实现看两个模型家族的代码差异。你会发现RMSNorm、RoPE、GQA这些核心组件几乎一样区别只在于超参数和细节处理。这个对比能帮你建立大模型架构其实是同一套底子的变形的认知以后再学其他模型会轻松很多。另外一个很有效的做法是去读它的分词器源码。Qwen2之后用tiktoken格式分词逻辑是全英文字符串压缩算法和之前常用的BPE实现有很多细节差异。分词是LLM输入的第一道关口搞懂它你对模型的理解会从网络结构延伸到数据链路。最后如果有余力可以试着把MiniQwen的模块结构复刻成ONNX格式用Netron可视化工具看看计算图。看到自己写的代码变成一张清晰的数据流图很多模块间的连接关系会比看代码更直观这种感觉是纯读文档给不了的。学习模型架构最忌讳的就是脑子里懂了手上一跑就废。虽然看起来只是复现几个模块但里面涉及的超参数对齐、矩阵维度推导、权重命名映射每一个环节都会逼着你把所有概念彻底消化一遍。手动跑通几个版本之后再去看任何一个模型的开源实现你都会有一种里面每一个数字我都知道它从哪里来的踏实感。这个系列的第一篇先讲到这下一篇我计划展开KV Cache的完整实现和推理加速的部分继续从零手搓我正在做的那个小模型。
返回列表