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

当前笔记顺序

Engine->Layers(当前embed_head.py和linear.py)

Qwen3的DecoderLayer概览(主要关注稠密架构设计与张量并行):


这两个代码文件的实现可以结合上面的Qwen模型架构图来理解,就能看明白VLLM的并行是怎么实现的

embed_head.py

作用解析:实现词嵌入层和输出头,这里开始实现并行,是本框架的核心之一

python 复制代码
import torch
from torch import nn
import torch.nn.functional as F
import torch.distributed as dist

from nanovllm.utils.context import get_context


class VocabParallelEmbedding(nn.Module):

    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
    ):
        super().__init__()
        #获取GPU编号
        self.tp_rank = dist.get_rank()
        #获取并行数量
        self.tp_size = dist.get_world_size()
        assert num_embeddings % self.tp_size == 0
        self.num_embeddings = num_embeddings
        self.num_embeddings_per_partition = self.num_embeddings // self.tp_size
        #每张卡获取负责的词表对应索引
        #这里选择切分词表而不是token的原因看问题1
        self.vocab_start_idx = self.num_embeddings_per_partition * self.tp_rank
        self.vocab_end_idx = self.vocab_start_idx + self.num_embeddings_per_partition
        self.weight = nn.Parameter(torch.empty(self.num_embeddings_per_partition, embedding_dim))
        self.weight.weight_loader = self.weight_loader

    #将完整权重切分到当前 GPU 的逻辑。
    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
        param_data = param.data
        #每个GPU需要负责嵌入的词数量
        shard_size = param_data.size(0)
        start_idx = self.tp_rank * shard_size
        #从0维,切分start到start+GPU需要负责的词数
        loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
        param_data.copy_(loaded_weight)

    def forward(self, x: torch.Tensor):
        if self.tp_size > 1:
            #注意我们切分的是词表,因此只保留我们能查到的词(词索引介于start和end之间)
            mask = (x >= self.vocab_start_idx) & (x < self.vocab_end_idx)
            #因为嵌入这一操作本质上是查表,因此我们需要将x里的token_id映射0-size的区间
            x = mask * (x - self.vocab_start_idx)
        y = F.embedding(x, self.weight)
        if self.tp_size > 1:
            #这里利用了广播机制
            y = mask.unsqueeze(1) * y
            #这里是求和,通过 all-reduce 把每个 token 在所有 GPU 上查到的那一份加起来,得到完整的 embedding。
            #并且这行会阻塞,直到所有 GPU 都完成
            dist.all_reduce(y)
        return y


#注意继承了上面的类
class ParallelLMHead(VocabParallelEmbedding):

    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
        bias: bool = False,
    ):
        assert not bias
        super().__init__(num_embeddings, embedding_dim)

    def forward(self, x: torch.Tensor):
        context = get_context()
        #此处解析见问题8
        if context.is_prefill:
            #只取每个序列最后一个token的输出
            last_indices = context.cu_seqlens_q[1:] - 1
            x = x[last_indices].contiguous()
        #linear这里的权重使用时是会转置的,所以上面类是y = F.embedding(x, self.weight),这里反过来
        logits = F.linear(x, self.weight)
        if self.tp_size > 1:
            #主进程创建
            all_logits = [torch.empty_like(logits) for _ in range(self.tp_size)] if self.tp_rank == 0 else None
            #logits_0 -> all_logits[0], logits_1 -> all_logits[1], ...
            dist.gather(logits, all_logits, 0)
            #这里按照最后一维进行合并,比如[[1,2],[3,4]] -> [1,2,3,4]
            logits = torch.cat(all_logits, -1) if self.tp_rank == 0 else None
        return logits
python 复制代码
​
1.为什么这里是把词汇表num_embeddings拆分,而不是拆分embedding_dim,还是说这不是词汇表,是模型的输入token

num_embeddings:就是词汇表大小(vocab_size),代表所有可被编码的 token 总数;
embedding_dim:每个 token 对应的向量维度(比如 768、4096)。
拆分逻辑的核心是并行策略的选择:大模型的张量并行(TP)有两种常见拆分方式:

拆分维度	适用场景	特点
词汇表(num_embeddings)	Embedding/LM Head 层	称为「词汇表并行」,每个进程只存部分 token 的 Embedding 权重,解决 Vocab×Dim 的参数量爆炸问题
向量维度(embedding_dim)	Transformer 的 FFN/Attention 层	称为「维度并行」,每个进程只存向量维度的一部分,解决 Dim×Dim 的参数量问题
为什么 Embedding 层选词汇表拆分:

