張量並行線性層: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 那一段。