当前笔记顺序
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的全部代码就讲解完毕了,后面会更新一点我自己的改造,敬请期待。
📚本系列文章(待写完修正)
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