Embedding 层的核心是 "token ID→向量",每个 token 对应唯一的向量,拆分词汇表能让每个进程只负责部分 token 的映射,符合 Embedding 的语义;
如果拆分embedding_dim,每个进程只存向量的一部分,需要额外的all_gather聚合向量维度,通信成本更高;而拆分词汇表仅需一次all_reduce,更高效;
补充:你提到的 "模型的输入 token" 本质是词汇表的索引(比如 token=100 对应词汇表第 100 个词),拆分num_embeddings就是拆分这些 token 的映射关系。

2.为什么self.weight = nn.Parameter(torch.empty(self.num_embeddings_per_partition, embedding_dim))而不是用nn.Embedding()

nn.Embedding()会自动初始化完整的权重矩阵(num_embeddings × embedding_dim),但我们需要的是权重分片(num_embeddings_per_partition × embedding_dim),手动用nn.Parameter可以精准控制权重形状;
我们需要自定义权重加载逻辑(self.weight.weight_loader = self.weight_loader):加载预训练权重时,只截取当前进程的分片,而nn.Embedding()的weight参数是封装的,无法直接绑定自定义加载函数;
ParallelLMHead要复用这个权重做线性变换(F.linear),如果用nn.Embedding(),需要额外提取weight参数,手动用nn.Parameter更直接;


3.我没看懂if context.is_prefill: # prefill阶段(首次处理长序列)这里的逻辑

在大模型推理的 Prefill 阶段,无论用户的 Prompt 有多长,语言模型(LM Head)的任务永远只有一个:预测接下来的那 1 个字。


4.LM Head 和 Embedding 层共享权重是怎么实现的,是logits = F.linear(x, self.weight)吗,但这权重不就反了吗

步骤 1:Embedding 层的逻辑(token→向量)
Embedding 层的前向是:

y = F.embedding(x, self.weight)

self.weight的 shape:[num_embeddings_per_partition, D](D=embedding_dim);
逻辑:根据 token ID(局部)查权重表,得到 shape=[B,S,D] 的向量;
步骤 2:LM Head 的逻辑(向量→logits)
LM Head 的前向是:

logits = F.linear(x, self.weight)

F.linear(x, W)的数学公式:x @ W.T(x 的 shape=[N,D],W 的 shape=[V,D],输出 shape=[N,V]);
这里的self.weight和 Embedding 层是同一个参数,shape=[V_part, D](V_part 是当前进程的词汇表分片大小);
步骤 3:"权重反转" 的本质
大模型中,Embedding 层的权重W_emb(shape=[V,D])和 LM Head 的权重W_lm(shape=[D,V])本应是转置关系,但为了共享权重,直接复用W_emb作为W_lm的权重,让F.linear自动做转置(x @ W_emb.T),等价于:

Embedding:token ID → W_emb[token_id](取行);
LM Head:向量x → x @ W_emb.T(乘权重的转置);
这是大模型的通用优化:无需维护两份权重(W_emb 和 W_lm),减少 50% 的参数量,且不影响模型效果(因为转置是线性变换,模型训练时会适配)。

5.为什么cu_seqlens_q[1:]从1开始而不是0

跟model_runner.py中cu_seqlens_q有关,初始化就这样第 0 个元素永远是 0

6.为什么用gather而不是all_reduce?
嵌入层:需要每个卡都有完整结果(后续计算需要)

输出层:只需要rank0有完整logits(用于采样下一个词)

7.为什么 Embedding 要手动切分,但 LM Head 不用?
Embedding 层:需要手动 mask

F.embedding(x, self.weight)  # x 是 token IDs
输入是 token IDs(离散的整数)
需要从切分后的权重中正确查表
如果 token 在另一个 GPU 的范围内,查到的索引会越界或错误
所以需要手动 mask + 重映射:让每个 GPU 只处理自己范围内的 token
LM Head:矩阵乘法天然就是分片计算

F.linear(x, self.weight)  # x 是 hidden states
输入是 hidden states(连续的向量)
矩阵乘法 x @ weight.T 天然就是分片计算
每个 GPU 的 weight 只有 vocab/tp_size 行
输出自然是 vocab/tp_size 维(只包含当前 GPU 负责的词)
不需要 mask:矩阵乘法自动产生正确结果


8.context.is_prefill是用在刚prefill完,需要decode第一个token的时候嘛?但问题是decode的时候不也是只要最后一个token就好了嘛

# Prefill阶段
hidden = [token0, token1, token2, token3]  # 4个token
# 需要筛选出最后一个: hidden[-1]来计算decode阶段的第一个token

# Decode阶段  
hidden = [token4]  # 只有1个token!
# 不需要筛选,因为它本身就是最后一个
​

linear.py

作用解析:实现线性层的张量并行,一个线性层 y = x @ W^T,W 形状 (out, in)。切分方法是Attention 里 qkv_proj(Column(按输出维度切))+ o_proj(Row(按输入维度切))、MLP 里 gate_up_proj(Column)+ down_proj(Row)

