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

当前笔记顺序

Engine->Layers->Models(当前Qwen3-0.6B与Qwen3-30B-A3B)


(Qwen3-0.6B)qwen3.py

Qwen3-0.6B是稠密架构,没有什么好讲的,跟transformer的decoder写法差不多,定义attention和mlp层,用attention和mlp层构建decoderlayer,用decoderlayer叠加构建Model,再加上vocab_embed和输出头就是完整的qwen3模型。唯二值得关注的点是:1.attention层里面还有RMSNorm。2.数据变换的维度不一致(见attention.py的问题6

python 复制代码
import 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,
            bias=qkv_bias,
        )
        self.o_proj = RowParallelLinear(
            self.total_num_heads * self.head_dim,
            hidden_size,
            bias=False,
        )
        if isinstance(rope_scaling, dict):
            rope_theta = rope_scaling.get("rope_theta", rope_theta)
        self.rotary_emb = get_rope(
            self.head_dim,
            rotary_dim=self.head_dim,
            max_position=max_position,
            base=rope_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, eps=rms_norm_eps)
            self.k_norm = RMSNorm(self.head_dim, eps=rms_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,
            bias=False,
        )
        self.down_proj = RowParallelLinear(
            intermediate_size,
            hidden_size,
            bias=False,
        )
        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_size=config.hidden_size,
            num_heads=config.num_attention_heads,
            num_kv_heads=config.num_key_value_heads,
            max_position=config.max_position_embeddings,
            rms_norm_eps=config.rms_norm_eps,
            qkv_bias=getattr(config, 'attention_bias', True),
            head_dim=getattr(config, 'head_dim', None),
            rope_theta=getattr(config, "rope_theta", 1000000),
            rope_scaling=getattr(config, "rope_scaling", None),
        )
        self.mlp = Qwen3MLP(
            hidden_size=config.hidden_size,
            intermediate_size=config.intermediate_size,
            hidden_act=config.hidden_act,
        )
        self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.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, eps=config.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)
