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

发布时间:2026/8/27 10:27:43

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

相关新闻

PLC编码器测速:中心差分法与自适应滤波算法解决低速跳变难题

PLC编码器测速:中心差分法与自适应滤波算法解决低速跳变难题

2026/8/27 10:27:43

1. 项目缘起:从“跳变”到“稳定”的测速挑战 在工业自动化现场,尤其是伺服电机、主轴驱动这类对速度反馈精度和实时性要求极高的场景,编码器测速的稳定性是控制系统能否平稳运行的基石。然而,任何一个在现场摸爬滚打过的工程师&a…

Node系列 · ORM:MD5 加密

Node系列 · ORM:MD5 加密

2026/8/27 10:27:43

Node系列 ORM:MD5 加密MD5 是 Node 后端最早接触的"加密"工具——给密码做哈希、给文件生成指纹。但 MD5 在 2004 年已被攻破,不再适合用于安全场景。本章讲清楚 MD5 的本质、它在哪些场景能用、哪些场景必须换方案。一、MD5 是什么 MD5&…

三步装好Notepad++ Markdown实时预览:MarkdownViewer++插件完整指南

三步装好Notepad++ Markdown实时预览:MarkdownViewer++插件完整指南

2026/8/27 10:17:43

三步装好Notepad Markdown实时预览:MarkdownViewer插件完整指南 【免费下载链接】MarkdownViewerPlusPlus A Notepad Plugin to view a Markdown file rendered on-the-fly 项目地址: https://gitcode.com/gh_mirrors/ma/MarkdownViewerPlusPlus 写完表格的最…

从物理动力学到策略优化:自行车运动员能量建模与MATLAB实现

从物理动力学到策略优化:自行车运动员能量建模与MATLAB实现

2026/8/27 11:27:56

1. 从一道赛题到一套方法论:自行车运动员能量特征建模的深度复盘去年带学生备赛美赛,A题“自行车运动员的能量特征”让不少队伍直呼“物理劝退”。题目本身并不复杂,核心是建立一个数学模型,描述运动员在给定功率输出下&#xff0…

三分钟让 Axure RP 全变中文:axure-cn 中文语言包安装、验收与避坑完整指南

三分钟让 Axure RP 全变中文:axure-cn 中文语言包安装、验收与避坑完整指南

2026/8/27 11:27:56

三分钟让 Axure RP 全变中文:axure-cn 中文语言包安装、验收与避坑完整指南 【免费下载链接】axure-cn Chinese language file for Axure RP. Axure RP 简体中文语言包。支持 Axure 11、10、9。不定期更新。 项目地址: https://gitcode.com/gh_mirrors/ax/axure-c…

ncmdump 使用教程:拖一个 NCM 文件就能转出 MP3,单首或整张专辑都行

ncmdump 使用教程:拖一个 NCM 文件就能转出 MP3,单首或整张专辑都行

2026/8/27 11:27:56

ncmdump 使用教程:拖一个 NCM 文件就能转出 MP3,单首或整张专辑都行 【免费下载链接】ncmdump 项目地址: https://gitcode.com/gh_mirrors/ncmd/ncmdump 你从网易云音乐下载的歌曲,不少是 NCM 格式的——这是网易的加密音频格式&…

Go语言方法值 vs 方法表达式:区别、应用与最佳实践

Go语言方法值 vs 方法表达式:区别、应用与最佳实践

2026/8/27 11:27:56

相信绝大多数 Go 开发者都写过这样的代码:把一个方法赋值给一个变量,然后像普通函数一样调用它。比如 f : obj.Method ,这个 f 在 Go 里叫 方法值(method value) 。但很多人第一次看到 T.Method(obj) 这种写法…

模型评测标准化:告别“跑几个样例”的不确定性

模型评测标准化:告别“跑几个样例”的不确定性

2026/8/27 11:27:56

你在两个模型之间做选型,用同一个测试集各跑了一遍。第一天模型 A 领先,第二天换了个提示词模板,模型 B 反超。你准备把结果写进汇报,但心里清楚,这个结论大概率经不起复测。这个场景很多做过模型评测的人都经历过。问…