python 复制代码
import torch
from torch import nn
import torch.nn.functional as F
import torch.distributed as dist


def divide(numerator, denominator):
    assert numerator % denominator == 0
    return numerator // denominator


class LinearBase(nn.Module):

    def __init__(
        self,
        input_size: int,
        output_size: int,
        bias: bool = False,
        tp_dim: int | None = None,
    ):
        super().__init__()
        #获取张量切分的维度
        self.tp_dim = tp_dim
        #获取当前显卡编号
        self.tp_rank = dist.get_rank()
        #获取进程总数
        self.tp_size = dist.get_world_size()
        #这里是nn.Linear的固定约定
        #好处是每一行对应一个输出神经元。 权重形状为 (out, in) 时,weight[i] 就是计算第 i 个输出所需的全部权重,是一段连续内存。
        self.weight = nn.Parameter(torch.empty(output_size, input_size))
        #这里和下面的bias都是强制子类实现weight_loader方法
        self.weight.weight_loader = self.weight_loader
        if bias:
            self.bias = nn.Parameter(torch.empty(output_size))
            self.bias.weight_loader = self.weight_loader
        else:
            self.register_parameter("bias", None)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        #强制子类实现线性层
        raise NotImplementedError


#复制线性层,所有TP rank持有相同的完整权重,用于不需要TP切分的场景(如输出层、embedding 之后的首层)。
class ReplicatedLinear(LinearBase):

    def __init__(
        self,
        input_size: int,
        output_size: int,
        bias: bool = False,
    ):
        super().__init__(input_size, output_size, bias)

    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
        param.data.copy_(loaded_weight)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return F.linear(x, self.weight, self.bias)

#一个线性层 y = x @ W^T,W 形状 (out, in)
#按输出维度切(Column):切 W 的第 0 维。每个 rank 拿到完整的输入,算出输出的一部分特征。
class ColumnParallelLinear(LinearBase):

    def __init__(
        self,
        input_size: int,
        output_size: int,
        bias: bool = False,
    ):
        tp_size = dist.get_world_size()
        #注意最后的0就是表示按照输出维度切
        super().__init__(input_size, divide(output_size, tp_size), bias, 0)

    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
        param_data = param.data
        shard_size = param_data.size(self.tp_dim)
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(self.tp_dim, start_idx, shard_size)
        param_data.copy_(loaded_weight)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        #这里没有通信将输出合并,而是推迟到下游
        return F.linear(x, self.weight, self.bias)


#这里的前提是Qwen3的MLP层里的gate矩阵和up矩阵合并成一个矩阵做计算,以减少计算量(kernel 启动次数减半、吞吐更高)
#但必须注意,gate 和 up 是分开加载的,各带一个 shard_id,是先切分gate和up,然后把他们合并
class MergedColumnParallelLinear(ColumnParallelLinear):

    def __init__(
        self,
        input_size: int,
        output_sizes: list[int],
        bias: bool = False,
    ):
        self.output_sizes = output_sizes
        super().__init__(input_size, sum(output_sizes), bias)

    #loaded_shard_id是大模型架构中的某一层的索引id
    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: int):
        param_data = param.data
        shard_offset = sum(self.output_sizes[:loaded_shard_id]) // self.tp_size
        shard_size = self.output_sizes[loaded_shard_id] // self.tp_size
        #这里先narrow后面又copy覆盖的原因是靠narrow匹配维度,不然copy会有问题
        param_data = param_data.narrow(self.tp_dim, shard_offset, shard_size)
        #沿 tp_dim 维度把张量切成 tp_size 份,取出对应的那份
        loaded_weight = loaded_weight.chunk(self.tp_size, self.tp_dim)[self.tp_rank]
        param_data.copy_(loaded_weight)


#跟上一个类同样的动机,反正qkv都是用x做投影,与其做三次,不如合并一个大矩阵一次计算完
#这个类在也是在切分QKV,但必须注意,是先切分再合并向量
class QKVParallelLinear(ColumnParallelLinear):

    def __init__(
        self,
        hidden_size: int,
        head_size: int,
        total_num_heads: int,
        total_num_kv_heads: int | None = None,
        bias: bool = False,
    ):
        tp_size = dist.get_world_size()
        total_num_kv_heads = total_num_kv_heads or total_num_heads
        self.head_size = head_size
        self.num_heads = divide(total_num_heads, tp_size)
        self.num_kv_heads = divide(total_num_kv_heads, tp_size)
        # * 2是因为k+v
        output_size = (total_num_heads + 2 * total_num_kv_heads) * self.head_size
        super().__init__(hidden_size, output_size, bias)

    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: str):
        param_data = param.data
        assert loaded_shard_id in ["q", "k", "v"]
        if loaded_shard_id == "q":
            shard_size = self.num_heads * self.head_size
            shard_offset = 0
        elif loaded_shard_id == "k":
            shard_size = self.num_kv_heads * self.head_size
            shard_offset = self.num_heads * self.head_size
        else:
            shard_size = self.num_kv_heads * self.head_size
            shard_offset = self.num_heads * self.head_size + self.num_kv_heads * self.head_size
        param_data = param_data.narrow(self.tp_dim, shard_offset, shard_size)
        loaded_weight = loaded_weight.chunk(self.tp_size, self.tp_dim)[self.tp_rank]
        param_data.copy_(loaded_weight)


