Skip to content

张量并行线性层:ColumnParallelLinear / QKVParallelLinear / RowParallelLinear

源码版本v0.25.1

职责

vLLM 的张量并行 (tensor parallelism, TP) 把单个 nn.Linear 拆成多个 GPU 各算一部分,前后向用 all-gather / all-reduce 拼回完整结果。这一层抽象落在 vllm/model_executor/layers/linear.py 的三个核心类上:ColumnParallelLinear 沿输出维切、RowParallelLinear 沿输入维切、QKVParallelLinear 处理 attention 里 Q/K/V 三段不同的 head 数(尤其是 GQA / MQA)。它们都继承自 LinearBase,真正的 GEMM 调用委托给 quant_method.apply(linear.py:551-569),所以同一套并行逻辑能跑未量化、fp8、AWQ、GPTQ 等各种后端。

每个 TP 线性层实例化时都会算出本 rank 的 tp_rank / tp_size,把 output_size 除以 tp_size 得到 output_size_per_partition(linear.py:437-442),并把这个分区信息写到参数的 output_dim / input_dim 属性上,后续 weight_loader 凭这个属性决定切哪一段。QKVParallelLinear 在此基础上额外处理 GQA 的 KV head 复制——当 tp_size >= total_num_kv_heads 时,每个 rank 持有一整份 KV head 的副本(num_kv_head_replicas = tp_size / total_num_kv_heads),(linear.py:994-1001)。

设计动机

为什么要在 LinearBase 之上引入 weight_loader / weight_loader_v2 两套机制?

  • 切法不同:ColumnParallelLinear 在输出维切(每个 rank 拿 [A_1, A_2, ...] 的一段),RowParallelLinear 在输入维切(A = [A_1; A_2; ...],X = [X_1, X_2, ...]),所以两类的 weight_loader 分别按 output_diminput_dimnarrow(linear.py:530-533)(linear.py:1650-1653)。
  • v2 走参数对象:新版 weight_loader_v2 把张量包装成 BasevLLMParameter,直接调 param.load_column_parallel_weight / param.load_row_parallel_weight,把"切哪段"封装在参数对象内部(linear.py:543-549)(linear.py:1663-1670),量化方法可以选择走 v2(WEIGHT_LOADER_V2_SUPPORTED)或老路。
  • Q/K/V 三段融合:QKVParallelLinear 把 Q/K/V 的权重在输出维拼成一个大矩阵,output_sizes = [q, k, v] 每段单独算 divide(..., tp_size)(linear.py:1003-1008),MergedColumnParallelLinear 是它的通用版,任何几个列并行的层都能融合。
  • GQA KV 复制:QKVParallelLinear.__init__if tp_size >= total_num_kv_heads: self.num_kv_heads = 1; self.num_kv_head_replicas = divide(tp_size, total_num_kv_heads)(linear.py:996-1001),保证每个 rank 都能算完整 KV head,后续 forward 不需要额外同步。
  • Shard 元信息:ColumnParallelLinearweight_loader_v2 还要识别 shard_id(QKV 里是 "q" / "k" / "v"),根据 shard id 决定写入 param 的哪一段 offset(linear.py:1023-1047)。
  • bias 只在 rank 0 加:RowParallelLinear.forwardbias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias(linear.py:1685-1688),reduce 之后再加 bias 会重复,所以只在 rank 0 加。
  • 可关闭 TP:disable_tp=True 时让 tp_rank=0 / tp_size=1,该层就退化成普通 Linear,MoE 的某些非并行 expert 用得上(linear.py:438-439)。

关键文件

数据流

下面这段是 ColumnParallelLinear.weight_loader 的核心逻辑——拿到磁盘上的完整张量后,按 tp_rank * shard_size 偏移切出本 rank 需要的那一段,再 copy_ 进参数:

python
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
    output_dim = getattr(param, "output_dim", None)

    is_sharded_weight = getattr(param, "is_sharded_weight", False)
    use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
    # bitsandbytes loads the weights of the specific portion
    # no need to narrow
    is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit

    param_data = param.data
    if output_dim is not None and not is_sharded_weight:
        shard_size = param_data.shape[output_dim]
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)

    # Special case for loading scales off disk, which often do not
    # have a shape (such as in the case of AutoFP8).
    if len(loaded_weight.shape) == 0:
        loaded_weight = loaded_weight.reshape(1)

    assert param_data.shape == loaded_weight.shape
    param_data.copy_(loaded_weight)

(linear.py:520-541)RowParallelLinear.weight_loader 镜像这段,只是 narrow 的维度从 output_dim 换成 input_dim(linear.py:1650-1653)。这两个函数被 DefaultModelLoader 在加载权重时通过 model.load_weights 间接调用——每个模型实现里都有类似 loaded_weights = param.weight_loader(param, loaded_weight) 的写法。

前向时,ColumnParallelLinear 走 all-gather(gather_output=True 时把每个 rank 的 output_parallel 拼成完整 output),RowParallelLinear 走 all-reduce(reduce_results=True 时把分块 GEMM 结果求和);两者搭配使用时正好可以省掉一次同步:QKVParallelLinear(column)→ attention → o_proj(row)是一次经典的 column-then-row 组合。

边界与失败

  • QKV shard_id 非法:validate_shard_id 限制 QKV 的 shard id 只能是 "q" / "k" / "v"None,其他值直接 raise ValueError(linear.py:1023-1029)。
  • RowParallelLinear reduce_results=False + 加 bias:不 reduce 又在 GEMM 里加 bias 会导致每个 rank 都加了一份,结果错误,代码里 raise ValueError(linear.py:1622-1626)。
  • FP8 块形状不匹配:_maybe_allow_fp8_block_shape_mismatchweight_block_size 不能整除某个 output_partition_sizes 时,放开 allow_fp8_block_shape_mismatch(linear.py:493-516),否则 kernel 会拒绝。
  • bitsandbytes 跳过 narrow:use_bitsandbytes_4bit 时 checkpoint 已经是按 rank 分好的,直接 copy_ 不能 narrow(linear.py:524-527)。
  • scales 是零维标量:AutoFP8 等 checkpoint 把 scale 存成零维张量,weight_loaderreshape(1) 兜底(linear.py:537-538)。
  • input_is_parallel=False 的 row 层:输入未切分时,RowParallelLinear.forward 自己 split_tensor_along_last_dim 取本 rank 那一份(linear.py:1676-1682)。

小结

三个并行线性类构成了 vLLM TP 的"权重切片语言",具体怎么算交给 quant_method。跟 量化层加载 的衔接是:ColumnParallelLinear.__init__ 里调 self.quant_method.create_weights(...),把 weight_loader(v1 或 v2)塞进每个参数;跟 DefaultModelLoader 的衔接是 model.load_weights 里逐张量调 weight_loader,把磁盘上的完整权重切成本 rank 那一段。

对照官方资料:vLLM 文档 · README