python 复制代码
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 内外出现三个 RMSNorm(input_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 进行跨卡求和。

Qwen3-30B-A3B(MOE支持)

相较于之前的模型实现,Qwen3-30B-A3B在模型架构上的主要变化是对MLP层进行了修改,增加了专家路由。MOE修改参考了GitHub - gogongxt/nano-vllm: Nano vLLM · GitHub

根据仓库架构图可知,我们需要修改的主要是三个代码文件,其中对现有一份代码文件进行了修改,并增加了两份代码文件。

仓库架构

models.py(新增)

作用解析:之前model_runner.py导入qwen3-0.6b是直接限定了模型,现在增加模型需要添加一个统一的路口

python 复制代码
from .qwen3 import Qwen3ForCausalLM
from .qwen3_moe import Qwen3MoeForCausalLM

model_dict = {
    "qwen3": Qwen3ForCausalLM,
    "qwen3_moe": Qwen3MoeForCausalLM,
}

model_runner.py

更改说明:就是把原来单模型入口改为多模型入口,并把默认的torch.dtype修改了一下,兼容不同版本transformer,其他一样

python 复制代码
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_size=self.world_size, rank=rank)
        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(name="nanovllm", create=True, size=2**20)
                dist.barrier()
            else:
                dist.barrier()
                self.shm = SharedMemory(name="nanovllm")
                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:n+4])
        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:n+4] = 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, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        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, dtype=torch.int64, pin_memory=True).cuda(non_blocking=True)
        positions = torch.tensor(positions, dtype=torch.int64, pin_memory=True).cuda(non_blocking=True)
        cu_seqlens_q = torch.tensor(cu_seqlens_q, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        cu_seqlens_k = torch.tensor(cu_seqlens_k, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        slot_mapping = torch.tensor(slot_mapping, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        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, dtype=torch.int64, pin_memory=True).cuda(non_blocking=True)
        positions = torch.tensor(positions, dtype=torch.int64, pin_memory=True).cuda(non_blocking=True)
        slot_mapping = torch.tensor(slot_mapping, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        context_lens = torch.tensor(context_lens, dtype=torch.int32, pin_memory=True).cuda(non_blocking=True)
        block_tables = self.prepare_block_tables(seqs)
        set_context(False, slot_mapping=slot_mapping, context_lens=context_lens, block_tables=block_tables)
        return input_ids, positions

    def prepare_sample(self, seqs: list[Sequence]):
        temperatures = [seq.temperature for seq in seqs]
        temperatures = torch.tensor(temperatures, dtype=torch.float32, pin_memory=True).cuda(non_blocking=True)
        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, dtype=torch.int64)
        positions = torch.zeros(max_bs, dtype=torch.int64)
        slot_mapping = torch.zeros(max_bs, dtype=torch.int32)
        context_lens = torch.zeros(max_bs, dtype=torch.int32)
        block_tables = torch.zeros(max_bs, max_num_blocks, dtype=torch.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_mapping=slot_mapping[:bs], context_lens=context_lens[:bs], block_tables=block_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_ids=input_ids,
            positions=positions,
            slot_mapping=slot_mapping,
            context_lens=context_lens,
            block_tables=block_tables,
            outputs=outputs,
        )

qwen3_moe.py

更改说明:在qwen3.py的基础上除了类名,只添加了MOE层,并稍微修改了Decoder块的MLP层的代码,这里仅展示不同的代码

MOE架构概览(图来自知乎作者北方的郎),MOE与MLP最大的区别就是MOE是拆分MLP后路由到TOP_K个子MLP进行计算

Qwen3MoeSparseMoeBlock
python 复制代码
class 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, bias=False)
        self.experts = nn.ModuleList(
            [
                Qwen3MoeMLP(
                    hidden_size=config.hidden_size,
                    intermediate_size=config.moe_intermediate_size,
                    hidden_act=config.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, dim=1, dtype=torch.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, keepdim=True)
        # 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,
            dtype=hidden_states.dtype,
            device=hidden_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_classes=self.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],注意N!=token总数,这里就是选出来的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_states
Qwen3MoeDecoderLayer(把原先该是MLP层的代码,改为了MLP OR MOE)
python 复制代码
class Qwen3MoeDecoderLayer(nn.Module):

    def __init__(
        self,
        config: Qwen3MoeConfig,
        layer_idx: int = -1,
    ) -> None:
        super().__init__()
        self.self_attn = Qwen3MoeAttention(
            hidden_size=config.hidden_size,
            num_heads=config.num_attention_heads,
            num_kv_heads=config.num_key_value_heads,
            max_position=config.max_position_embeddings,
            rms_norm_eps=config.rms_norm_eps,
            qkv_bias=getattr(config, 'attention_bias', False),
            head_dim=getattr(config, 'head_dim', None),
            rope_theta=getattr(config, "rope_theta", 1000000),
            rope_scaling=getattr(config, "rope_scaling", None),
        )
        ##只有这部分不同
        #Qwen3-30B-A3B的decoder_sparse_step=1,指的是每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(config=config)
        else:
            self.mlp = Qwen3MoeMLP(
                hidden_size=config.hidden_size,
                intermediate_size=config.intermediate_size,
                hidden_act=config.hidden_act,
            )
        ##
        self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.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的全部代码就讲解完毕了,后面会更新一点我自己的改造,敬请期待。

📚本系列文章(待写完修正)

1Nano-VLLM全代码解析笔记(1)-sequence

2Nano-VLLM全代码解析笔记(2)-block_manager

3Nano-VLLM全代码解析笔记(3)-llm_engine和scheduler

4Nano-VLLM全代码解析笔记(4)-model_runner

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

6Nano-VLLM全代码解析笔记(6)-embed_head和linear

7Nano-VLLM全代码解析笔记(7)-rotary_embedding

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

🔗上一篇:7Nano-VLLM全代码解析笔记(7)-rotary_embedding

相关推荐
智购科技无人售货机厂家1 小时前
2026自动售货机远程运维平台设计:从设备诊断到预测性维护的工程实践~YH
运维·python·物联网·架构·django·scikit-learn
JacksonMx1 小时前
Java 线程池:复用、Spring 管理、监控与线上故障排查全指南
开发语言·python
Lab_AI1 小时前
从代理到自研,从工具到平台:创腾科技的AI for Science破局之路
人工智能·ai·ai for science·ai4s·ai+材料创新·ai+药物发现·ai+药物研发
Zach_菠萝侠1 小时前
【deepseek harness研究】进化方向7:分布式与远程执行 思考、设计与实现
分布式·深度学习·deepseek
脉动数据行情1 小时前
Python WebSocket 实现融通金实时行情监听
开发语言·python·websocket
only-qi1 小时前
Python Agent 开发速通清单
开发语言·python
风云1 小时前
从实战到生产:Acl.Excel 八大场景与避坑指南(终篇)
性能优化·实战·最佳实践·nuget·避坑·acl.excel
AI服务老曹1 小时前
人流量统计线配置性能优化指南:从方向判定到资源调优实战
性能优化
Wang's Blog1 小时前
PostgreSQL笔记64:分区插件生态与扩展查询协议深度解析
数据库·笔记·postgresql