Nano-VLLM全代码解析笔记(5)-laynorm和attention

发布时间:2026/8/24 23:59:07
Nano-VLLM全代码解析笔记(5)-laynorm和attention 当前笔记顺序Engine-Layers(当前laynorm.py和attention.py)Qwen3的DecoderLayer概览(主要关注稠密架构设计与张量并行)layernorm.py作用解析就是正常的RMSNorm实现不过值得注意的是在qwen3的实现中会多一层RMSNorm的使用​import torch from torch import nn class RMSNorm(nn.Module): def __init__( self, hidden_size: int, eps: float 1e-6, ) - None: super().__init__() self.eps eps self.weight nn.Parameter(torch.ones(hidden_size)) torch.compile def rms_forward( self, x: torch.Tensor, ) - torch.Tensor: #这里先转float32再转回来转float32是为了计算正确转回来是为了显存和速度成本 #因为大模型训练 / 推理时为了节省显存、提升速度几乎都会用 低精度张量比如 float16 或 bfloat16但低精度有个致命问题取值范围太小计算易出错。 orig_dtype x.dtype x x.float() var x.pow(2).mean(dim-1, keepdimTrue) x.mul_(torch.rsqrt(var self.eps)) x x.to(orig_dtype).mul_(self.weight) return x torch.compile def add_rms_forward( self, x: torch.Tensor, residual: torch.Tensor, ) - tuple[torch.Tensor, torch.Tensor]: orig_dtype x.dtype x x.float().add_(residual.float()) residual x.to(orig_dtype) var x.pow(2).mean(dim-1, keepdimTrue) x.mul_(torch.rsqrt(var self.eps)) x x.to(orig_dtype).mul_(self.weight) return x, residual def forward( self, x: torch.Tensor, residual: torch.Tensor | None None, ) - torch.Tensor | tuple[torch.Tensor, torch.Tensor]: if residual is None: return self.rms_forward(x) else: return self.add_rms_forward(x, residual) ​attention.py作用解析实现硬件级别的KV缓存管理和注意力计算两大核心功能。KV_CACHE管理使用TRITON实现注意力计算使用FLASH ATTENION相关库实现同时在问题6中申明了一下模型输入的数据的变换维度这很重要因为这里的处理不同于正常transformerimport torch from torch import nn import triton import triton.language as tl from flash_attn import flash_attn_varlen_func, flash_attn_with_kvcache from nanovllm.utils.context import get_context #这个装饰器介绍看问题1 triton.jit #用 Triton JIT 编译的 kernel用于将 KV tensor也是显存中 存储到 GPU 缓存中。 def store_kvcache_kernel( key_ptr, key_stride, value_ptr, value_stride, k_cache_ptr, v_cache_ptr, slot_mapping_ptr, D: tl.constexpr, #Triton中的编译期常量标记用于标记核函数中必须在编译阶段确定值的参数。其值见下个函数 ): #当前线程块block的 ID并行处理每个 token idx tl.program_id(0) #加载该 token 对应的缓存槽 slot tl.load(slot_mapping_ptr idx) #该位置无效比如padding if slot -1: return key_offsets idx * key_stride tl.arange(0, D) value_offsets idx * value_stride tl.arange(0, D) #加载K/V数据GPU 显存 - GPU 寄存器 key tl.load(key_ptr key_offsets) value tl.load(value_ptr value_offsets) cache_offsets slot * D tl.arange(0, D) #写入缓存GPU 显存 - GPU 寄存器 tl.store(k_cache_ptr cache_offsets, key) tl.store(v_cache_ptr cache_offsets, value) def store_kvcache(key: torch.Tensor, value: torch.Tensor, k_cache: torch.Tensor, v_cache: torch.Tensor, slot_mapping: torch.Tensor): N, num_heads, head_dim key.shape #D是单个token在某一层多头合并后的总维度。 D num_heads * head_dim #最后一维步长为1确保最后一维连续 assert key.stride(-1) 1 and value.stride(-1) 1 #确保头的维度步长正确也是确保连续 assert key.stride(1) head_dim and value.stride(1) head_dim #确保Cache的步长正确 assert k_cache.stride(1) D and v_cache.stride(1) D #numel表示元素总数确保映射表大小匹配 assert slot_mapping.numel() N #[(N,)]是Triton的网格配置grid表示启动N个线程块。 store_kvcache_kernel[(N,)](key, key.stride(0), value, value.stride(0), k_cache, v_cache, slot_mapping, D) class Attention(nn.Module): def __init__( self, num_heads, head_dim, scale, num_kv_heads, ): super().__init__() self.num_heads num_heads self.head_dim head_dim #scale就是注意力计算中Softmax的缩放因子核心作用是避免注意力分数Q・K^T过大导致 Softmax 饱和梯度消失。 self.scale scale self.num_kv_heads num_kv_heads self.k_cache self.v_cache torch.tensor([]) def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor): context get_context() k_cache, v_cache self.k_cache, self.v_cache if k_cache.numel() and v_cache.numel(): store_kvcache(k, v, k_cache, v_cache, context.slot_mapping) if context.is_prefill: if context.block_tables is not None: # prefix cache k, v k_cache, v_cache o flash_attn_varlen_func(q, k, v, max_seqlen_qcontext.max_seqlen_q, cu_seqlens_qcontext.cu_seqlens_q, max_seqlen_kcontext.max_seqlen_k, cu_seqlens_kcontext.cu_seqlens_k, softmax_scaleself.scale, causalTrue, block_tablecontext.block_tables) else: # decode #unsqueeze是因为decode时用上一回生成的单个新token输入需要增加token序列维度 o flash_attn_with_kvcache(q.unsqueeze(1), k_cache, v_cache, cache_seqlenscontext.context_lens, block_tablecontext.block_tables, softmax_scaleself.scale, causalTrue) return o一些问题1.请介绍一下triton.jit triton.jit 是 Triton 框架的核心装饰器用于将Python 编写的函数编译为高性能的 GPU 核函数Kernel替代手动编写 CUDA C 核函数的繁琐过程。 核心特性 自动优化Triton 会自动处理 GPU 线程调度、内存访问优化、寄存器分配等底层细节无需开发者关注 CUDA 网格 / 块Grid/Block的手动配置 跨架构兼容编译后的核函数可在不同代际的 NVIDIA GPU如 Ampere、Hopper上高效运行无需针对不同架构适配 Python 语法友好用 Python 语法编写 GPU 逻辑降低异构编程门槛 动态生成代码支持编译期常量、动态形状等特性兼顾灵活性与性能。 在示例代码中store_kvcache_kernel 被该装饰器修饰后会被编译为 GPU 核函数负责将 K/V 数据写入缓存的核心逻辑。 2.为什么说idx tl.program_id(0)是获取当前核函数的线程ID它后面不是用来算第几个偏移吗难道线程和存储位置是对应的 1tl.program_id(0) 的含义 Triton 核函数的执行模型是「Grid-Program」网格 - 程序 tl.program_id(dim) 获取当前核函数实例在 dim 维度上的索引可理解为「线程 ID」更准确的是「Grid 维度的索引」 示例中 tl.program_id(0) 是一维 Grid 的索引取值范围是 0 ~ N-1因为启动核函数时指定了 [(N, )]即 Grid 大小为 N。 2线程与存储位置的对应关系 示例中核函数的设计逻辑是每个线程Program负责处理一个 Token 的 K/V 数据存储N 个 Token 对应 N 个线程。 idx 是第 idx 个线程对应处理第 idx 个 Token 的 K/V 数据 线程通过 idx 计算该 Token 的 K/V 数据在原始张量中的偏移key_offsets idx * key_stride tl.arange(0,D)再计算该 Token 要存入 Cache 的位置cache_offsets。 简言之线程 IDidx与 Token 索引一一对应而 Token 索引又对应其存储位置的偏移因此线程和存储位置是强绑定的每个线程只处理一个 Token 的存储。 3.解析cache_offsets slot * D tl.arange(0, D) 这是在计算要把数据写进物理显存KV Cache的具体内存地址。 slot物理槽位号可以理解为大楼里的“房间号”。 D一个 Token 的所有 Attention 头加起来的总数据量即 num_heads * head_dim可以理解为“房间的面积”。 slot * D这就走到了要找的那个房间的门口基础偏移量。 tl.arange(0, D)推开门给房间里 $0$ 到 $D-1$ 的每一个地砖内存单元都打上编号。 合在一起就是准确算出这 $D$ 个数据要存进这栋大楼的具体哪些绝对地址中。 4.store_kvcache_kernel[(N, )]的[(N, )]是什么意思 这是 Triton 核函数的启动配置表示 以「一维 Grid」启动核函数Grid 的大小为 N即启动 N 个并行的 Program / 线程 (N, ) 是 tuple 类型对应 Grid 的维度一维若为 (N, M) 则是二维 GridN 行 M 列。 示例中 N 是 Token 数量key.shape[0]启动 N 个线程每个线程处理一个 Token 的 K/V 存储与问题 3 的线程 - 存储位置对应逻辑一致。 5.为什么decode阶段q要unsqueeze unsqueeze(1) 是为了匹配 FlashAttention 对 Decode 阶段输入维度的要求核心是补充「序列长度」维度 1Prefill 与 Decode 阶段的 Q 维度差异 Prefill 阶段预填充处理完整的输入序列Q 的 shape 通常是 [N, num_heads, head_dim]N 是总 Token 数隐含序列长度维度 Decode 阶段逐 token 生成每次只处理一个 Token自回归生成Q 的原始 shape 是 [batch_size, num_heads, head_dim]缺少「序列长度」维度seq_len1。 2flash_attn_with_kvcache 的输入要求 该函数针对 Decode 阶段设计期望 Q 的 shape 包含 seq_len 维度即使 seq_len1即 [batch_size, seq_len, num_heads, head_dim]或简化为 [N, 1, D]。 示例中 q.unsqueeze(1) 是在第 1 维插入 seq_len1让 Q 的 shape 从 [N, num_heads*head_dim] 变为 [N, 1, num_heads*head_dim]匹配函数的输入维度要求确保 KV Cache 能正确对齐计算。 6.区别于transformer的数据维度这里是flash_attention的实现导致的我们需要进行补充不然会导致后面的代码理解错误 模型的输入从来不是 (batch, seq)而是所有序列的 token 拼接成的一维张量 (N,)模型内部自始至终保持 2D 的 (N, hidden_size)哪些 token 属于哪条序列不放在张量形状里而是放在 cu_seqlens / slot_mapping / block_tables / context_lens 这些元数据里。这就是 FlashAttention varlen 格式 分页 KV cache 的标准做法vLLM 也是这么干的。 1. 数据源头引擎层的一维拼接 引擎根本不构造 3D 张量。prepare_prefill 里nanovllm/engine/model_runner.py:129 input_ids.extend(seq[start:end]) —— 把本次 step 里所有序列要处理的 token 顺序拼进一个 list最终 torch.tensor 形状是 (N,)N ∑ 各序列 num_scheduled_tokens 同时用 cu_seqlens_q/k 累计每条序列的边界[0, len1, len1len2, ...]这就是 varlen 格式的目录 positions 也是拼接的但取值是每条序列内部的绝对位置这样旋转位置编码才正确。 decode 阶段model_runner.py:172更简单每条序列只取 seq.last_tokeninput_ids 形状 (num_seqs,)——每条序列一个 token仍然是 1D。 2. 模型内部为什么一直是 (N, hidden) embed_tokens(input_ids)nanovllm/models/qwen3_moe.py:269对 1D 的 id 做 embedding直接得到 (N, hidden_size)——2D没有任何时刻出现 batch/split 维。此后 48 层里所有算子都是逐 token 的 RMSNorm对最后一维求 mean/var(N, hidden) → (N, hidden) 注意力见下节用的是 varlen 内核 MLP / MoE纯逐 token 线性变换。 所以 qwen3_moe.py:149 解包出的两个值其实是sequence_length N本 step 里跨所有序列的 token 总数变量名有误导性见第 6 节、hidden_dim 2048。MoE 路由本来就是逐 token 独立决策每个 token 选出自己的 top-8 专家完全不需要序列边界信息后续 index_add_(0, top_x, ...)qwen3_moe.py:187按 flat 行号把各专家的贡献写回对应 token 行——一维扁平布局恰好就是 index_add_ 想要的形态。 3. 注意力varlen 内核 分页 KV cache q/k/v 在模型侧被 view 成 (N, num_heads, head_dim)qwen3_moe.py:81-83然后 prefillnanovllm/layers/attention.py:64-70flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)。flat 张量 cu_seqlens内核内部按序列边界做 causal mask每条序列只attend自己的历史 token。命中前缀缓存时block_tables is not Nonemodel_runner.py:162k/v 直接换成缓存里的整段张量cu_seqlens_k cu_seqlens_q 表示query 是新的、key 有更多历史。 decodeattention.py:71-74flash_attn_with_kvcache(q.unsqueeze(1), k_cache, v_cache, cache_seqlenscontext_lens, block_tableblock_tables)。q 变成 (bs, 1, heads, dim)内核按每条的 context_lens 从分页缓存里取它自己的历史 k/v 做 attention输出 (bs, heads, dim)。 KV cache 的形状是 (2, num_layers, num_blocks, block_size, num_kv_heads, head_dim)model_runner.py:115每个 token 有固定槽位 slot block编号 × block_size 块内偏移。attention.py:10-30 的 triton kernel 把每步的 k/v 写进这些槽位if slot -1: return 是给 warmup/CUDA-graph 静态布局用的哨兵。这样 decode 时任何 token 都能按 block_table 找回自己之前所有 token 的 KV——这就是分页。 4. 收尾logits 和采样 最后 norm 完还是 (N, hidden)过 lm_head 得 (N, vocab)。但 prefill 阶段我们只需要每条序列最后一个 token 的 logits 来采下一个 token所以 ParallelLMHeadnanovllm/layers/embed_head.py:56-66用 cu_seqlens_q[1:] - 1 把 (N, hidden) 收窄成 (num_seqs, hidden)decode 时 N num_seqs天然就是 (num_seqs, vocab)。Sampler 出 (num_seqs,) 的 token idscheduler.postprocessnanovllm/engine/scheduler.py:81逐条 append_token进入下一轮 decode。 5. 为什么非要用这种布局三个原因 零 padding 浪费序列长度不一如果按 (batch, max_len) 打包短序列要补 pad token多余的矩阵乘法全部白算扁平拼接只对真实 token 付费。 chunked prefill 的天然载体scheduler.schedulescheduler.py:42-46允许一条超长序列被拆成多次 stepnum_scheduled_tokens 截断每次 step 的 flat N 都不同batch的形状每步可变但永远是 1D——扁平 layout 使每次 step 的输入构造变成无脑拼接这正是上游 #218 重构后的设计。 与 flash-attn / 分页 cache 的存储格式配套序列结构信息谁是谁的历史、KV 放哪全在元数据里张量本身不用携带这些维。这也是为什么你能看到 prepare_prefill 里对 slot_mapping、block_tables、context_lens 的精心构造——它们才是真正的batch 结构。本系列文章(待写完修正)[1]Nano-VLLM全代码解析笔记(1)-sequence[2]Nano-VLLM全代码解析笔记(2)-block_manager[3]Nano-VLLM全代码解析笔记(3)-llm_engine和scheduler[4]Nano-VLLM全代码解析笔记(4)-model_runner[5]Nano-VLLM全代码解析笔记(5)-laynorm和attention[6]Nano-VLLM全代码解析笔记(6)-embed_head和linear[7]Nano-VLLM全代码解析笔记(7)-rotary_embedding[8]Nano-VLLM全代码解析笔记(8)-qwen3与qwen3_moe上一篇[4]Nano-VLLM全代码解析笔记(4)-model_runner下一篇[6]Nano-VLLM全代码解析笔记(6)-embed_head和linear