信息几何视角下的GFlowNets前向策略训练:自然梯度提升收敛稳定性

信息几何视角下的GFlowNets前向策略训练:自然梯度提升收敛稳定性

2026/8/27 11:17:56

这次我们看一个 GFlowNets 训练方向的新思路:Information-Geometric Forward Policy Training。它要解决的是生成流网络(Generative Flow Networks)中前向策略(forward policy)训练不稳定、分布偏移明显、在复杂离散图…

[光学原理与应用-521]:对光的错误理解与纠偏

[光学原理与应用-521]:对光的错误理解与纠偏

2026/8/27 11:10:02

首先光是一种能量的载体和形态,宏观上观察到的光是由无数个微观的光量子组成的,每个光子在产生的瞬间,其在真空的空间中以确定不变的速度沿着一个初始的方向一直向前,在微观层面,每个光量子的运动轨迹是以波函数所展现…

SIP通话转接原理与REFER方法实战解析

SIP通话转接原理与REFER方法实战解析

2026/8/27 7:25:23

1. 通话转接不是“挂断再拨号”,而是SIP会话的动态重定向你有没有遇到过这样的场景:客服坐席A正在和客户通电话,突然需要把这通对话无缝转给专家坐席B,客户完全感知不到中间的断连——既没听到忙音,也没被要求重新拨号…

Kolla-ansible单节点OpenStack部署实战:从环境准备到排坑指南

Kolla-ansible单节点OpenStack部署实战:从环境准备到排坑指南

2026/8/26 17:50:58

1. 为什么选择Kolla-ansible来部署单节点OpenStack?如果你正在寻找一种能把OpenStack从“概念”快速变成“可用的实验环境”的方法,那么Kolla-ansible几乎是当前最主流、最省心的选择。我见过太多人卡在手动编译依赖、配置服务、处理版本冲突的泥潭里&am…

Go语言构建企业级AI服务网关:统一管理英伟达等AI接口调用

Go语言构建企业级AI服务网关:统一管理英伟达等AI接口调用

2026/8/27 0:07:12

1. 项目概述:从零构建一个企业级的AI服务网关 最近在帮一个做内容审核的团队做技术架构升级,他们原来的业务里,每天有几十万张图片和短视频需要过审,最初是接了几个开源的AI模型自己部署,但效果和性能一直不太稳定。后…

LeetCode Hot100(51-60)算法精解与面试技巧

LeetCode Hot100(51-60)算法精解与面试技巧

2026/8/27 0:07:12

1. 题目背景与核心价值"hot100(51-60)"这个标题看起来像是某个编程题库或算法练习集中的一组题目编号。在技术社区中,类似命名通常指向LeetCode、牛客网等平台的热门题目集合。作为刷过300题的算法老手,我理解这类题目的核心价值在于&#xff…

CRC校验实战:从模2除法到HJ212协议排错

CRC校验实战:从模2除法到HJ212协议排错

2026/8/27 0:07:12

1. 为什么一个“校验码”能扛住工业现场90%的数据 corruption? 你有没有遇到过这样的场景:嵌入式设备通过RS-485上传温湿度数据,上位机偶尔收到一帧乱码——温度显示成-273℃,湿度跳到999%,但串口波形看起来完全正常&a…

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

2026/8/22 2:02:26

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…

导师推荐!2026最新AI论文工具测评与实用推荐

导师推荐!2026最新AI论文工具测评与实用推荐

2026/8/26 18:07:30

2026年真正好用的AI论文工具,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

告别游戏崩溃:XCOM 2模组管理器的智能革命

告别游戏崩溃:XCOM 2模组管理器的智能革命

2026/8/26 17:57:52

告别游戏崩溃:XCOM 2模组管理器的智能革命 【免费下载链接】xcom2-launcher The Alternative Mod Launcher (AML) is a replacement for the default game launchers from XCOM 2 and XCOM Chimera Squad. 项目地址: https://gitcode.com/gh_mirrors/xc/xcom2-lau…