#一个线性层 y = x @ W^T,W 形状 (out, in)
#按输入维度切(Row):切 W 的第 1 维。每个 rank 拿到输入的一部分,算出的是完整输出的部分和,需要 all_reduce 相加。
class RowParallelLinear(LinearBase):

    def __init__(
        self,
        input_size: int,
        output_size: int,
        bias: bool = False,
    ):
        tp_size = dist.get_world_size()
        super().__init__(divide(input_size, tp_size), output_size, bias, 1)

    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
        param_data = param.data
        #ndim表示张量是几维的
        if param_data.ndim == 1:
            param_data.copy_(loaded_weight)
            return
        shard_size = param_data.size(self.tp_dim)
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(self.tp_dim, start_idx, shard_size)
        param_data.copy_(loaded_weight)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        y = F.linear(x, self.weight, self.bias if self.tp_rank == 0 else None)
        if self.tp_size > 1:
            #累加后,阻塞,等待所有显卡累加完
            dist.all_reduce(y)
        return y

一些问题:

python 复制代码
1.loaded_weight = loaded_weight.chunk(self.tp_size, self.tp_dim)[self.tp_rank]是怎么拆分的?我没看懂
tensor.chunk(chunks, dim) 的作用是:将张量 tensor 沿指定维度 dim 拆分成 chunks 个近似等大的子张量(若总长度无法被 chunks 整除,最后一个子张量会稍小),返回这些子张量的列表。

2.重点解析MergedColumnParallelLinear和QKVParallelLinear的逻辑,比如MLP中GATE和UP向量,都是[3072, 1024]他们拼接后是[6144, 1024]还是[3072, 2048]?如果是[6144, 1024],这里MergedColumnParallelLinear对向量的切分不是按行吗?那岂不是又拆分回去了

合并后的权重是 [6144, 1024](沿输出维拼接),而不是 [3072, 2048]。你的疑惑点在于:既然按行(输出维)切分,那合并了岂不是又被拆回去了?答案是:TP 切分确实仍然发生,但切法不是"先拼成完整 [6144, 1024] 再按行连续均分",而是每个 rank 各取 gate 和 up 中"属于自己的一段",再拼在一起。所以合并的收益(一次 GEMM、SiluAndMul 无需通信)没有被破坏。

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

系列(1) 开篇(https://blog.csdn.net/xxx/article/details/xxxxxx)

系列(2) 核心原理(https://blog.csdn.net/xxx/article/details/xxxxxx)

🔗上一篇:系列(1)开篇(https://blog.csdn.net/xxx/article/details/xxxxxx)

🔗下一篇:系列(3)实战演练(https://blog.csdn.net/xxx/article/details/xxxxxx)

相关推荐
JTaoX1 小时前
PyCharm 连接本地虚拟机完整指南
ide·python·pycharm
Hi_Amos1 小时前
记一次 yfinance 源码调试:SOCKS5 代理下 Chart API 正常,历史数据却一直超时
python·k线·yfinance
晴天162 小时前
LLM 与世界模型:从“会说话“到“会理解世界“-Day27
人工智能·机器学习
青少儿编程课堂2 小时前
用图形化编程做一个“少年探险闯关”小游戏:方向键控制、碰撞检测与多关卡串起完整项目
c++·python·算法·bfs·信息学竞赛
雨田言炎2 小时前
STM32专题之内部FLASH详解
笔记·stm32·单片机
liulilittle2 小时前
长上下文的成本结构与「甜点区间」——从推理引擎的物理约束看 200K/256K/400K
c++·人工智能·ai·llm·注意力·qkv
evans在进步2 小时前
Spring Boot 工程化核心详解:Parent、Starter、热部署、事务与多数据源
spring boot·后端·python
这张生成的图像能检测吗2 小时前
即插即用模块 + 改进思路汇总目录
图像处理·人工智能·深度学习·目标检测·机器学习
xiaoxiangsiyan2 小时前
全网IPv6规模化改造实战指南
运维·网络·笔记·云原生·自动化