张量并行线性层:ColumnParallelLinear / QKVParallelLinear / RowParallelLinear
职责
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_dim或input_dim做narrow(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 元信息:
ColumnParallelLinear的weight_loader_v2还要识别shard_id(QKV 里是"q" / "k" / "v"),根据 shard id 决定写入 param 的哪一段 offset(linear.py:1023-1047)。 - bias 只在 rank 0 加:
RowParallelLinear.forward里bias_ = 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)。
关键文件
LinearMethodBase:141-180— 量化方法的抽象基类,定义create_weights/apply/process_weights_after_loading三个接口。LinearBase:231-288— 所有线性层的父类,update_param_tp_status把tp_rank/tp_size写到每个BasevLLMParameter上。ColumnParallelLinear.__init__:423-491— 按输出维切,output_size_per_partition = divide(output_size, tp_size)。ColumnParallelLinear weight_loader:520-549— v1 按output_dimnarrow,v2 委托给param.load_column_parallel_weight。ColumnParallelLinear.forward:551-569—quant_method.apply→ 可选tensor_model_parallel_all_gather。QKVParallelLinear.__init__:942-1021— 处理 GQA 的 KV head 复制,把 Q/K/V 三段output_sizes拼成output_size。QKVParallelLinear shard mapping:1023-1047—_get_shard_offset_mapping/_get_shard_size_mapping,告诉 weight_loader 写到 param 的哪一段。QKVParallelLinear fused checkpoint:1049-1097— Phi-3 那种磁盘上已融合的 QKV,自动切分后逐段调weight_loader_v2。RowParallelLinear.__init__:1572-1626— 按输入维切,input_size_per_partition = divide(input_size, tp_size)。RowParallelLinear.forward:1672-1698— 输入若未并行先 split,GEMM 后tensor_model_parallel_all_reduce。MergedColumnParallelLinear:580-636— 多个列并行层融合(QKV 的通用版),output_sizes列表决定每段。
数据流
下面这段是 ColumnParallelLinear.weight_loader 的核心逻辑——拿到磁盘上的完整张量后,按 tp_rank * shard_size 偏移切出本 rank 需要的那一段,再 copy_ 进参数:
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_mismatch在weight_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_loader用reshape(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 那一段。