Nano-VLLM全代码解析笔记(10)-GemmaRMSNorm和MRoPE

Nano-VLLM全代码解析笔记(10)-GemmaRMSNorm和MRoPE 前言本节参考资料VLLM源码Transformers源码(49 封私信) Qwen3.5 架构最全拆解Linear Attention 源码配图解析、Gated DeltaRule 公式源码逻辑介绍、Full Attention与 MoE 模块算子流程解析 - 知乎我的代码仓库(迭代更新ing)WilliamPockey/Nano-Vllm-Qwen-FitTODO LIST已完成Qwen3.5架构介绍已完成GemmaRMSNorm和MRoPE正在完成GDN线性注意力层待完成多模态支持视觉塔/多模态预处理/多模态条件生成类/权重加载待完成引擎LLMEngine/Sequence/Scheduler/ModelRunner修改待完成MTP支持待完成流式输出与多请求支持Qwen3.5架构图​本节修改的代码layernorm.py更改说明增加了GemmaRMSNorm类和RMSNormGated类前者是Qwen3.5-0.8B使用的对修改后的RMSNorm主要是把yx/RMS(x)​×γ变为了yx/RMS(x)​×(γ1)简单修改即可。而后者是Qwen3.5-0.8B的Gated DeltaNet所需的模块跟正常的RMSNorm区别不大会在介绍这个网络的时候用到他这里的silu区别于activation.py中的silu因为这里的gate和hidden_state不是像之前可以一起得到的class GemmaRMSNorm(nn.Module): RMS normalization for Gemma. difference from the above RMSNorm: 1. x * (1 w) instead of x * w. def __init__( self, hidden_size: int, eps: float 1e-6, ) - None: super().__init__() self.weight nn.Parameter(torch.zeros(hidden_size)) self.eps eps torch.compile def rms_forward( self, x: torch.Tensor, weight: torch.Tensor, ) - torch.Tensor: 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(weight.dtype).mul_(weight) return x.to(orig_dtype) torch.compile def add_rms_forward( self, x: torch.Tensor, residual: torch.Tensor, weight: 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(weight.dtype).mul_(weight) return x.to(orig_dtype), residual def forward( self, x: torch.Tensor, residual: torch.Tensor | None None, ) - torch.Tensor | tuple[torch.Tensor, torch.Tensor]: PyTorch-native implementation equivalent to forward(). weight self.weight.float() 1.0 if residual is None: return self.rms_forward(x, weight) return self.add_rms_forward(x, residual, weight) class RMSNormGated(nn.Module): def __init__( self, hidden_size: int, eps: float 1e-6, **kwargs ) - None: super().__init__() self.weight nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon eps self.activation silu def forward( self, hidden_states: torch.Tensor, gate: torch.Tensor ) - torch.Tensor: input_dtype hidden_states.dtype hidden_states hidden_states.to(torch.float32) variance hidden_states.pow(2).mean(-1, keepdimTrue) # Norm before gate hidden_states hidden_states * torch.rsqrt(variance self.variance_epsilon) hidden_states self.weight * hidden_states.to(input_dtype) hidden_states hidden_states * nn.functional.silu(gate.to(torch.float32)) return hidden_states.to(input_dtype)rotary_embedding.py更改说明Qwen3.5系列采用了部分旋转位置编码(partial RoPE)和MRoPE的策略因此我们需要修改原来的代码主要是改动了apply_rotary_emb函数和RotaryEmbedding类以适配部分旋转位置编码新增了MRotaryEmbedding类以支持MRoPE。这里没有对MRoPE进行缓存主要是考虑到多维的旋转位置编码会导致位置编码的变化。MRoPE介绍【面试高频】M-RoPE 多模态位置编码全解简单而言就是对token的最后一维hidden_state之前每一个地方的向量旋转角度只有位置决定我们现在要把它拆成3块每一块的向量旋转角度由对应的3个维度和维度中的位置决定(T时间\H高度\W宽度)from functools import lru_cache import torch from torch import nn # def apply_rotary_emb( # x: torch.Tensor, # cos: torch.Tensor, # sin: torch.Tensor, # ) - torch.Tensor: # x1, x2 torch.chunk(x.float(), 2, dim-1) # y1 x1 * cos - x2 * sin # y2 x2 * cos x1 * sin # return torch.cat((y1, y2), dim-1).to(x.dtype) #新版本部分位置旋转兼容旧版本全位置旋转 def apply_rotary_emb( x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor ) - torch.Tensor: rotary_dim cos.size(-1) * 2 # 现有 cache/现算路径都给 32 → 64 x_rot, x_pass x[..., :rotary_dim], x[..., rotary_dim:] x1, x2 torch.chunk(x_rot.float(), 2, dim-1) y1 x1 * cos - x2 * sin y2 x2 * cos x1 * sin return torch.cat((y1, y2, x_pass), dim-1).to(x.dtype) class RotaryEmbedding(nn.Module): def __init__( self, head_size: int, rotary_dim: int, max_position_embeddings: int, base: float, ) - None: super().__init__() self.head_size head_size # assert rotary_dim head_size assert rotary_dim head_size and rotary_dim % 2 0 inv_freq 1.0 / (base**(torch.arange(0, rotary_dim, 2, dtypetorch.float) / rotary_dim)) t torch.arange(max_position_embeddings, dtypetorch.float) freqs torch.einsum(i,j - ij, t, inv_freq) # 外积[位置数, 频率数] cos freqs.cos() sin freqs.sin() cache torch.cat((cos, sin), dim-1).unsqueeze_(1) self.register_buffer(cos_sin_cache, cache, persistentFalse) torch.compile def forward( self, positions: torch.Tensor, query: torch.Tensor, key: torch.Tensor, ) - tuple[torch.Tensor, torch.Tensor]: cos_sin self.cos_sin_cache[positions] cos, sin cos_sin.chunk(2, dim-1) query apply_rotary_emb(query, cos, sin) key apply_rotary_emb(key, cos, sin) return query, key class MRotaryEmbedding(nn.Module): def __init__( self, head_size: int, rotary_dim: int, base: float, mrope_section(11, 11, 10) ) - None: super().__init__() self.head_size head_size self.rotary_dim rotary_dim inv_freq 1.0 / (base ** (torch.arange(0, rotary_dim, 2, dtypetorch.float) / rotary_dim)) self.register_buffer(inv_freq, inv_freq, persistentFalse) # 预计算每个频率列取自 T/H/W 哪一路默认 0(T)H 占 slice(1, s1*3, 3)W 占 slice(2, s2*3, 3) sel torch.zeros(rotary_dim // 2, dtypetorch.long) sel[1 : mrope_section[1] * 3 : 3] 1 sel[2 : mrope_section[2] * 3 : 3] 2 self.register_buffer(freq_axis_sel, sel, persistentFalse) def forward( self, positions: torch.Tensor, query: torch.Tensor, key: torch.Tensor ) - tuple[torch.Tensor, torch.Tensor]: # positions: (3, N) freqs: (3, N, 32) freqs positions.float().unsqueeze(-1) * self.inv_freq #freqs.transpose(0, 1): (N, 3, 32) #indexsel (32,) → view → (1, 1, 32) → expand → (N, 1, 32) #freqs_t: (N, 32) freqs_t freqs.transpose(0, 1).gather(1, self.freq_axis_sel.view(1, 1, -1).expand(freqs.size(1), 1, -1)).squeeze(1) #cos: (N, 1, 32) cos freqs_t.cos().unsqueeze_(1) sin freqs_t.sin().unsqueeze_(1) return apply_rotary_emb(query, cos, sin), apply_rotary_emb(key, cos, sin)几个问题1.self.register_buffer(freq_axis_sel, sel, persistentFalse)这里最后的参数是什么意思 ❌ 保存模型时不会保存这个 buffer ❌ 加载模型时不会恢复这个 buffer ❌ model.state_dict() 不包含这个 buffer ✅ model.to(device) 仍然会移动这个 buffer因为存在内存中 但如果设为True前三个则反过来最后一个不变 2.为什么foward前不能加torch.compile 因为torch.compile对gather不友好 3.为什么不将MRoPE进行缓存 这里主要是考虑到输入的不同比如我输入文本、输入文本和图片、输入文本和视频都会导致位置编码变化。 而且图片的输入长宽会变MRoPE就算考虑设置最大长宽也会导致要保存的cos_sin太大 所以缓存不一定会命中所以暂时不考虑等后续优化时再看这里可能把缓存队列开大一些可以支持特定情况下的缓存 4.mrope_section为什么是(11,11,10) 这里是考虑到rotary_dim是32维而MRoPE简单来说就是让32维向量拆分成3个维度每个维度单独使用ROPE 因此这里(11,11,10)是人为设计 5.解释freqs positions.float().unsqueeze(-1) * self.inv_freq freqs_t freqs.gather(0, self.freq_axis_sel.unsqueeze(0).expand(freqs[0].shape)) MRoPE的三个维度是时间高度宽度token也被要求给出自己这三个维度对应的数值 因此freqs - freqs_t可以理解为根据一个 token的每一个 hidden_state 中的每一个位置的值选 T/H/C 中的一个频率 gather介绍 # 如果 dim0 out[i][j][k] input[ index[i][j][k] ][ j ][ k ] # 如果 dim1 out[i][j][k] input[ i ][ index[i][j][k] ][ k ] # 如果 dim2 out[i][j][k] input[ i ][ j ][ index[i][j][k] ] 这个公式看起来可能有点抽象它的具体工作过程是这样的 确定输出位置输出张量 out 的形状和 index 一模一样。我们会遍历 out 中的每一个位置比如 (i, j, k)。 读取索引值看 index 在同样位置 (i, j, k) 上的数值是多少。 替换对应维度的索引 如果 dim0意味着我们要替换的是第一个维度的索引。所以输出位置 (i, j, k) 的值来自于 input 中位置 ( index[i][j][k], j, k ) 的元素。 如果 dim1就替换第二个维度的索引取值位置变为 ( i, index[i][j][k], k )。以此类推。 gather的算法等价于 # freqs_t: [N,32] 输出 freqs_t torch.empty(N, 32) for n in range(N): # 遍历每个token for i in range(32): # 遍历每个频率对应一对hidden axis freq_axis_sel[i] # 0T,1H,2W只看i不看n freqs_t[n,i] positions[axis, n] * inv_freq[i] 因此如果是单模态比如说纯文本让THC这样就退化成了ROPE 6.我的疑惑是对于 hidden_state他选了通道 [0,1,2,0,1,2...,0,1,2]apply 里面是 chunk 拆成了两份那就成了 [0,1,2,0,1,2,0,1,2,0,1,2,0,1,2,0] 和 [1,2,0,1,2,0,1,2,0,1,2,0,1,2,0,1]然后再是旋转问题是比如说第一个分量分别是 0 和 1这都不是一个通道怎么能进行旋转我的意思是 x_rot 里面存的是 [0,1,2,0,1,2] 对应的不同通道的频率但我们 apply_rotary_embed 相当于对不同通道的频率进行计算 1.i 0,1,2,3,…31**频率索引** - freq_axis_sel[i] ∈ {0,1,2}决定**第 i 号频率**用 T/H/W 哪一套位置算出角度。 - 输出cos[...,i]、sin[...,i]这是**第 i 号频率的旋转角度系数**。 2. x_rothidden 向量shape […,64] 物理排布硬规则RoPE 原生约定M‑RoPE 不改 x_rot[..., 0], x_rot[..., 1] ← 配对使用 i0 的cos/sin x_rot[..., 2], x_rot[..., 3] ← 配对使用 i1 的cos/sin x_rot[..., 4], x_rot[..., 5] ← 配对使用 i2 的cos/sin x_rot[..., 6], x_rot[..., 7] ← 配对使用 i3 的cos/sin …… x_rot[...,2*i], x_rot[...,2*i1] ← 配对使用 i号频率的cos/sin 重点 hidden 上**位置 2i、2i1 这一对固定绑定第 i 号频率的 cos/sin**。 这个绑定关系是按数组下标位置硬绑定**不会因为 sel [i]0/1/2 发生任何改变**。 sel[i] 仅仅改变**i 号频率的 cos/sin 值是拿 T 算出来还是 H、还是 W 算出来**。它绝不调换 hidden 配对关系。 举极简小例子F3i0,i1,i2 - freq_axis_sel [0, 1, 2] - i0选 T 位置算角度 → cos0, sin0 - i1选 H 位置算角度 → cos1, sin1 - i2选 W 位置算角度 → cos2, sin2 x_rot [a0,a1, b0,b1, c0,c1] - a0,a1hidden 的 0、1 号位**强制使用 i0 的 (cos0,sin0)**T 角度 - b0,b1hidden 的 2、3 号位**强制使用 i1 的 (cos1,sin1)**H 角度 - c0,c1hidden 的 4、5 号位**强制使用 i2 的 (cos2,sin2)**W 角度本系列文章(待写完修正)[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[9]Nano-VLLM全代码解析笔记(9)-qwen3.5介绍上一篇[9]Nano-VLLM全代码解析笔记(9)-qwen3.5介绍