Nano-VLLM全代码解析笔记(8)-qwen3与qwen3_moe

发布时间:2026/8/27 10:27:47
Nano-VLLM全代码解析笔记(8)-qwen3与qwen3_moe 当前笔记顺序Engine-Layers-Models(当前Qwen3-0.6B与Qwen3-30B-A3B)Qwen3-0.6Bqwen3.pyQwen3-0.6B是稠密架构没有什么好讲的跟transformer的decoder写法差不多定义attention和mlp层用attention和mlp层构建decoderlayer用decoderlayer叠加构建Model再加上vocab_embed和输出头就是完整的qwen3模型。唯二值得关注的点是1.attention层里面还有RMSNorm。2.数据变换的维度不一致见attention.py的问题6import torch from torch import nn import torch.distributed as dist from transformers import Qwen3Config from nanovllm.layers.activation import SiluAndMul from nanovllm.layers.attention import Attention from nanovllm.layers.layernorm import RMSNorm from nanovllm.layers.linear import QKVParallelLinear, MergedColumnParallelLinear, RowParallelLinear from nanovllm.layers.rotary_embedding import get_rope from nanovllm.layers.embed_head import VocabParallelEmbedding, ParallelLMHead class Qwen3Attention(nn.Module): def __init__( self, hidden_size: int, num_heads: int, num_kv_heads: int, max_position: int 4096 * 32, head_dim: int | None None, rms_norm_eps: float 1e-06, qkv_bias: bool False, rope_theta: float 10000, rope_scaling: dict | None None, ) - None: super().__init__() tp_size dist.get_world_size() self.total_num_heads num_heads assert self.total_num_heads % tp_size 0 self.num_heads self.total_num_heads // tp_size self.total_num_kv_heads num_kv_heads assert self.total_num_kv_heads % tp_size 0 self.num_kv_heads self.total_num_kv_heads // tp_size self.head_dim head_dim or hidden_size // self.total_num_heads self.q_size self.num_heads * self.head_dim self.kv_size self.num_kv_heads * self.head_dim self.scaling self.head_dim ** -0.5 self.qkv_bias qkv_bias self.qkv_proj QKVParallelLinear( hidden_size, self.head_dim, self.total_num_heads, self.total_num_kv_heads, biasqkv_bias, ) self.o_proj RowParallelLinear( self.total_num_heads * self.head_dim, hidden_size, biasFalse, ) if isinstance(rope_scaling, dict): rope_theta rope_scaling.get(rope_theta, rope_theta) self.rotary_emb get_rope( self.head_dim, rotary_dimself.head_dim, max_positionmax_position, baserope_theta, ) self.attn Attention( self.num_heads, self.head_dim, self.scaling, self.num_kv_heads, ) if not self.qkv_bias: self.q_norm RMSNorm(self.head_dim, epsrms_norm_eps) self.k_norm RMSNorm(self.head_dim, epsrms_norm_eps) def forward( self, positions: torch.Tensor, hidden_states: torch.Tensor, ) - torch.Tensor: qkv self.qkv_proj(hidden_states) q, k, v qkv.split([self.q_size, self.kv_size, self.kv_size], dim-1) q q.view(-1, self.num_heads, self.head_dim) k k.view(-1, self.num_kv_heads, self.head_dim) v v.view(-1, self.num_kv_heads, self.head_dim) if not self.qkv_bias: q self.q_norm(q) k self.k_norm(k) q, k self.rotary_emb(positions, q, k) o self.attn(q, k, v) output self.o_proj(o.flatten(1, -1)) return output class Qwen3MLP(nn.Module): def __init__( self, hidden_size: int, intermediate_size: int, hidden_act: str, ) - None: super().__init__() self.gate_up_proj MergedColumnParallelLinear( hidden_size, [intermediate_size] * 2, biasFalse, ) self.down_proj RowParallelLinear( intermediate_size, hidden_size, biasFalse, ) assert hidden_act silu self.act_fn SiluAndMul() def forward(self, x): gate_up self.gate_up_proj(x) x self.act_fn(gate_up) x self.down_proj(x) return x class Qwen3DecoderLayer(nn.Module): def __init__( self, config: Qwen3Config, ) - None: super().__init__() self.self_attn Qwen3Attention( hidden_sizeconfig.hidden_size, num_headsconfig.num_attention_heads, num_kv_headsconfig.num_key_value_heads, max_positionconfig.max_position_embeddings, rms_norm_epsconfig.rms_norm_eps, qkv_biasgetattr(config, attention_bias, True), head_dimgetattr(config, head_dim, None), rope_thetagetattr(config, rope_theta, 1000000), rope_scalinggetattr(config, rope_scaling, None), ) self.mlp Qwen3MLP( hidden_sizeconfig.hidden_size, intermediate_sizeconfig.intermediate_size, hidden_actconfig.hidden_act, ) self.input_layernorm RMSNorm(config.hidden_size, epsconfig.rms_norm_eps) self.post_attention_layernorm RMSNorm(config.hidden_size, epsconfig.rms_norm_eps) def forward( self, positions: torch.Tensor, hidden_states: torch.Tensor, residual: torch.Tensor | None, ) - tuple[torch.Tensor, torch.Tensor]: if residual is None: hidden_states, residual self.input_layernorm(hidden_states), hidden_states else: hidden_states, residual self.input_layernorm(hidden_states, residual) hidden_states self.self_attn(positions, hidden_states) hidden_states, residual self.post_attention_layernorm(hidden_states, residual) hidden_states self.mlp(hidden_states) return hidden_states, residual class Qwen3Model(nn.Module): def __init__( self, config: Qwen3Config, ) - None: super().__init__() self.embed_tokens VocabParallelEmbedding(config.vocab_size, config.hidden_size) self.layers nn.ModuleList([Qwen3DecoderLayer(config) for _ in range(config.num_hidden_layers)]) self.norm RMSNorm(config.hidden_size, epsconfig.rms_norm_eps) def forward( self, input_ids: torch.Tensor, positions: torch.Tensor, ) - torch.Tensor: hidden_states self.embed_tokens(input_ids) residual None for layer in self.layers: hidden_states, residual layer(positions, hidden_states, residual) hidden_states, _ self.norm(hidden_states, residual) return hidden_states class Qwen3ForCausalLM(nn.Module): packed_modules_mapping { q_proj: (qkv_proj, q), k_proj: (qkv_proj, k), v_proj: (qkv_proj, v), gate_proj: (gate_up_proj, 0), up_proj: (gate_up_proj, 1), } def __init__( self, config: Qwen3Config ) - None: super().__init__() self.model Qwen3Model(config) self.lm_head ParallelLMHead(config.vocab_size, config.hidden_size) if config.tie_word_embeddings: self.lm_head.weight.data self.model.embed_tokens.weight.data def forward( self, input_ids: torch.Tensor, positions: torch.Tensor, ) - torch.Tensor: return self.model(input_ids, positions) def compute_logits( self, hidden_states: torch.Tensor, ) - torch.Tensor: return self.lm_head(hidden_states)1.貌似这里的模型设计是每次注意力计算前添加位置信息而不是传统transformer那样只在最开头添加位置信息对吗llama2也是这样的设计吗 Qwen3 确实是每层注意力计算前应用位置编码RoPE而非传统 Transformer 仅在 embedding 阶段加一次位置编码Llama2 的设计和 Qwen3 一致也是每层注意力前对 Q/K 应用 RoPE而非开头仅加一次。 RoPE旋转位置编码的核心是对注意力的 Query/Key 做旋转编码而非将位置编码直接加到 embedding 上。Qwen3 的每个 Decoder Layer 的自注意力模块中都会对 Q/K 执行 RoPE而非仅在 embedding 后加一次位置信息。传统 Transformer原始论文是将正弦位置编码直接加到 token embedding 上全程仅加一次。 2.output self.o_proj(o.flatten(1, -1))的变换维度是什么 o.flatten(1, -1) 的作用是将第 1 维也就是 num_heads及之后的所有维度展平。 变换前的维度是[batch_size * seq_len, num_heads, head_dim] 变换后的维度是[batch_size * seq_len, num_heads * head_dim]。 3.一个decoder层里面注意力前一个RMSNorm注意力模块里面一个RMSNorm注意力后面一个RMSNorm对吗为什么这样设计 如果配置中 qkv_bias 为 False确实会在 Attention 内外出现三个 RMSNorminput_layernorm, q_norm, k_norm这不仅是对的而且是较新模型如 Qwen2/3, Gemma, ViT-22B 等的重要改进。 注意力前的 input_layernorm这是标准的 Pre-Norm 架构。为了保证深层网络的梯度稳定归一化放在计算模块之前。 注意力模块内的 q_norm 和 k_norm这是施加在生成 Q 和 K 之后、施加 RoPE 之前的归一化。 为什么这样设计 随着模型参数量和学习率的增大模型在训练时容易出现 Logit Explosion逻辑值爆炸 现象即 Q 和 K 的点积变得异常巨大导致 Attention 权重趋近于 One-Hot 分布梯度消失。对 Q 和 K 进行归一化强制限制了它们的数值范围能够极大地稳定大规模集群上的模型训练。 注你提到的“注意力后面的一个 RMSNorm”实际上是 post_attention_layernorm它在逻辑上是属于下一个模块MLP层的 Pre-Norm。 4.if residual is None: # 第一层residual初始化为原始输入hidden_states归一化 hidden_states, residual self.input_layernorm(hidden_states), hidden_states else: # 非第一层归一化带残差更新hidden_states和residual hidden_states, residual self.input_layernorm(hidden_states, residual)我没看懂这里的逻辑 第一层residual is None 此时输入刚从 Embedding 出来没有前置的残差。代码执行正常的 RMSNorm 得到 hidden_states并把原始输入赋值给 residual 并向后传递。 后续层residual 不为空 在 Attention 计算完毕后传出来的 hidden_states 实际上是增量Attention的输出。 调用 self.input_layernorm(hidden_states, residual) 时内部执行了 x 增量 residual 计算出了真实的当前层输出 更新后的 residual x 保存下来供下一次跨层连接使用 返回 norm(x) 直接进入下一个模块如 MLP 这种设计让“残差相加”和“RMSNorm”在一个 GPU Kernel 内一次性算完大幅提高了运行速度。 5.请结合Linear.py讲解qwen3.py中使用的几个模块的维度是怎么拆分和组合的 A. Attention 部分的拆分与组合 QKVParallelLinear (列并行 - Column Parallel) 作用并行计算 Q、K、V 的投影。 拆分它把输出维度沿着卡切开了。所有的卡收到完全一样的输入 [N, hidden_size]。 维度每张卡独立运算只输出自己分配到的那几个头的 QKV。单卡输出维度为 [N, (local_q_heads 2 * local_kv_heads) * head_dim]。此时无需跨卡通信。 RowParallelLinear (行并行 - Row Parallel) - 对应 o_proj 作用将多卡上计算完毕的局部注意力结果整合回完整的 hidden_size。 拆分由于上一层的列并行现在每张卡上的结果 o 维度是 [N, local_heads * head_dim]。这正好对应了 o_proj 权重被按输入维度切分行切分。 组合每张卡用局部的 o 乘以局部的权重得到维度为 [N, hidden_size] 的部分和Partial Sum。最后通过底层调用的 dist.all_reduce(y) 把所有卡的矩阵加起来得到最终的完整输出。 B. MLP 部分的拆分与组合 MergedColumnParallelLinear (合并列并行) - 对应 gate_up_proj 作用并行计算 MLP 的升维部分Gate 和 Up 投影。 拆分同样是切割输出维度。每张卡收到相同的输入 [N, hidden_size]输出中间层大小的一小部分。单卡输出维度是 [N, 2 * (intermediate_size / TP)]。无通信。 RowParallelLinear (行并行) - 对应 down_proj 作用将 MLP 激活后的结果降维并汇总。 组合每张卡利用局部中间层变量 [N, intermediate_size / TP] 进行线性变换得到 [N, hidden_size] 的部分和再次使用 All-Reduce 进行跨卡求和。对比VLLM运行Qwen3-0.6B硬件单卡4090Nano-VLLMVLLMQwen3-30B-A3BMOE支持相较于之前的模型实现Qwen3-30B-A3B在模型架构上的主要变化是对MLP层进行了修改增加了专家路由。MOE修改参考了GitHub - gogongxt/nano-vllm: Nano vLLM · GitHub根据仓库架构图可知我们需要修改的主要是三个代码文件其中对现有一份代码文件进行了修改并增加了两份代码文件。值得注意的是这里是个粗略地实现性能与VLLM版本完全不能比。性能差异来自1.MoE专家逐个用Python循环跑小GEMM没有 grouped GEMM通用矩阵乘法2.专家没有按tensor parallel切分且每个专家都触发一次all-reduce3.VLLM可以对模型的非MOE部分进行CUDA Graph捕获但Nano-VLLM不支持仓库架构models.py新增作用解析之前model_runner.py导入qwen3-0.6b是直接限定了模型现在增加模型需要添加一个统一的路口from .qwen3 import Qwen3ForCausalLM from .qwen3_moe import Qwen3MoeForCausalLM model_dict { qwen3: Qwen3ForCausalLM, qwen3_moe: Qwen3MoeForCausalLM, }model_runner.py更改说明就是把原来单模型入口改为多模型入口并把默认的torch.dtype修改了一下兼容不同版本transformer其他一样import pickle import torch import torch.distributed as dist from multiprocessing.synchronize import Event from multiprocessing.shared_memory import SharedMemory from nanovllm.config import Config from nanovllm.engine.sequence import Sequence ###from nanovllm.models.qwen3 import Qwen3ForCausalLM #修改模型调用入口 from nanovllm.models.models import model_dict from nanovllm.layers.sampler import Sampler from nanovllm.utils.context import set_context, get_context, reset_context from nanovllm.utils.loader import load_model class ModelRunner: def __init__(self, config: Config, rank: int, event: Event | list[Event]): self.config config hf_config config.hf_config self.block_size config.kvcache_block_size self.enforce_eager config.enforce_eager ##新增 # MoE 动态专家路由不适合CUDA-graph捕获 (python loop index_add_) if hf_config.model_type qwen3_moe: self.enforce_eager True ## self.world_size config.tensor_parallel_size self.rank rank self.event event ##此处增加不同版本适配transformers 4.6x renamed torch_dtype to dtype self.dtype getattr(hf_config, dtype, getattr(hf_config, torch_dtype, torch.float16)) ## dist.init_process_group(nccl, tcp://localhost:2333, world_sizeself.world_size, rankrank) torch.cuda.set_device(rank) default_dtype torch.get_default_dtype() ###torch.set_default_dtype(hf_config.dtype) #适配上面的修改 torch.set_default_dtype(self.dtype) torch.set_default_device(cuda) ###self.model Qwen3ForCausalLM(hf_config) #修改为多模型适配 self.model model_dict[hf_config.model_type](hf_config) load_model(self.model, config.model) self.sampler Sampler() self.warmup_model() self.allocate_kv_cache() if not self.enforce_eager: self.capture_cudagraph() torch.set_default_device(cpu) torch.set_default_dtype(default_dtype) if self.world_size 1: if rank 0: self.shm SharedMemory(namenanovllm, createTrue, size2**20) dist.barrier() else: dist.barrier() self.shm SharedMemory(namenanovllm) self.loop() def exit(self): if self.world_size 1: self.shm.close() dist.barrier() if self.rank 0: self.shm.unlink() if not self.enforce_eager: del self.graphs, self.graph_pool torch.cuda.synchronize() dist.destroy_process_group() def loop(self): while True: method_name, args self.read_shm() self.call(method_name, *args) if method_name exit: break def read_shm(self): assert self.world_size 1 and self.rank 0 self.event.wait() n int.from_bytes(self.shm.buf[0:4], little) method_name, *args pickle.loads(self.shm.buf[4:n4]) self.event.clear() return method_name, args def write_shm(self, method_name, *args): assert self.world_size 1 and self.rank 0 data pickle.dumps([method_name, *args]) n len(data) self.shm.buf[0:4] n.to_bytes(4, little) self.shm.buf[4:n4] data for event in self.event: event.set() def call(self, method_name, *args): if self.world_size 1 and self.rank 0: self.write_shm(method_name, *args) method getattr(self, method_name, None) return method(*args) def warmup_model(self): torch.cuda.empty_cache() torch.cuda.reset_peak_memory_stats() max_num_batched_tokens, max_model_len self.config.max_num_batched_tokens, self.config.max_model_len seq_len min(max_num_batched_tokens, max_model_len) num_seqs min(max_num_batched_tokens // seq_len, self.config.max_num_seqs) seqs [Sequence([0] * seq_len) for _ in range(num_seqs)] for seq in seqs: seq.num_scheduled_tokens seq_len self.run(seqs, True) torch.cuda.empty_cache() def allocate_kv_cache(self): config self.config hf_config config.hf_config free, total torch.cuda.mem_get_info() used total - free peak torch.cuda.memory_stats()[allocated_bytes.all.peak] current torch.cuda.memory_stats()[allocated_bytes.all.current] num_kv_heads hf_config.num_key_value_heads // self.world_size head_dim getattr(hf_config, head_dim, hf_config.hidden_size // hf_config.num_attention_heads) ###block_bytes 2 * hf_config.num_hidden_layers * self.block_size * num_kv_heads * head_dim * hf_config.dtype.itemsize #适配上面的修改 block_bytes 2 * hf_config.num_hidden_layers * self.block_size * num_kv_heads * head_dim * self.dtype.itemsize config.num_kvcache_blocks int(total * config.gpu_memory_utilization - used - peak current) // block_bytes assert config.num_kvcache_blocks 0 self.kv_cache torch.empty(2, hf_config.num_hidden_layers, config.num_kvcache_blocks, self.block_size, num_kv_heads, head_dim) layer_id 0 for module in self.model.modules(): if hasattr(module, k_cache) and hasattr(module, v_cache): module.k_cache self.kv_cache[0, layer_id] module.v_cache self.kv_cache[1, layer_id] layer_id 1 def prepare_block_tables(self, seqs: list[Sequence]): max_len max(len(seq.block_table) for seq in seqs) block_tables [seq.block_table [-1] * (max_len - len(seq.block_table)) for seq in seqs] block_tables torch.tensor(block_tables, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) return block_tables def prepare_prefill(self, seqs: list[Sequence]): input_ids [] positions [] cu_seqlens_q [0] cu_seqlens_k [0] max_seqlen_q 0 max_seqlen_k 0 slot_mapping [] block_tables None for seq in seqs: start seq.num_cached_tokens seqlen_q seq.num_scheduled_tokens end start seqlen_q seqlen_k end input_ids.extend(seq[start:end]) positions.extend(range(start, end)) cu_seqlens_q.append(cu_seqlens_q[-1] seqlen_q) cu_seqlens_k.append(cu_seqlens_k[-1] seqlen_k) max_seqlen_q max(seqlen_q, max_seqlen_q) max_seqlen_k max(seqlen_k, max_seqlen_k) if not seq.block_table: # warmup continue start_block start // self.block_size end_block (end self.block_size - 1) // self.block_size for i in range(start_block, end_block): slot_start seq.block_table[i] * self.block_size if i start_block: slot_start start % self.block_size if i ! end_block - 1: slot_end seq.block_table[i] * self.block_size self.block_size else: slot_end seq.block_table[i] * self.block_size end - i * self.block_size slot_mapping.extend(range(slot_start, slot_end)) if cu_seqlens_k[-1] cu_seqlens_q[-1]: # prefix cache block_tables self.prepare_block_tables(seqs) input_ids torch.tensor(input_ids, dtypetorch.int64, pin_memoryTrue).cuda(non_blockingTrue) positions torch.tensor(positions, dtypetorch.int64, pin_memoryTrue).cuda(non_blockingTrue) cu_seqlens_q torch.tensor(cu_seqlens_q, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) cu_seqlens_k torch.tensor(cu_seqlens_k, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) slot_mapping torch.tensor(slot_mapping, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) set_context(True, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, slot_mapping, None, block_tables) return input_ids, positions def prepare_decode(self, seqs: list[Sequence]): input_ids [] positions [] slot_mapping [] context_lens [] for seq in seqs: input_ids.append(seq.last_token) positions.append(len(seq) - 1) context_lens.append(len(seq)) slot_mapping.append(seq.block_table[-1] * self.block_size seq.last_block_num_tokens - 1) input_ids torch.tensor(input_ids, dtypetorch.int64, pin_memoryTrue).cuda(non_blockingTrue) positions torch.tensor(positions, dtypetorch.int64, pin_memoryTrue).cuda(non_blockingTrue) slot_mapping torch.tensor(slot_mapping, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) context_lens torch.tensor(context_lens, dtypetorch.int32, pin_memoryTrue).cuda(non_blockingTrue) block_tables self.prepare_block_tables(seqs) set_context(False, slot_mappingslot_mapping, context_lenscontext_lens, block_tablesblock_tables) return input_ids, positions def prepare_sample(self, seqs: list[Sequence]): temperatures [seq.temperature for seq in seqs] temperatures torch.tensor(temperatures, dtypetorch.float32, pin_memoryTrue).cuda(non_blockingTrue) return temperatures torch.inference_mode() def run_model(self, input_ids: torch.Tensor, positions: torch.Tensor, is_prefill: bool): if is_prefill or self.enforce_eager or input_ids.size(0) 512: return self.model.compute_logits(self.model(input_ids, positions)) else: bs input_ids.size(0) context get_context() graph self.graphs[next(x for x in self.graph_bs if x bs)] graph_vars self.graph_vars graph_vars[input_ids][:bs] input_ids graph_vars[positions][:bs] positions graph_vars[slot_mapping].fill_(-1) graph_vars[slot_mapping][:bs] context.slot_mapping graph_vars[context_lens].zero_() graph_vars[context_lens][:bs] context.context_lens graph_vars[block_tables][:bs, :context.block_tables.size(1)] context.block_tables graph.replay() return self.model.compute_logits(graph_vars[outputs][:bs]) def run(self, seqs: list[Sequence], is_prefill: bool) - list[int]: input_ids, positions self.prepare_prefill(seqs) if is_prefill else self.prepare_decode(seqs) temperatures self.prepare_sample(seqs) if self.rank 0 else None logits self.run_model(input_ids, positions, is_prefill) token_ids self.sampler(logits, temperatures).tolist() if self.rank 0 else None reset_context() return token_ids torch.inference_mode() def capture_cudagraph(self): config self.config hf_config config.hf_config max_bs min(self.config.max_num_seqs, 512) max_num_blocks (config.max_model_len self.block_size - 1) // self.block_size input_ids torch.zeros(max_bs, dtypetorch.int64) positions torch.zeros(max_bs, dtypetorch.int64) slot_mapping torch.zeros(max_bs, dtypetorch.int32) context_lens torch.zeros(max_bs, dtypetorch.int32) block_tables torch.zeros(max_bs, max_num_blocks, dtypetorch.int32) outputs torch.zeros(max_bs, hf_config.hidden_size) self.graph_bs [1, 2, 4, 8] list(range(16, max_bs 1, 16)) self.graphs {} self.graph_pool None for bs in reversed(self.graph_bs): graph torch.cuda.CUDAGraph() set_context(False, slot_mappingslot_mapping[:bs], context_lenscontext_lens[:bs], block_tablesblock_tables[:bs]) outputs[:bs] self.model(input_ids[:bs], positions[:bs]) # warmup with torch.cuda.graph(graph, self.graph_pool): outputs[:bs] self.model(input_ids[:bs], positions[:bs]) # capture if self.graph_pool is None: self.graph_pool graph.pool() self.graphs[bs] graph torch.cuda.synchronize() reset_context() self.graph_vars dict( input_idsinput_ids, positionspositions, slot_mappingslot_mapping, context_lenscontext_lens, block_tablesblock_tables, outputsoutputs, )qwen3_moe.py更改说明在qwen3.py的基础上除了类名只添加了MOE层并稍微修改了Decoder块的MLP层的代码这里仅展示不同的代码MOE架构概览图来自知乎作者北方的郎MOE与MLP最大的区别就是MOE是拆分MLP后路由到TOP_K个子MLP进行计算​Qwen3MoeSparseMoeBlockclass Qwen3MoeSparseMoeBlock(nn.Module): def __init__( self, config: Qwen3MoeConfig, ) - None: super().__init__() self.hidden_size config.hidden_size #没用到 self.intermediate_size config.intermediate_size self.hidden_act config.hidden_act self.num_experts config.num_experts self.top_k config.num_experts_per_tok # gating #专家做了切分但是gate没有因此每张卡都有完整副本并进行相同计算 self.gate nn.Linear(self.hidden_size, self.num_experts, biasFalse) self.experts nn.ModuleList( [ Qwen3MoeMLP( hidden_sizeconfig.hidden_size, intermediate_sizeconfig.moe_intermediate_size, hidden_actconfig.hidden_act, ) for _ in range(self.num_experts) ] ) def forward(self, hidden_states: torch.Tensor): #sequence_length是当前batch中所有token的数量 #这与Flash_attention实现有关 sequence_length, hidden_dim hidden_states.shape router_logits self.gate(hidden_states) # [seq_len, num_experts] routing_weights F.softmax(router_logits, dim1, dtypetorch.float) # [seq_len, num_experts] routing_weights, selected_experts torch.topk( routing_weights, self.top_k, dim-1 ) #都是[seq_len, top_k] routing_weights / routing_weights.sum(dim-1, keepdimTrue) # we cast back to the input dtype routing_weights routing_weights.to(hidden_states.dtype) #初始化输出形状 [seq_len, hidden_dim]用于累加各专家输出。 final_hidden_states torch.zeros( hidden_states.shape, dtypehidden_states.dtype, devicehidden_states.device, ) #构造专家掩码 #one_hot 形状[seq_len, top_k, num_experts] #permute(2,1,0) 后[num_experts, top_k, seq_len] #expert_mask[e][t][k] 表示第 e 个专家是否被第 t 个 token 的第 k 个选择选中0/1。 expert_mask torch.nn.functional.one_hot( selected_experts, num_classesself.num_experts ).permute(2, 1, 0) #选择所有被选中的专家只要至少被一个选中就行 #expert_mask.sum(dim(-1, -2))对 top_k 和 seq_len 求和得到每个专家被选中的总次数标量。 #greater(..., 0) 得到布尔向量nonzero() 返回被至少一个 token 选中的专家索引列表。 expert_hitted torch.greater(expert_mask.sum(dim(-1, -2)), 0).nonzero() for expert_idx in expert_hitted: expert_idx expert_idx.item() expert_layer self.experts[expert_idx] #expert_mask[expert_idx]形状 [top_k, seq_len]因为 permute 后第一维是专家维度 #squeeze(0) 去掉第一维因为 expert_idx 是标量第一维大小为 1得到 [top_k, seq_len] #idx[N]表示排名0 或 1对应 top-1 或 top-2。top_x[N]表示选中了改专家的token的索引序列中的位置 idx, top_x torch.where(expert_mask[expert_idx].squeeze(0)) #hidden_states[None, top_x] 形状[1, N, hidden_dim]注意Ntoken总数这里就是选出来的token数 #reshape(-1, hidden_dim) → [N, hidden_dim]取出所有需要喂给当前专家的 token 的隐状态。None用于维度拓展等价于unsqueeze()这里这种先unsqueeze再reshape的写法是一种统一接口的写法 current_state hidden_states[None, top_x].reshape(-1, hidden_dim) #这里的None是为了广播 current_hidden_states ( expert_layer(current_state) * routing_weights[top_x, idx, None] ) #index_add_ 在维度 0序列维度上按照 top_x 中的索引将 current_hidden_states 加到 final_hidden_states 对应位置。这里会累加专家的贡献 final_hidden_states.index_add_( 0, top_x, current_hidden_states.to(hidden_states.dtype) ) return final_hidden_statesQwen3MoeDecoderLayer把原先该是MLP层的代码改为了MLP OR MOEclass Qwen3MoeDecoderLayer(nn.Module): def __init__( self, config: Qwen3MoeConfig, layer_idx: int -1, ) - None: super().__init__() self.self_attn Qwen3MoeAttention( hidden_sizeconfig.hidden_size, num_headsconfig.num_attention_heads, num_kv_headsconfig.num_key_value_heads, max_positionconfig.max_position_embeddings, rms_norm_epsconfig.rms_norm_eps, qkv_biasgetattr(config, attention_bias, False), head_dimgetattr(config, head_dim, None), rope_thetagetattr(config, rope_theta, 1000000), rope_scalinggetattr(config, rope_scaling, None), ) ##只有这部分不同 #Qwen3-30B-A3B的decoder_sparse_step1指的是每decoder_sparse_step个层出现一层稀疏层在这里除了指定的MLP层其余都是MOE #关于为什么用了layer_idx not in mlp_only_layers还要右边的判断条件应该是为了实验用途比如关闭某几层的MOE特性 mlp_only_layers getattr(config, mlp_only_layers, []) if (layer_idx not in mlp_only_layers) and ( config.num_experts 0 and (layer_idx 1) % config.decoder_sparse_step 0 ): self.mlp Qwen3MoeSparseMoeBlock(configconfig) else: self.mlp Qwen3MoeMLP( hidden_sizeconfig.hidden_size, intermediate_sizeconfig.intermediate_size, hidden_actconfig.hidden_act, ) ## self.input_layernorm RMSNorm(config.hidden_size, epsconfig.rms_norm_eps) self.post_attention_layernorm RMSNorm(config.hidden_size, epsconfig.rms_norm_eps) def forward( self, positions: torch.Tensor, hidden_states: torch.Tensor, residual: torch.Tensor | None, ) - tuple[torch.Tensor, torch.Tensor]: if residual is None: hidden_states, residual self.input_layernorm(hidden_states), hidden_states else: hidden_states, residual self.input_layernorm(hidden_states, residual) hidden_states self.self_attn(positions, hidden_states) hidden_states, residual self.post_attention_layernorm(hidden_states, residual) hidden_states self.mlp(hidden_states) return hidden_states, residual到这里Nano-VLLM的全部代码就讲解完毕了后面会更新一点我自己的改造敬请期待。本系列文章(待写完修正)[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上一篇[7]Nano-VLLM全代码解析笔记(7)-rotary_embedding